Akahsizrr commited on
Commit
267ee66
·
verified ·
1 Parent(s): 6eefb37

Bake runtime fixes into model code: SwiGLU clamp, router stability, from_pretrained auto-fix

Browse files
Files changed (1) hide show
  1. fuse2_model.py +75 -2
fuse2_model.py CHANGED
@@ -66,6 +66,11 @@ class SwiGLUExpert(nn.Module):
66
  down_proj: (hidden, intermediate)
67
  """
68
 
 
 
 
 
 
69
  def __init__(self, hidden_size: int, intermediate_size: int):
70
  super().__init__()
71
  self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
@@ -73,7 +78,9 @@ class SwiGLUExpert(nn.Module):
73
  self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
74
 
75
  def forward(self, x: torch.Tensor) -> torch.Tensor:
76
- return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
 
 
77
 
78
 
79
  class Fuse2Router(nn.Module):
@@ -115,8 +122,11 @@ class Fuse2Router(nn.Module):
115
  router_logits: (batch*seq, num_experts) — raw logits for load balancing
116
  """
117
  # sqrtsoftplus scoring (from DeepSeek V4)
 
 
 
118
  logits = self.gate(hidden_states) # (tokens, num_experts)
119
- scores = F.softplus(logits).sqrt()
120
 
121
  # Top-k selection
122
  topk_weights, topk_indices = scores.topk(self.top_k, dim=-1)
@@ -514,6 +524,69 @@ class Fuse2ForCausalLM(Qwen3ForCausalLM):
514
  **kwargs,
515
  )
516
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
517
 
518
  def load_expert_weights(
519
  model: Fuse2ForCausalLM,
 
66
  down_proj: (hidden, intermediate)
67
  """
68
 
69
+ # DeepSeek V4 Flash uses swiglu_limit=10.0 to clamp intermediate
70
+ # activations. Without this, outlier values grow exponentially across
71
+ # 36 layers and produce NaN by layer 6.
72
+ SWIGLU_LIMIT = 10.0
73
+
74
  def __init__(self, hidden_size: int, intermediate_size: int):
75
  super().__init__()
76
  self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
 
78
  self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
79
 
80
  def forward(self, x: torch.Tensor) -> torch.Tensor:
81
+ gate_up = F.silu(self.gate_proj(x)) * self.up_proj(x)
82
+ gate_up = gate_up.clamp(-self.SWIGLU_LIMIT, self.SWIGLU_LIMIT)
83
+ return self.down_proj(gate_up)
84
 
85
 
86
  class Fuse2Router(nn.Module):
 
122
  router_logits: (batch*seq, num_experts) — raw logits for load balancing
123
  """
124
  # sqrtsoftplus scoring (from DeepSeek V4)
125
+ # Clamp softplus to min=1e-6 before sqrt to prevent NaN gradients
126
+ # when logits are very negative (softplus → 0 → sqrt(0) = 0, but
127
+ # gradient sqrt'(0) = inf).
128
  logits = self.gate(hidden_states) # (tokens, num_experts)
129
+ scores = F.softplus(logits).clamp(min=1e-6).sqrt()
130
 
131
  # Top-k selection
132
  topk_weights, topk_indices = scores.topk(self.top_k, dim=-1)
 
524
  **kwargs,
525
  )
526
 
527
+ @classmethod
528
+ def from_pretrained(cls, *args, **kwargs):
529
+ """Load from HuggingFace Hub with automatic runtime fixes.
530
+
531
+ This overrides the default from_pretrained to apply three critical
532
+ fixes after weight loading:
533
+
534
+ 1. Initialize coding_gate and coding_norm if they're still on meta
535
+ device (these params are not in the safetensors checkpoint).
536
+ 2. Cast all RMSNorm/LayerNorm weights from float32 to bfloat16 to
537
+ enable fused SDPA kernel dispatch (otherwise falls back to slow
538
+ Python loops).
539
+ 3. Ensure coding_enabled is True (config default).
540
+
541
+ With these fixes, from_pretrained produces a working model without
542
+ any manual post-load patching.
543
+ """
544
+ model = super().from_pretrained(*args, **kwargs)
545
+ model._apply_runtime_fixes()
546
+ return model
547
+
548
+ def _apply_runtime_fixes(self):
549
+ """Apply runtime fixes after weight loading.
550
+
551
+ Called automatically by from_pretrained. Can also be called manually
552
+ if the model was loaded via a custom path (e.g., init_empty_weights +
553
+ manual safetensors loading).
554
+ """
555
+ device = next(self.parameters()).device
556
+ fixed_meta = 0
557
+ fixed_norms = 0
558
+
559
+ for layer in self.model.layers:
560
+ if not isinstance(layer, Fuse2AugmentedLayer):
561
+ continue
562
+
563
+ # Fix 1: coding_gate on meta device → initialize to -2.0
564
+ if hasattr(layer, 'coding_gate'):
565
+ if layer.coding_gate.device.type == 'meta':
566
+ layer.coding_gate = nn.Parameter(
567
+ torch.tensor(-2.0, device=device))
568
+ fixed_meta += 1
569
+
570
+ # Fix 2: coding_norm on meta device → create fresh RMSNorm
571
+ if hasattr(layer, 'coding_norm'):
572
+ if hasattr(layer.coding_norm, 'weight') and \
573
+ layer.coding_norm.weight.device.type == 'meta':
574
+ layer.coding_norm = nn.RMSNorm(
575
+ layer.coding_norm.weight.shape[0], eps=1e-6).to(device)
576
+ fixed_meta += 1
577
+
578
+ # Fix 3: Cast float32 norm weights to bfloat16 for fused kernels
579
+ for module in self.modules():
580
+ if hasattr(module, 'weight') and hasattr(module, 'eps'):
581
+ if module.weight.dtype == torch.float32:
582
+ module.weight.data = module.weight.data.to(torch.bfloat16)
583
+ fixed_norms += 1
584
+
585
+ # Ensure coding is enabled
586
+ self.set_coding_enabled(True)
587
+
588
+ return {"meta_params_fixed": fixed_meta, "norms_cast_to_bf16": fixed_norms}
589
+
590
 
591
  def load_expert_weights(
592
  model: Fuse2ForCausalLM,