fix: RoPE inv_freq fp32 construction + clean-copy self-heal (NPU state-dict load corruption)

#4
Files changed (1) hide show
  1. modeling_moss_vl.py +23 -5
modeling_moss_vl.py CHANGED
@@ -209,11 +209,14 @@ class MossVLVisionRotaryEmbedding(nn.Module):
209
 
210
  def __init__(self, dim: int, theta: float = 10000.0) -> None:
211
  super().__init__()
212
- inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float) / dim))
213
  self.register_buffer("inv_freq", inv_freq, persistent=False)
 
214
 
215
  def forward(self, seqlen: int) -> torch.Tensor:
216
- seq = torch.arange(seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype)
 
 
217
  freqs = torch.outer(seq, self.inv_freq)
218
  return freqs
219
 
@@ -445,12 +448,23 @@ class MossVLTextRotaryEmbedding(nn.Module):
445
  self.original_max_seq_len = config.max_position_embeddings
446
 
447
  self.config = config
448
- self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]
449
-
450
- inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
 
 
 
 
 
 
451
  self.register_buffer("inv_freq", inv_freq, persistent=False)
452
  self.original_inv_freq = self.inv_freq
453
 
 
 
 
 
 
454
 
455
  if hasattr(config, "rope_scaling") and config.rope_scaling is not None:
456
  self.mrope_section = config.rope_scaling.get("mrope_section", [24, 20, 20])
@@ -477,6 +491,10 @@ class MossVLTextRotaryEmbedding(nn.Module):
477
  @torch.no_grad()
478
  @dynamic_rope_update
479
  def forward(self, x, position_ids):
 
 
 
 
480
 
481
  if position_ids.ndim == 2:
482
  position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1)
 
209
 
210
  def __init__(self, dim: int, theta: float = 10000.0) -> None:
211
  super().__init__()
212
+ inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float, device="cpu") / dim))
213
  self.register_buffer("inv_freq", inv_freq, persistent=False)
214
+ self._inv_freq_clean = inv_freq.clone()
215
 
216
  def forward(self, seqlen: int) -> torch.Tensor:
217
+ if torch.isnan(self.inv_freq).any() or torch.isinf(self.inv_freq).any():
218
+ self.inv_freq = self._inv_freq_clean.to(self.inv_freq.device)
219
+ seq = torch.arange(seqlen, dtype=self.inv_freq.dtype, device=self.inv_freq.device)
220
  freqs = torch.outer(seq, self.inv_freq)
221
  return freqs
222
 
 
448
  self.original_max_seq_len = config.max_position_embeddings
449
 
450
  self.config = config
451
+ if self.rope_type == "default" or self.rope_type not in ROPE_INIT_FUNCTIONS:
452
+ import torch as _torch
453
+ base = config.rope_theta if hasattr(config, "rope_theta") else 10000.0
454
+ head_dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads
455
+ inv_freq = 1.0 / (base ** (_torch.arange(0, head_dim, 2, device="cpu").float() / head_dim))
456
+ self.attention_scaling = 1.0
457
+ else:
458
+ self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]
459
+ inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
460
  self.register_buffer("inv_freq", inv_freq, persistent=False)
461
  self.original_inv_freq = self.inv_freq
462
 
463
+ # Guard: from_pretrained may corrupt inv_freq with garbage from the
464
+ # state dict (persistent=False buffers can still be overwritten during
465
+ # weight loading on some transformers versions). Re-validate on forward.
466
+ self._inv_freq_computed = inv_freq.clone()
467
+
468
 
469
  if hasattr(config, "rope_scaling") and config.rope_scaling is not None:
470
  self.mrope_section = config.rope_scaling.get("mrope_section", [24, 20, 20])
 
491
  @torch.no_grad()
492
  @dynamic_rope_update
493
  def forward(self, x, position_ids):
494
+ # Guard: from_pretrained can corrupt inv_freq — restore from clean copy
495
+ if torch.isnan(self.inv_freq).any() or torch.isinf(self.inv_freq).any():
496
+ self.inv_freq = self._inv_freq_computed.to(self.inv_freq.device)
497
+ self.original_inv_freq = self._inv_freq_computed.to(self.inv_freq.device)
498
 
499
  if position_ids.ndim == 2:
500
  position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1)