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

Function loss_func

examples/linear_llama3/pretrain_llama.py:186–215  ·  view source on GitHub ↗

Loss function. Args: loss_mask (torch.Tensor): Used to mask out some portions of the loss output_tensor (torch.Tensor): The tensor with the losses

(loss_mask: torch.Tensor, output_tensor: torch.Tensor)

Source from the content-addressed store, hash-verified

184 return batch.values()
185
186def loss_func(loss_mask: torch.Tensor, output_tensor: torch.Tensor):
187 """Loss function.
188
189 Args:
190 loss_mask (torch.Tensor): Used to mask out some portions of the loss
191 output_tensor (torch.Tensor): The tensor with the losses
192 """
193 args = get_args()
194
195 losses = output_tensor.float()
196 loss_mask = loss_mask.view(-1).float()
197 if args.context_parallel_size > 1:
198 loss = torch.cat([torch.sum(losses.view(-1) * loss_mask).view(1), loss_mask.sum().view(1)])
199 torch.distributed.all_reduce(loss, group=mpu.get_context_parallel_group())
200 loss = loss[0] / loss[1]
201 else:
202 loss = torch.sum(losses.view(-1) * loss_mask) / loss_mask.sum()
203
204 # Check individual rank losses are not NaN prior to DP all-reduce.
205 if args.check_for_nan_in_loss_and_grad:
206 global_rank = torch.distributed.get_rank()
207 assert not loss.isnan(), (
208 f'Rank {global_rank}: found NaN in local forward loss calculation. '
209 f'Device: {torch.cuda.current_device()}, node: {os.uname()[1]}'
210 )
211
212 # Reduce loss for logging.
213 averaged_loss = average_losses_across_data_parallel_group([loss])
214
215 return loss * args.context_parallel_size, {'lm loss': averaged_loss[0]}
216
217
218def forward_step(data_iterator, model: GPTModel):

Callers 2

cross_entropy_loss_funcFunction · 0.50
apply_aux_lossMethod · 0.50

Calls 1

get_argsFunction · 0.50

Tested by

no test coverage detected