Skip to content

Commit 2e35b0c

Browse files
SolitaryThinkerjzhang38RandNMR73
authored
[refactor]: linear/mlp FP4 path additions for Wan-2.1 (Attn-QAT 6/12) (#1390)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com> Co-authored-by: Matthew Noto <notomatthew31@gmail.com>
1 parent 1c627a3 commit 2e35b0c

4 files changed

Lines changed: 148 additions & 27 deletions

File tree

fastvideo/layers/linear.py

Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -219,6 +219,15 @@ class ReplicatedLinear(LinearBase):
219219
(e.g. model.layers.0.qkv_proj)
220220
"""
221221

222+
# Opt-in instrumentation: when ``enable_shape_tracking`` is set to True,
223+
# ``forward`` records every unique ``(input_shape, output_shape)`` pair
224+
# observed across all ``ReplicatedLinear`` instances, along with the
225+
# subclass name that produced it. Used by upcoming QAT-aware backends
226+
# to discover which GEMM shapes need quantized kernels. Defaults to
227+
# False; default forward path is bit-identical to pre-slice behavior.
228+
enable_shape_tracking = False
229+
_shape_to_layer_types: dict[tuple[torch.Size, torch.Size], set[str]] = {}
230+
222231
def __init__(
223232
self,
224233
input_size: int,
@@ -285,6 +294,8 @@ def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
285294
bias = self.bias if not self.skip_bias_add else None
286295
assert self.quant_method is not None
287296
output = self.quant_method.apply(self, x, bias)
297+
if self.enable_shape_tracking:
298+
self._track_shape(x.shape, output.shape)
288299
output_bias = self.bias if self.skip_bias_add else None
289300
return output, output_bias
290301

@@ -294,6 +305,41 @@ def extra_repr(self) -> str:
294305
s += f", bias={self.bias is not None}"
295306
return s
296307

308+
@classmethod
309+
def get_shape_mapping(cls) -> dict:
310+
"""Get the mapping from (input_shape, output_shape) to layer types."""
311+
return cls._shape_to_layer_types.copy()
312+
313+
@classmethod
314+
def reset_shape_tracking(cls) -> None:
315+
"""Clear tracked shapes and layer type mappings."""
316+
cls._shape_to_layer_types.clear()
317+
318+
def _track_shape(self, input_shape: torch.Size, output_shape: torch.Size) -> None:
319+
shape_key = (input_shape, output_shape)
320+
if shape_key not in self._shape_to_layer_types:
321+
self._shape_to_layer_types[shape_key] = set()
322+
logger.debug("Layer: %s | input shape: %s --> output shape: %s, Quant Method: %s", self.prefix, input_shape,
323+
output_shape, self.quant_method.__class__.__name__)
324+
self._shape_to_layer_types[shape_key].add(self.__class__.__name__)
325+
326+
@classmethod
327+
def print_shape_summary(cls) -> None:
328+
"""Log a summary of all unique shapes and their layer types."""
329+
if not cls._shape_to_layer_types:
330+
logger.info("No shapes have been processed yet.")
331+
return
332+
333+
lines = [
334+
"=== Matrix Multiplication Shape Summary ===",
335+
f"Total unique shapes: {len(cls._shape_to_layer_types)}",
336+
]
337+
for i, (shape_key, layer_types) in enumerate(cls._shape_to_layer_types.items(), 1):
338+
input_shape, output_shape = shape_key
339+
lines.append(f"{i}. Input: {input_shape} → Output: {output_shape}")
340+
lines.append(f" Layer types: {', '.join(sorted(layer_types))}")
341+
logger.info("\n".join(lines))
342+
297343

298344
class ColumnParallelLinear(LinearBase):
299345
"""Linear layer with column parallelism.

fastvideo/layers/mlp.py

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55

66
from fastvideo.layers.activation import get_act_fn
77
from fastvideo.layers.linear import ReplicatedLinear
8+
from fastvideo.layers.quantization import QuantizationConfig
89

910

1011
class MLP(nn.Module):
@@ -21,18 +22,27 @@ def __init__(
2122
act_type: str = "gelu_pytorch_tanh",
2223
dtype: torch.dtype | None = None,
2324
prefix: str = "",
25+
quant_config: QuantizationConfig | None = None,
2426
):
2527
super().__init__()
2628
self.fc_in = ReplicatedLinear(
2729
input_dim,
2830
mlp_hidden_dim, # For activation func like SiLU that need 2x width
2931
bias=bias,
30-
params_dtype=dtype)
32+
params_dtype=dtype,
33+
quant_config=quant_config,
34+
prefix=f"{prefix}.fc_in",
35+
)
3136

3237
self.act = get_act_fn(act_type)
3338
if output_dim is None:
3439
output_dim = input_dim
35-
self.fc_out = ReplicatedLinear(mlp_hidden_dim, output_dim, bias=bias, params_dtype=dtype)
40+
self.fc_out = ReplicatedLinear(mlp_hidden_dim,
41+
output_dim,
42+
bias=bias,
43+
params_dtype=dtype,
44+
quant_config=quant_config,
45+
prefix=f"{prefix}.fc_out")
3646

3747
def forward(self, x: torch.Tensor) -> torch.Tensor:
3848
x, _ = self.fc_in(x)

fastvideo/models/dits/wanvideo.py

Lines changed: 41 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
from fastvideo.logger import init_logger
2727
from fastvideo.models.dits.base import BaseDiT
2828
from fastvideo.platforms import AttentionBackendEnum, current_platform
29+
from fastvideo.layers.quantization import QuantizationConfig
2930

3031
from fastvideo.distributed.parallel_state import get_sp_world_size
3132

@@ -106,7 +107,9 @@ def __init__(self,
106107
window_size=(-1, -1),
107108
qk_norm=True,
108109
eps=1e-6,
109-
parallel_attention=False) -> None:
110+
parallel_attention=False,
111+
quant_config: QuantizationConfig | None = None,
112+
prefix: str = "") -> None:
110113
assert dim % num_heads == 0
111114
super().__init__()
112115
self.dim = dim
@@ -118,10 +121,10 @@ def __init__(self,
118121
self.parallel_attention = parallel_attention
119122

120123
# layers
121-
self.to_q = ReplicatedLinear(dim, dim)
122-
self.to_k = ReplicatedLinear(dim, dim)
123-
self.to_v = ReplicatedLinear(dim, dim)
124-
self.to_out = ReplicatedLinear(dim, dim)
124+
self.to_q = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_q")
125+
self.to_k = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_k")
126+
self.to_v = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_v")
127+
self.to_out = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.to_out")
125128
self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
126129
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
127130

@@ -194,13 +197,15 @@ def __init__(
194197
qk_norm=True,
195198
eps=1e-6,
196199
supported_attention_backends: tuple[AttentionBackendEnum, ...]
197-
| None = None
200+
| None = None,
201+
quant_config: QuantizationConfig | None = None,
202+
prefix: str = "",
198203
) -> None:
199204
super().__init__(dim, num_heads, window_size, qk_norm, eps,
200-
supported_attention_backends)
205+
supported_attention_backends, quant_config=quant_config, prefix=prefix)
201206

202-
self.add_k_proj = ReplicatedLinear(dim, dim)
203-
self.add_v_proj = ReplicatedLinear(dim, dim)
207+
self.add_k_proj = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.add_k_proj")
208+
self.add_v_proj = ReplicatedLinear(dim, dim, quant_config=quant_config, prefix=f"{prefix}.add_v_proj")
204209
self.norm_added_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
205210
self.norm_added_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
206211

@@ -246,16 +251,17 @@ def __init__(self,
246251
added_kv_proj_dim: int | None = None,
247252
supported_attention_backends: tuple[AttentionBackendEnum, ...]
248253
| None = None,
254+
quant_config: QuantizationConfig | None = None,
249255
prefix: str = ""):
250256
super().__init__()
251257

252258
# 1. Self-attention
253259
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
254-
self.to_q = ReplicatedLinear(dim, dim, bias=True)
255-
self.to_k = ReplicatedLinear(dim, dim, bias=True)
256-
self.to_v = ReplicatedLinear(dim, dim, bias=True)
260+
self.to_q = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_q")
261+
self.to_k = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_k")
262+
self.to_v = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_v")
257263

258-
self.to_out = ReplicatedLinear(dim, dim, bias=True)
264+
self.to_out = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_out")
259265
self.attn1 = DistributedAttention(
260266
num_heads=num_heads,
261267
head_size=dim // num_heads,
@@ -290,13 +296,17 @@ def __init__(self,
290296
self.attn2 = WanI2VCrossAttention(dim,
291297
num_heads,
292298
qk_norm=qk_norm,
293-
eps=eps)
299+
eps=eps,
300+
quant_config=quant_config,
301+
prefix=f"{prefix}.attn2")
294302
else:
295303
# T2V
296304
self.attn2 = WanT2VCrossAttention(dim,
297305
num_heads,
298306
qk_norm=qk_norm,
299-
eps=eps)
307+
eps=eps,
308+
quant_config=quant_config,
309+
prefix=f"{prefix}.attn2")
300310
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
301311
dim,
302312
norm_type="layer",
@@ -306,7 +316,7 @@ def __init__(self,
306316
compute_dtype=torch.float32)
307317

308318
# 3. Feed-forward
309-
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
319+
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh", quant_config=quant_config, prefix=f"{prefix}.ffn")
310320
self.mlp_residual = ScaleResidual()
311321

312322
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
@@ -406,17 +416,17 @@ def __init__(self,
406416
added_kv_proj_dim: int | None = None,
407417
supported_attention_backends: tuple[AttentionBackendEnum, ...]
408418
| None = None,
419+
quant_config: QuantizationConfig | None = None,
409420
prefix: str = ""):
410421
super().__init__()
411422

412423
# 1. Self-attention
413424
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
414-
self.to_q = ReplicatedLinear(dim, dim, bias=True)
415-
self.to_k = ReplicatedLinear(dim, dim, bias=True)
416-
self.to_v = ReplicatedLinear(dim, dim, bias=True)
417-
self.to_gate_compress = ReplicatedLinear(dim, dim, bias=True)
418-
419-
self.to_out = ReplicatedLinear(dim, dim, bias=True)
425+
self.to_q = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_q")
426+
self.to_k = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_k")
427+
self.to_v = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_v")
428+
self.to_gate_compress = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_gate_compress")
429+
self.to_out = ReplicatedLinear(dim, dim, bias=True, quant_config=quant_config, prefix=f"{prefix}.to_out")
420430
self.attn1 = DistributedAttention_VSA(
421431
num_heads=num_heads,
422432
head_size=dim // num_heads,
@@ -451,13 +461,17 @@ def __init__(self,
451461
self.attn2 = WanI2VCrossAttention(dim,
452462
num_heads,
453463
qk_norm=qk_norm,
454-
eps=eps)
464+
eps=eps,
465+
quant_config=quant_config,
466+
prefix=f"{prefix}.attn2")
455467
else:
456468
# T2V
457469
self.attn2 = WanT2VCrossAttention(dim,
458470
num_heads,
459471
qk_norm=qk_norm,
460-
eps=eps)
472+
eps=eps,
473+
quant_config=quant_config,
474+
prefix=f"{prefix}.attn2")
461475
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
462476
dim,
463477
norm_type="layer",
@@ -467,7 +481,7 @@ def __init__(self,
467481
compute_dtype=torch.float32)
468482

469483
# 3. Feed-forward
470-
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
484+
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh", quant_config=quant_config, prefix=f"{prefix}.ffn")
471485
self.mlp_residual = ScaleResidual()
472486

473487
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
@@ -556,6 +570,7 @@ class WanTransformer3DModel(BaseDiT):
556570
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
557571
Any]) -> None:
558572
super().__init__(config=config, hf_config=hf_config)
573+
self.quant_config = config.quant_config
559574

560575
inner_dim = config.num_attention_heads * config.attention_head_dim
561576
self.hidden_size = config.hidden_size
@@ -594,6 +609,7 @@ def __init__(self, config: WanVideoConfig, hf_config: dict[str,
594609
config.eps,
595610
config.added_kv_proj_dim,
596611
self._supported_attention_backends,
612+
quant_config=config.quant_config,
597613
prefix=f"{config.prefix}.blocks.{i}")
598614
for i in range(config.num_layers)
599615
])
Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,49 @@
1+
# SPDX-License-Identifier: Apache-2.0
2+
"""Regression coverage for PR #1390's S2-1 plumbing finding: the dormant FP4 shape-tracking path must stay
3+
gated off unless explicitly enabled, and the MLP quant_config=None default path must keep using ReplicatedLinear's
4+
unquantized fallback.
5+
"""
6+
from __future__ import annotations
7+
8+
import torch
9+
10+
from fastvideo.layers.linear import ReplicatedLinear, UnquantizedLinearMethod
11+
from fastvideo.layers.mlp import MLP
12+
13+
14+
def test_replicated_linear_shape_tracking_default_off() -> None:
15+
ReplicatedLinear.reset_shape_tracking()
16+
assert ReplicatedLinear.enable_shape_tracking is False
17+
18+
linear = ReplicatedLinear(input_size=8, output_size=4)
19+
linear(torch.randn(2, 8))
20+
21+
assert len(ReplicatedLinear._shape_to_layer_types) == 0
22+
23+
24+
def test_replicated_linear_shape_tracking_enabled_records_unique_shapes() -> None:
25+
ReplicatedLinear.reset_shape_tracking()
26+
ReplicatedLinear.enable_shape_tracking = True
27+
try:
28+
linear = ReplicatedLinear(input_size=8, output_size=4)
29+
linear(torch.randn(2, 8))
30+
linear(torch.randn(3, 8))
31+
32+
assert len(ReplicatedLinear._shape_to_layer_types) == 2
33+
for layer_types in ReplicatedLinear._shape_to_layer_types.values():
34+
assert "ReplicatedLinear" in layer_types
35+
36+
ReplicatedLinear.reset_shape_tracking()
37+
assert len(ReplicatedLinear._shape_to_layer_types) == 0
38+
finally:
39+
ReplicatedLinear.enable_shape_tracking = False
40+
41+
42+
def test_mlp_quant_config_none_uses_unquantized_path() -> None:
43+
mlp = MLP(input_dim=8, mlp_hidden_dim=16)
44+
45+
assert isinstance(mlp.fc_in.quant_method, UnquantizedLinearMethod)
46+
assert isinstance(mlp.fc_out.quant_method, UnquantizedLinearMethod)
47+
48+
output = mlp.forward(torch.randn(2, 8))
49+
assert output.shape == (2, 8)

0 commit comments

Comments
 (0)