reworked iop_corr, added explainability tools and plotting tools, cleanup codebase

This commit is contained in:
rpotter6298
2026-03-12 16:12:59 +01:00
parent 13ad32683f
commit 6d0698d0fb
35 changed files with 7765 additions and 5723 deletions
+24
View File
@@ -219,6 +219,15 @@ def _set_single_phase(model: SingleEyeHT, phase: str) -> None:
# Ablation modes have no fusion bridge; fused_warmup is meaningless — treat as tower_warmup
if bridge_mode in ("image_only", "metadata_only") and phase == "fused_warmup":
phase = "tower_warmup"
if phase == "md_warmup":
_set_requires_grad(model.img_tower, False)
_set_requires_grad(model.md_tower, True)
_set_requires_grad(model.bridge.classifier_img, False)
_set_requires_grad(model.bridge.classifier_md, True)
_set_requires_grad(model.bridge.W_img, False)
_set_requires_grad(model.bridge.W_md, False)
_set_requires_grad(model.bridge.classifier_fused, False)
return
if phase == "tower_warmup":
_set_requires_grad(model.img_tower, bridge_mode != "metadata_only")
_set_requires_grad(model.md_tower, bridge_mode != "image_only")
@@ -282,12 +291,27 @@ def train_single_epoch(
x = batch.get("image_1")
m = batch.get("matrix_1")
y = batch.get("label_1")
if phase == "md_warmup":
if not torch.is_tensor(m):
continue
m = m.to(device)
y = _to_label_tensor(y, device)
md_feats = model.md_tower(m)
logits = model.bridge.classifier_md(md_feats)
loss = F.cross_entropy(logits, y)
opt.zero_grad(); loss.backward(); opt.step()
bs = y.shape[0]
total_loss += float(loss.item()) * bs
total_correct += int((logits.argmax(1) == y).sum())
total_n += bs
continue
if not torch.is_tensor(x) or not torch.is_tensor(m):
continue
x = x.to(device)
m = m.to(device)
y = _to_label_tensor(y, device)
bridge_mode = model.bridge.mode
img_feats = None if bridge_mode == "metadata_only" else model.img_tower(x)
md_feats = None if bridge_mode == "image_only" else model.md_tower(m)