"""C7 group-risk-only branch: C6 architecture plus smooth worst-group loss.""" from __future__ import annotations import numpy as np import torch from .c6 import C6 class C7Group(C6): """Inference uses C6; training adds smooth worst-group risk.""" def smooth_group_risk(losses: torch.Tensor, group_ids: np.ndarray, lambda_group: float = 0.1, group_temperature: float = 0.05) -> torch.Tensor: """Match the selected group penalty from the Q2 training protocol.""" if group_temperature <= 0 or not 0 <= lambda_group <= 1: raise ValueError("invalid group risk parameters") groups = torch.as_tensor(group_ids, device=losses.device, dtype=torch.long) if groups.shape != losses.shape: raise ValueError("group_ids must match per-sample losses") group_losses, priors = [], [] for group in torch.unique(groups): selected = groups == group group_losses.append(losses[selected].mean()) priors.append(selected.float().mean()) values = torch.stack(group_losses) prior = torch.stack(priors).clamp_min(1e-8) expected = (prior * values).sum() worst = group_temperature * torch.logsumexp(torch.log(prior) + values / group_temperature, dim=0) return (1.0 - lambda_group) * expected + lambda_group * worst