@@ -27,10 +27,19 @@ def is_fp4_marlin_supported():
2727 return current_platform .has_device_capability (75 )
2828
2929
30- def _nvfp4_compute_scale_factor (marlin_scales : torch .Tensor ) -> float :
30+ def _nvfp4_compute_scale_factor (
31+ marlin_scales : torch .Tensor ,
32+ a_dtype : torch .dtype | None = None ,
33+ ) -> float :
3134 """Compute the power-of-2 scale_factor needed so that all non-zero
3235 values in marlin_scales * 2^7 are >= 2 after rescaling.
3336 Returns a Python float (power of 2, >= 1.0)."""
37+
38+ # Since half has a smaller dynamic range compared to bfloat16,
39+ # no rescaling is applied here if active dtype is half.
40+ if a_dtype is not None and a_dtype == torch .half :
41+ return 1.0
42+
3443 ws_float = marlin_scales .float () * (2 ** 7 )
3544 nonzero_mask = ws_float > 0
3645 if nonzero_mask .any ():
@@ -44,6 +53,7 @@ def _nvfp4_compute_scale_factor(marlin_scales: torch.Tensor) -> float:
4453def nvfp4_marlin_process_scales (
4554 marlin_scales : torch .Tensor ,
4655 scale_factor : float | None = None ,
56+ a_dtype : torch .dtype | None = None ,
4757) -> tuple [torch .Tensor , float ]:
4858 """Process NVFP4 weight scales into the special S0E5M3 format for Marlin.
4959
@@ -91,7 +101,7 @@ def nvfp4_marlin_process_scales(
91101 # to fully utilize the E4M3 dynamic range (e.g., global_scale=1).
92102 # The caller must compensate by dividing global_scale by scale_factor.
93103 if scale_factor is None :
94- scale_factor = _nvfp4_compute_scale_factor (marlin_scales )
104+ scale_factor = _nvfp4_compute_scale_factor (marlin_scales , a_dtype )
95105 if scale_factor > 1.0 :
96106 marlin_scales = (marlin_scales .float () * scale_factor ).to (torch .half )
97107
@@ -119,12 +129,14 @@ def mxfp4_marlin_process_scales(marlin_scales, input_dtype=None):
119129 return marlin_scales
120130
121131
122- def nvfp4_marlin_process_global_scale (global_scale ):
123- assert global_scale .dtype in [torch .half , torch .bfloat16 ]
132+ def nvfp4_marlin_process_global_scale (global_scale , a_dtype : torch .dtype | None = None ):
133+ if a_dtype is None :
134+ a_dtype = global_scale .dtype
135+ assert a_dtype in [torch .half , torch .bfloat16 ]
124136 fp4_exponent = 2
125- if global_scale . dtype == torch .half :
137+ if a_dtype == torch .half :
126138 target_exponent = 5
127- elif global_scale . dtype == torch .bfloat16 :
139+ elif a_dtype == torch .bfloat16 :
128140 target_exponent = 8
129141 # exponent_bias_fp16 = 2 ** 4 - 2 ** 1 = 14
130142 # exponent_bias_bf16 = 2 ** 7 - 2 ** 1 = 126
@@ -244,11 +256,15 @@ def prepare_fp4_layer_for_marlin(
244256 )
245257
246258 if is_nvfp4 :
247- weight_scale , scale_factor = nvfp4_marlin_process_scales (weight_scale )
259+ weight_scale , scale_factor = nvfp4_marlin_process_scales (
260+ weight_scale , a_dtype = param_dtype
261+ )
248262 layer .weight_scale = torch .nn .Parameter (weight_scale , requires_grad = False )
249263
250- weight_global_scale = layer .weight_global_scale .to (param_dtype )
251- weight_global_scale = nvfp4_marlin_process_global_scale (weight_global_scale )
264+ weight_global_scale = layer .weight_global_scale .to (torch .float32 )
265+ weight_global_scale = nvfp4_marlin_process_global_scale (
266+ weight_global_scale , param_dtype
267+ )
252268 weight_global_scale = weight_global_scale / scale_factor
253269 layer .weight_global_scale = torch .nn .Parameter (
254270 weight_global_scale , requires_grad = False
@@ -339,7 +355,6 @@ def premute_scales(
339355 scales : torch .Tensor , g_scales : torch .Tensor , name : str
340356 ) -> tuple [torch .Tensor , torch .Tensor ]:
341357 scales = scales .to (param_dtype )
342- g_scales = g_scales .to (param_dtype )
343358
344359 tensor_list = []
345360 num_shards = 2 if is_act_and_mul else 1
@@ -350,7 +365,7 @@ def premute_scales(
350365
351366 # All experts share one global_scale, so compute the max
352367 # scale_factor across all experts first, then apply uniformly.
353- combined_scale_factor = _nvfp4_compute_scale_factor (scales )
368+ combined_scale_factor = _nvfp4_compute_scale_factor (scales , param_dtype )
354369
355370 for i in range (E ):
356371 scale = scales [i ].T
@@ -362,12 +377,12 @@ def premute_scales(
362377 is_a_8bit = is_a_8bit ,
363378 )
364379 marlin_scales , _ = nvfp4_marlin_process_scales (
365- marlin_scales , scale_factor = combined_scale_factor
380+ marlin_scales , scale_factor = combined_scale_factor , a_dtype = param_dtype
366381 )
367382 tensor_list .append (marlin_scales )
368383
369384 scales = torch .cat ([x .unsqueeze (0 ) for x in tensor_list ], 0 )
370- g_scales = nvfp4_marlin_process_global_scale (g_scales )
385+ g_scales = nvfp4_marlin_process_global_scale (g_scales , param_dtype )
371386 g_scales = g_scales / combined_scale_factor
372387 return scales , g_scales
373388
@@ -438,7 +453,7 @@ def prepare_moe_fp4_layer_for_marlin(
438453 scales = scales .view (torch .float8_e8m0fnu )
439454 scales = scales .to (param_dtype )
440455 if is_nvfp4 :
441- global_scale = getattr (layer , name + "_weight_scale_2" ). to ( param_dtype )
456+ global_scale = getattr (layer , name + "_weight_scale_2" )
442457
443458 tensor_list = []
444459 if "w13" in name :
@@ -449,7 +464,7 @@ def prepare_moe_fp4_layer_for_marlin(
449464 # For NVFP4: compute unified scale_factor across all experts
450465 combined_scale_factor = None
451466 if is_nvfp4 :
452- combined_scale_factor = _nvfp4_compute_scale_factor (scales )
467+ combined_scale_factor = _nvfp4_compute_scale_factor (scales , param_dtype )
453468
454469 for i in range (e ):
455470 scale = scales [i ].T
@@ -463,7 +478,9 @@ def prepare_moe_fp4_layer_for_marlin(
463478 )
464479 if is_nvfp4 :
465480 marlin_scales , _ = nvfp4_marlin_process_scales (
466- marlin_scales , scale_factor = combined_scale_factor
481+ marlin_scales ,
482+ scale_factor = combined_scale_factor ,
483+ a_dtype = param_dtype ,
467484 )
468485 else :
469486 marlin_scales = mxfp4_marlin_process_scales (
@@ -477,7 +494,7 @@ def prepare_moe_fp4_layer_for_marlin(
477494
478495 if is_nvfp4 :
479496 assert combined_scale_factor is not None
480- global_scale = nvfp4_marlin_process_global_scale (global_scale )
497+ global_scale = nvfp4_marlin_process_global_scale (global_scale , param_dtype )
481498 global_scale = global_scale / combined_scale_factor
482499 global_scale = torch .nn .Parameter (global_scale , requires_grad = False )
483500 setattr (layer , name + "_weight_scale_2" , global_scale )
@@ -665,7 +682,7 @@ def rand_marlin_weight_nvfp4_like(weight, group_size, input_dtype=None):
665682 )
666683 marlin_scales , scale_factor = nvfp4_marlin_process_scales (marlin_scales )
667684
668- global_scale = nvfp4_marlin_process_global_scale (global_scale )
685+ global_scale = nvfp4_marlin_process_global_scale (global_scale ). to ( torch . float32 )
669686 global_scale = global_scale / scale_factor
670687
671688 return weight_ref .T , marlin_qweight , marlin_scales , global_scale
0 commit comments