Download code/training/src/training_validation/loss.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 15.2 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/loss.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/training_validation/loss.py
-
curl -L -o loss.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/loss.py
15.2 kB
| 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}") | |