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

Function loss_func

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

Source from the content-addressed store, hash-verified

155
156
157def 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
182def valid_loss_func(loss_mask, output_tensor):

Callers 1

forward_stepFunction · 0.50

Calls 3

get_argsFunction · 0.90
compute_lm_lossFunction · 0.85

Tested by

no test coverage detected