MCPcopy Create free account
hub / github.com/pytorch/examples / get_data

Function get_data

language_translation/src/data.py:19–96  ·  view source on GitHub ↗
(opts)

Source from the content-addressed store, hash-verified

17# Get data, tokenizer, text transform, vocab objs, etc. Everything we
18# need to start training the model
19def 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

Callers 3

inferenceFunction · 0.90
mainFunction · 0.90
data.pyFile · 0.85

Calls 2

_yield_tokensFunction · 0.85
_seq_transformFunction · 0.85

Tested by

no test coverage detected