1. Per row FP8 wrongly fell to NVFP4 conversion, as NVFP4 checked on
   dims alone. Also check on dtype for FP8
2. Need to reshape FP8 QKV projections scales in the same way that
   weights are reshaped
This commit is contained in:
Oliver Simons
2026-09-30 15:55:51 +02:00
parent 6541a7508d
commit 60dcb35afe
2 changed files with 27 additions and 2 deletions
+9 -2
View File
@@ -595,6 +595,9 @@ class ModelBase:
for name, value in new_tensors.items():
self.model_tensors[name] = value
def _transform_fp8_scale(self, name: str, scale: Tensor) -> Tensor:
return scale
def _prepare_fp8_e4m3_tensors(self):
if self._fp8_as_q8:
return
@@ -623,6 +626,8 @@ class ModelBase:
if scale.numel() != 1:
continue
scale = self._transform_fp8_scale(weight_name, scale)
weight_prefix = weight_name.removesuffix(".weight")
# Transformers fine-grained FP8 uses activation_scale while ModelOpt uses input_scale.
# Ref: https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/models/mistral3.py#L357-L361
@@ -971,9 +976,11 @@ class ModelBase:
weight = LazyTorchTensor.to_eager(self.model_tensors[name]())
scale = LazyTorchTensor.to_eager(self.model_tensors[scale_name]())
# Skip non-NVFP4 tensors (e.g. FP8 with per-channel 1D scales)
if scale.ndim < 2:
# Leave FP8 weights and their scales for _prepare_fp8_e4m3_tensors.
if weight.dtype in (torch.float8_e4m3fn, torch.float8_e5m2) or scale.ndim < 2:
continue
if weight.ndim != 2 or scale.shape[0] != weight.shape[0] or scale.shape[1] * 8 != weight.shape[1]:
raise ValueError(f"NVFP4 weight {name!r} has incompatible shapes: weight {list(weight.shape)}, scale {list(scale.shape)}")
scale2 = LazyTorchTensor.to_eager(self.model_tensors.get(scale2_name, lambda: torch.tensor(1.0))())
input_scale = LazyTorchTensor.to_eager(self.model_tensors.get(input_scale_name, lambda: torch.tensor(1.0))())
+18
View File
@@ -475,6 +475,24 @@ class _LinearAttentionVReorderBase(Qwen3NextModel):
perm[dim], perm[dim + 1] = perm[dim + 1], perm[dim]
return tensor.permute(*perm).contiguous().reshape(*shape)
def _transform_fp8_scale(self, name: str, scale: Tensor) -> Tensor:
num_k_heads = self.hparams.get("linear_num_key_heads", 0)
num_v_heads = self.hparams.get("linear_num_value_heads", 0)
if scale.numel() == 1 or num_k_heads == 0 or num_v_heads == 0 or num_k_heads == num_v_heads:
return scale
num_v_per_k = num_v_heads // num_k_heads
head_v_dim = self.hparams["linear_value_head_dim"]
if name.endswith(".linear_attn.in_proj_qkv.weight"):
qk_dim = 2 * self.hparams["linear_key_head_dim"] * num_k_heads
v_scale = self._reorder_v_heads(scale[qk_dim:], 0, num_k_heads, num_v_per_k, head_v_dim)
return torch.cat([scale[:qk_dim], v_scale])
if name.endswith(".linear_attn.in_proj_z.weight"):
return self._reorder_v_heads(scale, 0, num_k_heads, num_v_per_k, head_v_dim)
if name.endswith((".linear_attn.in_proj_a.weight", ".linear_attn.in_proj_b.weight")):
return self._reorder_v_heads(scale, 0, num_k_heads, num_v_per_k, 1)
return scale
def _transform_nvfp4_weight(self, name: str, weight: Tensor, scale: Tensor) -> tuple[Tensor, Tensor]:
if not name.endswith((
".linear_attn.in_proj_qkv.weight",