From f02ae874f3f51e0ed1a391112155c40944f454b3 Mon Sep 17 00:00:00 2001 From: 51616 Date: Fri, 5 Sep 2025 17:51:37 +0900 Subject: [PATCH] sparse kl --- src/ctx_to_lora/trainer.py | 18 ++++++++++++------ 1 file changed, 12 insertions(+), 6 deletions(-) diff --git a/src/ctx_to_lora/trainer.py b/src/ctx_to_lora/trainer.py index c499c81..43d9f38 100644 --- a/src/ctx_to_lora/trainer.py +++ b/src/ctx_to_lora/trainer.py @@ -164,13 +164,19 @@ class DistillationTrainer(ModulatedModelTrainer): ##### KL loss outputs_logits = outputs.logits[label_pos[0], label_pos[1] - 1] # shift back 1 - teacher_logp = torch.full_like(outputs_logits, -torch.inf) - teacher_logp.scatter_(1, indices, target_logp) - # reduction = "batchmean" if num_items_in_batch is None else "sum" + logq_full_denom = torch.logsumexp(outputs_logits, dim=-1, keepdim=True) # (N,1) + selected_logits = outputs_logits.gather(1, indices) # (N,K) + # log softmax at selected indices + logq_selected = selected_logits - logq_full_denom + p = target_logp.exp() + loss = -(p * logq_selected).sum(dim=-1) - p = teacher_logp.exp() - logq = nn.functional.log_softmax(outputs_logits, dim=-1) - loss = -torch.sum(p * logq, dim=-1) + # teacher_logp = torch.full_like(outputs_logits, -torch.inf) + # teacher_logp.scatter_(1, indices, target_logp) + # # reduction = "batchmean" if num_items_in_batch is None else "sum" + # p = teacher_logp.exp() + # logq = nn.functional.log_softmax(outputs_logits, dim=-1) + # loss = -torch.sum(p * logq, dim=-1) if self.use_per_ctx_average_loss: loss = per_ctx_loss_kl(inputs, labels, loss)