(opts)
| 17 | # Get data, tokenizer, text transform, vocab objs, etc. Everything we |
| 18 | # need to start training the model |
| 19 | def get_data(opts): |
| 20 | |
| 21 | src_lang = opts.src |
| 22 | tgt_lang = opts.tgt |
| 23 | |
| 24 | multi30k.URL["train"] = "https://raw.githubusercontent.com/neychev/small_DL_repo/master/datasets/Multi30k/training.tar.gz" |
| 25 | multi30k.URL["valid"] = "https://raw.githubusercontent.com/neychev/small_DL_repo/master/datasets/Multi30k/validation.tar.gz" |
| 26 | |
| 27 | # Define a token "unkown", "padding", "beginning of sentence", and "end of sentence" |
| 28 | special_symbols = { |
| 29 | "<unk>":0, |
| 30 | "<pad>":1, |
| 31 | "<bos>":2, |
| 32 | "<eos>":3 |
| 33 | } |
| 34 | |
| 35 | # Get training examples from torchtext (the multi30k dataset) |
| 36 | train_iterator = Multi30k(split="train", language_pair=(src_lang, tgt_lang)) |
| 37 | valid_iterator = Multi30k(split="valid", language_pair=(src_lang, tgt_lang)) |
| 38 | |
| 39 | # Grab a tokenizer for these languages |
| 40 | src_tokenizer = get_tokenizer("spacy", src_lang) |
| 41 | tgt_tokenizer = get_tokenizer("spacy", tgt_lang) |
| 42 | |
| 43 | # Build a vocabulary object for these languages |
| 44 | src_vocab = build_vocab_from_iterator( |
| 45 | _yield_tokens(train_iterator, src_tokenizer, True), |
| 46 | min_freq=1, |
| 47 | specials=list(special_symbols.keys()), |
| 48 | special_first=True |
| 49 | ) |
| 50 | |
| 51 | tgt_vocab = build_vocab_from_iterator( |
| 52 | _yield_tokens(train_iterator, tgt_tokenizer, False), |
| 53 | min_freq=1, |
| 54 | specials=list(special_symbols.keys()), |
| 55 | special_first=True |
| 56 | ) |
| 57 | |
| 58 | src_vocab.set_default_index(special_symbols["<unk>"]) |
| 59 | tgt_vocab.set_default_index(special_symbols["<unk>"]) |
| 60 | |
| 61 | # Helper function to sequentially apply transformations |
| 62 | def _seq_transform(*transforms): |
| 63 | def func(txt_input): |
| 64 | for transform in transforms: |
| 65 | txt_input = transform(txt_input) |
| 66 | return txt_input |
| 67 | return func |
| 68 | |
| 69 | # Function to add BOS/EOS and create tensor for input sequence indices |
| 70 | def _tensor_transform(token_ids): |
| 71 | return torch.cat( |
| 72 | (torch.tensor([special_symbols["<bos>"]]), |
| 73 | torch.tensor(token_ids), |
| 74 | torch.tensor([special_symbols["<eos>"]])) |
| 75 | ) |
| 76 |
no test coverage detected