MCPcopy Create free account
hub / github.com/pytorch/tutorials / get_dataloader

Function get_dataloader

intermediate_source/seq2seq_translation_tutorial.py:562–582  ·  view source on GitHub ↗
(batch_size)

Source from the content-addressed store, hash-verified

560 return (input_tensor, target_tensor)
561
562def get_dataloader(batch_size):
563 input_lang, output_lang, pairs = prepareData('eng', 'fra', True)
564
565 n = len(pairs)
566 input_ids = np.zeros((n, MAX_LENGTH), dtype=np.int32)
567 target_ids = np.zeros((n, MAX_LENGTH), dtype=np.int32)
568
569 for idx, (inp, tgt) in enumerate(pairs):
570 inp_ids = indexesFromSentence(input_lang, inp)
571 tgt_ids = indexesFromSentence(output_lang, tgt)
572 inp_ids.append(EOS_token)
573 tgt_ids.append(EOS_token)
574 input_ids[idx, :len(inp_ids)] = inp_ids
575 target_ids[idx, :len(tgt_ids)] = tgt_ids
576
577 train_data = TensorDataset(torch.LongTensor(input_ids).to(device),
578 torch.LongTensor(target_ids).to(device))
579
580 train_sampler = RandomSampler(train_data)
581 train_dataloader = DataLoader(train_data, sampler=train_sampler, batch_size=batch_size)
582 return input_lang, output_lang, train_dataloader
583
584
585######################################################################

Callers 1

Calls 2

prepareDataFunction · 0.85
indexesFromSentenceFunction · 0.70

Tested by

no test coverage detected