Files
modeling_zhaocui/submit/final/model/c7_distill.py
T

42 lines
1.8 KiB
Python

"""C7 distillation-only branch: C6 architecture plus teacher loss."""
from __future__ import annotations
import math
import numpy as np
import torch
from torch.nn import functional as F
from .c6 import C6
DISTILL_TEMPERATURE = 2.0
DISTILL_WEIGHT = 0.1
class C7Distill(C6):
"""Inference uses C6; training adds weighted teacher distillation."""
def distillation_per_sample(student: dict, teacher: dict,
original: np.ndarray, current: np.ndarray) -> torch.Tensor:
"""Entropy/retention-weighted KL and score term for Q2 distillation."""
temp = DISTILL_TEMPERATURE
p_teacher = teacher["tempered_probs_by_path"].mean(dim=0).detach().clamp_min(1e-8)
p_student = student["tempered_probs_by_path"].mean(dim=0).clamp_min(1e-8)
entropy = -(p_teacher * p_teacher.log()).sum(dim=-1)
confidence_weight = (1.0 - entropy / math.log(3.0)).clamp(0.0, 1.0)
orig_t = torch.as_tensor(original, device=p_teacher.device, dtype=torch.float32)
curr_t = torch.as_tensor(current, device=p_teacher.device, dtype=torch.float32)
retained = []
for modality in range(3):
denominator = orig_t[:, :, modality].sum(dim=1)
ratio = (orig_t[:, :, modality] * curr_t[:, :, modality]).sum(dim=1) / denominator.clamp_min(1.0)
retained.append(torch.where(denominator > 0, ratio, torch.ones_like(ratio)))
weight = confidence_weight * torch.stack(retained, dim=-1).mean(dim=-1)
kl = (p_teacher * (p_teacher.log() - p_student.log())).sum(dim=-1) * temp * temp
teacher_score = teacher["mixed_score"].detach()
student_score = student["mixed_score"]
regression = F.huber_loss((teacher_score - student_score) / 3.0,
torch.zeros_like(teacher_score), reduction="none", delta=0.25)
return weight * (kl + regression)