lsh9034's picture
Add files using upload-large-folder tool
76d61a0 verified
Raw History Blame Contribute Delete
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}")