pre-refactor 041426

This commit is contained in:
rpotter6298
2026-04-14 19:42:16 +02:00
parent eb9eafe715
commit 13290575d5
75 changed files with 8900 additions and 77 deletions
+19 -16
View File
@@ -14,6 +14,7 @@ class Bridge(nn.Module):
num_classes,
fusion_dim=256,
mode="fused",
dropout: float = 0.5,
use_se: bool = True,
se_reduction: int = 16,
se_pre_norm: bool = True,
@@ -37,7 +38,7 @@ class Bridge(nn.Module):
# heads
self.classifier_fused = nn.Sequential(
nn.ReLU(),
nn.Dropout(0.5),
nn.Dropout(dropout),
nn.Linear(fusion_dim, num_classes),
)
self.classifier_img = nn.Linear(img_dim, num_classes)
@@ -54,26 +55,28 @@ class Bridge(nn.Module):
return self.se_log.get(reset=reset)
return None
def _compute_fused(self, img_feats, md_feats):
"""Return z_fused embedding (before classifier_fused). Used by encode() and forward()."""
hi = self.ln_img(self.W_img(img_feats))
hm = self.ln_md(self.W_md(md_feats))
fused = hi * hm
if self.se is not None:
fused, gates = self.se(fused)
if self.se_log.enabled:
self.se_log.accumulate(gates)
return fused
def encode(self, img_feats, md_feats) -> torch.Tensor:
"""Return z_fused embedding without applying the classifier head."""
assert self.mode == "fused", "encode() only valid in fused mode"
return self._compute_fused(img_feats, md_feats)
def forward(self, img_feats, md_feats):
out_img = None if self.mode == "clinical_only" else self.classifier_img(img_feats)
out_md = None if self.mode == "image_only" else self.classifier_cd(md_feats)
if self.mode == "fused":
hi = self.ln_img(self.W_img(img_feats)) # image features
hm = self.ln_md(self.W_md(md_feats)) # clinical data features
fused = hi * hm # elementwise product
# apply SE gates
if self.se is not None:
fused, gates = self.se(fused)
if self.se_log.enabled:
self.se_log.accumulate(gates)
if self.se is not None and self.training and self.se_log.enabled:
if not hasattr(self, "_dbg_seen"):
self._dbg_seen = 0
if self._dbg_seen < 3: # print only a few times
print("[SE] gate mean this batch:", gates.mean().item())
self._dbg_seen += 1
fused = self._compute_fused(img_feats, md_feats)
out_f = self.classifier_fused(fused)
return out_f, out_img, out_md
# if ablation modes: