(batch)
| 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) |
nothing calls this directly
no outgoing calls
no test coverage detected