42 lines
1.8 KiB
Python
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)
|