Spaces:
Running on Zero
Running on Zero
Fix model architecture: decoder_embed_dim=384, add norm layer, correct forward method
Browse files- inference_engine.py +28 -20
inference_engine.py
CHANGED
|
@@ -69,15 +69,16 @@ TASKS = {
|
|
| 69 |
}
|
| 70 |
|
| 71 |
# 模型默认参数(与训练时一致)
|
|
|
|
| 72 |
DEFAULT_MODEL_ARGS = {
|
| 73 |
'img_size': 128,
|
| 74 |
'patch_size': 16,
|
| 75 |
'embed_dim': 768,
|
| 76 |
'depth': 12,
|
| 77 |
'num_heads': 12,
|
| 78 |
-
'decoder_embed_dim':
|
| 79 |
-
'decoder_depth':
|
| 80 |
-
'decoder_num_heads':
|
| 81 |
'pool': 'mean',
|
| 82 |
'dropout': 0.5,
|
| 83 |
}
|
|
@@ -106,7 +107,11 @@ class MultiMAE3DForDownstream(nn.Module):
|
|
| 106 |
super().__init__()
|
| 107 |
self.encoder = encoder
|
| 108 |
self.pool = pool
|
| 109 |
-
self.
|
|
|
|
|
|
|
|
|
|
|
|
|
| 110 |
|
| 111 |
# 预测头
|
| 112 |
self.head = nn.Sequential(
|
|
@@ -114,30 +119,33 @@ class MultiMAE3DForDownstream(nn.Module):
|
|
| 114 |
nn.Linear(embed_dim, num_outputs)
|
| 115 |
)
|
| 116 |
|
| 117 |
-
def forward(self, images, observed
|
| 118 |
"""
|
| 119 |
Args:
|
| 120 |
images: [B, 4, D, H, W] - 4 modalities
|
| 121 |
observed: [B, 4] - 0/1 mask for available modalities
|
| 122 |
-
mc: [B] - modality combination index (optional)
|
| 123 |
-
|
| 124 |
Returns:
|
| 125 |
logits: [B, num_outputs]
|
| 126 |
"""
|
| 127 |
-
#
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 135 |
else:
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
logits = self.head(
|
| 140 |
-
|
| 141 |
return logits
|
| 142 |
|
| 143 |
|
|
|
|
| 69 |
}
|
| 70 |
|
| 71 |
# 模型默认参数(与训练时一致)
|
| 72 |
+
# IMPORTANT: decoder_embed_dim must be divisible by 6 for 3D sincos position embedding
|
| 73 |
DEFAULT_MODEL_ARGS = {
|
| 74 |
'img_size': 128,
|
| 75 |
'patch_size': 16,
|
| 76 |
'embed_dim': 768,
|
| 77 |
'depth': 12,
|
| 78 |
'num_heads': 12,
|
| 79 |
+
'decoder_embed_dim': 384, # Must be divisible by 6 (384/6=64) ✓
|
| 80 |
+
'decoder_depth': 2, # Pretrain default
|
| 81 |
+
'decoder_num_heads': 12,
|
| 82 |
'pool': 'mean',
|
| 83 |
'dropout': 0.5,
|
| 84 |
}
|
|
|
|
| 107 |
super().__init__()
|
| 108 |
self.encoder = encoder
|
| 109 |
self.pool = pool
|
| 110 |
+
self.num_patches_per_modality = encoder.num_patches
|
| 111 |
+
self.num_global_tokens = encoder.num_global_tokens
|
| 112 |
+
|
| 113 |
+
# LayerNorm before head (important for checkpoint compatibility)
|
| 114 |
+
self.norm = nn.LayerNorm(embed_dim)
|
| 115 |
|
| 116 |
# 预测头
|
| 117 |
self.head = nn.Sequential(
|
|
|
|
| 119 |
nn.Linear(embed_dim, num_outputs)
|
| 120 |
)
|
| 121 |
|
| 122 |
+
def forward(self, images: torch.Tensor, observed: torch.Tensor) -> torch.Tensor:
|
| 123 |
"""
|
| 124 |
Args:
|
| 125 |
images: [B, 4, D, H, W] - 4 modalities
|
| 126 |
observed: [B, 4] - 0/1 mask for available modalities
|
|
|
|
|
|
|
| 127 |
Returns:
|
| 128 |
logits: [B, num_outputs]
|
| 129 |
"""
|
| 130 |
+
# encode() returns [B, 1 + 4*num_patches, embed_dim]
|
| 131 |
+
encoder_out = self.encoder.encode(images, observed)
|
| 132 |
+
|
| 133 |
+
if self.pool == 'cls':
|
| 134 |
+
features = encoder_out[:, 0] # CLS token -> [B, D]
|
| 135 |
+
elif self.pool == 'mean':
|
| 136 |
+
# Mean pool over modality tokens with masking for missing modalities
|
| 137 |
+
tokens = encoder_out[:, self.num_global_tokens:] # [B, 4*N_p, D]
|
| 138 |
+
B, _, D = tokens.shape
|
| 139 |
+
N = self.num_patches_per_modality
|
| 140 |
+
# Build per-token mask: repeat each modality's observed flag N times
|
| 141 |
+
mask = observed.unsqueeze(-1).expand(-1, -1, N) # [B, 4, N]
|
| 142 |
+
mask = mask.reshape(B, 4 * N).unsqueeze(-1) # [B, 4*N, 1]
|
| 143 |
+
features = (tokens * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1.0)
|
| 144 |
else:
|
| 145 |
+
raise ValueError(f"Unknown pool type: {self.pool}")
|
| 146 |
+
|
| 147 |
+
features = self.norm(features)
|
| 148 |
+
logits = self.head(features)
|
|
|
|
| 149 |
return logits
|
| 150 |
|
| 151 |
|