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)
| 184 | return batch.values() |
| 185 | |
| 186 | def 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 | |
| 218 | def forward_step(data_iterator, model: GPTModel): |
no test coverage detected