Simmonstt commited on
Commit
dc32e0a
·
verified ·
1 Parent(s): 17541bf

Fix model architecture: decoder_embed_dim=384, add norm layer, correct forward method

Browse files
Files changed (1) hide show
  1. 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': 512,
79
- 'decoder_depth': 8,
80
- 'decoder_num_heads': 16,
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.num_outputs = num_outputs
 
 
 
 
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, mc=None):
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
- # Encode
128
- x = self.encoder.forward_encoder(images, observed)
129
-
130
- # Pool
131
- if self.pool == 'mean':
132
- x = x.mean(dim=1) # [B, embed_dim]
133
- elif self.pool == 'cls':
134
- x = x[:, 0] # Use CLS token
 
 
 
 
 
 
135
  else:
136
- x = x[:, 0]
137
-
138
- # Head
139
- logits = self.head(x) # [B, num_outputs]
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