diff --git a/conversion/base.py b/conversion/base.py index f332613334..5c02b847e8 100644 --- a/conversion/base.py +++ b/conversion/base.py @@ -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))()) diff --git a/conversion/qwen.py b/conversion/qwen.py index 64d606176e..9b987da260 100644 --- a/conversion/qwen.py +++ b/conversion/qwen.py @@ -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",