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

Function _collate_fn

language_translation/src/data.py:82–90  ·  view source on GitHub ↗
(batch)

Source from the content-addressed store, hash-verified

80 # Now we want to convert the torchtext data pipeline to a dataloader. We
81 # will need to collate batches
82 def _collate_fn(batch):
83 src_batch, tgt_batch = [], []
84 for src_sample, tgt_sample in batch:
85 src_batch.append(src_lang_transform(src_sample.rstrip("\n")))
86 tgt_batch.append(tgt_lang_transform(tgt_sample.rstrip("\n")))
87
88 src_batch = pad_sequence(src_batch, padding_value=special_symbols["<pad>"])
89 tgt_batch = pad_sequence(tgt_batch, padding_value=special_symbols["<pad>"])
90 return src_batch, tgt_batch
91
92 # Create the dataloader
93 train_dataloader = DataLoader(train_iterator, batch_size=opts.batch, collate_fn=_collate_fn)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected