(loss_mask, output_tensor)
| 180 | |
| 181 | |
| 182 | def valid_loss_func(loss_mask, output_tensor): |
| 183 | args = get_args() |
| 184 | |
| 185 | def compute_lm_loss(losses: torch.Tensor, loss_mask: torch.Tensor): |
| 186 | loss_mask = loss_mask.view(-1).float() |
| 187 | loss = torch.sum(losses.view(-1) * loss_mask) / torch.clamp_min(loss_mask.sum(), 1e-8) |
| 188 | |
| 189 | return loss |
| 190 | |
| 191 | losses = output_tensor.float() |
| 192 | loss = compute_lm_loss(losses, loss_mask) |
| 193 | |
| 194 | # Reduce loss for logging. |
| 195 | averaged_loss = average_losses_across_data_parallel_group([loss]) |
| 196 | |
| 197 | return loss, {"lm loss": averaged_loss[0]} |
| 198 | |
| 199 | |
| 200 | def forward_step(data_iterator, model): |
nothing calls this directly
no test coverage detected