(opts)
| 51 | |
| 52 | # Opens an user interface where users can translate an arbitrary sentence |
| 53 | def inference(opts): |
| 54 | |
| 55 | # Get training data, tokenizer and vocab |
| 56 | # objects as well as any special symbols we added to our dataset |
| 57 | _, _, src_vocab, tgt_vocab, src_transform, _, special_symbols = get_data(opts) |
| 58 | |
| 59 | src_vocab_size = len(src_vocab) |
| 60 | tgt_vocab_size = len(tgt_vocab) |
| 61 | |
| 62 | # Create model |
| 63 | model = Translator( |
| 64 | num_encoder_layers=opts.enc_layers, |
| 65 | num_decoder_layers=opts.dec_layers, |
| 66 | embed_size=opts.embed_size, |
| 67 | num_heads=opts.attn_heads, |
| 68 | src_vocab_size=src_vocab_size, |
| 69 | tgt_vocab_size=tgt_vocab_size, |
| 70 | dim_feedforward=opts.dim_feedforward, |
| 71 | dropout=opts.dropout |
| 72 | ).to(DEVICE) |
| 73 | |
| 74 | # Load in weights |
| 75 | model.load_state_dict(torch.load(opts.model_path)) |
| 76 | |
| 77 | # Set to inference |
| 78 | model.eval() |
| 79 | |
| 80 | # Accept input and keep translating until they quit |
| 81 | while True: |
| 82 | print("> ", end="") |
| 83 | |
| 84 | sentence = input() |
| 85 | |
| 86 | # Convert to tokens |
| 87 | src = src_transform(sentence).view(-1, 1) |
| 88 | num_tokens = src.shape[0] |
| 89 | |
| 90 | src_mask = (torch.zeros(num_tokens, num_tokens)).type(torch.bool) |
| 91 | |
| 92 | # Decode |
| 93 | tgt_tokens = greedy_decode( |
| 94 | model, src, src_mask, max_len=num_tokens+5, start_symbol=special_symbols["<bos>"], end_symbol=special_symbols["<eos>"] |
| 95 | ).flatten() |
| 96 | |
| 97 | # Convert to list of tokens |
| 98 | output_as_list = list(tgt_tokens.cpu().numpy()) |
| 99 | |
| 100 | # Convert tokens to words |
| 101 | output_list_words = tgt_vocab.lookup_tokens(output_as_list) |
| 102 | |
| 103 | # Remove special tokens and convert to string |
| 104 | translation = " ".join(output_list_words).replace("<bos>", "").replace("<eos>", "") |
| 105 | |
| 106 | print(translation) |
| 107 | |
| 108 | # Train the model for 1 epoch |
| 109 | def train(model, train_dl, loss_fn, optim, special_symbols, opts): |
no test coverage detected