(data_iterator, args, timers)
| 48 | |
| 49 | from transformers import DataCollatorForSeq2Seq |
| 50 | def get_batch(data_iterator, args, timers): |
| 51 | # Items and their type. |
| 52 | keys = ['input_ids', 'labels'] |
| 53 | datatype = torch.int64 |
| 54 | |
| 55 | # Broadcast data. |
| 56 | timers('data loader').start() |
| 57 | if data_iterator is not None: |
| 58 | data = next(data_iterator) |
| 59 | else: |
| 60 | data = None |
| 61 | timers('data loader').stop() |
| 62 | data_b = mpu.broadcast_data(keys, data, datatype) |
| 63 | # Unpack. |
| 64 | tokens = data_b['input_ids'].long() |
| 65 | labels = data_b['labels'].long() |
| 66 | |
| 67 | return tokens, labels |
| 68 | |
| 69 | |
| 70 | from torch.nn import CrossEntropyLoss |
no test coverage detected