MCPcopy Create free account
hub / github.com/OpenSparseLLMs/Linear-MoE / get_batch

Function get_batch

examples/linear_moe_deepseek_v2/pretrain_deepseek.py:165–189  ·  view source on GitHub ↗

Generate a batch.

(data_iterator)

Source from the content-addressed store, hash-verified

163
164
165def 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
192def loss_func(loss_mask: torch.Tensor, output_tensor: torch.Tensor):

Callers 1

forward_stepFunction · 0.70

Calls 2

get_argsFunction · 0.50

Tested by

no test coverage detected