Compare commits

...
Author SHA1 Message Date
SolitaryThinker 5c13867381 [bugfix]: keep _rms_norm_dispatch output in the input dtype under autocast 2026-07-05 13:11:42 -07:00
Raghav e19cc78271 [bugfix] nvfp4: reject non-float inputs, extract dtype coercion + test
Address review feedback on the fp32-cast: instead of silently casting any
non-bf16/fp16 dtype (which would also accept int/bool), reject non-floating
inputs fast with TypeError and only cast floating dtypes (fp32/fp64) to bf16.

Extract the logic into a module-level _coerce_fp4_input_dtype() shared by
quantize_input() and apply(), and add a CPU-only unit test (no flashinfer/CUDA)
covering bf16/fp16 passthrough, fp32/fp64 -> bf16, and int/bool rejection.
2026-07-05 13:11:11 -07:00
Raghav 10b858345e [bugfix] nvfp4: cast fp32 inputs to bf16 instead of asserting
NVFP4QuantizeMethod.quantize_input() and apply() asserted bf16/fp16 inputs
and crashed on the fp32 pre-attention norm output in eager mode (without the
torch.compile fusion that keeps it bf16):

  AssertionError: only allow bf16/fp16 inputs to fp4 linear, got torch.float32

The FP4 linear emits bf16 regardless (mm_fp4 out dtype is bf16), and the
sibling fastvideo/layers/fp4linear.py already casts non-bf16/fp16 inputs, so
cast fp32 -> bf16 here too instead of failing. Lets NVFP4Config inference run
in eager mode (repro: LTX2-distilled on DGX Spark / sm_121, no compile).
2026-07-05 13:11:11 -07:00
3 changed files with 50 additions and 5 deletions
+18 -4
View File
@@ -261,6 +261,22 @@ def _mm_fp4(
)
def _coerce_fp4_input_dtype(x: torch.Tensor) -> torch.Tensor:
"""Coerce an activation to a dtype the FP4 linear accepts.
The pre-attention norm can emit fp32 (e.g. in eager mode, without the
torch.compile fusion that keeps it bf16). The FP4 linear emits bf16
regardless (see _mm_fp4 out dtype), so cast fp32 -> bf16 rather than
failing, matching the sibling fastvideo/layers/fp4linear.py. Non-floating
inputs (e.g. int/bool) are a genuine error and are rejected fast.
"""
if not x.is_floating_point():
raise TypeError(f"fp4 linear expects floating-point inputs, got {x.dtype}")
if x.dtype not in (torch.bfloat16, torch.float16):
x = x.to(torch.bfloat16)
return x
class NVFP4QuantizeMethod(QuantizeMethodBase):
def __init__(self, layer_prefix: str = ""):
@@ -285,8 +301,7 @@ class NVFP4QuantizeMethod(QuantizeMethodBase):
def quantize_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
SfLayout, _, _ = _require_flashinfer()
assert x.dtype == torch.bfloat16 or x.dtype == torch.float16, (
f"only allow bf16/fp16 inputs to fp4 linear, got {x.dtype}")
x = _coerce_fp4_input_dtype(x)
x_2d = x.view(-1, x.shape[-1])
x_fp4, x_scale = _nvfp4_quantize(
x_2d,
@@ -332,8 +347,7 @@ class NVFP4QuantizeMethod(QuantizeMethodBase):
if x_scale.dim() > 2:
x_scale = x_scale.view(-1, x_scale.shape[-1])
else:
assert x.dtype == torch.bfloat16 or x.dtype == torch.float16, (
f"only allow bf16/fp16 inputs to fp4 linear, got {x.dtype}")
x = _coerce_fp4_input_dtype(x)
x = x.view(-1, x.shape[-1])
x_global_sf = self.x_global_sf
x_fp4, x_scale = _nvfp4_quantize(
+5 -1
View File
@@ -72,8 +72,12 @@ def _rms_norm_dispatch(
"""Use QuACK RMSNorm only for LTX-2 refine stage, else Torch RMSNorm."""
if _is_ltx2_refine_stage():
return _quack_rmsnorm(x, weight=weight, eps=eps)
# torch 2.12 added rms_norm to autocast's float32 cast policy, so under
# autocast(bfloat16) this norm now returns float32 where torch <= 2.11
# returned bfloat16. Restore the input dtype so autocast-blind consumers
# downstream don't receive upcast activations.
return torch.nn.functional.rms_norm(
x, (x.shape[-1], ), weight=weight, eps=eps)
x, (x.shape[-1], ), weight=weight, eps=eps).to(x.dtype)
class StageAwareRMSNorm(nn.RMSNorm):
@@ -63,3 +63,30 @@ def _raise_module_on_import(name: str) -> types.ModuleType:
raise ImportError(f"No module named '{name}.{item}'")
return _RaisingModule(name)
def test_coerce_fp4_input_dtype_casts_and_rejects():
"""The FP4 input-dtype coercion (CPU-only, no flashinfer/CUDA): bf16/fp16
pass through, other floats (the fp32 pre-attention norm in eager mode) are
cast to bf16, and non-floating inputs are rejected fast."""
import torch
from fastvideo.layers.quantization.nvfp4_config import (
_coerce_fp4_input_dtype)
# bf16 / fp16 pass through untouched.
bf16 = torch.zeros(4, 8, dtype=torch.bfloat16)
assert _coerce_fp4_input_dtype(bf16) is bf16
fp16 = torch.zeros(4, 8, dtype=torch.float16)
assert _coerce_fp4_input_dtype(fp16) is fp16
# Other floating dtypes (fp32 from an unfused norm, fp64) -> bf16.
assert _coerce_fp4_input_dtype(
torch.zeros(4, 8, dtype=torch.float32)).dtype is torch.bfloat16
assert _coerce_fp4_input_dtype(
torch.zeros(4, 8, dtype=torch.float64)).dtype is torch.bfloat16
# Non-floating inputs are a real error, not silently cast.
for bad in (torch.int32, torch.int64, torch.bool):
with pytest.raises(TypeError, match="floating-point"):
_coerce_fp4_input_dtype(torch.zeros(4, 8, dtype=bad))