Fix dtype mismatches: cast entire model to bfloat16 in init
Browse files- 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 |
-
|
| 165 |
-
|
| 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=
|
| 170 |
-
res_outputs = torch.empty(len(x_list), N, self.hidden_size, device=device, 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 |
-
|
| 197 |
-
|
| 198 |
-
|
|
|
|
| 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,
|
| 430 |
else:
|
| 431 |
config = AutoConfig.from_pretrained(model_id, trust_remote_code=True)
|
| 432 |
self.model = AutoModelForCausalLM.from_config(
|
| 433 |
-
config,
|
| 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 |
-
|
|
|
|
| 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"):
|