CD-Models / utils /model_adapters.py
Dineth Perera
Publish tested dataset winners and benchmark rankings
ce209f5
Raw
History Blame Contribute Delete
46.1 kB
from __future__ import annotations
import importlib
import sys
from abc import ABC, abstractmethod
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from types import SimpleNamespace
from typing import Any
import torch
import torch.nn as nn
import torch.nn.functional as F
from utils.metrics import normalize_binary_prediction
ROOT = Path(__file__).resolve().parents[1]
@dataclass
class NormalizedOutput:
binary: torch.Tensor
score: torch.Tensor | None
metric_tensor: torch.Tensor
class AdapterUnavailable(RuntimeError):
pass
@contextmanager
def repo_import_context(repo: Path):
repo_str = str(repo)
old_path = list(sys.path)
upstream_prefixes = ("models", "model", "network", "cd", "DSIFN", "changedetection")
for name in list(sys.modules):
if any(name == prefix or name.startswith(f"{prefix}.") for prefix in upstream_prefixes):
del sys.modules[name]
sys.path.insert(0, repo_str)
try:
yield
finally:
sys.path = old_path
class BaseModelAdapter(ABC):
model_name = "base"
supports_inprocess_eval = False
supports_unified_eval = False
supports_unified_training = False
supports_flops = False
supports_amp_training = True
notes_or_failure_reason = "Adapter is not implemented."
model_class_path = ""
input_format = ""
output_format = ""
checkpoint_format = ""
final_output_for_metrics = ""
def __init__(self, model_name: str | None = None) -> None:
if model_name:
self.model_name = model_name
def _missing(self, method: str) -> NotImplementedError:
return NotImplementedError(f"{self.model_name} adapter does not implement {method}: {self.notes_or_failure_reason}")
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
raise self._missing("build_model")
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
raise self._missing("load_checkpoint")
def forward(self, model: nn.Module, batch: tuple, device: torch.device) -> Any:
raise self._missing("forward")
def normalize_output(self, raw_output: Any, batch: tuple, dataset_config: dict) -> NormalizedOutput:
raise self._missing("normalize_output")
def get_dummy_inputs(self, dataset_config: dict, device: torch.device) -> tuple:
raise self._missing("get_dummy_inputs")
def compute_loss(self, raw_output: Any, batch: tuple, model_config: dict, dataset_config: dict, device: torch.device) -> dict[str, torch.Tensor]:
output = raw_output[-1] if isinstance(raw_output, (list, tuple)) else raw_output
output = output.float()
mask = batch[2].to(device, non_blocking=True).float()
if output.ndim == 4 and output.shape[1] == 2:
target = mask.squeeze(1).long()
loss = F.cross_entropy(output, target)
elif output.ndim == 4 and output.shape[1] == 1:
target = mask
if output.shape[-2:] != target.shape[-2:]:
output = F.interpolate(output, size=target.shape[-2:], mode="bilinear", align_corners=True)
if output.detach().min() >= 0 and output.detach().max() <= 1:
loss = F.binary_cross_entropy(output.clamp(1e-6, 1 - 1e-6), target)
else:
loss = F.binary_cross_entropy_with_logits(output, target)
else:
raise RuntimeError(f"{self.model_name} default loss cannot handle output shape {tuple(output.shape)}.")
return {"loss": loss, "main_loss": loss.detach()}
def save_checkpoint(
self,
model: nn.Module,
optimizer: torch.optim.Optimizer,
scheduler: Any,
path: Path,
metadata: dict,
) -> None:
payload = dict(metadata)
payload["model_state_dict"] = model.state_dict()
payload["optimizer_state_dict"] = optimizer.state_dict()
if scheduler is not None:
payload["scheduler_state_dict"] = scheduler.state_dict()
torch.save(payload, path)
def build_optimizer(self, model: nn.Module, model_config: dict) -> torch.optim.Optimizer:
lr = float(model_config.get("lr", 1e-4))
weight_decay = float(model_config.get("weight_decay", 0.0) or 0.0)
optimizer_name = str(model_config.get("optimizer", "adam")).lower()
if optimizer_name == "sgd":
return torch.optim.SGD(model.parameters(), lr=lr, momentum=0.9, weight_decay=weight_decay)
if optimizer_name == "adamw":
return torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)
return torch.optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay)
def build_scheduler(self, optimizer: torch.optim.Optimizer, model_config: dict) -> Any:
scheduler_name = str(model_config.get("scheduler", "")).lower()
epochs = int(model_config.get("num_epochs", 200))
if scheduler_name == "cosine":
return torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
if scheduler_name == "step":
return torch.optim.lr_scheduler.StepLR(optimizer, step_size=50, gamma=0.5)
return None
def get_primary_metric(self) -> str:
return "f1"
def get_threshold(self, dataset_config: dict) -> float:
return float(dataset_config.get("eval", {}).get("threshold", 0.5))
def get_threshold_mode(self) -> str:
return "argmax" if "two-channel" in self.output_format else "threshold"
class UnsupportedAdapter(BaseModelAdapter):
def __init__(self, model_name: str, reason: str) -> None:
super().__init__(model_name)
self.notes_or_failure_reason = reason
class TwoTensorLogitAdapter(BaseModelAdapter):
supports_inprocess_eval = True
supports_unified_eval = True
supports_unified_training = True
supports_flops = True
input_format = "two tensors (A, B), each [B, 3, H, W]"
output_format = "two-channel logits [B, 2, H, W]"
final_output_for_metrics = "raw two-channel logits"
def forward(self, model: nn.Module, batch: tuple, device: torch.device) -> torch.Tensor:
a, b, _mask, _names = batch
return model(a.to(device, non_blocking=True), b.to(device, non_blocking=True))
def normalize_output(self, raw_output: Any, batch: tuple, dataset_config: dict) -> NormalizedOutput:
if isinstance(raw_output, (list, tuple)):
raw_output = raw_output[-1]
threshold = float(dataset_config.get("eval", {}).get("threshold", 0.5))
binary, score = normalize_binary_prediction(raw_output.detach().cpu(), threshold=threshold)
return NormalizedOutput(binary=binary, score=score, metric_tensor=raw_output.detach().cpu())
def get_dummy_inputs(self, dataset_config: dict, device: torch.device) -> tuple:
size = int(dataset_config.get("img_size", dataset_config.get("image_size", 256)))
return (
torch.zeros(1, 3, size, size, device=device),
torch.zeros(1, 3, size, size, device=device),
)
class FCAdapter(TwoTensorLogitAdapter):
notes_or_failure_reason = "Direct FC model from fully_convolutional_change_detection."
checkpoint_format = "CD-Models checkpoint dict with model_state_dict, or raw state dict"
def __init__(self, model_name: str) -> None:
super().__init__(model_name)
self.model_class_path = {
"fc_ef": "model_repos/fully_convolutional_change_detection/unet.py:Unet",
"fc_siam_conc": "model_repos/fully_convolutional_change_detection/siamunet_conc.py:SiamUnet_conc",
"fc_siam_diff": "model_repos/fully_convolutional_change_detection/siamunet_diff.py:SiamUnet_diff",
}[model_name]
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "model_repos" / "fully_convolutional_change_detection"
with repo_import_context(repo):
if self.model_name == "fc_ef":
module = importlib.import_module("unet")
model = module.Unet(input_nbr=6, label_nbr=2)
elif self.model_name == "fc_siam_conc":
module = importlib.import_module("siamunet_conc")
model = module.SiamUnet_conc(input_nbr=3, label_nbr=2)
else:
module = importlib.import_module("siamunet_diff")
model = module.SiamUnet_diff(input_nbr=3, label_nbr=2)
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"{self.model_name} checkpoint {checkpoint_path} has no recognized state dict.")
model.load_state_dict(state, strict=True)
def compute_loss(self, raw_output: Any, batch: tuple, model_config: dict, dataset_config: dict, device: torch.device) -> dict[str, torch.Tensor]:
target = batch[2].squeeze(1).long().to(device, non_blocking=True)
weight = None
if self.model_name == "fc_siam_diff":
pos = target.eq(1).sum().float()
neg = target.eq(0).sum().float()
if bool(pos.item() > 0):
pos_weight = (neg / pos.clamp_min(1.0)).clamp(min=1.0, max=50.0)
weight = torch.stack([torch.ones_like(pos_weight), pos_weight]).to(device)
loss = F.nll_loss(raw_output.float(), target, weight=weight)
result = {"loss": loss, "nll_loss": loss.detach()}
if weight is not None:
result["positive_class_weight"] = weight[1].detach()
return result
class BITAdapter(TwoTensorLogitAdapter):
model_name = "bit_cd"
notes_or_failure_reason = "BIT_CD models.networks.define_G with base_transformer_pos_s4_dd8."
model_class_path = "BIT_CD/models/networks.py:define_G"
checkpoint_format = "raw state dict, CD-Models model_state_dict, or upstream model_G_state_dict"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "BIT_CD"
with repo_import_context(repo):
networks = importlib.import_module("models.networks")
args = SimpleNamespace(net_G="base_transformer_pos_s4_dd8", lr_policy="linear", max_epochs=1)
model = networks.define_G(args, gpu_ids=[])
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_G_state_dict" in checkpoint:
state = checkpoint["model_G_state_dict"]
elif isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"BIT_CD checkpoint {checkpoint_path} has no recognized state dict.")
model.load_state_dict(state, strict=True)
class ChangeFormerAdapter(BITAdapter):
model_name = "changeformer"
supports_amp_training = False
notes_or_failure_reason = "ChangeFormer models.networks.define_G with ChangeFormerV6 embed_dim=64."
model_class_path = "ChangeFormer/models/networks.py:define_G"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "ChangeFormer"
with repo_import_context(repo):
networks = importlib.import_module("models.networks")
args = SimpleNamespace(net_G="ChangeFormerV6", embed_dim=64, lr_policy="linear", max_epochs=1)
model = networks.define_G(args, gpu_ids=[])
return model.to(device)
class BiFAAdapter(TwoTensorLogitAdapter):
model_name = "bifa"
supports_amp_training = False
notes_or_failure_reason = "BiFA models.bifa.BiFA with mit_b0 backbone; returns two-channel logits."
model_class_path = "BiFA/models/bifa.py:BiFA"
output_format = "two-channel logits [B, 2, H, W]"
checkpoint_format = "raw state dict or CD-Models model_state_dict"
final_output_for_metrics = "raw two-channel logits"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "BiFA"
with repo_import_context(repo):
module = importlib.import_module("models.bifa")
model = module.BiFA(backbone="mit_b0")
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"BiFA checkpoint {checkpoint_path} has no recognized state dict.")
model.load_state_dict(state, strict=True)
class SiamNestedUNetAdapter(TwoTensorLogitAdapter):
model_name = "siam_nestedunet"
supports_amp_training = False
notes_or_failure_reason = "Siam-NestedUNet SNUNet_ECAM returns a tuple; final output is tuple[-1]."
model_class_path = "Siam-NestedUNet/models/Models.py:SNUNet_ECAM"
output_format = "tuple containing two-channel logits [B, 2, H, W]"
checkpoint_format = "raw state dict or CD-Models model_state_dict"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "Siam-NestedUNet"
with repo_import_context(repo):
module = importlib.import_module("models.Models")
model = module.SNUNet_ECAM(3, 2)
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"Siam-NestedUNet checkpoint {checkpoint_path} has no recognized state dict.")
model.load_state_dict(state, strict=True)
class STANetModule(nn.Module):
def __init__(self, net_f: nn.Module, net_a: nn.Module) -> None:
super().__init__()
self.netF = net_f
self.netA = net_a
def forward(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
feat_a = self.netF(a)
feat_b = self.netF(b)
feat_a, feat_b = self.netA(feat_a, feat_b)
dist = torch.norm(feat_a - feat_b, p=2, dim=1, keepdim=True)
return F.interpolate(dist, size=a.shape[-2:], mode="bilinear", align_corners=True)
class STANetAdapter(BaseModelAdapter):
model_name = "stanet"
supports_inprocess_eval = True
supports_unified_eval = True
supports_unified_training = True
supports_flops = True
notes_or_failure_reason = "STANet netF + CDSA(PAM), distance map thresholded at 1.0 as in train_wildfire.py."
model_class_path = "STANet/models/backbone.py:define_F + CDSA"
input_format = "two tensors (A, B), each [B, 3, H, W]"
output_format = "one-channel L2 distance map [B, 1, H, W]"
checkpoint_format = "dict with netF and netA state dicts"
final_output_for_metrics = "binary distance > 1.0"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "STANet"
with repo_import_context(repo):
backbone = importlib.import_module("models.backbone")
net_f = backbone.define_F(in_c=3, f_c=64, type="mynet3")
net_a = backbone.CDSA(in_c=64, ds=1, mode="PAM")
return STANetModule(net_f, net_a).to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
model.load_state_dict(checkpoint["model_state_dict"], strict=True)
return
if not isinstance(checkpoint, dict) or "netF" not in checkpoint or "netA" not in checkpoint:
raise RuntimeError(f"STANet checkpoint {checkpoint_path} must contain netF and netA keys.")
model.netF.load_state_dict(checkpoint["netF"], strict=True)
model.netA.load_state_dict(checkpoint["netA"], strict=True)
def forward(self, model: nn.Module, batch: tuple, device: torch.device) -> torch.Tensor:
a, b, _mask, _names = batch
return model(a.to(device, non_blocking=True), b.to(device, non_blocking=True))
def normalize_output(self, raw_output: Any, batch: tuple, dataset_config: dict) -> NormalizedOutput:
dist = raw_output.detach().cpu()
binary = (dist > 1.0).float()
return NormalizedOutput(binary=binary, score=None, metric_tensor=binary)
def get_dummy_inputs(self, dataset_config: dict, device: torch.device) -> tuple:
size = int(dataset_config.get("img_size", dataset_config.get("image_size", 256)))
return (
torch.zeros(1, 3, size, size, device=device),
torch.zeros(1, 3, size, size, device=device),
)
def compute_loss(self, raw_output: Any, batch: tuple, model_config: dict, dataset_config: dict, device: torch.device) -> dict[str, torch.Tensor]:
repo = ROOT / "STANet"
with repo_import_context(repo):
module = importlib.import_module("models.loss")
criterion = module.BCL().to(device)
target = batch[2].to(device, non_blocking=True).float()
label = target.clone()
label[target == 1] = -1
label[target == 0] = 1
loss = criterion(raw_output, label)
return {"loss": loss, "bcl_loss": loss.detach()}
class TinyCDAdapter(TwoTensorLogitAdapter):
model_name = "tinycd"
supports_inprocess_eval = True
supports_unified_eval = True
supports_unified_training = True
supports_flops = True
supports_amp_training = False
notes_or_failure_reason = "TinyCD ChangeClassifier, sigmoid one-channel probability output."
model_class_path = "model_repos/Tiny_model_4_CD/models/change_classifier.py:ChangeClassifier"
output_format = "one-channel sigmoid probability map [B, 1, H, W]"
checkpoint_format = "raw state dict or CD-Models model_state_dict"
final_output_for_metrics = "sigmoid probability thresholded with dataset eval threshold"
def get_threshold_mode(self) -> str:
return "threshold"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "model_repos" / "Tiny_model_4_CD"
with repo_import_context(repo):
module = importlib.import_module("models.change_classifier")
model = module.ChangeClassifier()
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"TinyCD checkpoint {checkpoint_path} has no recognized state dict.")
model.load_state_dict(state, strict=True)
def compute_loss(self, raw_output: Any, batch: tuple, model_config: dict, dataset_config: dict, device: torch.device) -> dict[str, torch.Tensor]:
target = batch[2].to(device, non_blocking=True).float()
output = raw_output.float()
if output.shape[-2:] != target.shape[-2:]:
output = F.interpolate(output, size=target.shape[-2:], mode="bilinear", align_corners=True)
loss = F.binary_cross_entropy(output.clamp(1e-6, 1 - 1e-6), target)
return {"loss": loss, "bce_loss": loss.detach()}
class DSAMNetAdapter(BaseModelAdapter):
model_name = "dsamnet"
supports_inprocess_eval = True
supports_unified_eval = True
supports_unified_training = True
supports_flops = True
notes_or_failure_reason = "DSAMNet returns distance map plus deep supervision outputs; distance > 1 is change."
model_class_path = "model_repos/DSAMNet/model/dsamnet.py:DSAMNet"
input_format = "two tensors (A, B), each [B, 3, H, W]"
output_format = "tuple(dist [B,1,H,W], ds2, ds3)"
checkpoint_format = "raw state dict or CD-Models model_state_dict"
final_output_for_metrics = "dist > 1.0"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "model_repos" / "DSAMNet"
with repo_import_context(repo):
module = importlib.import_module("model.dsamnet")
model = module.DSAMNet(2)
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"DSAMNet checkpoint {checkpoint_path} has no recognized state dict.")
model.load_state_dict(state, strict=True)
def forward(self, model: nn.Module, batch: tuple, device: torch.device) -> Any:
a, b, _mask, _names = batch
return model(a.to(device, non_blocking=True), b.to(device, non_blocking=True))
def normalize_output(self, raw_output: Any, batch: tuple, dataset_config: dict) -> NormalizedOutput:
dist = raw_output[0].detach().cpu() if isinstance(raw_output, (list, tuple)) else raw_output.detach().cpu()
binary = (dist > 1.0).float()
return NormalizedOutput(binary=binary, score=None, metric_tensor=binary)
def get_dummy_inputs(self, dataset_config: dict, device: torch.device) -> tuple:
size = int(dataset_config.get("img_size", dataset_config.get("image_size", 256)))
return (
torch.zeros(1, 3, size, size, device=device),
torch.zeros(1, 3, size, size, device=device),
)
def compute_loss(self, raw_output: Any, batch: tuple, model_config: dict, dataset_config: dict, device: torch.device) -> dict[str, torch.Tensor]:
if not isinstance(raw_output, (list, tuple)) or len(raw_output) < 3:
raise RuntimeError("DSAMNet training expects (dist, ds2, ds3).")
dist, ds2, ds3 = raw_output[:3]
repo = ROOT / "model_repos" / "DSAMNet"
with repo_import_context(repo):
bcl_module = importlib.import_module("loss.BCL")
dice_module = importlib.import_module("loss.DiceLoss")
bcl = bcl_module.BCL().to(device)
dice = dice_module.DiceLoss().to(device)
mask = batch[2].to(device, non_blocking=True).float()
one_hot = torch.cat([1.0 - mask, mask], dim=1)
dice_loss = 0.5 * (dice(ds2, one_hot) + dice(ds3, one_hot))
ct_loss = bcl(dist, mask.clone())
total = ct_loss + float(model_config.get("wDice", 0.5)) * dice_loss
return {"loss": total, "contrastive_loss": ct_loss.detach(), "dice_loss": dice_loss.detach()}
class CGNetAdapter(TwoTensorLogitAdapter):
model_name = "cgnet"
notes_or_failure_reason = "CGNet returns (change_map, final_map); final_map is used for metrics."
model_class_path = "model_repos/CGNet-CD/network/CGNet.py:CGNet"
output_format = "tuple(one-channel coarse logits, one-channel final logits)"
checkpoint_format = "raw state dict or CD-Models model_state_dict"
final_output_for_metrics = "tuple[-1] final_map"
def get_threshold_mode(self) -> str:
return "threshold"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "model_repos" / "CGNet-CD"
with repo_import_context(repo):
module = importlib.import_module("network.CGNet")
model = module.CGNet()
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"CGNet checkpoint {checkpoint_path} has no recognized state dict.")
model.load_state_dict(state, strict=True)
def compute_loss(self, raw_output: Any, batch: tuple, model_config: dict, dataset_config: dict, device: torch.device) -> dict[str, torch.Tensor]:
if not isinstance(raw_output, (list, tuple)) or len(raw_output) < 2:
raise RuntimeError("CGNet training expects a tuple of two prediction maps.")
target = batch[2].to(device, non_blocking=True).float()
loss1 = F.binary_cross_entropy_with_logits(raw_output[0], target)
loss2 = F.binary_cross_entropy_with_logits(raw_output[1], target)
loss = loss1 + loss2
return {"loss": loss, "coarse_loss": loss1.detach(), "final_loss": loss2.detach()}
class HANetAdapter(TwoTensorLogitAdapter):
model_name = "hanet"
supports_amp_training = False
notes_or_failure_reason = "HANet HAN returns a tuple; tuple[-1] is used for metrics."
model_class_path = "model_repos/HANet-CD/models/HANet.py:HAN"
output_format = "tuple containing two-channel logits [B, 2, H, W]"
checkpoint_format = "full torch-saved model, raw state dict, or CD-Models model_state_dict"
final_output_for_metrics = "tuple[-1]"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "model_repos" / "HANet-CD"
with repo_import_context(repo):
module = importlib.import_module("models.HANet")
model = module.HAN(3, 2)
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, nn.Module):
model.load_state_dict(checkpoint.state_dict(), strict=True)
return
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"HANet checkpoint {checkpoint_path} has no recognized model/state dict.")
model.load_state_dict(state, strict=True)
def compute_loss(self, raw_output: Any, batch: tuple, model_config: dict, dataset_config: dict, device: torch.device) -> dict[str, torch.Tensor]:
logits = raw_output[-1] if isinstance(raw_output, (list, tuple)) else raw_output
logits = logits.float()
target = batch[2].long().to(device, non_blocking=True)
if target.ndim == 4:
target = target.squeeze(1)
ce = F.cross_entropy(logits, target)
probas = F.softmax(logits, dim=1)
one_hot = F.one_hot(target, num_classes=logits.shape[1]).permute(0, 3, 1, 2).float()
dims = (0, 2, 3)
dice = 1.0 - ((2.0 * torch.sum(probas * one_hot, dims) + 1e-7) / (torch.sum(probas + one_hot, dims) + 1e-7)).mean()
loss = ce + dice
return {"loss": loss, "cross_entropy": ce.detach(), "dice_loss": dice.detach()}
class ELGCNetAdapter(TwoTensorLogitAdapter):
model_name = "elgcnet"
notes_or_failure_reason = "ELGCNet models.networks.define_G with net_G=ELGCNet and dec_embed_dim=256."
model_class_path = "model_repos/elgcnet/models/networks.py:define_G"
output_format = "two-channel logits [B, 2, H, W]"
checkpoint_format = "upstream model_G_state_dict, raw state dict, or CD-Models model_state_dict"
final_output_for_metrics = "raw two-channel logits"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "model_repos" / "elgcnet"
with repo_import_context(repo):
networks = importlib.import_module("models.networks")
args = SimpleNamespace(net_G="ELGCNet", dec_embed_dim=256)
model = networks.define_G(args, gpu_ids=[])
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_G_state_dict" in checkpoint:
state = checkpoint["model_G_state_dict"]
elif isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"ELGCNet checkpoint {checkpoint_path} has no recognized state dict.")
model.load_state_dict(state, strict=True)
class DSIFNAdapter(BaseModelAdapter):
supports_inprocess_eval = True
supports_unified_eval = True
supports_unified_training = True
supports_flops = True
supports_amp_training = False
notes_or_failure_reason = "IFNet/pytorch version DSIFN with two VGG16 feature towers; upstream validation uses preds[-1]."
model_class_path = "IFNet/pytorch version/DSIFN.py:DSIFN"
input_format = "two tensors (A, B), each [B, 3, H, W]"
output_format = "tuple of five one-channel sigmoid maps"
checkpoint_format = "raw state dict, CD-Models model_state_dict, or upstream state_dict"
final_output_for_metrics = "tuple[-1], resized to target mask then thresholded"
def __init__(self, model_name: str) -> None:
super().__init__(model_name)
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "IFNet" / "pytorch version"
with repo_import_context(repo):
import torchvision.models as tv_models
original_vgg16 = tv_models.vgg16
def vgg16_no_download(*args: Any, **kwargs: Any) -> nn.Module:
kwargs.pop("pretrained", None)
kwargs["weights"] = None
return original_vgg16(*args, **kwargs)
tv_models.vgg16 = vgg16_no_download
try:
module = importlib.import_module("DSIFN")
model = module.DSIFN(module.vgg16_base(), module.vgg16_base())
finally:
tv_models.vgg16 = original_vgg16
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
if any(str(key).startswith("net.") for key in state):
model.load_state_dict(state, strict=True)
return
elif isinstance(checkpoint, dict) and "state_dict" in checkpoint:
state = checkpoint["state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"{self.model_name} checkpoint {checkpoint_path} has no recognized state dict.")
model.load_state_dict(state, strict=True)
def forward(self, model: nn.Module, batch: tuple, device: torch.device) -> Any:
a, b, _mask, _names = batch
return model(a.to(device, non_blocking=True), b.to(device, non_blocking=True))
def normalize_output(self, raw_output: Any, batch: tuple, dataset_config: dict) -> NormalizedOutput:
if not isinstance(raw_output, (list, tuple)) or not raw_output:
raise RuntimeError(f"{self.model_name} expected DSIFN tuple output, got {type(raw_output)!r}.")
score = raw_output[-1].detach().cpu()
target = batch[2]
if score.shape[-2:] != target.shape[-2:]:
score = F.interpolate(score, size=target.shape[-2:], mode="bilinear", align_corners=True)
threshold = float(dataset_config.get("eval", {}).get("threshold", 0.5))
binary, score = normalize_binary_prediction(score, threshold=threshold)
return NormalizedOutput(binary=binary, score=score, metric_tensor=score)
def get_dummy_inputs(self, dataset_config: dict, device: torch.device) -> tuple:
size = int(dataset_config.get("img_size", dataset_config.get("image_size", 256)))
return (
torch.zeros(1, 3, size, size, device=device),
torch.zeros(1, 3, size, size, device=device),
)
def get_threshold_mode(self) -> str:
return "threshold"
def compute_loss(self, raw_output: Any, batch: tuple, model_config: dict, dataset_config: dict, device: torch.device) -> dict[str, torch.Tensor]:
if not isinstance(raw_output, (list, tuple)) or not raw_output:
raise RuntimeError(f"{self.model_name} training expects five DSIFN outputs.")
target = batch[2].to(device, non_blocking=True).float()
repo = ROOT / "IFNet" / "pytorch version"
with repo_import_context(repo):
loss_module = importlib.import_module("loss")
cd_loss = loss_module.cd_loss
total = torch.zeros((), device=device)
for pred in raw_output:
pred = pred.float()
if pred.shape[-2:] != target.shape[-2:]:
pred = F.interpolate(pred, size=target.shape[-2:], mode="bilinear", align_corners=True)
total = total + cd_loss(pred.squeeze(1), target.squeeze(1))
loss = total / len(raw_output)
return {"loss": loss, "dsifn_loss": loss.detach()}
class SChangerModule(nn.Module):
def __init__(self, net: nn.Module) -> None:
super().__init__()
self.net = net
def forward(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
return self.net((a, b))
class SChangerAdapter(TwoTensorLogitAdapter):
model_name = "schanger"
notes_or_failure_reason = "SChanger-base from SChanger/cd/SChanger.py; eval forward returns sigmoid probability map."
model_class_path = "SChanger/cd/SChanger.py:SChanger"
output_format = "one-channel sigmoid probability map [B, 1, H, W] in eval mode"
checkpoint_format = "raw state dict, CD-Models model_state_dict, or upstream state_dict"
final_output_for_metrics = "eval-mode sigmoid probability thresholded with dataset eval threshold"
def get_threshold_mode(self) -> str:
return "threshold"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "SChanger"
with repo_import_context(repo):
module = importlib.import_module("cd.SChanger")
net = module.SChanger(
num_classes=1,
input_channels=3,
c_list=[8 * 3, 8 * 4, 8 * 6, 8 * 8, 8 * 13, 8 * 15],
dropout=0.2,
)
return SChangerModule(net).to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
if any(str(key).startswith(("net.", "trainer.")) for key in state):
model.load_state_dict(state, strict=True)
return
elif isinstance(checkpoint, dict) and "state_dict" in checkpoint:
state = checkpoint["state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"SChanger checkpoint {checkpoint_path} has no recognized state dict.")
model.net.load_state_dict(state, strict=True)
def compute_loss(self, raw_output: Any, batch: tuple, model_config: dict, dataset_config: dict, device: torch.device) -> dict[str, torch.Tensor]:
target = batch[2].to(device, non_blocking=True).float()
output = raw_output
if output.ndim != 4:
raise RuntimeError(f"SChanger training expected [B,C,H,W] output, got {tuple(output.shape)}.")
total = torch.zeros((), device=device)
for idx in range(output.shape[1]):
total = total + F.binary_cross_entropy_with_logits(output[:, idx:idx + 1], target)
loss = total / output.shape[1]
return {"loss": loss, "deep_supervision_bce": loss.detach()}
class Change3DBCDModule(nn.Module):
def __init__(self, trainer: nn.Module) -> None:
super().__init__()
self.trainer = trainer
def forward(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
return self.trainer.update_bcd(a, b)
class Change3DAdapter(TwoTensorLogitAdapter):
model_name = "change3d"
notes_or_failure_reason = "Change3D Trainer.update_bcd with X3D-L encoder and ChangeDecoder sigmoid output."
model_class_path = "Change3D/model/trainer.py:Trainer"
output_format = "one-channel sigmoid probability map [B, 1, H, W]"
checkpoint_format = "raw Trainer state dict, CD-Models model_state_dict, or upstream state_dict"
final_output_for_metrics = "update_bcd sigmoid probability thresholded with dataset eval threshold"
def get_threshold_mode(self) -> str:
return "threshold"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "Change3D"
size = int(dataset_config.get("img_size", dataset_config.get("image_size", 256)))
pretrained = repo / "model" / "X3D_L.pyth"
args = SimpleNamespace(
dataset=str(dataset_config.get("name", "")).upper() + "_CD",
in_height=size,
in_width=size,
num_perception_frame=1,
num_class=1,
pretrained=str(pretrained) if pretrained.exists() else None,
batch_size=int(dataset_config.get("batch_size", 1)),
lr=float(model_config.get("lr", 0.0002)),
lr_mode="poly",
max_steps=1,
max_epochs=1,
step_loss=100,
)
with repo_import_context(repo):
module = importlib.import_module("model.trainer")
trainer = module.Trainer(args)
return Change3DBCDModule(trainer).to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint:
state = checkpoint["model_state_dict"]
elif isinstance(checkpoint, dict) and "state_dict" in checkpoint:
state = checkpoint["state_dict"]
elif isinstance(checkpoint, dict) and all(torch.is_tensor(v) for v in checkpoint.values()):
state = checkpoint
else:
raise RuntimeError(f"Change3D checkpoint {checkpoint_path} has no recognized state dict.")
model.trainer.load_state_dict(state, strict=True)
def compute_loss(self, raw_output: Any, batch: tuple, model_config: dict, dataset_config: dict, device: torch.device) -> dict[str, torch.Tensor]:
target = batch[2].to(device, non_blocking=True).float()
output = raw_output.float()
if output.shape[-2:] != target.shape[-2:]:
output = F.interpolate(output, size=target.shape[-2:], mode="bilinear", align_corners=True)
bce = F.binary_cross_entropy(output.clamp(1e-6, 1 - 1e-6), target)
smooth = 1.0
intersection = torch.sum(output * target)
dice = 1.0 - ((2.0 * intersection + smooth) / (torch.sum(output) + torch.sum(target) + smooth))
loss = bce + dice
return {"loss": loss, "bce_loss": bce.detach(), "dice_loss": dice.detach()}
class ChangeMambaAdapter(TwoTensorLogitAdapter):
model_name = "changemamba"
supports_unified_training = False
notes_or_failure_reason = "ChangeMambaBCD with VMamba backbone; upstream training checkpoint stores weights under the model key."
model_class_path = "model_repos/ChangeMamba/changedetection/models/ChangeMambaBCD.py:ChangeMambaBCD"
output_format = "two-channel logits [B, 2, H, W]"
checkpoint_format = "ChangeMamba training checkpoint with model key, state_dict key, or raw state dict"
final_output_for_metrics = "raw two-channel logits"
def build_model(self, model_config: dict, dataset_config: dict, device: torch.device) -> nn.Module:
repo = ROOT / "model_repos" / "ChangeMamba"
variant = str(model_config.get("vmamba_variant", "small")).lower()
cfg_name = {
"tiny": "vssm_tiny_224_0229flex.yaml",
"small": "vssm_small_224.yaml",
"base": "vssm_base_224.yaml",
}.get(variant)
if cfg_name is None:
raise ValueError(f"Unsupported ChangeMamba VMamba variant: {variant}")
cfg_path = repo / "changedetection" / "configs" / "vssm1" / cfg_name
args = SimpleNamespace(
cfg=str(cfg_path),
opts=None,
batch_size=int(dataset_config.get("batch_size", 8)),
data_path=None,
zip=None,
cache_mode=None,
pretrained=None,
encoder_pretrained_path=None,
model_checkpoint_path=None,
resume=None,
resume_training_path=None,
accumulation_steps=None,
use_checkpoint=None,
disable_amp=None,
output=None,
tag=None,
eval=None,
throughput=None,
enable_amp=None,
fused_layernorm=None,
optim=None,
)
with repo_import_context(repo):
config_module = importlib.import_module("changedetection.configs.config")
utils_module = importlib.import_module("changedetection.script.script_utils")
model_module = importlib.import_module("changedetection.models.ChangeMambaBCD")
config = config_module.get_config(args)
model = model_module.ChangeMambaBCD(pretrained=None, **utils_module.get_vssm_kwargs(config))
return model.to(device)
def load_checkpoint(self, model: nn.Module, checkpoint_path: Path, device: torch.device) -> None:
# Load large upstream ChangeMamba training archives on CPU first.
# Direct CUDA deserialization can fail with PytorchStreamReader read
# errors on otherwise valid zip checkpoints, while load_state_dict will
# copy the extracted tensors into the already-placed CUDA model.
checkpoint = torch.load(checkpoint_path, map_location="cpu")
with repo_import_context(ROOT / "model_repos" / "ChangeMamba"):
checkpoint_module = importlib.import_module("changedetection.checkpoints")
state = checkpoint_module.extract_model_state_dict(checkpoint)
model.load_state_dict(state, strict=True)
_SUPPORTED = {
"fc_ef": FCAdapter,
"fc_siam_conc": FCAdapter,
"fc_siam_diff": FCAdapter,
"bifa": BiFAAdapter,
"bit_cd": BITAdapter,
"changeformer": ChangeFormerAdapter,
"changemamba": ChangeMambaAdapter,
"change3d": Change3DAdapter,
"siam_nestedunet": SiamNestedUNetAdapter,
"stanet": STANetAdapter,
"tinycd": TinyCDAdapter,
"dsamnet": DSAMNetAdapter,
"dsifn": DSIFNAdapter,
"ifnet": DSIFNAdapter,
"cgnet": CGNetAdapter,
"hanet": HANetAdapter,
"elgcnet": ELGCNetAdapter,
"schanger": SChangerAdapter,
}
_UNSUPPORTED_REASONS = {
"cdmamba": "CDMamba requires the cd-mamba-ssm environment with CUDA-visible mamba_ssm/Triton; cd-mamba-ssm is not available in this shell, and Mamba reports torch.cuda.is_available() == False.",
"changer": "Changer command construction is wired through Open-CD/OpenMMLab; verify runtime dependencies in cd-openmmlab before treating it as trainable.",
"rsm_cd": "RSM-CD requires the cd-mamba-ssm environment and VMamba/selective-scan CUDA extensions; cd-mamba-ssm is not available in this shell.",
}
def get_model_adapter(model_name: str) -> BaseModelAdapter:
if model_name in _SUPPORTED:
cls = _SUPPORTED[model_name]
return cls(model_name) if cls in {FCAdapter, DSIFNAdapter} else cls()
return UnsupportedAdapter(model_name, _UNSUPPORTED_REASONS.get(model_name, "No adapter registered."))