from __future__ import annotations import torch import torch.nn as nn from v3.classes.SE_attention import SEBlock, SEGateLogger class Bridge(nn.Module): def __init__( self, img_dim, meta_dim, num_classes, fusion_dim=256, mode="fused", dropout: float = 0.5, use_se: bool = True, se_reduction: int = 16, se_pre_norm: bool = True, ): super().__init__() self.mode = mode self.use_se = use_se # project towers to equal width self.W_img = nn.Linear(img_dim, fusion_dim) self.W_md = nn.Linear(meta_dim, fusion_dim) # optional: layernorm before SE self.ln_img = nn.LayerNorm(fusion_dim) if se_pre_norm else nn.Identity() self.ln_md = nn.LayerNorm(fusion_dim) if se_pre_norm else nn.Identity() # SE gate on the fused vector self.se = SEBlock(fusion_dim, reduction=se_reduction, residual=True) if use_se else None self.se_log = SEGateLogger(enabled=use_se, track_channels=False, dim=fusion_dim) # heads self.classifier_fused = nn.Sequential( nn.ReLU(), nn.Dropout(dropout), nn.Linear(fusion_dim, num_classes), ) self.classifier_img = nn.Linear(img_dim, num_classes) self.classifier_cd = nn.Linear(meta_dim, num_classes) def reset_se_stats(self): """Call at epoch start.""" if getattr(self, "se_log", None): self.se_log.reset() def get_se_stats(self, reset: bool = True): """Call after eval. Returns dict or None.""" if getattr(self, "se_log", None) and self.se_log.enabled: 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": fused = self._compute_fused(img_feats, md_feats) out_f = self.classifier_fused(fused) return out_f, out_img, out_md # if ablation modes: if self.mode == "image_only": return out_img, out_img, None if self.mode == "clinical_only": return out_md, None, out_md class VoteBridge(nn.Module): def __init__(self, num_classes): super().__init__() self.vote_combiner = nn.Linear(num_classes * 2, num_classes) # two sets of logits def forward(self, out_img, out_md): votes = torch.cat([out_img, out_md], dim=1) return self.vote_combiner(votes)