"""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)