MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / valid_loss_func

Function valid_loss_func

codegeex/megatron/tools/finetune_codegeex.py:182–197  ·  view source on GitHub ↗
(loss_mask, output_tensor)

Source from the content-addressed store, hash-verified

180
181
182def 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
200def forward_step(data_iterator, model):

Callers

nothing calls this directly

Calls 3

get_argsFunction · 0.90
compute_lm_lossFunction · 0.85

Tested by

no test coverage detected