logging rework temp save

This commit is contained in:
rpotter6298
2026-03-04 09:40:09 +01:00
parent 080b5999fa
commit d36b9508ff
41 changed files with 2182 additions and 383 deletions
+209 -32
View File
@@ -34,6 +34,7 @@ class SingleEyeHT(nn.Module):
num_classes: int,
md_hidden_dim: int = 128,
fusion_dim: int = 256,
bridge_mode: str = "fused",
):
super().__init__()
self.img_tower = ImageTower(
@@ -52,7 +53,7 @@ class SingleEyeHT(nn.Module):
meta_dim=self.md_tower.out_dim,
num_classes=num_classes,
fusion_dim=fusion_dim,
mode="fused",
mode=bridge_mode,
use_se=False,
)
@@ -61,8 +62,8 @@ class SingleEyeHT(nn.Module):
return self.img_tower.transform
def forward(self, x: torch.Tensor, meta: torch.Tensor) -> torch.Tensor:
img_feats = self.img_tower(x)
md_feats = self.md_tower(meta)
img_feats = None if self.bridge.mode == "metadata_only" else self.img_tower(x)
md_feats = None if self.bridge.mode == "image_only" else self.md_tower(meta)
out_f, _, _ = self.bridge(img_feats, md_feats)
return out_f
@@ -214,11 +215,15 @@ def _set_requires_grad(module: nn.Module, enabled: bool) -> None:
def _set_single_phase(model: SingleEyeHT, phase: str) -> None:
bridge_mode = model.bridge.mode
# 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 == "tower_warmup":
_set_requires_grad(model.img_tower, True)
_set_requires_grad(model.md_tower, True)
_set_requires_grad(model.bridge.classifier_img, True)
_set_requires_grad(model.bridge.classifier_md, True)
_set_requires_grad(model.img_tower, bridge_mode != "metadata_only")
_set_requires_grad(model.md_tower, bridge_mode != "image_only")
_set_requires_grad(model.bridge.classifier_img, bridge_mode != "metadata_only")
_set_requires_grad(model.bridge.classifier_md, bridge_mode != "image_only")
_set_requires_grad(model.bridge.W_img, False)
_set_requires_grad(model.bridge.W_md, False)
_set_requires_grad(model.bridge.classifier_fused, False)
@@ -282,19 +287,33 @@ def train_single_epoch(
x = x.to(device)
m = m.to(device)
y = _to_label_tensor(y, device)
img_feats = model.img_tower(x)
md_feats = model.md_tower(m)
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)
if phase == "tower_warmup":
logits_i = model.bridge.classifier_img(img_feats)
logits_m = model.bridge.classifier_md(md_feats)
loss = 0.5 * (F.cross_entropy(logits_i, y) + F.cross_entropy(logits_m, y))
logits = 0.5 * (F.softmax(logits_i, dim=1) + F.softmax(logits_m, dim=1))
if bridge_mode == "metadata_only":
logits = model.bridge.classifier_md(md_feats)
loss = F.cross_entropy(logits, y)
elif bridge_mode == "image_only":
logits = model.bridge.classifier_img(img_feats)
loss = F.cross_entropy(logits, y)
else:
logits_i = model.bridge.classifier_img(img_feats)
logits_m = model.bridge.classifier_md(md_feats)
loss = 0.5 * (F.cross_entropy(logits_i, y) + F.cross_entropy(logits_m, y))
logits = 0.5 * (F.softmax(logits_i, dim=1) + F.softmax(logits_m, dim=1))
elif phase == "fused_warmup":
logits, _, _ = model.bridge(img_feats, md_feats)
loss = F.cross_entropy(logits, y)
else:
if random() < bcd_prob:
if bridge_mode == "metadata_only":
logits = model.bridge.classifier_md(md_feats)
loss = F.cross_entropy(logits, y)
elif bridge_mode == "image_only":
logits = model.bridge.classifier_img(img_feats)
loss = F.cross_entropy(logits, y)
elif random() < bcd_prob:
if random() < 0.5:
logits = model.bridge.classifier_img(img_feats)
else:
@@ -447,6 +466,79 @@ def collect_probs_classic(
return np.concatenate(y_chunks), np.concatenate(p_chunks, axis=0)
def collect_probs_ensemble_pereye(
model: "SingleEyeHT",
loader: DataLoader,
device: torch.device,
*,
return_ids: bool = False,
):
"""
Per-patient, per-eye probs for all 3 heads from a bilateral loader (ensemble mode).
OD corresponds to image_1/matrix_1; OS to image_2/matrix_2.
Arrays are in patient order (not interleaved at sample level).
Returns:
(y, pf_od, pi_od, pm_od, pf_os, pi_os, pm_os)
or, when return_ids=True:
(y, pf_od, pi_od, pm_od, pf_os, pi_os, pm_os, patient_ids)
Patient-level averaged ensemble probs can be recovered as:
p_en = 0.5 * (pf_od + pf_os)
"""
model.eval()
y_chunks: list = []
pf_od_c, pi_od_c, pm_od_c = [], [], []
pf_os_c, pi_os_c, pm_os_c = [], [], []
id_chunks: list[str] = []
with torch.no_grad():
for batch in loader:
x1 = batch.get("image_1"); m1 = batch.get("matrix_1")
x2 = batch.get("image_2"); m2 = batch.get("matrix_2")
y = batch.get("label_1")
if not (torch.is_tensor(x1) and torch.is_tensor(m1) and
torch.is_tensor(x2) and torch.is_tensor(m2)):
continue
y_t = _to_label_tensor(y, device)
def _fwd(x, m):
img_feats = None if model.bridge.mode == "metadata_only" else model.img_tower(x.to(device))
md_feats = None if model.bridge.mode == "image_only" else model.md_tower(m.to(device))
out_f, out_i, out_m = model.bridge(img_feats, md_feats)
pf = F.softmax(out_f, dim=1)
pi = F.softmax(out_i, dim=1) if out_i is not None else pf
pm = F.softmax(out_m, dim=1) if out_m is not None else pf
return pf, pi, pm
pf_od, pi_od, pm_od = _fwd(x1, m1)
pf_os, pi_os, pm_os = _fwd(x2, m2)
y_chunks.append(y_t.cpu().numpy())
pf_od_c.append(pf_od.cpu().numpy()); pi_od_c.append(pi_od.cpu().numpy()); pm_od_c.append(pm_od.cpu().numpy())
pf_os_c.append(pf_os.cpu().numpy()); pi_os_c.append(pi_os.cpu().numpy()); pm_os_c.append(pm_os.cpu().numpy())
if return_ids:
ids = batch.get("id_1", [""] * len(y_t))
if torch.is_tensor(ids):
ids = ids.tolist()
id_chunks.extend([str(i) for i in ids])
if not y_chunks:
z = np.zeros((0, 0), dtype=np.float32)
empty_i = np.array([], dtype=np.int64)
base = (empty_i, z, z, z, z, z, z)
return base + (np.array([], dtype=object),) if return_ids else base
y = np.concatenate(y_chunks)
pf_od = np.concatenate(pf_od_c, axis=0); pi_od = np.concatenate(pi_od_c, axis=0); pm_od = np.concatenate(pm_od_c, axis=0)
pf_os = np.concatenate(pf_os_c, axis=0); pi_os = np.concatenate(pi_os_c, axis=0); pm_os = np.concatenate(pm_os_c, axis=0)
if return_ids:
return y, pf_od, pi_od, pm_od, pf_os, pi_os, pm_os, np.array(id_chunks, dtype=object)
return y, pf_od, pi_od, pm_od, pf_os, pi_os, pm_os
def collect_probs_ensemble(
model: SingleEyeHT,
loader: DataLoader,
@@ -532,15 +624,22 @@ def collect_probs_single_components(
device: torch.device,
*,
aggregate_patient: bool,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
return_logits: bool = False,
):
"""
Collect fused/img/md probabilities for SingleEyeHT.
Collect fused/img/md probabilities (and optionally raw logits) for SingleEyeHT.
- aggregate_patient=False: eye-level (OD/OS as independent samples)
- aggregate_patient=True : patient-level (average OD/OS per head)
- return_logits=False: returns (y, probs_f, probs_i, probs_m)
- return_logits=True: returns (y, probs_f, probs_i, probs_m,
logits_f, logits_i, logits_m)
Note: logits are averaged across eyes when aggregate_patient=True,
which is equivalent to averaging in logit space (before softmax).
"""
model.eval()
y_chunks = []
pf_chunks, pi_chunks, pm_chunks = [], [], []
lf_chunks, li_chunks, lm_chunks = [], [], []
with torch.no_grad():
for batch in loader:
x1 = batch.get("image_1"); m1 = batch.get("matrix_1")
@@ -550,40 +649,118 @@ def collect_probs_single_components(
continue
y_t = _to_label_tensor(y, device)
def _per_eye_probs(x, m):
img_feats = model.img_tower(x.to(device))
md_feats = model.md_tower(m.to(device))
def _per_eye(x, m):
img_feats = None if model.bridge.mode == "metadata_only" else model.img_tower(x.to(device))
md_feats = None if model.bridge.mode == "image_only" else model.md_tower(m.to(device))
out_f, out_i, out_m = model.bridge(img_feats, md_feats)
return (
F.softmax(out_f, dim=1),
F.softmax(out_i, dim=1),
F.softmax(out_m, dim=1),
)
pf = F.softmax(out_f, dim=1)
pi = F.softmax(out_i, dim=1) if out_i is not None else pf
pm = F.softmax(out_m, dim=1) if out_m is not None else pf
lf = out_f
li = out_i if out_i is not None else out_f
lm = out_m if out_m is not None else out_f
return pf, pi, pm, lf, li, lm
pf_od, pi_od, pm_od = _per_eye_probs(x1, m1)
pf_os, pi_os, pm_os = _per_eye_probs(x2, m2)
pf_od, pi_od, pm_od, lf_od, li_od, lm_od = _per_eye(x1, m1)
pf_os, pi_os, pm_os, lf_os, li_os, lm_os = _per_eye(x2, m2)
if aggregate_patient:
y_chunks.append(y_t.cpu().numpy())
pf_chunks.append((0.5 * (pf_od + pf_os)).cpu().numpy())
pi_chunks.append((0.5 * (pi_od + pi_os)).cpu().numpy())
pm_chunks.append((0.5 * (pm_od + pm_os)).cpu().numpy())
lf_chunks.append((0.5 * (lf_od + lf_os)).cpu().numpy())
li_chunks.append((0.5 * (li_od + li_os)).cpu().numpy())
lm_chunks.append((0.5 * (lm_od + lm_os)).cpu().numpy())
else:
y_np = y_t.cpu().numpy()
y_chunks += [y_np, y_np]
pf_chunks += [pf_od.cpu().numpy(), pf_os.cpu().numpy()]
pi_chunks += [pi_od.cpu().numpy(), pi_os.cpu().numpy()]
pm_chunks += [pm_od.cpu().numpy(), pm_os.cpu().numpy()]
lf_chunks += [lf_od.cpu().numpy(), lf_os.cpu().numpy()]
li_chunks += [li_od.cpu().numpy(), li_os.cpu().numpy()]
lm_chunks += [lm_od.cpu().numpy(), lm_os.cpu().numpy()]
if not y_chunks:
z = np.zeros((0, 0), dtype=np.float32)
if return_logits:
return np.array([], dtype=np.int64), z, z, z, z, z, z
return np.array([], dtype=np.int64), z, z, z
return (
np.concatenate(y_chunks),
np.concatenate(pf_chunks, axis=0),
np.concatenate(pi_chunks, axis=0),
np.concatenate(pm_chunks, axis=0),
)
y = np.concatenate(y_chunks)
pf = np.concatenate(pf_chunks, axis=0)
pi = np.concatenate(pi_chunks, axis=0)
pm = np.concatenate(pm_chunks, axis=0)
if return_logits:
lf = np.concatenate(lf_chunks, axis=0)
li = np.concatenate(li_chunks, axis=0)
lm = np.concatenate(lm_chunks, axis=0)
return y, pf, pi, pm, lf, li, lm
return y, pf, pi, pm
def collect_probs_eye_level(
model: "SingleEyeHT",
loader: DataLoader,
device: torch.device,
*,
return_ids: bool = False,
):
"""
Collect fused/img/md probabilities from a single-eye loader (image_1/matrix_1 only).
Used for eval-mode passes over the training set.
Returns (y, probs_f, probs_i, probs_m) or, when return_ids=True,
(y, probs_f, probs_i, probs_m, sample_ids) where sample_ids is an
array of strings like "2OD", "4OS".
"""
model.eval()
y_chunks, pf_chunks, pi_chunks, pm_chunks, id_chunks = [], [], [], [], []
with torch.no_grad():
for batch in loader:
x = batch.get("image_1")
m = batch.get("matrix_1")
y = batch.get("label_1")
if not (torch.is_tensor(x) and torch.is_tensor(m)):
continue
y_t = _to_label_tensor(y, device)
img_feats = None if model.bridge.mode == "metadata_only" else model.img_tower(x.to(device))
md_feats = None if model.bridge.mode == "image_only" else model.md_tower(m.to(device))
out_f, out_i, out_m = model.bridge(img_feats, md_feats)
pf = F.softmax(out_f, dim=1)
pi = F.softmax(out_i, dim=1) if out_i is not None else pf
pm = F.softmax(out_m, dim=1) if out_m is not None else pf
y_chunks.append(y_t.cpu().numpy())
pf_chunks.append(pf.cpu().numpy())
pi_chunks.append(pi.cpu().numpy())
pm_chunks.append(pm.cpu().numpy())
if return_ids:
ids = batch.get("id_1", [""] * len(y_t))
eyes = batch.get("eye_id_1", [""] * len(y_t))
# ids/eyes may be tensors (int) or lists of strings
if torch.is_tensor(ids):
ids = ids.tolist()
if torch.is_tensor(eyes):
eyes = eyes.tolist()
id_chunks.extend(
[f"{pid}{eye}" for pid, eye in zip(ids, eyes)]
)
if not y_chunks:
z = np.zeros((0, 0), dtype=np.float32)
empty_ids = np.array([], dtype=object)
if return_ids:
return np.array([], dtype=np.int64), z, z, z, empty_ids
return np.array([], dtype=np.int64), z, z, z
y = np.concatenate(y_chunks)
pf = np.concatenate(pf_chunks, axis=0)
pi = np.concatenate(pi_chunks, axis=0)
pm = np.concatenate(pm_chunks, axis=0)
if return_ids:
return y, pf, pi, pm, np.array(id_chunks, dtype=object)
return y, pf, pi, pm
def collect_probs_bilateral_components(