(loss_mask, output_tensor)
| 155 | |
| 156 | |
| 157 | def loss_func(loss_mask, output_tensor): |
| 158 | args = get_args() |
| 159 | |
| 160 | def compute_lm_loss(losses: torch.Tensor, loss_mask: torch.Tensor): |
| 161 | if args.gold: |
| 162 | losses_ = losses.detach() |
| 163 | prob = torch.exp(-losses_) # Pθ(s) |
| 164 | torch.sqrt_(prob) # Pθ(s)ᵃ |
| 165 | torch.clamp_min_(prob, args.gold_beta) # max(Pθ(s)ᵃ,β) |
| 166 | losses = prob * losses |
| 167 | |
| 168 | loss_mask = loss_mask.view(-1).float() |
| 169 | loss = torch.sum(losses.view(-1) * loss_mask) / torch.clamp_min(loss_mask.sum(), 1e-8) |
| 170 | |
| 171 | return loss |
| 172 | |
| 173 | losses = output_tensor.float() |
| 174 | loss = compute_lm_loss(losses, loss_mask) |
| 175 | |
| 176 | # Reduce loss for logging. |
| 177 | averaged_loss = average_losses_across_data_parallel_group([loss]) |
| 178 | |
| 179 | return loss, {"lm loss": averaged_loss[0]} |
| 180 | |
| 181 | |
| 182 | def valid_loss_func(loss_mask, output_tensor): |
no test coverage detected