pre-refactor 041426
This commit is contained in:
+19
-16
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user