import torch import sentencepiece as spm from pathlib import Path from model import Translation, TranslationConfig BASE_DIR = Path(__file__).resolve().parent sp = spm.SentencePieceProcessor() sp.load(str(BASE_DIR / "spm.model")) BOS = sp.bos_id() EOS = sp.eos_id() PAD = 0 MAX_LEN = 32 checkpoint = torch.load( BASE_DIR / "translator_checkpoint.pth", map_location="cpu", weights_only=True, ) config = TranslationConfig() config.vocab_size = 8000 model = Translation(config) model.load_state_dict( checkpoint["model_state_dict"] ) model.eval() def get_model_parameters(): parameters = model.parameter_summary() total_parameters = sum(item["num_parameters"] for item in parameters) return { "total_parameters": total_parameters, "parameters": parameters, } @torch.no_grad() def translate(sentence, max_new_tokens=32): # Encode source sentence src = sp.encode(sentence)[:MAX_LEN - 1] + [EOS] # Pad src += [PAD] * (MAX_LEN - len(src)) # Tensor src = torch.tensor(src).unsqueeze(0) # Generate translation output_tokens = model.generate( src, max_new_tokens=max_new_tokens ) # Convert tensor to list output_tokens = output_tokens[0].tolist() # Remove BOS if output_tokens[0] == BOS: output_tokens = output_tokens[1:] # Stop at EOS if EOS in output_tokens: output_tokens = output_tokens[ :output_tokens.index(EOS) ] # Decode translated_text = sp.decode(output_tokens) return translated_text if __name__ == "__main__": for item in get_model_parameters()["parameters"]: print( f'{item["name"]}: shape={item["shape"]}, ' f'parameters={item["num_parameters"]}' ) print(translate("How are you?"))