Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
87737d10fd | ||
|
|
91f2acbbc9 | ||
|
|
2e8eb75c11 |
@@ -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(
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user