33 lines
1.3 KiB
Python
33 lines
1.3 KiB
Python
"""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
|