temp / src /models /fusion.py
minhquang47's picture
Upload 34 files
0c7d9c9 verified
Raw
History Blame Contribute Delete
10.1 kB
from dataclasses import dataclass
from typing import Dict, Optional, Tuple
import torch
import torch.nn as nn
@dataclass
class StructuralFusionConfig:
"""
Cấu hình cho StructuralVisualFusion.
Args:
d_model: Hidden dimension của visual tokens từ ViT/VLM.
topo_feature_dim: Số chiều feature topology/morphology mỗi patch.
hidden_dim: Hidden dimension của MLP chiếu topo_features -> d_model.
prior_strength: Hệ số spatial modulation ban đầu cho lesion prior.
topo_strength: Hệ số enrichment ban đầu cho topo embedding.
dropout: Dropout trong topo MLP.
use_layer_norm: Chuẩn hóa visual/topo branch để ổn định training.
learnable_gates: Nếu True, alpha/beta là tham số học được.
normalize_prior: Nếu True, normalize prior_mask về [0, 1] theo từng sample.
residual_scale: Scale residual cuối để tránh làm lệch mạnh pretrained tokens.
"""
d_model: int
topo_feature_dim: int
hidden_dim: Optional[int] = None
prior_strength: float = 0.30
topo_strength: float = 0.10
dropout: float = 0.10
use_layer_norm: bool = True
learnable_gates: bool = True
normalize_prior: bool = True
residual_scale: float = 1.0
def __post_init__(self):
if self.d_model <= 0:
raise ValueError("d_model phải > 0")
if self.topo_feature_dim <= 0:
raise ValueError("topo_feature_dim phải > 0")
if self.hidden_dim is None:
self.hidden_dim = max(self.d_model // 2, self.topo_feature_dim * 4)
if self.hidden_dim <= 0:
raise ValueError("hidden_dim phải > 0")
class StructuralVisualFusion(nn.Module):
"""
Fusion module nối ViT visual tokens với lesion prior và topology features.
Forward logic:
V_fused = V * (1 + alpha * P) + beta * MLP(F_topo)
Trong đó:
V: visual_tokens, shape [B, N, D]
P: prior_mask, shape [B, N] hoặc [B, H, W]
F_topo: topo_features, shape [B, N, F] hoặc [B, H, W, F]
Output chính:
fused_tokens, shape [B, N, D]
"""
def __init__(self, config: StructuralFusionConfig):
super().__init__()
self.config = config
self.visual_norm = nn.LayerNorm(config.d_model) if config.use_layer_norm else nn.Identity()
self.topo_norm = nn.LayerNorm(config.topo_feature_dim) if config.use_layer_norm else nn.Identity()
self.topo_projector = nn.Sequential(
nn.Linear(config.topo_feature_dim, config.hidden_dim),
nn.GELU(),
nn.Dropout(config.dropout),
nn.Linear(config.hidden_dim, config.d_model),
nn.Dropout(config.dropout),
)
if config.learnable_gates:
self.prior_gate = nn.Parameter(torch.tensor(float(config.prior_strength)))
self.topo_gate = nn.Parameter(torch.tensor(float(config.topo_strength)))
else:
self.register_buffer("prior_gate", torch.tensor(float(config.prior_strength)))
self.register_buffer("topo_gate", torch.tensor(float(config.topo_strength)))
self.output_norm = nn.LayerNorm(config.d_model) if config.use_layer_norm else nn.Identity()
def forward(
self,
visual_tokens: torch.Tensor,
prior_mask: Optional[torch.Tensor] = None,
topo_features: Optional[torch.Tensor] = None,
return_diagnostics: bool = True,
):
"""
Args:
visual_tokens: Tensor [B, N, D].
prior_mask: Optional tensor [B, N] hoặc [B, H, W].
topo_features: Optional tensor [B, N, F] hoặc [B, H, W, F].
return_diagnostics: Nếu True trả về dict debug; nếu False chỉ trả fused_tokens.
Returns:
dict hoặc Tensor:
{
"fused_tokens": [B, N, D],
"prior_weights": [B, N, 1] | None,
"topo_embedding": [B, N, D] | None,
"prior_gate": scalar,
"topo_gate": scalar,
}
"""
self._validate_visual_tokens(visual_tokens)
batch_size, num_tokens, d_model = visual_tokens.shape
fused_tokens = self.visual_norm(visual_tokens)
prior_weights = None
topo_embedding = None
if prior_mask is not None:
prior_weights = self._prepare_prior_mask(
prior_mask,
batch_size=batch_size,
num_tokens=num_tokens,
device=visual_tokens.device,
dtype=visual_tokens.dtype,
)
fused_tokens = fused_tokens * (1.0 + self.prior_gate * prior_weights)
if topo_features is not None:
topo_features = self._prepare_topo_features(
topo_features,
batch_size=batch_size,
num_tokens=num_tokens,
device=visual_tokens.device,
dtype=visual_tokens.dtype,
)
topo_features = self.topo_norm(topo_features)
topo_embedding = self.topo_projector(topo_features)
fused_tokens = fused_tokens + self.topo_gate * topo_embedding
if self.config.residual_scale != 1.0:
fused_tokens = visual_tokens + self.config.residual_scale * (fused_tokens - visual_tokens)
fused_tokens = self.output_norm(fused_tokens)
if not return_diagnostics:
return fused_tokens
return {
"fused_tokens": fused_tokens,
"prior_weights": prior_weights,
"topo_embedding": topo_embedding,
"prior_gate": self.prior_gate.detach().clone(),
"topo_gate": self.topo_gate.detach().clone(),
}
def _validate_visual_tokens(self, visual_tokens: torch.Tensor) -> None:
if not torch.is_tensor(visual_tokens):
raise TypeError("visual_tokens phải là torch.Tensor")
if visual_tokens.ndim != 3:
raise ValueError(
f"visual_tokens phải có shape [B, N, D], nhận {tuple(visual_tokens.shape)}"
)
if visual_tokens.shape[-1] != self.config.d_model:
raise ValueError(
f"visual_tokens D={visual_tokens.shape[-1]} không khớp d_model={self.config.d_model}"
)
def _prepare_prior_mask(
self,
prior_mask: torch.Tensor,
batch_size: int,
num_tokens: int,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor:
if not torch.is_tensor(prior_mask):
prior_mask = torch.as_tensor(prior_mask)
prior_mask = prior_mask.to(device=device, dtype=dtype)
if prior_mask.ndim == 3:
prior_mask = prior_mask.flatten(start_dim=1)
elif prior_mask.ndim == 2:
pass
else:
raise ValueError(
f"prior_mask phải có shape [B, N] hoặc [B, H, W], nhận {tuple(prior_mask.shape)}"
)
if prior_mask.shape[0] != batch_size:
raise ValueError(
f"prior_mask batch={prior_mask.shape[0]} không khớp visual batch={batch_size}"
)
if prior_mask.shape[1] != num_tokens:
raise ValueError(
f"prior_mask N={prior_mask.shape[1]} không khớp visual tokens N={num_tokens}"
)
if self.config.normalize_prior:
prior_mask = self._minmax_normalize_per_sample(prior_mask)
return prior_mask.unsqueeze(-1)
def _prepare_topo_features(
self,
topo_features: torch.Tensor,
batch_size: int,
num_tokens: int,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor:
if not torch.is_tensor(topo_features):
topo_features = torch.as_tensor(topo_features)
topo_features = topo_features.to(device=device, dtype=dtype)
if topo_features.ndim == 4:
topo_features = topo_features.flatten(start_dim=1, end_dim=2)
elif topo_features.ndim == 3:
pass
else:
raise ValueError(
"topo_features phải có shape [B, N, F] hoặc [B, H, W, F], "
f"nhận {tuple(topo_features.shape)}"
)
if topo_features.shape[0] != batch_size:
raise ValueError(
f"topo_features batch={topo_features.shape[0]} không khớp visual batch={batch_size}"
)
if topo_features.shape[1] != num_tokens:
raise ValueError(
f"topo_features N={topo_features.shape[1]} không khớp visual tokens N={num_tokens}"
)
if topo_features.shape[2] != self.config.topo_feature_dim:
raise ValueError(
f"topo_features F={topo_features.shape[2]} không khớp "
f"topo_feature_dim={self.config.topo_feature_dim}"
)
return topo_features
@staticmethod
def _minmax_normalize_per_sample(values: torch.Tensor, eps: float = 1e-8) -> torch.Tensor:
min_values = values.amin(dim=1, keepdim=True)
max_values = values.amax(dim=1, keepdim=True)
norm = (values - min_values) / (max_values - min_values + eps)
return torch.nan_to_num(norm, nan=0.0, posinf=0.0, neginf=0.0)
def build_structural_visual_fusion(
d_model: int,
topo_feature_dim: int,
**kwargs,
) -> StructuralVisualFusion:
"""Factory tiện dụng cho training scripts."""
config = StructuralFusionConfig(
d_model=d_model,
topo_feature_dim=topo_feature_dim,
**kwargs,
)
return StructuralVisualFusion(config)
__all__ = [
"StructuralFusionConfig",
"StructuralVisualFusion",
"build_structural_visual_fusion",
]