Generate a batch.
(data_iterator)
| 163 | |
| 164 | |
| 165 | def get_batch(data_iterator): |
| 166 | """Generate a batch.""" |
| 167 | |
| 168 | # TODO: this is pretty hacky, find a better way |
| 169 | if (not mpu.is_pipeline_first_stage()) and (not mpu.is_pipeline_last_stage()): |
| 170 | return None, None, None, None, None |
| 171 | |
| 172 | args = get_args() |
| 173 | |
| 174 | if "-Raw" in args.dataset: |
| 175 | # get batches based on the TP rank you are on |
| 176 | batch = get_batch_on_this_tp_rank_original(data_iterator) |
| 177 | # slice batch along sequence dimension for context parallelism |
| 178 | batch = get_batch_on_this_cp_rank(batch) |
| 179 | |
| 180 | elif "-Idxmap" in args.dataset: |
| 181 | # get batches based on the TP rank you are on |
| 182 | batch = get_batch_on_this_tp_rank(data_iterator) |
| 183 | # slice batch along sequence dimension for context parallelism |
| 184 | batch = get_batch_on_this_cp_rank(batch) |
| 185 | |
| 186 | else: |
| 187 | raise ValueError("please set correct --dataset ") |
| 188 | |
| 189 | return batch.values() |
| 190 | |
| 191 | |
| 192 | def loss_func(loss_mask: torch.Tensor, output_tensor: torch.Tensor): |
no test coverage detected