leochen085 commited on
Commit
bcd17bc
·
verified ·
1 Parent(s): 3b29d80

Fix dtype mismatches: cast entire model to bfloat16 in init

Browse files
Files changed (1) hide show
  1. modeling_slip.py +23 -13
modeling_slip.py CHANGED
@@ -161,18 +161,18 @@ class MultiSizePatchEmbed(nn.Module):
161
  return new_w, base_b, new_res_w, res_b
162
 
163
  def forward(self, x_list, attention_mask, time_idx):
164
- amp_dtype = torch.get_autocast_gpu_dtype() if torch.is_autocast_enabled() else torch.float32
165
- device = torch.device("cuda", torch.cuda.current_device()) if torch.cuda.is_available() else torch.device("cpu")
166
  sizes = torch.tensor([x.shape[-1] for x in x_list])
167
  unique_sizes = sizes.unique(sorted=True)
168
  N = x_list[0].shape[0]
169
- outputs = torch.empty(len(x_list), N, self.intermediate_size, device=device, dtype=amp_dtype)
170
- res_outputs = torch.empty(len(x_list), N, self.hidden_size, device=device, dtype=amp_dtype)
171
  for psize in unique_sizes.tolist():
172
  idxs = (sizes == psize).nonzero(as_tuple=True)[0]
173
- xs = torch.stack([x_list[i] for i in idxs]).to(device=device, non_blocking=True)
174
- mask = torch.stack([attention_mask[i] for i in idxs]).to(device=device, non_blocking=True)
175
- ti = torch.stack([time_idx[i] for i in idxs]).to(device=device, non_blocking=True)
176
  xs = torch.cat([xs, mask, ti], dim=-1)
177
  w, b, r_w, r_b = self.resize_weight(psize * 3)
178
  res_outputs[idxs] = F.linear(xs, r_w, r_b)
@@ -193,9 +193,10 @@ class PatchEmbedding(nn.Module):
193
  self.residual_layer = nn.Linear(patch_size * 3, hidden_size)
194
 
195
  def forward(self, x, mask, time_idx):
196
- x = rearrange(x, 'bs nvar (nump ps) -> (bs nvar) nump ps', ps=self.patch_size)
197
- mask = rearrange(mask, 'bs nvar (nump ps) -> (bs nvar) nump ps', ps=self.patch_size)
198
- time_idx = rearrange(time_idx, 'bs nvar (nump ps) -> (bs nvar) nump ps', ps=self.patch_size)
 
199
  x = torch.cat([x, mask, time_idx], dim=-1)
200
  return self.dropout(self.output_layer(self.act(self.hidden_layer(x)))) + self.residual_layer(x)
201
 
@@ -243,6 +244,9 @@ class CrossAttention(nn.Module):
243
  self.o_proj = nn.Linear(dim, dim, bias=False)
244
 
245
  def forward(self, query, context, attention_mask=None, **kwargs):
 
 
 
246
  bsz, q_len, _ = query.size()
247
  k_len = context.size(1)
248
  query = self.norm(query)
@@ -426,11 +430,11 @@ class Gemma3MultimodalModel(nn.Module):
426
  # Only init from pretrained when training from scratch.
427
  if init_from_pretrained:
428
  self.model = AutoModelForCausalLM.from_pretrained(
429
- model_id, dtype=torch.bfloat16, trust_remote_code=True)
430
  else:
431
  config = AutoConfig.from_pretrained(model_id, trust_remote_code=True)
432
  self.model = AutoModelForCausalLM.from_config(
433
- config, torch_dtype=torch.bfloat16, trust_remote_code=True)
434
 
435
  self.split_layer = split_layer
436
  hidden_size = self.model.config.hidden_size
@@ -442,10 +446,13 @@ class Gemma3MultimodalModel(nn.Module):
442
  dim=hidden_size, context_dim=hidden_size, num_heads=num_heads, dropout_rate=0.1)
443
  self.model.model.layers[i] = Gemma3MultimodalLayer(
444
  self.model.model.layers[i], Residual(cross_attn))
 
 
445
  self.to(torch.bfloat16)
446
 
447
  def condition_image(self, image_embeds):
448
- self.image_embeds = image_embeds.to(next(self.parameters()).device, dtype=torch.bfloat16)
 
449
  for layer in self.model.model.layers:
450
  if isinstance(layer, Gemma3MultimodalLayer):
451
  layer.condition_vis_x(self.image_embeds)
@@ -536,6 +543,9 @@ class SLIPModel(SLIPPreTrainedModel):
536
  self.temperature = nn.Parameter(torch.tensor(math.log(1 / 0.07)))
537
  self.temperature_max = math.log(1 / 0.07)
538
 
 
 
 
539
  def embed_sensor(self, sensors, sensor_attn_mask=None, time_index=None):
540
  sensor_tokens, attn_mask = self.sensor_encoder(sensors, sensor_attn_mask, time_index=time_index)
541
  if hasattr(self, "img_attn_pool"):
 
161
  return new_w, base_b, new_res_w, res_b
162
 
163
  def forward(self, x_list, attention_mask, time_idx):
164
+ device = self.shared_linear.weight.device
165
+ dtype = self.shared_linear.weight.dtype
166
  sizes = torch.tensor([x.shape[-1] for x in x_list])
167
  unique_sizes = sizes.unique(sorted=True)
168
  N = x_list[0].shape[0]
169
+ outputs = torch.empty(len(x_list), N, self.intermediate_size, device=device, dtype=dtype)
170
+ res_outputs = torch.empty(len(x_list), N, self.hidden_size, device=device, dtype=dtype)
171
  for psize in unique_sizes.tolist():
172
  idxs = (sizes == psize).nonzero(as_tuple=True)[0]
173
+ xs = torch.stack([x_list[i] for i in idxs]).to(device=device, dtype=dtype, non_blocking=True)
174
+ mask = torch.stack([attention_mask[i] for i in idxs]).to(device=device, dtype=dtype, non_blocking=True)
175
+ ti = torch.stack([time_idx[i] for i in idxs]).to(device=device, dtype=dtype, non_blocking=True)
176
  xs = torch.cat([xs, mask, ti], dim=-1)
177
  w, b, r_w, r_b = self.resize_weight(psize * 3)
178
  res_outputs[idxs] = F.linear(xs, r_w, r_b)
 
193
  self.residual_layer = nn.Linear(patch_size * 3, hidden_size)
194
 
195
  def forward(self, x, mask, time_idx):
196
+ dtype = self.hidden_layer.weight.dtype
197
+ x = rearrange(x, 'bs nvar (nump ps) -> (bs nvar) nump ps', ps=self.patch_size).to(dtype=dtype)
198
+ mask = rearrange(mask, 'bs nvar (nump ps) -> (bs nvar) nump ps', ps=self.patch_size).to(dtype=dtype)
199
+ time_idx = rearrange(time_idx, 'bs nvar (nump ps) -> (bs nvar) nump ps', ps=self.patch_size).to(dtype=dtype)
200
  x = torch.cat([x, mask, time_idx], dim=-1)
201
  return self.dropout(self.output_layer(self.act(self.hidden_layer(x)))) + self.residual_layer(x)
202
 
 
244
  self.o_proj = nn.Linear(dim, dim, bias=False)
245
 
246
  def forward(self, query, context, attention_mask=None, **kwargs):
247
+ dtype = self.q_proj.weight.dtype
248
+ query = query.to(dtype=dtype)
249
+ context = context.to(dtype=dtype)
250
  bsz, q_len, _ = query.size()
251
  k_len = context.size(1)
252
  query = self.norm(query)
 
430
  # Only init from pretrained when training from scratch.
431
  if init_from_pretrained:
432
  self.model = AutoModelForCausalLM.from_pretrained(
433
+ model_id, trust_remote_code=True)
434
  else:
435
  config = AutoConfig.from_pretrained(model_id, trust_remote_code=True)
436
  self.model = AutoModelForCausalLM.from_config(
437
+ config, trust_remote_code=True)
438
 
439
  self.split_layer = split_layer
440
  hidden_size = self.model.config.hidden_size
 
446
  dim=hidden_size, context_dim=hidden_size, num_heads=num_heads, dropout_rate=0.1)
447
  self.model.model.layers[i] = Gemma3MultimodalLayer(
448
  self.model.model.layers[i], Residual(cross_attn))
449
+
450
+ # Ensure all parameters (including new cross-attention modules) are bfloat16
451
  self.to(torch.bfloat16)
452
 
453
  def condition_image(self, image_embeds):
454
+ param = next(self.parameters())
455
+ self.image_embeds = image_embeds.to(device=param.device, dtype=param.dtype)
456
  for layer in self.model.model.layers:
457
  if isinstance(layer, Gemma3MultimodalLayer):
458
  layer.condition_vis_x(self.image_embeds)
 
543
  self.temperature = nn.Parameter(torch.tensor(math.log(1 / 0.07)))
544
  self.temperature_max = math.log(1 / 0.07)
545
 
546
+ # Cast entire model to bfloat16 (matches training checkpoint dtype)
547
+ self.to(torch.bfloat16)
548
+
549
  def embed_sensor(self, sensors, sensor_attn_mask=None, time_index=None):
550
  sensor_tokens, attn_mask = self.sensor_encoder(sensors, sensor_attn_mask, time_index=time_index)
551
  if hasattr(self, "img_attn_pool"):