from __future__ import annotations import torch from torch import nn import torch.nn.functional as F class CIBCELoss(nn.Module): def __init__( self, output_key: str = "ci", label_key: str = "ci", weight: float = 1.0, eps: float = 1e-6, ): super().__init__() self.output_key = str(output_key) self.label_key = str(label_key) self.weight = float(weight) self.eps = float(eps) def forward(self, outputs: dict[str, torch.Tensor], labels: dict[str, torch.Tensor]) -> torch.Tensor: preds = torch.clamp(outputs[self.output_key].float(), self.eps, 1.0 - self.eps) targets = labels[self.label_key].float() if targets.ndim == 4 and targets.shape[1] == 1: targets = targets[:, 0] loss = self.weight * F.binary_cross_entropy(preds, targets) self.last_components = {"bce": float(loss.detach().cpu())} return loss class CIFocalLoss(nn.Module): def __init__( self, output_key: str = "ci", label_key: str = "ci", weight: float = 1.0, alpha: float = 0.25, gamma: float = 2.0, eps: float = 1e-6, ): super().__init__() self.output_key = str(output_key) self.label_key = str(label_key) self.weight = float(weight) self.alpha = float(alpha) self.gamma = float(gamma) self.eps = float(eps) def forward(self, outputs: dict[str, torch.Tensor], labels: dict[str, torch.Tensor]) -> torch.Tensor: preds = torch.clamp(outputs[self.output_key].float(), self.eps, 1.0 - self.eps) targets = labels[self.label_key].float().clamp(0.0, 1.0) if targets.ndim == 4 and targets.shape[1] == 1: targets = targets[:, 0] bce = F.binary_cross_entropy(preds, targets, reduction="none") pt = targets * preds + (1.0 - targets) * (1.0 - preds) alpha_t = targets * self.alpha + (1.0 - targets) * (1.0 - self.alpha) loss = self.weight * (alpha_t * (1.0 - pt).pow(self.gamma) * bce).mean() self.last_components = {"focal": float(loss.detach().cpu())} return loss class CIFocalTverskyLoss(nn.Module): def __init__( self, output_key: str = "ci", label_key: str = "ci", weight: float = 1.0, focal_weight: float = 0.5, tversky_weight: float = 0.5, fp_weight: float = 0.7, fn_weight: float = 0.3, focal_gamma: float = 2.0, tversky_gamma: float = 1.0, positive_alpha: float = 0.35, target_threshold: float = 0.5, eps: float = 1e-6, ): super().__init__() self.output_key = str(output_key) self.label_key = str(label_key) self.weight = float(weight) self.focal_weight = float(focal_weight) self.tversky_weight = float(tversky_weight) self.fp_weight = float(fp_weight) self.fn_weight = float(fn_weight) self.focal_gamma = float(focal_gamma) self.tversky_gamma = float(tversky_gamma) self.positive_alpha = float(positive_alpha) self.target_threshold = float(target_threshold) self.eps = float(eps) def forward(self, outputs: dict[str, torch.Tensor], labels: dict[str, torch.Tensor]) -> torch.Tensor: probs = torch.clamp(outputs[self.output_key].float(), self.eps, 1.0 - self.eps) targets = labels[self.label_key].float().clamp(0.0, 1.0) if targets.ndim == 4 and targets.shape[1] == 1: targets = targets[:, 0] bce = F.binary_cross_entropy(probs, targets, reduction="none") pt = targets * probs + (1.0 - targets) * (1.0 - probs) alpha_t = targets * self.positive_alpha + (1.0 - targets) * (1.0 - self.positive_alpha) focal = alpha_t * (1.0 - pt).pow(self.focal_gamma) * bce focal_loss = focal.mean() reduce_dims = tuple(range(1, probs.dim())) y_hard = (targets >= self.target_threshold).float() tp = (probs * y_hard).sum(dim=reduce_dims) fp = (probs * (1.0 - y_hard)).sum(dim=reduce_dims) fn = ((1.0 - probs) * y_hard).sum(dim=reduce_dims) tversky = (tp + self.eps) / (tp + self.fp_weight * fp + self.fn_weight * fn + self.eps) tversky_loss = (1.0 - tversky).pow(self.tversky_gamma).mean() loss = self.focal_weight * focal_loss + self.tversky_weight * tversky_loss total = self.weight * loss self.last_components = { "focal": float((self.weight * self.focal_weight * focal_loss).detach().cpu()), "tversky": float((self.weight * self.tversky_weight * tversky_loss).detach().cpu()), "total": float(total.detach().cpu()), } return total class BTMaskedSmoothL1Loss(nn.Module): def __init__( self, output_key: str = "bt", label_key: str = "bt", weight: float = 1.0, lead_weights: list[float] | None = None, beta: float = 1.0, ): super().__init__() self.output_key = str(output_key) self.label_key = str(label_key) self.weight = float(weight) self.beta = float(beta) self.register_buffer("lead_weights", torch.tensor(lead_weights or [], dtype=torch.float32), persistent=False) def forward(self, outputs: dict[str, torch.Tensor], labels: dict[str, torch.Tensor]) -> torch.Tensor: preds = outputs[self.output_key].float() targets = labels[self.label_key].float() valid = torch.isfinite(targets) if not bool(valid.any()): zero = preds.sum() * 0.0 self.last_components = {"bt_smooth_l1": 0.0} return zero loss = F.smooth_l1_loss(preds, torch.nan_to_num(targets), beta=self.beta, reduction="none") if self.lead_weights.numel() > 0: if self.lead_weights.numel() != preds.shape[1]: raise ValueError(f"lead_weights length {self.lead_weights.numel()} != T {preds.shape[1]}") view_shape = (1, preds.shape[1]) + (1,) * (preds.ndim - 2) loss = loss * self.lead_weights.to(preds.device).view(view_shape) loss = self.weight * loss[valid].mean() self.last_components = {"bt_smooth_l1": float(loss.detach().cpu())} return loss class BTMaskedMSELoss(nn.Module): def __init__( self, output_key: str = "bt", label_key: str = "bt", weight: float = 1.0, lead_weights: list[float] | None = None, ): super().__init__() self.output_key = str(output_key) self.label_key = str(label_key) self.weight = float(weight) self.register_buffer("lead_weights", torch.tensor(lead_weights or [], dtype=torch.float32), persistent=False) def forward(self, outputs: dict[str, torch.Tensor], labels: dict[str, torch.Tensor]) -> torch.Tensor: preds = outputs[self.output_key].float() targets = labels[self.label_key].float() valid = torch.isfinite(targets) if not bool(valid.any()): zero = preds.sum() * 0.0 self.last_components = {"bt_mse": 0.0} return zero loss = (preds - torch.nan_to_num(targets)) ** 2 if self.lead_weights.numel() > 0: if self.lead_weights.numel() != preds.shape[1]: raise ValueError(f"lead_weights length {self.lead_weights.numel()} != T {preds.shape[1]}") view_shape = (1, preds.shape[1]) + (1,) * (preds.ndim - 2) loss = loss * self.lead_weights.to(preds.device).view(view_shape) loss = self.weight * loss[valid].mean() self.last_components = {"bt_mse": float(loss.detach().cpu())} return loss class CIBTLoss(nn.Module): def __init__( self, ci_loss: dict, bt_loss: dict | None = None, normalize_terms: bool = False, ref_batches: int = 100, eps: float = 1e-8, ): super().__init__() ci_cfg = dict(ci_loss) bt_cfg = dict(bt_loss or {"name": "bt_mse", "weight": 0.05}) self.ci_weight = float(ci_cfg.pop("weight", 1.0)) self.bt_weight = float(bt_cfg.pop("weight", 1.0)) self.ci_loss = _build_one_loss(ci_cfg, default_name=str(ci_cfg.get("name", "bce"))) self.bt_loss = _build_one_loss(bt_cfg, default_name="bt_mse") self.normalize_terms = bool(normalize_terms) self.ref_batches = int(ref_batches) self.eps = float(eps) self.register_buffer("ref_count", torch.tensor(0.0), persistent=True) self.register_buffer("ci_ref_sum", torch.tensor(0.0), persistent=True) self.register_buffer("bt_ref_sum", torch.tensor(0.0), persistent=True) self.register_buffer("ci_ref", torch.tensor(1.0), persistent=True) self.register_buffer("bt_ref", torch.tensor(1.0), persistent=True) self.last_components: dict[str, float] = {} def forward(self, outputs: dict[str, torch.Tensor], labels: dict[str, torch.Tensor]) -> torch.Tensor: ci = self.ci_loss(outputs, labels) bt = self.bt_loss(outputs, labels) if self.normalize_terms: self._update_refs(ci, bt) ci_term = ci / self.ci_ref.clamp_min(self.eps) bt_term = bt / self.bt_ref.clamp_min(self.eps) else: ci_term = ci bt_term = bt ci_weighted = self.ci_weight * ci_term bt_weighted = self.bt_weight * bt_term total = ci_weighted + bt_weighted components = { "ci": float(ci_weighted.detach().cpu()), "bt": float(bt_weighted.detach().cpu()), "total": float(total.detach().cpu()), "ci_raw": float(ci.detach().cpu()), "bt_raw": float(bt.detach().cpu()), "ci_norm": float(ci_term.detach().cpu()), "bt_norm": float(bt_term.detach().cpu()), "ci_weight": float(self.ci_weight), "bt_weight": float(self.bt_weight), "ci_ref": float(self.ci_ref.detach().cpu()), "bt_ref": float(self.bt_ref.detach().cpu()), "ref_count": float(self.ref_count.detach().cpu()), "ref_ready": float(self._ref_ready()), } for name, value in getattr(self.ci_loss, "last_components", {}).items(): components[f"ci_{name}"] = float(value) for name, value in getattr(self.bt_loss, "last_components", {}).items(): component_name = str(name) if str(name).startswith("bt_") else f"bt_{name}" components[component_name] = float(value) self.last_components = components return total def _ref_ready(self) -> bool: return bool(self.normalize_terms and self.ref_count.item() >= max(self.ref_batches, 1)) def _update_refs(self, ci: torch.Tensor, bt: torch.Tensor) -> None: if self.ref_batches <= 0: if self.ref_count.item() < 1: self.ci_ref.copy_(ci.detach().abs().clamp_min(self.eps)) self.bt_ref.copy_(bt.detach().abs().clamp_min(self.eps)) self.ref_count.fill_(1.0) return if self.ref_count.item() >= self.ref_batches: return self.ci_ref_sum.add_(ci.detach().abs()) self.bt_ref_sum.add_(bt.detach().abs()) self.ref_count.add_(1.0) denom = self.ref_count.clamp_min(1.0) self.ci_ref.copy_((self.ci_ref_sum / denom).clamp_min(self.eps)) self.bt_ref.copy_((self.bt_ref_sum / denom).clamp_min(self.eps)) class CombinedLoss(nn.Module): def __init__(self, losses: dict[str, nn.Module]): super().__init__() if not losses: raise ValueError("CombinedLoss requires at least one loss") self.losses = nn.ModuleDict(losses) self.last_components: dict[str, float] = {} def forward(self, outputs: dict[str, torch.Tensor], labels: dict[str, torch.Tensor]) -> torch.Tensor: total = None components = {} for name, loss_fn in self.losses.items(): value = loss_fn(outputs, labels) components[name] = float(value.detach().cpu()) for sub_name, sub_value in getattr(loss_fn, "last_components", {}).items(): component_name = str(sub_name) if name == "loss" else f"{name}_{sub_name}" components[component_name] = float(sub_value) total = value if total is None else total + value self.last_components = components if total is None: raise RuntimeError("no losses were evaluated") return total def build_loss(config: dict) -> nn.Module: config = dict(config) if str(config.get("name", "")).lower() == "ci_bt": return _build_one_loss(config, default_name="ci_bt") if "losses" in config: return CombinedLoss( { str(component_name): _build_one_loss(dict(component_cfg), default_name=str(component_name)) for component_name, component_cfg in dict(config["losses"]).items() } ) return CombinedLoss({"loss": _build_one_loss(config, default_name=str(config.get("name", "binary_focal")))}) def _build_one_loss(config: dict, default_name: str) -> nn.Module: name = str(config.get("name", default_name)).lower() params = dict(config.get("params", {})) output_key = str(config.get("output_key", "bt" if name.startswith("bt") or "smooth_l1" in name else "ci")) label_key = str(config.get("label_key", config.get("target_label", output_key))) weight = float(config.get("weight", 1.0)) if name in {"bce", "binary_bce", "binary_cross_entropy", "ci_bce"}: return CIBCELoss(output_key=output_key, label_key=label_key, weight=weight, **params) if name in {"binary_focal", "focal", "binary_focal_loss", "ci", "ci_focal"}: return CIFocalLoss(output_key=output_key, label_key=label_key, weight=weight, **params) if name in {"focal_tversky", "tversky_focal", "ci_focal_tversky", "far_aware_tversky_focal", "far_aware", "ci_far_aware"}: return CIFocalTverskyLoss(output_key=output_key, label_key=label_key, weight=weight, **params) if name in {"masked_lead_smooth_l1", "bt_smooth_l1", "bt", "bt_masked_smooth_l1"}: return BTMaskedSmoothL1Loss(output_key=output_key, label_key=label_key, weight=weight, **params) if name in {"bt_mse", "masked_bt_mse", "bt_masked_mse"}: return BTMaskedMSELoss(output_key=output_key, label_key=label_key, weight=weight, **params) if name == "ci_bt": if "ci_loss" not in config: raise ValueError("ci_bt loss requires a ci_loss config block") bt_loss = config.get("bt_loss") if bt_loss is None: bt_loss = { "name": "bt_mse", "output_key": config.get("bt_output_key", "bt"), "label_key": config.get("bt_label_key", "bt"), "weight": config.get("bt_weight", 0.05), } return CIBTLoss( ci_loss=dict(config["ci_loss"]), bt_loss=dict(bt_loss), **params, ) raise ValueError(f"unknown loss: {name}")