(loss_mask, output_tensor)
| 150 | |
| 151 | |
| 152 | def loss_func(loss_mask, output_tensor): |
| 153 | losses = output_tensor.float() |
| 154 | loss_mask = loss_mask.view(-1).float() |
| 155 | loss = torch.sum(losses.view(-1) * loss_mask) / loss_mask.sum() |
| 156 | |
| 157 | # Reduce loss for logging. |
| 158 | averaged_loss = average_losses_across_data_parallel_group([loss]) |
| 159 | |
| 160 | return loss, {"lm loss": averaged_loss[0]} |
| 161 | |
| 162 | |
| 163 | def forward_step(data_iterator, model): |
nothing calls this directly
no test coverage detected