Complete standalone final deliverable and unaligned Q2 results
This commit is contained in:
@@ -0,0 +1,32 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user