diff --git a/docs/en/attention.md b/docs/en/attention.md index 4e58877..0634ff0 100644 --- a/docs/en/attention.md +++ b/docs/en/attention.md @@ -71,6 +71,9 @@ class SparseAttentionConfig: sol_tau: float = 1.0 # Sol-Attn routing threshold sol_threshold_type: str = "diag" # "diag" or "exact" sol_kv_splits: int | str = "auto" # "auto", 1, 2, or 4 + sol_fp8: bool = False # FP8 Q/K/V Sol-Attn on SM90 + sol_fp8_layer_start: int = 0 # First layer using FP8 Sol-Attn + sol_fp8_layer_end: int | None = None # Exclusive end; None means all remaining layers ``` ## Calling Flow @@ -197,10 +200,20 @@ config = AttentionConfig.sol_attention() pipe_config.dit_config.attention_config = config ``` -Sol-Attn is used only for contiguous, noncausal BF16 self-attention with equal Q/K/V -shapes and head dimension 128. Unsupported calls, dense warmup layers or timesteps, -and kernel runtime failures fall back to the existing dense attention path. Ring/USP -also remains dense because its online merge requires log-sum-exp output. +Sol-Attn is used for contiguous, noncausal self-attention with equal Q/K/V shapes +and head dimension 128. BF16 is supported by the architecture-specific kernels; +SM90 additionally supports E4M3 Q/K/V with FP32 accumulation. Unsupported calls, +dense warmup layers or timesteps, and kernel runtime failures fall back to the +existing dense attention path. Ring/USP remains dense because its online merge +requires log-sum-exp output. + +`sol_fp8_layer_start` and `sol_fp8_layer_end` restrict E4M3 Q/K/V to a half-open +transformer-layer range. Sparse layers outside that range continue to use BF16 +Sol-Attn. This controls accumulated FP8 routing error in diffusion models. +Setting `dense_timesteps=0`, `dense_layers=0`, and a negative `tau` forces all +KV blocks onto the exact route. The Wan optimized example exposes this as +`--attention fp8-dense`; `--attention fp8-sol` enables centroid routing with the +same FP8 Q/K/V and QK/PV kernel. ### QwenImagePipeline / ZImagePipeline diff --git a/docs/en/quantization.md b/docs/en/quantization.md index 0f56858..41a6f09 100644 --- a/docs/en/quantization.md +++ b/docs/en/quantization.md @@ -94,7 +94,7 @@ quant_config = QuantConfig( ) ``` -For MiniMax H3, use `quantization="tf-kernel-fp8"` with +For MiniMax H3, use `quantization="fp8"` (`"tf-kernel-fp8"` remains an alias) with `examples/minimax_h3/minimax_h3_fl2va_h100.py`. This backend is single-GPU only and keeps the FP8 weights resident after first-use conversion. It is distinct from the scaled-FP8 checkpoint path below: the latter expects weights and scales already serialized in the checkpoint. diff --git a/docs/zh/attention.md b/docs/zh/attention.md index 475d648..365cc41 100644 --- a/docs/zh/attention.md +++ b/docs/zh/attention.md @@ -71,6 +71,9 @@ class SparseAttentionConfig: sol_tau: float = 1.0 # Sol-Attn 路由阈值 sol_threshold_type: str = "diag" # "diag" 或 "exact" sol_kv_splits: int | str = "auto" # "auto"、1、2 或 4 + sol_fp8: bool = False # SM90 FP8 Q/K/V Sol-Attn + sol_fp8_layer_start: int = 0 # 启用 FP8 Sol-Attn 的首层 + sol_fp8_layer_end: int | None = None # 结束层(不包含);None 表示其余所有层 ``` ## 调用流程 @@ -197,10 +200,19 @@ config = AttentionConfig.sol_attention() pipe_config.dit_config.attention_config = config ``` -Sol-Attn 仅用于连续、非因果、BF16、Q/K/V 形状相同且 head dimension 为 128 的 -self-attention。其他调用、dense 预热层/时间步以及内核运行失败都会回退到现有密集路径。 +Sol-Attn 用于连续、非因果、Q/K/V 形状相同且 head dimension 为 128 的 +self-attention。各架构内核支持 BF16,SM90 还支持使用 FP32 累加的 E4M3 Q/K/V。 +其他调用、dense 预热层/时间步以及内核运行失败都会回退到现有密集路径。 Ring/USP 需要 LSE 做在线合并,因此仍使用支持 LSE 的密集后端。 +`sol_fp8_layer_start` 和 `sol_fp8_layer_end` 用半开区间限制使用 E4M3 Q/K/V +的 transformer 层,区间外的稀疏层继续使用 BF16 Sol-Attn,以控制扩散模型中 +逐层累积的 FP8 路由误差。 + +设置 `dense_timesteps=0`、`dense_layers=0` 和负数 `tau` 会强制所有 KV block +走 exact 路径。Wan 优化示例将其暴露为 `--attention fp8-dense`; +`--attention fp8-sol` 使用相同的 FP8 Q/K/V 与 QK/PV kernel 并启用质心路由。 + ### QwenImagePipeline / ZImagePipeline 仅支持密集注意力(图像生成没有时序维度): diff --git a/examples/minimax_h3/README.md b/examples/minimax_h3/README.md index 4269ad6..6084467 100644 --- a/examples/minimax_h3/README.md +++ b/examples/minimax_h3/README.md @@ -317,7 +317,7 @@ MiniMax H3 supports three single-GPU online quantization backends for the DiT tr | CLI value | Backend | Weight/activation path | |---|---|---| | torchao-fp8 | TorchAO | FP8 dynamic activation and FP8 weight when supported, otherwise TorchAO's FP8 weight-only path | -| tf-kernel-fp8 | TeleFuser tf-kernel | Per-token activation and per-output-channel weight FP8 (W8A8), BF16 output | +| fp8 (`tf-kernel-fp8` alias) | TeleFuser tf-kernel | Per-token activation and per-output-channel weight FP8 (W8A8), BF16 output | | bnb-nf4 | bitsandbytes | NF4 weight-only with BF16 compute | All three paths convert the 258 Linear layers in the main and token-refiner transformer blocks. The FP32 video/audio @@ -342,12 +342,13 @@ python examples/minimax_h3/minimax_h3_fl2va_h100.py \ --output outputs/minimax_h3_bnb_nf4.mp4 python examples/minimax_h3/minimax_h3_fl2va_h100.py \ --mode t2va \ - --quantization tf-kernel-fp8 \ + --quantization fp8 \ --duration 5 \ --output outputs/minimax_h3_tf_kernel_fp8.mp4 ~~~ -The FL2VA CLI accepts `--quantization` with `torchao-fp8`, `tf-kernel-fp8`, or `bnb-nf4`; omit it for BF16. The Python +The FL2VA CLI accepts `--quantization` with `fp8`, `torchao-fp8`, or `bnb-nf4`; `tf-kernel-fp8` remains a compatibility +alias and omitting the option keeps BF16. The Python loader accepts the same names: ~~~python @@ -356,7 +357,7 @@ from examples.minimax_h3.common import load_minimax_h3_pipeline pipeline = load_minimax_h3_pipeline( "/path/to/MiniMaxAI_MiniMax-H3", partition="FL2VA", - quantization="tf-kernel-fp8", + quantization="fp8", ) ~~~ @@ -364,12 +365,42 @@ Online quantization currently requires ulysses_degree=1, tp_degree=1, and FSDP d would invalidate those wrappers' BF16 parameter-sharding contract, so unsupported combinations fail before checkpoint loading. +### FP8 Sol-Attn + +The same FL2VA example exposes dense/Sol and BF16/FP8 as independent switches. `--quantization fp8` applies +tf-kernel W8A8 Linear GEMMs to the transformer blocks. `--attn-impl SOL_ATTN` enables the MiniMax-H3 Sol policy: +the first 10 denoising steps and first 2 DiT layers remain dense, the full condition prefix is an exact KV sink, and +prefix queries are recomputed with BF16 dense attention. Adding `--sol-fp8` quantizes post-RoPE Q/K/V in active sparse +layers and dispatches the SM90 CuTe FP8 Sol mainloop. + +~~~bash +# BF16 dense +python -m examples.minimax_h3.minimax_h3_fl2va_h100 --mode t2va --output outputs/h3_bf16.mp4 + +# BF16 Sol +python -m examples.minimax_h3.minimax_h3_fl2va_h100 \ + --mode t2va --attn-impl SOL_ATTN --output outputs/h3_bf16_sol.mp4 + +# FP8 Linear + dense attention +python -m examples.minimax_h3.minimax_h3_fl2va_h100 \ + --mode t2va --quantization fp8 --output outputs/h3_fp8.mp4 + +# FP8 Linear + FP8 Sol attention +python -m examples.minimax_h3.minimax_h3_fl2va_h100 \ + --mode t2va --quantization fp8 --attn-impl SOL_ATTN --sol-fp8 \ + --output outputs/h3_fp8_sol.mp4 +~~~ + +Use `--sol-dense-steps`, `--sol-dense-layers`, `--sol-tau`, `--sol-threshold-type`, +`--sol-fp8-layer-start`, and `--sol-fp8-layer-end` to override the policy for controlled ablations. The defaults +match the released H100 MiniMax-H3 Sol profile. + For matched BF16/TorchAO-FP8/tf-kernel-FP8/NF4 profiling, use the validation benchmark. It writes the synchronized MP4 plus a JSON report containing load time, end-to-end generation time, stage timings, and denoising allocator peaks: ~~~bash -python tools/validation/benchmark_minimax_h3_quantization.py \ - --backend tf-kernel-fp8 \ +python -m tools.validation.benchmark_minimax_h3_quantization \ + --backend fp8-sol \ --duration 5 \ --steps 50 \ --output outputs/minimax_h3_tf_kernel_fp8_50step.mp4 diff --git a/examples/minimax_h3/common.py b/examples/minimax_h3/common.py index 253f4c9..0883a67 100644 --- a/examples/minimax_h3/common.py +++ b/examples/minimax_h3/common.py @@ -138,6 +138,7 @@ def minimax_h3_quant_config(quantization: str | QuantType | None) -> QuantConfig if isinstance(quantization, str): normalized = quantization.strip().lower().replace("_", "-") names = { + "fp8": QuantType.FP8, "torchao-fp8": QuantType.TORCHAO_FP8, "bnb-nf4": QuantType.BNB_NF4, "tf-kernel-fp8": QuantType.FP8, @@ -145,7 +146,7 @@ def minimax_h3_quant_config(quantization: str | QuantType | None) -> QuantConfig try: quant_type = names[normalized] except KeyError as exc: - raise ValueError("quantization must be 'torchao-fp8', 'tf-kernel-fp8', 'bnb-nf4', or None") from exc + raise ValueError("quantization must be 'fp8', 'torchao-fp8', 'tf-kernel-fp8', 'bnb-nf4', or None") from exc elif isinstance(quantization, QuantType): quant_type = quantization else: @@ -172,6 +173,13 @@ def load_minimax_h3_pipeline( text_encoder_tp_degree: int | None = None, enable_fsdp: bool | None = None, attn_impl: AttnImplType | str = AttnImplType.FLASH_ATTN_4, + sol_fp8: bool = False, + sol_dense_steps: int = 10, + sol_dense_layers: int = 2, + sol_tau: float = 1.0, + sol_threshold_type: str = "exact", + sol_fp8_layer_start: int = 0, + sol_fp8_layer_end: int | None = None, feature_cache_config: FeatureCacheConfig | None = None, adaln_cache_path: str | Path | None = None, online_adaln_cache: bool = False, @@ -209,6 +217,8 @@ def load_minimax_h3_pipeline( attn_impl = AttnImplType[attn_impl] except KeyError as exc: raise ValueError(f"unsupported attention implementation: {attn_impl}") from exc + if sol_fp8 and attn_impl != AttnImplType.SOL_ATTN: + raise ValueError("sol_fp8 requires attn_impl=SOL_ATTN") component_root = Path(model_root) / partition if not component_root.is_dir(): raise FileNotFoundError(f"MiniMax H3 partition not found: {component_root}") @@ -248,12 +258,25 @@ def load_minimax_h3_pipeline( offload_config=resident_offload, parallel_config=text_parallel, ) + attention_config = ( + AttentionConfig.sol_attention( + dense_timesteps=sol_dense_steps, + dense_layers=sol_dense_layers, + tau=sol_tau, + threshold_type=sol_threshold_type, + sol_fp8=sol_fp8, + sol_fp8_layer_start=sol_fp8_layer_start, + sol_fp8_layer_end=sol_fp8_layer_end, + ) + if attn_impl == AttnImplType.SOL_ATTN + else AttentionConfig.dense_attention(attn_impl) + ) dit_runtime = ModelRuntimeConfig( device_type=runtime_device.type, device_id=runtime_device.index or 0, torch_dtype=torch.bfloat16, offload_config=dit_offload, - attention_config=AttentionConfig.dense_attention(attn_impl), + attention_config=attention_config, feature_cache_config=feature_cache_config or FeatureCacheConfig(), quant_config=quant_config, lora_configs=[LoraConfig(path=str(lora_path), strength=lora_strength)] if lora_path else [], diff --git a/examples/minimax_h3/minimax_h3_fl2va_h100.py b/examples/minimax_h3/minimax_h3_fl2va_h100.py index 57af404..b833507 100644 --- a/examples/minimax_h3/minimax_h3_fl2va_h100.py +++ b/examples/minimax_h3/minimax_h3_fl2va_h100.py @@ -35,6 +35,7 @@ "enable_fsdp": None, "online_adaln_cache": True, "attn_impl": AttnImplType.FLASH_ATTN_4, + "sol_fp8": False, "feature_cache_model_type": "MiniMax-H3-Base", "feature_cache_n_derivatives": 1, "feature_cache_taylor_threshold": 2, @@ -85,6 +86,13 @@ def get_pipeline( enable_fsdp: bool | None = PPL_CONFIG["enable_fsdp"], online_adaln_cache: bool = PPL_CONFIG["online_adaln_cache"], attn_impl: AttnImplType | str = PPL_CONFIG["attn_impl"], + sol_fp8: bool = PPL_CONFIG["sol_fp8"], + sol_dense_steps: int = 10, + sol_dense_layers: int = 2, + sol_tau: float = 1.0, + sol_threshold_type: str = "exact", + sol_fp8_layer_start: int = 0, + sol_fp8_layer_end: int | None = None, enable_feature_cache: bool = False, feature_cache_model_type: str = PPL_CONFIG["feature_cache_model_type"], feature_cache_n_derivatives: int = PPL_CONFIG["feature_cache_n_derivatives"], @@ -104,6 +112,13 @@ def get_pipeline( enable_fsdp=enable_fsdp, online_adaln_cache=online_adaln_cache, attn_impl=attn_impl, + sol_fp8=sol_fp8, + sol_dense_steps=sol_dense_steps, + sol_dense_layers=sol_dense_layers, + sol_tau=sol_tau, + sol_threshold_type=sol_threshold_type, + sol_fp8_layer_start=sol_fp8_layer_start, + sol_fp8_layer_end=sol_fp8_layer_end, feature_cache_config=FeatureCacheConfig( enabled=enable_feature_cache, model_type=feature_cache_model_type, @@ -270,16 +285,23 @@ def _main(default_quantization: str | None = PPL_CONFIG["quantization"]) -> None parser.add_argument("--device", default=PPL_CONFIG["device"]) parser.add_argument( "--quantization", - choices=("torchao-fp8", "tf-kernel-fp8", "bnb-nf4"), + choices=("fp8", "torchao-fp8", "tf-kernel-fp8", "bnb-nf4"), default=default_quantization, help="Online DiT Linear quantization backend (single GPU only).", ) parser.add_argument("--gpu-num", "--ulysses-degree", dest="gpu_num", type=int, choices=(1, 2, 4), default=1) parser.add_argument( "--attn-impl", - choices=("FLASH_ATTN_4", "SAGE_ATTN_2_8_8_SM90"), + choices=("FLASH_ATTN_4", "SAGE_ATTN_2_8_8_SM90", "SOL_ATTN"), default=PPL_CONFIG["attn_impl"].name, ) + parser.add_argument("--sol-fp8", action="store_true", help="Use FP8 Q/K/V in active Sol-Attn layers.") + parser.add_argument("--sol-dense-steps", type=int, default=10) + parser.add_argument("--sol-dense-layers", type=int, default=2) + parser.add_argument("--sol-tau", type=float, default=1.0) + parser.add_argument("--sol-threshold-type", choices=("exact", "diag"), default="exact") + parser.add_argument("--sol-fp8-layer-start", type=int, default=0) + parser.add_argument("--sol-fp8-layer-end", type=int) parser.add_argument("--enable-feature-cache", action="store_true") parser.add_argument("--feature-cache-model-type", default=PPL_CONFIG["feature_cache_model_type"]) parser.add_argument( @@ -324,6 +346,13 @@ def _main(default_quantization: str | None = PPL_CONFIG["quantization"]) -> None num_inference_steps=args.steps, enable_fsdp=args.enable_fsdp, attn_impl=args.attn_impl, + sol_fp8=args.sol_fp8, + sol_dense_steps=args.sol_dense_steps, + sol_dense_layers=args.sol_dense_layers, + sol_tau=args.sol_tau, + sol_threshold_type=args.sol_threshold_type, + sol_fp8_layer_start=args.sol_fp8_layer_start, + sol_fp8_layer_end=args.sol_fp8_layer_end, enable_feature_cache=args.enable_feature_cache, feature_cache_model_type=args.feature_cache_model_type, feature_cache_n_derivatives=args.feature_cache_n_derivatives, diff --git a/examples/minimax_h3/minimax_h3_ref2va_h100.py b/examples/minimax_h3/minimax_h3_ref2va_h100.py index a7de40f..ef6971e 100644 --- a/examples/minimax_h3/minimax_h3_ref2va_h100.py +++ b/examples/minimax_h3/minimax_h3_ref2va_h100.py @@ -263,7 +263,7 @@ def main() -> None: parser.add_argument("--flow-shift", type=float, default=PPL_CONFIG["flow_shift"]) parser.add_argument("--audio-flow-shift", type=float, default=PPL_CONFIG["audio_flow_shift"]) parser.add_argument("--device", default=PPL_CONFIG["device"]) - parser.add_argument("--quantization", choices=("torchao-fp8", "tf-kernel-fp8", "bnb-nf4")) + parser.add_argument("--quantization", choices=("fp8", "torchao-fp8", "tf-kernel-fp8", "bnb-nf4")) parser.add_argument("--gpu-num", "--ulysses-degree", dest="gpu_num", type=int, choices=(1, 2, 4), default=1) fsdp_group = parser.add_mutually_exclusive_group() fsdp_group.add_argument("--enable-fsdp", dest="enable_fsdp", action="store_true") diff --git a/examples/wan_video/README.md b/examples/wan_video/README.md index 9a7e03f..64852c9 100644 --- a/examples/wan_video/README.md +++ b/examples/wan_video/README.md @@ -89,7 +89,6 @@ python examples/wan_video/wan21_1_3b_text_to_video_h100.py --resolution 480p --a **Features:** - Video Frame Interpolation (VFI) with RIFE model for 30fps output - CFG parallel when cfg_scale > 1 - #### wan21_1_3b_text_to_video_hf.py T2V with HuggingFace format loading. @@ -173,6 +172,132 @@ pipe_config.dit_config.attention_config = AttentionConfig.sol_attention() Sol-Attn is built into TeleFuser. Eligible BF16 self-attention calls use the sparse kernel; unsupported calls automatically use the existing dense fallback. The defaults follow the official Wan2.1 profile: Morton3D token ordering, dense layer 0, and 10 dense warm-up steps for the standard 50-step schedule. +#### wan21_1_3b_text_to_video_optimized_h100.py + +Provides one entry point for independently enabling attention and quantization +optimizations. The defaults are `--attention dense --quantization none`, which +run the BF16 baseline. + +Attention choices: + +- `dense`: BF16 PyTorch SDPA +- `sol`: BF16 Sol-Attn with dense warm-up and fallback calls +- `fp8-dense`: E4M3 Q/K/V with FP8 QK/PV WGMMA; all KV blocks are exact +- `fp8-sol`: the same FP8 kernel with Sol routing enabled + +For both FP8 modes, `--fp8-layer-start` and `--fp8-layer-end` select the +half-open transformer-layer range that uses FP8 Q/K/V. Other layers use the +corresponding BF16 dense or Sol path. Restricting FP8 Q/K/V to middle layers +avoids accumulating small quantization changes across the full denoiser. + +Quantization choices: + +- `none`: BF16 DiT +- `tf-kernel-fp8`: TeleFuser dynamic W8A8 FP8 GEMM + +`--fp8-linear-scope all` quantizes every transformer-block Linear layer. +`--fp8-linear-scope ffn` keeps self/cross-attention projections in BF16 and +quantizes the 60 FFN Linear layers. The default `auto` selects `all`; generated +video validation shows that all-Linear FP8 preserves quality. Attention Q/K/V +are more sensitive, so their default FP8 layer range is 10-19 for this 30-layer +Wan2.1 model. + +Only DiT transformer-block Linear layers are quantized; the VAE and text encoder +remain BF16. Select the two optimization axes independently: + +```bash +# Dense + BF16 baseline +python examples/wan_video/wan21_1_3b_text_to_video_optimized_h100.py \ + --model-root /path/to/Wan2.1-T2V-1.3B \ + --attention dense --quantization none + +# Sol-Attn + BF16 +python examples/wan_video/wan21_1_3b_text_to_video_optimized_h100.py \ + --model-root /path/to/Wan2.1-T2V-1.3B \ + --attention sol --quantization none + +# FP8 Dense: exact FP8 attention + FP8 Linear +python examples/wan_video/wan21_1_3b_text_to_video_optimized_h100.py \ + --model-root /path/to/Wan2.1-T2V-1.3B \ + --attention fp8-dense \ + --quantization tf-kernel-fp8 \ + --fp8-layer-start 10 \ + --fp8-layer-end 20 + +# FP8 Sol: routed FP8 attention + FP8 Linear +python examples/wan_video/wan21_1_3b_text_to_video_optimized_h100.py \ + --model-root /path/to/Wan2.1-T2V-1.3B \ + --attention fp8-sol \ + --quantization tf-kernel-fp8 \ + --fp8-layer-start 10 \ + --fp8-layer-end 20 \ + --dense-timesteps 10 \ + --dense-layers 1 \ + --tau 1.0 \ + --threshold-type diag \ + --kv-splits auto +``` + +In this example, FP8 means the E4M3 attention implementation rather than a +BF16-attention run with only its Linear layers quantized. Post-RoPE Q/K/V are +quantized and QK/PV run through the CuTe SM90 WGMMA mainloop. Q/K use one scale +per 64-token block, V uses per-channel scales and a K-major layout, and FP32 +accumulators are used throughout. `fp8-dense` forces every routed KV block onto +the exact path, while `fp8-sol` permits centroid approximation. `auto` selects +two KV splits for long FP8 sequences, which is faster at Wan's sequence length +without changing the FP32 accumulation contract. +Partial tiles are physically padded while the original sequence length remains +masked in the kernel. FP8 split execution restores the represented N64 route +length before PV, matching the BF16 summed-centroid contract. The accompanying +FP8 Linear GEMMs use the tf-kernel backend. Self-attention Q/K/V projections +share one dynamic activation quantization instead of quantizing the same input +three times. With a partial FP8 layer range, FP8 Dense sends unquantized layers +to SDPA and FP8 Sol sends unquantized sparse layers to Triton, avoiding a second +CuTe specialization in the cold-start path. +The final log reports generation time, frames per second, and peak allocated and +reserved CUDA memory. + +##### H100 benchmark + +This clean-process cold-start benchmark runs each configuration in a separate process +with no other GPU processes on one H100 80GB. It uses the official Wan2.1 T2V-1.3B example prompt, `832x480`, +81 frames, 50 UniPC steps, CFG 5.0, sigma shift 5.0, and seed 42. Generation timing +starts after pipeline loading, so it includes first-execution kernel/JIT costs but +excludes model loading. Peak memory is `torch.cuda.max_memory_allocated()` over the +same generation interval. + +| Quantization | Attention | Throughput (frames/s) | Peak allocated (GiB) | +| --- | --- | ---: | ---: | +| BF16 | Dense | 0.8491 | 16.147 | +| BF16 | Sol-Attn | 1.1090 | 17.023 | +| FP8 | Dense (Q/K/V layers 10-19, exact) | 0.8739 | 15.730 | +| FP8 | Sol-Attn (Q/K/V layers 10-19) | 1.1565 | 15.730 | + +Both FP8 rows quantize all 300 transformer-block Linear layers and use the same +E4M3 attention layer range. FP8 Dense therefore measures this implementation's +exact QK/PV path, not BF16 SDPA with only Linear quantization. Against the +corresponding BF16 output, FP8 Dense measures 22.0257 dB PSNR / 0.828783 SSIM, +and FP8 Sol measures 20.8502 dB PSNR / 0.792656 SSIM. + +The benchmark prompt is: + +> Two anthropomorphic cats in comfy boxing gear and bright gloves fight intensely +> on a spotlighted stage. + +Example command (change `--attention` and `--quantization` for each ablation): + +```bash +python examples/wan_video/wan21_1_3b_text_to_video_optimized_h100.py \ + --model-root /path/to/Wan2.1-T2V-1.3B \ + --prompt "Two anthropomorphic cats in comfy boxing gear and bright gloves fight intensely on a spotlighted stage." \ + --attention fp8-sol \ + --quantization tf-kernel-fp8 \ + --fp8-linear-scope all \ + --fp8-layer-start 10 --fp8-layer-end 20 \ + --width 832 --height 480 \ + --num-frames 81 --num-inference-steps 50 \ + --sample-solver unipc --cfg-scale 5.0 --sigma-shift 5.0 --seed 42 +``` #### wan21_1_3b_text_to_video_cache_calibrate.py diff --git a/examples/wan_video/wan21_1_3b_text_to_video_optimized_h100.py b/examples/wan_video/wan21_1_3b_text_to_video_optimized_h100.py new file mode 100644 index 0000000..f9abafa --- /dev/null +++ b/examples/wan_video/wan21_1_3b_text_to_video_optimized_h100.py @@ -0,0 +1,307 @@ +"""Wan2.1 1.3B T2V with optional attention and quantization optimizations. + +Attention can use BF16 dense/Sol or FP8 dense/Sol kernels. DiT Linear layers can +remain BF16 or use tf-kernel FP8, TorchAO FP8, or bitsandbytes NF4. The example +keeps the DiT on CUDA so quantized Linear modules are not repeatedly rebuilt. +""" + +from __future__ import annotations + +import os +import time + +import click +import torch + +from telefuser.core.config import ( + AttentionConfig, + AttnImplType, + QuantConfig, + QuantKernelBackend, + QuantType, + WeightOffloadType, +) +from telefuser.core.module_manager import ModuleManager +from telefuser.pipelines.wan_video.wan21_video import Wan21VideoPipeline, Wan21VideoPipelineConfig +from telefuser.utils.utils import get_example_name +from telefuser.utils.video import get_target_video_size_from_ratio, save_video + +TF_MODEL_ZOO_PATH = os.environ.get("TF_MODEL_ZOO_PATH", "model_zoo") +PPL_CONFIG = { + "model_root": TF_MODEL_ZOO_PATH + "/Wan2.1-T2V-1.3B", + "negative_prompt": ( + "Camera shake, overly saturated colors, overexposed, static, blurry details, subtitles, " + "worst quality, low quality, JPEG compression artifacts, ugly, incomplete, deformed limbs" + ), + "num_inference_steps": 40, + "num_frames": 81, + "resolution": "480p", + "cfg_scale": 5.0, + "sigma_shift": 8.0, +} + +FP8_ATTENTION_MODES = ("fp8-dense", "fp8-sol") + + +def configure_attention_backends() -> None: + """Configure dense attention backends used directly or by SOL fallbacks.""" + if hasattr(torch.backends.cuda, "enable_cudnn_sdp"): + torch.backends.cuda.enable_cudnn_sdp(False) + if hasattr(torch.backends.cuda, "enable_flash_sdp"): + torch.backends.cuda.enable_flash_sdp(True) + if hasattr(torch.backends.cuda, "enable_math_sdp"): + torch.backends.cuda.enable_math_sdp(True) + if hasattr(torch.backends.cuda, "enable_mem_efficient_sdp"): + torch.backends.cuda.enable_mem_efficient_sdp(True) + + +def make_quant_config(quantization: str, *, fp8_linear_scope: str = "all") -> QuantConfig: + """Build the online quantization config used for the Wan DiT.""" + if fp8_linear_scope not in ("all", "ffn"): + raise ValueError("fp8_linear_scope must be 'all' or 'ffn'") + if quantization == "none": + return QuantConfig() + if quantization == "tf-kernel-fp8": + return QuantConfig( + enabled=True, + quant_type=QuantType.FP8, + kernel_backend=QuantKernelBackend.TF_KERNEL, + quantize_modules=(".ffn.",) if fp8_linear_scope == "ffn" else None, + ) + if quantization == "torchao-fp8": + return QuantConfig( + enabled=True, + quant_type=QuantType.TORCHAO_FP8, + kernel_backend=QuantKernelBackend.TORCHAO, + ) + if quantization == "bnb-nf4": + return QuantConfig( + enabled=True, + quant_type=QuantType.BNB_NF4, + kernel_backend=QuantKernelBackend.BITSANDBYTES, + ) + raise ValueError("quantization must be 'none', 'tf-kernel-fp8', 'torchao-fp8', or 'bnb-nf4'") + + +def resolve_fp8_linear_scope(attention: str, fp8_linear_scope: str) -> str: + """Resolve the Linear FP8 scope for the selected attention mode.""" + if fp8_linear_scope == "auto": + return "all" + if fp8_linear_scope not in ("all", "ffn"): + raise ValueError("fp8_linear_scope must be 'auto', 'all', or 'ffn'") + return fp8_linear_scope + + +def make_attention_config( + attention: str, + *, + dense_timesteps: int = 10, + dense_layers: int = 1, + tau: float = 1.0, + threshold_type: str = "diag", + kv_splits: int | str = "auto", + fp8_layer_start: int = 10, + fp8_layer_end: int | None = 20, +) -> AttentionConfig: + """Build the selected BF16 or FP8 dense/Sol attention configuration.""" + if attention == "dense": + return AttentionConfig.dense_attention(AttnImplType.TORCH_SDPA) + if attention not in ("sol", *FP8_ATTENTION_MODES): + raise ValueError("attention must be 'dense', 'sol', 'fp8-dense', or 'fp8-sol'") + + fp8_dense = attention == "fp8-dense" + return AttentionConfig.sol_attention( + dense_timesteps=0 if fp8_dense else dense_timesteps, + dense_layers=0 if fp8_dense else dense_layers, + # A negative threshold makes every routed KV block exact. This runs the + # FP8 CuTe mainloop without introducing Sol's centroid approximation. + tau=-1000.0 if fp8_dense else tau, + threshold_type=threshold_type, + kv_splits=kv_splits, + sol_fp8=attention in FP8_ATTENTION_MODES, + sol_fp8_layer_start=fp8_layer_start, + sol_fp8_layer_end=fp8_layer_end, + ) + + +def get_pipeline( + *, + model_root: str = PPL_CONFIG["model_root"], + attention: str = "dense", + quantization: str = "none", + fp8_linear_scope: str = "auto", + dense_timesteps: int = 10, + dense_layers: int = 1, + tau: float = 1.0, + threshold_type: str = "diag", + kv_splits: int | str = "auto", + fp8_layer_start: int = 10, + fp8_layer_end: int | None = 20, + sample_solver: str = "euler", +) -> Wan21VideoPipeline: + """Load Wan2.1 with independently selectable attention and quantization.""" + fp8_linear_scope = resolve_fp8_linear_scope(attention, fp8_linear_scope) + quant_config = make_quant_config(quantization, fp8_linear_scope=fp8_linear_scope) + module_manager = ModuleManager(torch_dtype=torch.bfloat16, device="cpu") + module_manager.load_model(f"{model_root}/Wan2.1_VAE.pth", device="cpu", torch_dtype=torch.bfloat16) + module_manager.load_model( + f"{model_root}/diffusion_pytorch_model.safetensors", + device="cuda", + torch_dtype=torch.bfloat16, + quant_config=quant_config, + ) + module_manager.load_model(f"{model_root}/models_t5_umt5-xxl-enc-bf16.pth", device="cpu", torch_dtype=torch.bfloat16) + + pipeline = Wan21VideoPipeline(device="cuda", torch_dtype=torch.bfloat16) + config = Wan21VideoPipelineConfig() + config.dit_config.attention_config = make_attention_config( + attention, + dense_timesteps=dense_timesteps, + dense_layers=dense_layers, + tau=tau, + threshold_type=threshold_type, + kv_splits=kv_splits, + fp8_layer_start=fp8_layer_start, + fp8_layer_end=fp8_layer_end, + ) + config.dit_config.quant_config = quant_config + config.dit_config.offload_config.offload_type = WeightOffloadType.NO_CPU_OFFLOAD + config.sample_solver = sample_solver + config.enable_metrics = True + pipeline.init(module_manager, config) + return pipeline + + +def run( + pipeline: Wan21VideoPipeline, + prompt: str, + *, + seed: int = 42, + resolution: str = "480p", + width: int | None = None, + height: int | None = None, + num_inference_steps: int = PPL_CONFIG["num_inference_steps"], + num_frames: int = PPL_CONFIG["num_frames"], + cfg_scale: float = PPL_CONFIG["cfg_scale"], + sigma_shift: float = PPL_CONFIG["sigma_shift"], +): + """Generate one deterministic validation video.""" + if (width is None) != (height is None): + raise ValueError("width and height must be provided together") + if width is None or height is None: + width, height = get_target_video_size_from_ratio( + "16:9", resolution=resolution, height_division_factor=2, width_division_factor=2 + ) + return pipeline( + prompt=prompt, + negative_prompt=PPL_CONFIG["negative_prompt"], + num_inference_steps=num_inference_steps, + num_frames=num_frames, + cfg_scale=cfg_scale, + seed=seed, + height=height, + width=width, + sigma_shift=sigma_shift, + tiled=True, + ) + + +@click.command() +@click.option("--prompt", default="A small paper boat floating down a sunlit stream.") +@click.option("--seed", default=42, type=int) +@click.option("--resolution", default="480p", type=click.Choice(["480p", "720p"])) +@click.option("--width", type=int) +@click.option("--height", type=int) +@click.option("--num-inference-steps", default=PPL_CONFIG["num_inference_steps"], type=int) +@click.option("--num-frames", default=PPL_CONFIG["num_frames"], type=int) +@click.option("--cfg-scale", default=PPL_CONFIG["cfg_scale"], type=float) +@click.option("--sigma-shift", default=PPL_CONFIG["sigma_shift"], type=float) +@click.option("--sample-solver", default="euler", type=click.Choice(["euler", "unipc"])) +@click.option("--model-root", default=PPL_CONFIG["model_root"]) +@click.option( + "--attention", + default="dense", + type=click.Choice(["dense", "sol", "fp8-dense", "fp8-sol"]), +) +@click.option( + "--quantization", + default="none", + type=click.Choice(["none", "tf-kernel-fp8", "torchao-fp8", "bnb-nf4"]), +) +@click.option("--fp8-linear-scope", default="auto", type=click.Choice(["auto", "all", "ffn"])) +@click.option("--dense-timesteps", default=10, type=int) +@click.option("--dense-layers", default=1, type=int) +@click.option("--tau", default=1.0, type=float) +@click.option("--threshold-type", default="diag", type=click.Choice(["diag", "exact"])) +@click.option("--kv-splits", default="auto", type=click.Choice(["auto", "1", "2", "4"])) +@click.option("--fp8-layer-start", "--sol-fp8-layer-start", default=10, type=int) +@click.option("--fp8-layer-end", "--sol-fp8-layer-end", default=20, type=int) +@click.option("--output", default=get_example_name(__file__, "mp4")) +def main( + prompt: str, + seed: int, + resolution: str, + width: int | None, + height: int | None, + num_inference_steps: int, + num_frames: int, + cfg_scale: float, + sigma_shift: float, + sample_solver: str, + model_root: str, + attention: str, + quantization: str, + fp8_linear_scope: str, + dense_timesteps: int, + dense_layers: int, + tau: float, + threshold_type: str, + kv_splits: str, + fp8_layer_start: int, + fp8_layer_end: int | None, + output: str, +) -> None: + """Run Wan2.1 with optional attention and quantization optimizations.""" + configure_attention_backends() + pipeline = get_pipeline( + model_root=model_root, + attention=attention, + quantization=quantization, + fp8_linear_scope=fp8_linear_scope, + dense_timesteps=dense_timesteps, + dense_layers=dense_layers, + tau=tau, + threshold_type=threshold_type, + kv_splits=kv_splits if kv_splits == "auto" else int(kv_splits), + fp8_layer_start=fp8_layer_start, + fp8_layer_end=fp8_layer_end, + sample_solver=sample_solver, + ) + torch.cuda.reset_peak_memory_stats() + start = time.perf_counter() + video = run( + pipeline, + prompt, + seed=seed, + resolution=resolution, + width=width, + height=height, + num_inference_steps=num_inference_steps, + num_frames=num_frames, + cfg_scale=cfg_scale, + sigma_shift=sigma_shift, + ) + torch.cuda.synchronize() + elapsed = time.perf_counter() - start + save_video(video, output, fps=16, quality=6) + peak_allocated = torch.cuda.max_memory_allocated() / 2**30 + peak_reserved = torch.cuda.max_memory_reserved() / 2**30 + click.echo( + f"attention={attention} quantization={quantization} elapsed_s={elapsed:.2f} " + f"throughput_fps={num_frames / elapsed:.4f} " + f"peak_allocated_gib={peak_allocated:.3f} peak_reserved_gib={peak_reserved:.3f} output={output}" + ) + + +if __name__ == "__main__": + main() diff --git a/telefuser/core/config.py b/telefuser/core/config.py index 505b481..7bd1f5f 100644 --- a/telefuser/core/config.py +++ b/telefuser/core/config.py @@ -150,7 +150,10 @@ class SparseAttentionConfig: use_sage_attention: bool = False # Use sage attention backend sol_tau: float = 1.0 # Sol-Attn routing threshold multiplier sol_threshold_type: str = "diag" # Sol-Attn threshold estimator: "diag" or "exact" - sol_kv_splits: int | str = "auto" # Auto selects split 4 for long SM90 sequences + sol_kv_splits: int | str = "auto" # Auto selects split 2 for long FP8 SM90 sequences + sol_fp8: bool = False # Quantize post-RoPE Q/K/V activations for FP8 Sol-Attn + sol_fp8_layer_start: int = 0 # First transformer layer using FP8 Sol-Attn + sol_fp8_layer_end: int | None = None # Exclusive end; None enables all remaining layers def __post_init__(self) -> None: if self.sparse_impl != "sol": @@ -159,6 +162,10 @@ def __post_init__(self) -> None: raise ValueError("Sol-Attn threshold type must be 'diag' or 'exact'") if self.sol_kv_splits not in ("auto", 1, 2, 4): raise ValueError("Sol-Attn KV splits must be 'auto', 1, 2, or 4") + if self.sol_fp8_layer_start < 0: + raise ValueError("Sol-Attn FP8 layer start must be non-negative") + if self.sol_fp8_layer_end is not None and self.sol_fp8_layer_end <= self.sol_fp8_layer_start: + raise ValueError("Sol-Attn FP8 layer end must be greater than its start") def should_use_dense(self, numeral_timestep: int, layer_idx: int) -> bool: """Check if dense attention should be used for current step/layer. @@ -229,6 +236,9 @@ def sol_attention( tau: float = 1.0, threshold_type: str = "diag", kv_splits: int | str = "auto", + sol_fp8: bool = False, + sol_fp8_layer_start: int = 0, + sol_fp8_layer_end: int | None = None, **kwargs: any, ) -> AttentionConfig: """Create a Sol-Attn config for dynamic sparse video self-attention.""" @@ -241,6 +251,9 @@ def sol_attention( sol_tau=tau, sol_threshold_type=threshold_type, sol_kv_splits=kv_splits, + sol_fp8=sol_fp8, + sol_fp8_layer_start=sol_fp8_layer_start, + sol_fp8_layer_end=sol_fp8_layer_end, ), **kwargs, ) diff --git a/telefuser/kernel/sol_attn/common/runtime.py b/telefuser/kernel/sol_attn/common/runtime.py index 5182d0c..5c93cdd 100644 --- a/telefuser/kernel/sol_attn/common/runtime.py +++ b/telefuser/kernel/sol_attn/common/runtime.py @@ -4,11 +4,14 @@ def to_cute_tensor(tensor): + leading_dim = tensor.ndim - 1 + if tensor.stride(leading_dim) != 1: + leading_dim = next(i for i, stride in enumerate(tensor.stride()) if stride == 1) return from_dlpack( tensor, assumed_align=16, enable_tvm_ffi=True, - ).mark_layout_dynamic(leading_dim=tensor.ndim - 1) + ).mark_layout_dynamic(leading_dim=leading_dim) __all__ = ["to_cute_tensor"] diff --git a/telefuser/kernel/sol_attn/interface.py b/telefuser/kernel/sol_attn/interface.py index 933513a..1843222 100644 --- a/telefuser/kernel/sol_attn/interface.py +++ b/telefuser/kernel/sol_attn/interface.py @@ -15,6 +15,22 @@ _compiled = {} +def _is_token_contiguous_bthd(x: torch.Tensor) -> bool: + batch, tokens, heads, head_dim = x.shape + return x.stride() == ( + heads * head_dim * tokens, + 1, + head_dim * tokens, + tokens, + ) + + +def _to_token_contiguous_bthd(x: torch.Tensor) -> torch.Tensor: + if _is_token_contiguous_bthd(x): + return x + return x.permute(0, 2, 3, 1).contiguous().permute(0, 3, 1, 2) + + def _validate_inputs( q, k, @@ -27,12 +43,13 @@ def _validate_inputs( raise ValueError("q, k, and v must share shape [B, T, H, 128]") if q.shape[1] == 0 or q.shape[3] != 128: raise ValueError("Sol-Attn requires T > 0 and head dimension 128") - if any(x.dtype != torch.bfloat16 for x in (q, k, v)): - raise TypeError("q, k, and v must use torch.bfloat16") + if any(x.dtype not in (torch.bfloat16, torch.float8_e4m3fn) for x in (q, k, v)): + raise TypeError("q, k, and v must use torch.bfloat16 or torch.float8_e4m3fn") if q.device.type != "cuda" or k.device != q.device or v.device != q.device: raise ValueError("q, k, and v must be on the same CUDA device") - if not (q.is_contiguous() and k.is_contiguous() and v.is_contiguous()): - raise ValueError("q, k, and v must be contiguous BTHD tensors") + v_layout_valid = v.is_contiguous() or (v.dtype == torch.float8_e4m3fn and _is_token_contiguous_bthd(v)) + if not (q.is_contiguous() and k.is_contiguous() and v_layout_valid): + raise ValueError("q and k must be contiguous BTHD; FP8 v may also be token-contiguous BTHD") if thresh_type not in ("diag", "exact"): raise ValueError("thresh_type must be 'diag' or 'exact'") if not isinstance(sink_tokens, int): @@ -70,17 +87,10 @@ def _backend_for_arch( """Select CuTe when specialized and available, otherwise Triton.""" if arch[0] < 8: - raise RuntimeError( - "Sol-Attn requires an NVIDIA GPU with compute capability >= 8.0; " - f"got SM{arch[0]}{arch[1]}" - ) + raise RuntimeError(f"Sol-Attn requires an NVIDIA GPU with compute capability >= 8.0; got SM{arch[0]}{arch[1]}") cute_backend = _CUTE_BACKENDS.get(arch) if cute_backend is not None: - available = ( - _cute_runtime_available() - if cute_available is None - else cute_available - ) + available = _cute_runtime_available() if cute_available is None else cute_available if available: return cute_backend return "triton" @@ -125,18 +135,20 @@ def _compile_sm90( kv_splits, sink_range, stream, + fp8_inputs, ): import cutlass.cute as cute from .sm90 import make_kernel - operator = make_kernel(tokens, kv_splits) + operator = make_kernel(tokens, kv_splits, fp8_inputs=fp8_inputs) args = _to_cute_tensors(tensors) compiled = cute.compile( operator, *args, scale, sink_range, + tokens, stream=stream, options="--enable-tvm-ffi", ) @@ -209,28 +221,68 @@ def _sol_attn_cute( kv_splits, sink_tokens, sink_start, + q_scale=None, + k_scale=None, + v_scale=None, ): - from .preprocess import prepare - batch, tokens, heads, _ = q.shape + fp8_inputs = q.dtype == torch.float8_e4m3fn with torch.cuda.device(q.device): - kc, vc, threshold = prepare( - q, - k, - v, - scale=scale, - tau=tau, - thresh_type=thresh_type, - ) - output = torch.empty_like(v) + if fp8_inputs: + from .triton_ref.preprocess import prepare_sm90_fp8 + + kc, vc, threshold, kc_scale = prepare_sm90_fp8( + q, + k, + v, + scale=scale, + tau=tau, + thresh_type=thresh_type, + tokens=tokens, + q_scale=q_scale, + k_scale=k_scale, + v_scale=v_scale, + ) + else: + from .preprocess import prepare + + kc, vc, threshold = prepare( + q, + k, + v, + scale=scale, + tau=tau, + thresh_type=thresh_type, + ) + dummy_scale = torch.ones((1,), device=q.device, dtype=torch.float32) + q_scale = k_scale = v_scale = kc_scale = dummy_scale + if fp8_inputs and tokens % BLOCK_SIZE: + padded_tokens = ((tokens + BLOCK_SIZE - 1) // BLOCK_SIZE) * BLOCK_SIZE + q_padded = torch.zeros( + (batch, padded_tokens, heads, q.shape[-1]), + device=q.device, + dtype=q.dtype, + ) + k_padded = torch.zeros_like(q_padded) + q_padded[:, :tokens].copy_(q) + k_padded[:, :tokens].copy_(k) + v_storage = torch.zeros( + (batch, heads, v.shape[-1], padded_tokens), + device=v.device, + dtype=v.dtype, + ) + v_storage[..., :tokens].copy_(v.permute(0, 2, 3, 1)) + q, k = q_padded, k_padded + v = v_storage.permute(0, 3, 1, 2) + output = torch.empty(v.shape, device=v.device, dtype=torch.bfloat16) lse = torch.empty( - (batch, tokens, heads), + (batch, q.shape[1], heads), device=q.device, dtype=torch.float32, ) stream = _stream(q.device) - key = (q.device.index, arch, batch, tokens, heads, kv_splits) + key = (q.device.index, arch, batch, tokens, heads, kv_splits, q.dtype) if arch == (9, 0): if sink_tokens: @@ -242,17 +294,17 @@ def _sol_attn_cute( sink_range = sink_start_block | (sink_end_block << 16) else: sink_range = 0 - tensors = [q, k, v, output, kc, vc, threshold, lse] + tensors = [q, k, v, output, kc, vc, threshold, lse, q_scale, k_scale, v_scale, kc_scale] if kv_splits > 1: tensors.extend( [ torch.empty( - (batch, tokens, kv_splits * heads, 128), + (batch, q.shape[1], kv_splits * heads, 128), device=q.device, dtype=torch.bfloat16, ), torch.empty( - (batch, tokens, kv_splits * heads), + (batch, q.shape[1], kv_splits * heads), device=q.device, dtype=torch.float32, ), @@ -268,6 +320,7 @@ def _sol_attn_cute( kv_splits, sink_range, stream, + fp8_inputs, ) else: args = _to_cute_tensors(tensors) @@ -275,6 +328,7 @@ def _sol_attn_cute( *args, scale, sink_range, + tokens, stream=stream, ) elif arch == (10, 0): @@ -329,7 +383,7 @@ def _sol_attn_cute( sink_end_block, stream=stream, ) - return output + return output[:, :tokens] def sol_attn( @@ -343,14 +397,81 @@ def sol_attn( kv_splits: int = 1, sink_tokens: int = 0, sink_start: int | None = None, + q_scale: torch.Tensor | None = None, + k_scale: torch.Tensor | None = None, + v_scale: torch.Tensor | None = None, + force_triton: bool = False, ) -> torch.Tensor: - """Compute noncausal Sol-Attn for contiguous BF16 BTHD tensors. + """Compute noncausal Sol-Attn for contiguous BF16 or FP8 BTHD tensors. ``sink_start`` and ``sink_tokens`` keep every KV block overlapping the corresponding contiguous token range exact for all queries. Omitting ``sink_start`` places the range at the token suffix. """ + fp8_inputs = any(x.dtype == torch.float8_e4m3fn for x in (q, k, v)) + if fp8_inputs: + if force_triton: + raise ValueError("force_triton is only supported for BF16 Sol-Attn") + if kv_splits not in (1, 2, 4): + raise ValueError("kv_splits must be 1, 2, or 4") + if not all(x.dtype == torch.float8_e4m3fn for x in (q, k, v)): + raise TypeError("q, k, and v must all use the same dtype") + if any(scale is None for scale in (q_scale, k_scale, v_scale)): + raise ValueError("FP8 Sol-Attn requires q_scale, k_scale, and v_scale") + _validate_inputs(q, k, v, thresh_type, sink_tokens, sink_start) + arch = tuple(torch.cuda.get_device_capability(q.device)) + native_sm90_fp8 = arch == (9, 0) and _cute_runtime_available() + blocks = (q.shape[1] + BLOCK_SIZE - 1) // BLOCK_SIZE + expected_scale_shape = (q.shape[0], blocks, q.shape[2]) + for name, tensor in (("q_scale", q_scale), ("k_scale", k_scale)): + if tensor.shape != expected_scale_shape or tensor.device != q.device: + raise ValueError(f"{name} must have shape {expected_scale_shape} on the Q/K/V device") + if not tensor.is_contiguous(): + raise ValueError(f"{name} must be contiguous") + if native_sm90_fp8: + expected_v_scale_shape = (q.shape[0], q.shape[2], q.shape[3]) + if v_scale.shape != expected_v_scale_shape or v_scale.device != q.device: + raise ValueError("SM90 FP8 Sol-Attn requires v_scale with shape [B, H, D]") + if not v_scale.is_contiguous(): + raise ValueError("v_scale must be contiguous") + v = _to_token_contiguous_bthd(v) + scale = q.shape[-1] ** -0.5 if scale is None else float(scale) + return _sol_attn_cute( + q, + k, + v, + arch=arch, + scale=scale, + tau=float(tau), + thresh_type=thresh_type, + kv_splits=kv_splits, + sink_tokens=sink_tokens, + sink_start=sink_start, + q_scale=q_scale, + k_scale=k_scale, + v_scale=v_scale, + ) + if v_scale.shape != expected_scale_shape or v_scale.device != q.device: + raise ValueError(f"v_scale must have shape {expected_scale_shape} on the Q/K/V device") + if not v_scale.is_contiguous(): + raise ValueError("v_scale must be contiguous") + from .triton_ref import sol_attn as triton_sol_attn + + return triton_sol_attn( + q, + k, + v, + scale=scale, + tau=tau, + thresh_type=thresh_type, + sink_tokens=sink_tokens, + sink_start=sink_start, + q_scale=q_scale, + k_scale=k_scale, + v_scale=v_scale, + ) + arch = _validate_inputs( q, k, @@ -361,7 +482,7 @@ def sol_attn( ) if kv_splits not in (1, 2, 4): raise ValueError("kv_splits must be 1, 2, or 4") - backend = _backend_for_arch(arch) + backend = "triton" if force_triton else _backend_for_arch(arch) scale = q.shape[-1] ** -0.5 if scale is None else float(scale) tau = float(tau) @@ -393,6 +514,9 @@ def sol_attn( kv_splits=kv_splits, sink_tokens=sink_tokens, sink_start=sink_start, + q_scale=q_scale, + k_scale=k_scale, + v_scale=v_scale, ) diff --git a/telefuser/kernel/sol_attn/sm90/atoms.py b/telefuser/kernel/sol_attn/sm90/atoms.py index e40d2a5..70add7b 100644 --- a/telefuser/kernel/sol_attn/sm90/atoms.py +++ b/telefuser/kernel/sol_attn/sm90/atoms.py @@ -7,15 +7,22 @@ from ._compat import sm90_utils -def make_pv_mma(tile_m: int = 64, tile_v: int = 128) -> cute.TiledMma: +def make_pv_mma( + tile_m: int = 64, + tile_v: int = 128, + a_dtype=cutlass.BFloat16, + b_dtype=cutlass.BFloat16, + source: str = "RS", +) -> cute.TiledMma: + b_major = "K" if b_dtype is cutlass.Float8E4M3FN else "MN" return sm90_utils.make_tiled_mma( - cutlass.BFloat16, + a_dtype, "K", - "MN", + b_major, tile_v, - source="RS", + source=source, atom_layout_mnk=(tile_m // 64, 1, 1), - b_dtype=cutlass.BFloat16, + b_dtype=b_dtype, acc_dtype=Float32, ) diff --git a/telefuser/kernel/sol_attn/sm90/fwd.py b/telefuser/kernel/sol_attn/sm90/fwd.py index 453ab94..b6fb585 100644 --- a/telefuser/kernel/sol_attn/sm90/fwd.py +++ b/telefuser/kernel/sol_attn/sm90/fwd.py @@ -1,5 +1,7 @@ """Hopper forward operators.""" +import math + import cuda.bindings.driver as cuda import cutlass import cutlass.cute as cute @@ -18,7 +20,7 @@ def __init__(self, *args, **kwargs): head_dim=self.tile_hdimv, tile_m=16, k_block_size=64, - log_max_splits=1 if self.sol_attn_num_splits == 2 else 2, + log_max_splits=int(math.log2(self.sol_attn_num_splits)), num_threads=128, stages=4, partial_dtype=cutlass.BFloat16, @@ -35,10 +37,15 @@ def __call__( vc: cute.Tensor, threshold: cute.Tensor, lse: cute.Tensor, + q_scale: cute.Tensor, + k_scale: cute.Tensor, + v_scale: cute.Tensor, + kc_scale: cute.Tensor, o_partial: cute.Tensor, lse_partial: cute.Tensor, softmax_scale: cutlass.Float32, sink_range: cutlass.Int32, + logical_tokens: cutlass.Int32, stream: cuda.CUstream = None, ): SolAttnMainloopSm90.__call__( @@ -51,8 +58,13 @@ def __call__( vc, threshold, lse_partial, + q_scale, + k_scale, + v_scale, + kc_scale, softmax_scale, sink_range, + logical_tokens, stream=stream, ) diff --git a/telefuser/kernel/sol_attn/sm90/kernel.py b/telefuser/kernel/sol_attn/sm90/kernel.py index c0b3fb7..007a41d 100644 --- a/telefuser/kernel/sol_attn/sm90/kernel.py +++ b/telefuser/kernel/sol_attn/sm90/kernel.py @@ -6,7 +6,7 @@ from .mainloop import SolAttnMainloopSm90 -def make_kernel(tokens: int, kv_splits: int): +def make_kernel(tokens: int, kv_splits: int, fp8_inputs: bool = False): blocks = (tokens + 63) // 64 full_groups, tail = divmod(blocks, 64) has_full_groups = tail == 0 @@ -25,22 +25,17 @@ def make_kernel(tokens: int, kv_splits: int): tile_n=64, num_stages=1, num_threads=128, - sol_attn_assume_lane_group_route_reduce=( - has_full_blocks and has_full_groups - ), + sol_attn_assume_lane_group_route_reduce=(has_full_blocks and has_full_groups), sol_attn_assume_full_k_exact_blocks=has_full_blocks, sol_attn_tail_exact_words1=0 < tail <= 8, sol_attn_assume_full_route_groups=has_full_groups, - sol_attn_static_num_full_route_groups=( - -1 if has_full_groups else full_groups - ), + sol_attn_static_num_full_route_groups=(-1 if has_full_groups else full_groups), sol_attn_static_tail_valid_count=(-1 if has_full_groups else tail), sol_attn_tail_physical_tile16=0 < tail <= 16, - sol_attn_exact_mask_seqlen_last_only=( - not has_full_blocks - ), + sol_attn_exact_mask_seqlen_last_only=(not has_full_blocks), sol_attn_tail16_lane_group_route_reduce=tail == 16, sol_attn_num_splits=kv_splits, + fp8_inputs=fp8_inputs, ) diff --git a/telefuser/kernel/sol_attn/sm90/mainloop.py b/telefuser/kernel/sol_attn/sm90/mainloop.py index 47ba69f..047dbb1 100644 --- a/telefuser/kernel/sol_attn/sm90/mainloop.py +++ b/telefuser/kernel/sol_attn/sm90/mainloop.py @@ -1,48 +1,47 @@ # Copyright (c) 2025, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao. # SM90 (Hopper) forward pass for flash attention, extracted from flash_fwd.py. +from functools import partial from types import SimpleNamespace from typing import Callable, Optional -from functools import partial import cuda.bindings.driver as cuda - import cutlass import cutlass.cute as cute -from cutlass import Float32, Int32, const_expr -from cutlass.cute.nvgpu import cpasync, warpgroup -from cutlass.utils import LayoutEnum import cutlass.utils.hopper_helpers as sm90_utils_basic -from cutlass import pipeline -from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait +from cutlass import Float32, Int32, const_expr, pipeline from cutlass.base_dsl.arch import Arch +from cutlass.cute.nvgpu import cpasync, warpgroup +from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait +from cutlass.utils import LayoutEnum -from ._compat import copy_utils -from ._compat import layout_utils -from ._compat import sm90_utils - -from telefuser.kernel.sol_attn._vendor.flash_attn.cute.cute_dsl_utils import assume_tensor_aligned +from telefuser.kernel.sol_attn._vendor.flash_attn.cute import pipeline as pipeline_custom from telefuser.kernel.sol_attn._vendor.flash_attn.cute import utils -from telefuser.kernel.sol_attn._vendor.flash_attn.cute.mask import AttentionMask -from telefuser.kernel.sol_attn._vendor.flash_attn.cute.softmax import Softmax, apply_score_mod_inner -from telefuser.kernel.sol_attn._vendor.flash_attn.cute.seqlen_info import SeqlenInfoQK from telefuser.kernel.sol_attn._vendor.flash_attn.cute.block_info import BlockInfo from telefuser.kernel.sol_attn._vendor.flash_attn.cute.block_sparsity import BlockSparseTensors -from telefuser.kernel.sol_attn._vendor.flash_attn.cute import pipeline as pipeline_custom -from telefuser.kernel.sol_attn._vendor.flash_attn.cute.pack_gqa import PackGQA, pack_gqa_layout, make_packgqa_tiled_tma_atom +from telefuser.kernel.sol_attn._vendor.flash_attn.cute.cute_dsl_utils import assume_tensor_aligned +from telefuser.kernel.sol_attn._vendor.flash_attn.cute.flash_fwd import FlashAttentionForwardBase +from telefuser.kernel.sol_attn._vendor.flash_attn.cute.mask import AttentionMask from telefuser.kernel.sol_attn._vendor.flash_attn.cute.named_barrier import NamedBarrierFwd -from ._compat.cute_dsl_utils import ParamsBase +from telefuser.kernel.sol_attn._vendor.flash_attn.cute.pack_gqa import ( + PackGQA, + make_packgqa_tiled_tma_atom, + pack_gqa_layout, +) +from telefuser.kernel.sol_attn._vendor.flash_attn.cute.seqlen_info import SeqlenInfoQK +from telefuser.kernel.sol_attn._vendor.flash_attn.cute.softmax import Softmax, apply_score_mod_inner from telefuser.kernel.sol_attn._vendor.flash_attn.cute.tile_scheduler import ( - TileSchedulerArguments, - SingleTileScheduler, SingleTileLPTScheduler, + SingleTileScheduler, SingleTileVarlenScheduler, + TileSchedulerArguments, ) -from telefuser.kernel.sol_attn._vendor.flash_attn.cute.flash_fwd import FlashAttentionForwardBase -from . import atoms as sol_attn_atoms -from . import exact as exact_stream from telefuser.kernel.sol_attn.common import selector as sol_attn_selector +from . import atoms as sol_attn_atoms +from . import exact as exact_stream +from ._compat import copy_utils, layout_utils, sm90_utils +from ._compat.cute_dsl_utils import ParamsBase SOL_ATTN_ROUTE_MASK_BARRIER_ID = 7 SOL_ATTN_ROUTE_SUM_BARRIER_ID = 8 @@ -62,14 +61,18 @@ def __init__( sol_attn_exact_mask_seqlen_last_only: bool = False, sol_attn_tail16_lane_group_route_reduce: bool = False, sol_attn_num_splits: int = 1, + fp8_inputs: bool = False, **kwargs, ): super().__init__(*args, **kwargs) - self.qk_dtype = cutlass.BFloat16 - self.pv_dtype = self.dtype + self.fp8_inputs = fp8_inputs + self.qk_dtype = cutlass.Float8E4M3FN if fp8_inputs else cutlass.BFloat16 + self.p_dtype = cutlass.Float8E4M3FN if fp8_inputs else self.dtype + self.v_dtype = cutlass.Float8E4M3FN if fp8_inputs else self.dtype + self.fp8_probability_scale = 1.0 self.sol_attn_group_size = 64 self.sol_attn_group_words = 2 - self.mma_pv_is_rs = True + self.mma_pv_is_rs = not fp8_inputs self.sol_attn_mma_regs_override = 128 self.sol_attn_warp_route_mask = True self.sol_attn_fast_route_lens = True @@ -82,9 +85,7 @@ def __init__( self.sol_attn_assume_full_route_groups = sol_attn_assume_full_route_groups self.sol_attn_static_num_full_route_groups = sol_attn_static_num_full_route_groups self.sol_attn_static_tail_valid_count = sol_attn_static_tail_valid_count - self.sol_attn_tail_exact_words1 = ( - sol_attn_tail_exact_words1 and 0 < self.sol_attn_static_tail_valid_count <= 32 - ) + self.sol_attn_tail_exact_words1 = sol_attn_tail_exact_words1 and 0 < self.sol_attn_static_tail_valid_count <= 32 self.sol_attn_tail_route_mask_words1 = False self.sol_attn_tail_physical_tile16 = ( sol_attn_tail_physical_tile16 and 0 < self.sol_attn_static_tail_valid_count <= 16 @@ -109,25 +110,22 @@ def __init__( def _get_smem_layout_atom(self): sQ_layout_atom = warpgroup.make_smem_layout_atom( - sm90_utils_basic.get_smem_layout_atom( - LayoutEnum.ROW_MAJOR, self.qk_dtype, self.tile_hdim - ), + sm90_utils_basic.get_smem_layout_atom(LayoutEnum.ROW_MAJOR, self.qk_dtype, self.tile_hdim), self.qk_dtype, ) sK_layout_atom = sQ_layout_atom sV_layout_atom = warpgroup.make_smem_layout_atom( - sm90_utils_basic.get_smem_layout_atom( - LayoutEnum.ROW_MAJOR, self.pv_dtype, self.tile_hdimv - ), - self.pv_dtype, + sm90_utils_basic.get_smem_layout_atom(LayoutEnum.ROW_MAJOR, self.v_dtype, self.tile_hdimv), + self.v_dtype, + ) + sO_layout_atom = warpgroup.make_smem_layout_atom( + sm90_utils_basic.get_smem_layout_atom(LayoutEnum.ROW_MAJOR, self.dtype, self.tile_hdimv), + self.dtype, ) - sO_layout_atom = sV_layout_atom if not self.mma_pv_is_rs: sP_layout_atom = warpgroup.make_smem_layout_atom( - sm90_utils_basic.get_smem_layout_atom( - LayoutEnum.ROW_MAJOR, self.pv_dtype, self.tile_n - ), - self.pv_dtype, + sm90_utils_basic.get_smem_layout_atom(LayoutEnum.ROW_MAJOR, self.p_dtype, self.tile_n), + self.p_dtype, ) else: sP_layout_atom = None @@ -135,18 +133,21 @@ def _get_smem_layout_atom(self): def _get_tiled_mma(self): tiled_mma_qk = sm90_utils.make_tiled_mma( - cutlass.BFloat16, + self.qk_dtype, "K", "K", self.tile_n, source="SS", atom_layout_mnk=(self.tile_m // 64, 1, 1), - b_dtype=cutlass.BFloat16, + b_dtype=self.qk_dtype, acc_dtype=Float32, ) tiled_mma_pv = sol_attn_atoms.make_pv_mma( tile_m=self.tile_m, tile_v=self.tile_hdimv, + a_dtype=self.p_dtype, + b_dtype=self.v_dtype, + source="RS" if self.mma_pv_is_rs else "SS", ) return tiled_mma_qk, tiled_mma_pv @@ -175,27 +176,191 @@ def sol_attn_qk_gemm_zero_init( swap_AB, ) + @cute.jit + def sol_attn_pv_gemm( + self, + tiled_mma: cute.TiledMma, + acc: cute.Tensor, + tile_acc: Optional[cute.Tensor], + tCrA: cute.Tensor, + tCrB: cute.Tensor, + zero_init, + B_idx: Int32, + wg_wait: cutlass.Constexpr[int], + ): + if const_expr(self.fp8_inputs and self.sol_attn_num_splits == 1): + sm90_utils.gemm_w_idx( + tiled_mma, + tile_acc, + tCrA, + tCrB, + zero_init=True, + B_idx=B_idx, + wg_wait=0, + ) + acc.store(acc.load() + tile_acc.load()) + elif const_expr(self.fp8_inputs): + sm90_utils.gemm_w_idx( + tiled_mma, + acc, + tCrA, + tCrB, + zero_init=False, + B_idx=B_idx, + wg_wait=wg_wait, + ) + else: + sm90_utils.gemm_w_idx( + tiled_mma, + acc, + tCrA, + tCrB, + zero_init=zero_init, + B_idx=B_idx, + wg_wait=wg_wait, + ) + + @cute.jit + def sol_attn_scale_exact_scores( + self, + acc_S: cute.Tensor, + q_block: Int32, + n_block: Int32, + q_scale: cute.Tensor, + k_scale: cute.Tensor, + ): + """Apply one Q/K scale per N64 tile to an FP8 QK accumulator.""" + + acc_S_mn = layout_utils.reshape_acc_to_mn(acc_S) + factor = Float32(q_scale[q_block]) * Float32(k_scale[n_block]) + for i in cutlass.range_constexpr(cute.size(acc_S_mn)): + acc_S_mn[i] = Float32(acc_S_mn[i]) * factor + + @cute.jit + def sol_attn_scale_route_scores( + self, + acc_S: cute.Tensor, + tScS_mn: cute.Tensor, + q_block: Int32, + route_n_block: Int32, + q_scale: cute.Tensor, + kc_scale: cute.Tensor, + ): + """Apply block-scaled Q and per-centroid K dequantization.""" + + acc_S_mn = layout_utils.reshape_acc_to_mn(acc_S) + q_factor = Float32(q_scale[q_block]) + for i in cutlass.range_constexpr(cute.size(acc_S_mn)): + col = tScS_mn[i][1] + factor = q_factor * Float32(kc_scale[route_n_block + col]) + acc_S_mn[i] = Float32(acc_S_mn[i]) * factor + + @cute.jit + def sol_attn_scale_route_probabilities( + self, + acc_S: cute.Tensor, + tScS_mn: cute.Tensor, + route_n_block: Int32, + seqlen: SeqlenInfoQK, + ): + """Convert centroid probabilities to equivalent block-sum weights.""" + + acc_S_mn = layout_utils.reshape_acc_to_mn(acc_S) + for i in cutlass.range_constexpr(cute.size(acc_S_mn)): + col = tScS_mn[i][1] + n_block = route_n_block + col + current_len = seqlen.seqlen_k - n_block * Int32(self.tile_n) + if current_len > Int32(self.tile_n): + current_len = Int32(self.tile_n) + if current_len < Int32(0): + current_len = Int32(0) + acc_S_mn[i] = Float32(acc_S_mn[i]) * Float32(current_len) + + @cute.jit + def sol_attn_convert_probability( + self, + src: cute.Tensor, + dst: cute.Tensor, + ): + """Convert post-softmax probabilities to the PV MMA operand dtype.""" + + if const_expr(self.fp8_inputs and self.mma_pv_is_rs): + for i in cutlass.range_constexpr(cute.size(src)): + dst[i] = cutlass.Float8E4M3FN(Float32(src[i]) * Float32(self.fp8_probability_scale)) + + tid = cute.arch.thread_idx()[0] % Int32(4) + values_u32 = cute.recast_tensor(dst, cutlass.Uint32) + for n in cutlass.range_constexpr(cute.size(values_u32, mode=[1])): + for k in cutlass.range_constexpr(cute.size(values_u32, mode=[2])): + for ii in cutlass.range_constexpr(0, 8, 4): + value0 = values_u32[ii // 2, n, k] + value1 = values_u32[ii // 2 + 1, n, k] + + send_high = 1 + if tid == Int32(1) or tid == Int32(2): + send_high = 0 + recv_lane = (Int32(0x3021) >> (tid * Int32(4))) & Int32(0xF) + value_a = value1 + if send_high == 0: + value_a = value0 + value_a = cute.arch.shuffle_sync_op(value_a, recv_lane, 0xFFFFFFFF, 7199) + + send_high = 1 - send_high + recv_lane = (Int32(0x2130) >> (tid * Int32(4))) & Int32(0xF) + value_b = value1 + if send_high == 0: + value_b = value0 + value_b = cute.arch.shuffle_sync_op(value_b, recv_lane, 0xFFFFFFFF, 7199) + + order0 = 0x5410 + order1 = 0x7632 + if send_high == 0: + order0 = 0x1054 + order1 = 0x3276 + values_u32[ii // 2, n, k] = cute.arch.prmt(value_a, value_b, order0) + values_u32[ii // 2 + 1, n, k] = cute.arch.prmt(value_a, value_b, order1) + elif const_expr(self.fp8_inputs): + for i in cutlass.range_constexpr(cute.size(src)): + dst[i] = cutlass.Float8E4M3FN(Float32(src[i]) * Float32(self.fp8_probability_scale)) + else: + utils.cvt_f16(src, dst) + + @cute.jit + def sol_attn_apply_v_scale( + self, + acc_O: cute.Tensor, + tiled_mma_pv: cute.TiledMma, + tidx: Int32, + v_scale: cute.Tensor, + ): + """Apply the per-channel V scale after all FP8 PV accumulations.""" + + thr_mma = tiled_mma_pv.get_slice(tidx) + cO = cute.make_identity_tensor((self.tile_m, self.tile_hdimv)) + taccOcO = layout_utils.reshape_acc_to_mn(thr_mma.partition_C(cO)) + acc_O_mn = layout_utils.reshape_acc_to_mn(acc_O) + for i in cutlass.range(cute.size(acc_O_mn), unroll_full=True): + col = taccOcO[i][1] + acc_O_mn[i] = Float32(acc_O_mn[i]) * Float32(v_scale[col]) / Float32(self.fp8_probability_scale) + def _get_shared_storage_cls(self): - sQ_struct, sK_struct = [ - cute.struct.Align[ - cute.struct.MemRange[self.qk_dtype, cute.cosize(layout)], self.buffer_align_bytes - ] - for layout in (self.sQ_layout, self.sK_layout) + sQ_elements = cute.cosize(self.sQ_layout) + if const_expr(self.fp8_inputs): + sQ_elements = max(sQ_elements, cute.cosize(self.sO_layout) * 2) + sQ_struct = cute.struct.Align[cute.struct.MemRange[self.qk_dtype, sQ_elements], self.buffer_align_bytes] + sK_struct = cute.struct.Align[ + cute.struct.MemRange[self.qk_dtype, cute.cosize(self.sK_layout)], self.buffer_align_bytes ] sV_struct = cute.struct.Align[ - cute.struct.MemRange[self.pv_dtype, cute.cosize(self.sV_layout)], + cute.struct.MemRange[self.v_dtype, cute.cosize(self.sV_layout)], self.buffer_align_bytes, ] cosize_sQV = max(cute.cosize(self.sQ_layout), cute.cosize(self.sV_layout)) - sQV_struct = cute.struct.Align[cute.struct.MemRange[self.pv_dtype, cosize_sQV], 1024] + sQV_struct = cute.struct.Align[cute.struct.MemRange[self.v_dtype, cosize_sQV], 1024] cosize_sP = cute.cosize(self.sP_layout) if const_expr(self.sP_layout is not None) else 0 - sP_struct = cute.struct.Align[cute.struct.MemRange[self.pv_dtype, cosize_sP], 1024] - route_mask_struct = cute.struct.Align[ - cute.struct.MemRange[Int32, 4], 16 - ] - route_sums_struct = cute.struct.Align[ - cute.struct.MemRange[Float32, 4 * self.tile_n], 16 - ] + sP_struct = cute.struct.Align[cute.struct.MemRange[self.p_dtype, cosize_sP], 1024] + route_mask_struct = cute.struct.Align[cute.struct.MemRange[Int32, 4], 16] + route_sums_struct = cute.struct.Align[cute.struct.MemRange[Float32, 4 * self.tile_n], 16] # 1 stage * 2 for Q pipeline (full + empty), self.num_stages*2 for K, self.num_stages*2 for V, mbar_ptr_Q_struct = cute.struct.MemRange[cutlass.Int64, 1 * 2] mbar_ptr_K_struct = cute.struct.MemRange[cutlass.Int64, self.num_stages * 2] @@ -264,7 +429,7 @@ def sol_attn_reduce_route_sums_guarded( ): """Fallback route-column reduction that ignores invalid q rows.""" - for off in cutlass.range_constexpr(self.tile_n): + for off in cutlass.range(self.tile_n, unroll_full=True): partial = Float32(0.0) for i in cutlass.range(cute.size(acc_S_mn), unroll_full=True): row = tScS_mn[i][0] @@ -502,12 +667,8 @@ def sol_attn_build_route_mask_from_acc( off1 = Int32(32) + lane route_col0 = route_col_offset + off0 route_col1 = route_col_offset + off1 - col_sum0 = Float32(route_sums[0, route_col0]) + Float32( - route_sums[1, route_col0] - ) - col_sum1 = Float32(route_sums[0, route_col1]) + Float32( - route_sums[1, route_col1] - ) + col_sum0 = Float32(route_sums[0, route_col0]) + Float32(route_sums[1, route_col0]) + col_sum1 = Float32(route_sums[0, route_col1]) + Float32(route_sums[1, route_col1]) col_sum0 += Float32(route_sums[2, route_col0]) col_sum1 += Float32(route_sums[2, route_col1]) col_sum0 += Float32(route_sums[3, route_col0]) @@ -547,12 +708,7 @@ def sol_attn_build_route_mask_from_acc( off = Int32(word * 32) + lane route_col = route_col_offset + off valid = True - if const_expr( - not ( - self.sol_attn_assume_full_route_groups - or assume_full_route_group - ) - ): + if const_expr(not (self.sol_attn_assume_full_route_groups or assume_full_route_group)): valid = off < valid_count exact = False if valid: @@ -572,10 +728,8 @@ def sol_attn_build_route_mask_from_acc( ) if sink_enabled: exact = exact or ( - group_start_n_block + off - >= sink_start_block - and group_start_n_block + off - < sink_end_block + group_start_n_block + off >= sink_start_block + and group_start_n_block + off < sink_end_block ) if const_expr(self.sol_attn_approx_colmask): column_mask = -Float32.inf @@ -588,21 +742,11 @@ def sol_attn_build_route_mask_from_acc( word_bits = Int32(0) if exact: word_bits = Int32(1) << lane - word_bits = word_bits | cute.arch.shuffle_sync_down( - word_bits, 16 - ) - word_bits = word_bits | cute.arch.shuffle_sync_down( - word_bits, 8 - ) - word_bits = word_bits | cute.arch.shuffle_sync_down( - word_bits, 4 - ) - word_bits = word_bits | cute.arch.shuffle_sync_down( - word_bits, 2 - ) - word_bits = word_bits | cute.arch.shuffle_sync_down( - word_bits, 1 - ) + word_bits = word_bits | cute.arch.shuffle_sync_down(word_bits, 16) + word_bits = word_bits | cute.arch.shuffle_sync_down(word_bits, 8) + word_bits = word_bits | cute.arch.shuffle_sync_down(word_bits, 4) + word_bits = word_bits | cute.arch.shuffle_sync_down(word_bits, 2) + word_bits = word_bits | cute.arch.shuffle_sync_down(word_bits, 1) if lane == Int32(0): if const_expr(word == 0): mask0 = word_bits @@ -619,9 +763,7 @@ def sol_attn_build_route_mask_from_acc( for off in cutlass.range_constexpr(self.sol_attn_group_size): route_col = route_col_offset + Int32(off) valid = True - if const_expr( - not (self.sol_attn_assume_full_route_groups or assume_full_route_group) - ): + if const_expr(not (self.sol_attn_assume_full_route_groups or assume_full_route_group)): valid = Int32(off) < valid_count col_sum = ( Float32(route_sums[0, route_col]) @@ -640,16 +782,12 @@ def sol_attn_build_route_mask_from_acc( ) if sink_enabled: exact = exact or ( - group_start_n_block + Int32(off) - >= sink_start_block - and group_start_n_block + Int32(off) - < sink_end_block + group_start_n_block + Int32(off) >= sink_start_block + and group_start_n_block + Int32(off) < sink_end_block ) if exact: - mask0, mask1, mask2, mask3 = ( - sol_attn_selector.sol_attn_set_exact_bit( - mask0, mask1, mask2, mask3, Int32(off) - ) + mask0, mask1, mask2, mask3 = sol_attn_selector.sol_attn_set_exact_bit( + mask0, mask1, mask2, mask3, Int32(off) ) return mask0, mask1, mask2, mask3 @@ -685,19 +823,14 @@ def sol_attn_mask_route_approx_columns( valid = group_col >= Int32(0) if valid: valid = group_col < valid_count - elif const_expr( - not (self.sol_attn_assume_full_route_groups or assume_full_route_group) - ): + elif const_expr(not (self.sol_attn_assume_full_route_groups or assume_full_route_group)): valid = col < valid_count exact = False if valid: route_mask_words = self.sol_attn_group_words if const_expr(route_mask_words_override != 0): route_mask_words = route_mask_words_override - if const_expr( - self.sol_attn_tail_route_mask_words1 - and not assume_full_route_group - ): + if const_expr(self.sol_attn_tail_route_mask_words1 and not assume_full_route_group): route_mask_words = 1 exact = sol_attn_selector.sol_attn_test_exact_bit_limited_words( mask0, mask1, mask2, mask3, group_col, route_mask_words @@ -771,9 +904,7 @@ def sol_attn_apply_route_current_lens_to_row_sum( """Correct route approx denominator for VC tiles that are block sums.""" acc_S_mn = layout_utils.reshape_acc_to_mn(acc_S) - last_n_block = ( - (seqlen.seqlen_k + Int32(self.tile_n - 1)) // Int32(self.tile_n) - ) - Int32(1) + last_n_block = ((seqlen.seqlen_k + Int32(self.tile_n - 1)) // Int32(self.tile_n)) - Int32(1) tail_len = seqlen.seqlen_k - last_n_block * Int32(self.tile_n) for r in cutlass.range(cute.size(softmax.row_sum), unroll_full=True): extra = Float32(0.0) @@ -807,9 +938,7 @@ def sol_attn_apply_route_current_lens_to_row_sum_fast( ): """Fast denominator correction for full-length route groups.""" - last_n_block = ( - (seqlen.seqlen_k + Int32(self.tile_n - 1)) // Int32(self.tile_n) - ) - Int32(1) + last_n_block = ((seqlen.seqlen_k + Int32(self.tile_n - 1)) // Int32(self.tile_n)) - Int32(1) tail_len = seqlen.seqlen_k - last_n_block * Int32(self.tile_n) group_end = group_start_n_block + valid_count full_len_group = (tail_len == Int32(self.tile_n)) or (group_end <= last_n_block) @@ -846,8 +975,13 @@ def __call__( mVC: cute.Tensor, mGlobalThresh: cute.Tensor, mLSE: Optional[cute.Tensor], + mQScale: cute.Tensor, + mKScale: cute.Tensor, + mVScale: cute.Tensor, + mKCScale: cute.Tensor, softmax_scale: Float32, sink_range: Int32, + logical_tokens: Int32, stream: cuda.CUstream = None, ): """Configure and launch the Hopper Sol-Attn kernel.""" @@ -866,34 +1000,27 @@ def __call__( aux_tensors = None self.varlen_q = mCuSeqlensQ is not None or mSeqUsedQ is not None - mQ, mK, mV, mO, mKC, mVC, mGlobalThresh = [ + mQ, mK, mV, mO, mKC, mVC, mGlobalThresh, mQScale, mKScale, mVScale, mKCScale = [ assume_tensor_aligned(t) - for t in (mQ, mK, mV, mO, mKC, mVC, mGlobalThresh) + for t in (mQ, mK, mV, mO, mKC, mVC, mGlobalThresh, mQScale, mKScale, mVScale, mKCScale) ] if const_expr(piecewise_k is not None): - piecewise_k, piecewise_v = [ - assume_tensor_aligned(t) for t in (piecewise_k, piecewise_v) - ] + piecewise_k, piecewise_v = [assume_tensor_aligned(t) for t in (piecewise_k, piecewise_v)] SOL_ATTN_BTHD_TRANSPOSE = [1, 3, 2, 0] SOL_ATTN_BNH_TRANSPOSE = [1, 2, 0] - mQ, mK, mV, mO, mKC, mVC = [ - layout_utils.select(t, SOL_ATTN_BTHD_TRANSPOSE) - for t in (mQ, mK, mV, mO, mKC, mVC) - ] - mGlobalThresh = layout_utils.select( - mGlobalThresh, SOL_ATTN_BNH_TRANSPOSE - ) + mQ, mK, mV, mO, mKC, mVC = [layout_utils.select(t, SOL_ATTN_BTHD_TRANSPOSE) for t in (mQ, mK, mV, mO, mKC, mVC)] + mGlobalThresh = layout_utils.select(mGlobalThresh, SOL_ATTN_BNH_TRANSPOSE) + if const_expr(self.fp8_inputs): + mQScale, mKScale, mKCScale = [ + layout_utils.select(t, SOL_ATTN_BNH_TRANSPOSE) for t in (mQScale, mKScale, mKCScale) + ] + mVScale = layout_utils.select(mVScale, [2, 1, 0]) if const_expr(piecewise_k is not None): piecewise_k, piecewise_v = [ - layout_utils.select(t, SOL_ATTN_BTHD_TRANSPOSE) - for t in (piecewise_k, piecewise_v) + layout_utils.select(t, SOL_ATTN_BTHD_TRANSPOSE) for t in (piecewise_k, piecewise_v) ] LSE_layout_transpose = [1, 2, 0] - mLSE = ( - layout_utils.select(mLSE, LSE_layout_transpose) - if const_expr(mLSE is not None) - else None - ) + mLSE = layout_utils.select(mLSE, LSE_layout_transpose) if const_expr(mLSE is not None) else None tiled_mma_qk, tiled_mma_pv = self._get_tiled_mma() self.num_mma_threads = tiled_mma_qk.size @@ -906,9 +1033,7 @@ def __call__( self.num_producer_threads = 32 self.num_Q_load_threads = self.num_threads_per_warp_group # If not TMA_Q self.num_epilogue_threads = self.num_mma_threads - self.num_mma_regs, self.num_producer_regs = {1: (256, 56), 2: (240, 24), 3: (160, 32)}[ - self.num_wg_mma - ] + self.num_mma_regs, self.num_producer_regs = {1: (256, 56), 2: (240, 24), 3: (160, 32)}[self.num_wg_mma] self.use_block_sparsity = cutlass.const_expr(blocksparse_tensors is not None) self.has_piecewise_kv = cutlass.const_expr(piecewise_k is not None) if const_expr(self.use_block_sparsity): @@ -917,17 +1042,13 @@ def __call__( raise NotImplementedError("one-warpgroup SOL_ATTN path does not support piecewise KV") self.use_scheduler_barrier = self.num_wg_mma == 2 - self.use_tma_Q = self.arch >= Arch.sm_90 and not ( - self.pack_gqa and self.tile_m % self.qhead_per_kvhead != 0 - ) + self.use_tma_Q = self.arch >= Arch.sm_90 and not (self.pack_gqa and self.tile_m % self.qhead_per_kvhead != 0) if const_expr(not self.use_tma_Q): raise NotImplementedError("one-warpgroup SOL_ATTN path requires TMA Q/O") # FP32 split partials require a direct register-to-global epilogue. # A BF16 split partial matches V/O dtype and can reuse the shared-memory # plus TMA-O epilogue. - self.use_tma_O = ( - self.sol_attn_num_splits == 1 or mO.element_type == self.dtype - ) + self.use_tma_O = self.sol_attn_num_splits == 1 or mO.element_type == self.dtype # Producer needs more registers when doing cp.async Q or KV loads if const_expr(self.num_wg_mma == 2 and (not self.use_tma_Q or not self.use_tma_KV)): self.num_mma_regs, self.num_producer_regs = 224, 40 @@ -936,18 +1057,26 @@ def __call__( self.rescale_O_before_gemm = False self._setup_attributes() # TODO: we prob don't need most of what's in _setup_attributes - self.sQ_layout, self.sK_layout, self.sV_layout, self.sO_layout = [ + self.sQ_layout, self.sK_layout = [ sm90_utils.make_smem_layout(mX.element_type, LayoutEnum.ROW_MAJOR, shape, stage) for mX, shape, stage in [ (mQ, (self.tile_m, self.tile_hdim), None), (mK, (self.tile_n, self.tile_hdim), self.num_stages), - (mV, (self.tile_n, self.tile_hdimv), self.num_stages), - # sO always holds the BF16 PV epilogue tile. Split-KV's - # global mO is an FP32 partial workspace, so derive this - # shared-memory layout from V instead of global O. - (mV, (self.tile_m, self.tile_hdimv), None), ] ] + v_layout = LayoutEnum.COL_MAJOR if const_expr(self.fp8_inputs) else LayoutEnum.ROW_MAJOR + self.sV_layout = sm90_utils.make_smem_layout( + mV.element_type, + v_layout, + (self.tile_n, self.tile_hdimv), + self.num_stages, + ) + self.sO_layout = sm90_utils.make_smem_layout( + self.dtype, + LayoutEnum.ROW_MAJOR, + (self.tile_m, self.tile_hdimv), + None, + ) self.sP_layout = None if const_expr(not self.mma_pv_is_rs): self.sP_layout = sm90_utils.make_smem_layout( @@ -1042,9 +1171,7 @@ def __call__( if const_expr(self.use_tma_O): mO_tma = mO_og if const_expr(self.pack_gqa) else mO if const_expr(self.varlen_q): - mO_tma = copy_utils.create_ragged_tensor_for_tma( - mO_tma, ragged_dim=0, ptr_shift=True - ) + mO_tma = copy_utils.create_ragged_tensor_for_tma(mO_tma, ragged_dim=0, ptr_shift=True) tma_atom_O, tma_tensor_O = make_tiled_tma_atom_fn( gmem_tiled_copy_O, mO_tma, @@ -1055,20 +1182,14 @@ def __call__( TileScheduler = SingleTileVarlenScheduler else: TileScheduler = ( - SingleTileScheduler - if const_expr(not self.is_causal or self.is_local) - else SingleTileLPTScheduler + SingleTileScheduler if const_expr(not self.is_causal or self.is_local) else SingleTileLPTScheduler ) tile_sched_args = TileSchedulerArguments( cute.ceil_div(cute.size(mQ.shape[0]), self.tile_m), cute.size(mQ.shape[2]), - cute.size(mQ.shape[3]) - if const_expr(mCuSeqlensQ is None) - else cute.size(mCuSeqlensQ.shape[0] - 1), + cute.size(mQ.shape[3]) if const_expr(mCuSeqlensQ is None) else cute.size(mCuSeqlensQ.shape[0] - 1), self.sol_attn_num_splits, - cute.size(mK.shape[0]) - if const_expr(mPageTable is None) - else mK.shape[0] * mPageTable.shape[1], + cute.size(mK.shape[0]) if const_expr(mPageTable is None) else mK.shape[0] * mPageTable.shape[1], mQ.shape[1], mV.shape[1], total_q=cute.size(mQ.shape[0]) @@ -1085,14 +1206,10 @@ def __call__( ) tile_sched_params = TileScheduler.to_underlying_arguments(tile_sched_args) grid_dim = TileScheduler.get_grid_shape(tile_sched_params) - softmax_scale_log2, softmax_scale = utils.compute_softmax_scale_log2( - softmax_scale, self.score_mod - ) + softmax_scale_log2, softmax_scale = utils.compute_softmax_scale_log2(softmax_scale, self.score_mod) window_size_left = Int32(window_size_left) if window_size_left is not None else None window_size_right = Int32(window_size_right) if window_size_right is not None else None - fastdiv_mods = utils.compute_fastdiv_mods( - mQ, mK, self.qhead_per_kvhead, self.pack_gqa, aux_tensors, mPageTable - ) + fastdiv_mods = utils.compute_fastdiv_mods(mQ, mK, self.qhead_per_kvhead, self.pack_gqa, aux_tensors, mPageTable) self.kernel( tma_tensor_Q if const_expr(self.use_tma_Q) else mQ, @@ -1105,6 +1222,10 @@ def __call__( tma_tensor_O if const_expr(self.use_tma_O) else mO, mGlobalThresh, mLSE, + mQScale, + mKScale, + mVScale, + mKCScale, mCuSeqlensQ, mCuSeqlensK, mSeqUsedQ, @@ -1121,6 +1242,7 @@ def __call__( softmax_scale_log2, softmax_scale, sink_range, + logical_tokens, window_size_left, window_size_right, learnable_sink, @@ -1160,6 +1282,10 @@ def kernel( mO: cute.Tensor, mGlobalThresh: cute.Tensor, mLSE: Optional[cute.Tensor], + mQScale: cute.Tensor, + mKScale: cute.Tensor, + mVScale: cute.Tensor, + mKCScale: cute.Tensor, mCuSeqlensQ: Optional[cute.Tensor], mCuSeqlensK: Optional[cute.Tensor], mSeqUsedQ: Optional[cute.Tensor], @@ -1176,6 +1302,7 @@ def kernel( softmax_scale_log2: Float32, softmax_scale: Optional[Float32], sink_range: Int32, + logical_tokens: Int32, window_size_left: Optional[Int32], window_size_right: Optional[Int32], learnable_sink: Optional[cute.Tensor], @@ -1290,9 +1417,7 @@ def kernel( if const_expr(not self.Q_in_regs): sV = storage.sV.get_tensor(sV_layout.outer, swizzle=sV_layout.inner) else: - sV = storage.sQ.get_tensor( - sV_layout.outer, swizzle=sV_layout.inner, dtype=mV.element_type - ) + sV = storage.sQ.get_tensor(sV_layout.outer, swizzle=sV_layout.inner, dtype=mV.element_type) # Transpose view of V to tensor with layout (head_dim_v, tile_n) for tiled mma sVt = layout_utils.transpose_view(sV) sP = None @@ -1300,9 +1425,7 @@ def kernel( sP = storage.sP.get_tensor(sP_layout.outer, swizzle=sP_layout.inner) # reuse sQ's data iterator sO = storage.sQ.get_tensor(sO_layout.outer, swizzle=sO_layout.inner, dtype=self.dtype) - route_mask = storage.route_mask.get_tensor( - cute.make_layout((4,)) - ) + route_mask = storage.route_mask.get_tensor(cute.make_layout((4,))) route_sums = storage.route_sums.get_tensor(cute.make_layout((4, self.tile_n))) block_info = BlockInfo( @@ -1317,10 +1440,8 @@ def kernel( ) SeqlenInfoCls = partial( SeqlenInfoQK.create, - seqlen_q_static=mQ.shape[0] if const_expr(not self.pack_gqa) else mQ.shape[0][1], - seqlen_k_static=mK.shape[0] - if const_expr(mPageTable is None) - else mK.shape[0] * mPageTable.shape[1], + seqlen_q_static=(logical_tokens if const_expr(not self.pack_gqa) else mQ.shape[0][1]), + seqlen_k_static=(logical_tokens if const_expr(mPageTable is None) else mK.shape[0] * mPageTable.shape[1]), mCuSeqlensQ=mCuSeqlensQ, mCuSeqlensK=mCuSeqlensK, mSeqUsedQ=mSeqUsedQ, @@ -1371,6 +1492,10 @@ def kernel( AttentionMaskCls, TileSchedulerCls, mGlobalThresh, + mQScale, + mKScale, + mVScale, + mKCScale, route_mask, route_sums, softmax_scale_log2, @@ -1403,9 +1528,7 @@ def epilogue_one_warpgroup_tma_o( barrier_id=int(NamedBarrierFwd.Epilogue), number_of_threads=self.num_epilogue_threads, ) - smem_copy_atom_O = utils.get_smem_store_atom( - self.arch.major * 10 + self.arch.minor, self.dtype - ) + smem_copy_atom_O = utils.get_smem_store_atom(self.arch.major * 10 + self.arch.minor, self.dtype) smem_thr_copy_O = cute.make_tiled_copy_C(smem_copy_atom_O, tiled_mma).get_slice(tidx) taccOrO = smem_thr_copy_O.retile(rO) taccOsO = smem_thr_copy_O.partition_D(sO) @@ -1415,9 +1538,7 @@ def epilogue_one_warpgroup_tma_o( if const_expr(mLSE is not None): mLSE_cur = mLSE[None, head_idx, batch_idx] gLSE = cute.local_tile(mLSE_cur, (self.tile_m,), (m_block,)) - gLSE_expanded_layout = cute.append( - gLSE.layout, cute.make_layout((self.tile_hdimv,), stride=(0,)) - ) + gLSE_expanded_layout = cute.append(gLSE.layout, cute.make_layout((self.tile_hdimv,), stride=(0,))) gLSE_expanded = cute.make_tensor(gLSE.iterator, gLSE_expanded_layout) thr_mma = tiled_mma.get_slice(tidx) taccOgLSE = layout_utils.reshape_acc_to_mn(thr_mma.partition_C(gLSE_expanded)) @@ -1425,10 +1546,7 @@ def epilogue_one_warpgroup_tma_o( t0accOcO = layout_utils.reshape_acc_to_mn(thr_mma.get_slice(0).partition_C(cO)) if taccOcO[0][1] == 0: for m in cutlass.range_constexpr(cute.size(taccOgLSE.shape[1])): - if ( - t0accOcO[m, 0][0] - < seqlen.seqlen_q - m_block * self.tile_m - taccOcO[0][0] - ): + if t0accOcO[m, 0][0] < seqlen.seqlen_q - m_block * self.tile_m - taccOcO[0][0]: taccOgLSE[m, 0] = lse[m] mO_cur = mO[None, None, head_idx, batch_idx] @@ -1438,9 +1556,7 @@ def epilogue_one_warpgroup_tma_o( number_of_threads=self.num_epilogue_threads, ) gO = cute.local_tile(mO_cur, (self.tile_m, self.tile_hdimv), (m_block, 0)) - store_O, _, _ = copy_utils.tma_get_copy_fn( - tma_atom_O, 0, cute.make_layout(1), sO, gO, single_stage=True - ) + store_O, _, _ = copy_utils.tma_get_copy_fn(tma_atom_O, 0, cute.make_layout(1), sO, gO, single_stage=True) warp_idx = cute.arch.make_warp_uniform(cute.arch.warp_idx()) if warp_idx == Int32(0): store_O() @@ -1469,9 +1585,7 @@ def epilogue_one_warpgroup_split_partial( """ mO_cur = mO[None, None, partial_head_idx, batch_idx] - gO = cute.local_tile( - mO_cur, (self.tile_m, self.tile_hdimv), (m_block, 0) - ) + gO = cute.local_tile(mO_cur, (self.tile_m, self.tile_hdimv), (m_block, 0)) copy_atom = cute.make_copy_atom( cute.nvgpu.CopyUniversalOp(), Float32, @@ -1486,29 +1600,16 @@ def epilogue_one_warpgroup_split_partial( mLSE_cur = mLSE[None, partial_head_idx, batch_idx] gLSE = cute.local_tile(mLSE_cur, (self.tile_m,), (m_block,)) - gLSE_expanded_layout = cute.append( - gLSE.layout, cute.make_layout((self.tile_hdimv,), stride=(0,)) - ) - gLSE_expanded = cute.make_tensor( - gLSE.iterator, gLSE_expanded_layout - ) + gLSE_expanded_layout = cute.append(gLSE.layout, cute.make_layout((self.tile_hdimv,), stride=(0,))) + gLSE_expanded = cute.make_tensor(gLSE.iterator, gLSE_expanded_layout) thr_mma = tiled_mma.get_slice(tidx) - taccOgLSE = layout_utils.reshape_acc_to_mn( - thr_mma.partition_C(gLSE_expanded) - ) + taccOgLSE = layout_utils.reshape_acc_to_mn(thr_mma.partition_C(gLSE_expanded)) cO = cute.make_identity_tensor((self.tile_m, self.tile_hdimv)) taccOcO = layout_utils.reshape_acc_to_mn(thr_mma.partition_C(cO)) - t0accOcO = layout_utils.reshape_acc_to_mn( - thr_mma.get_slice(0).partition_C(cO) - ) + t0accOcO = layout_utils.reshape_acc_to_mn(thr_mma.get_slice(0).partition_C(cO)) if taccOcO[0][1] == 0: for m in cutlass.range_constexpr(cute.size(taccOgLSE.shape[1])): - if ( - t0accOcO[m, 0][0] - < seqlen.seqlen_q - - m_block * self.tile_m - - taccOcO[0][0] - ): + if t0accOcO[m, 0][0] < seqlen.seqlen_q - m_block * self.tile_m - taccOcO[0][0]: taccOgLSE[m, 0] = lse[m] @cute.jit @@ -1543,6 +1644,10 @@ def mma_one_warpgroup_sol_attn_route_tma( AttentionMaskCls: Callable, TileSchedulerCls: cutlass.Constexpr[Callable], mGlobalThresh: cute.Tensor, + mQScale: cute.Tensor, + mKScale: cute.Tensor, + mVScale: cute.Tensor, + mKCScale: cute.Tensor, route_mask: cute.Tensor, route_sums: cute.Tensor, softmax_scale_log2: Float32, @@ -1560,38 +1665,32 @@ def mma_one_warpgroup_sol_attn_route_tma( else: q_producer_phase = Int32(1) q_consumer_phase = Int32(0) - kv_producer_state = pipeline.make_pipeline_state( - pipeline.PipelineUserType.Producer, self.num_stages - ) - kv_consumer_state = pipeline.make_pipeline_state( - pipeline.PipelineUserType.Consumer, self.num_stages - ) + kv_producer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Producer, self.num_stages) + kv_consumer_state = pipeline.make_pipeline_state(pipeline.PipelineUserType.Consumer, self.num_stages) tile_scheduler = TileSchedulerCls() work_tile = tile_scheduler.initial_work_tile_info() if work_tile.is_valid_tile: m_block, head_idx, batch_idx, split_idx = work_tile.tile_idx partial_head_idx = ( - head_idx - + split_idx * mQ.shape[2] - if const_expr(self.sol_attn_num_splits > 1) - else head_idx + head_idx + split_idx * mQ.shape[2] if const_expr(self.sol_attn_num_splits > 1) else head_idx ) seqlen = SeqlenInfoCls(batch_idx) - head_idx_kv = ( - head_idx // self.qhead_per_kvhead - if const_expr(not self.pack_gqa) - else head_idx - ) + head_idx_kv = head_idx // self.qhead_per_kvhead if const_expr(not self.pack_gqa) else head_idx mQ_cur = seqlen.offset_batch_Q(mQ, batch_idx, dim=3)[None, None, head_idx] - mK_cur = seqlen.offset_batch_K(mK, batch_idx, dim=3)[ - None, None, head_idx_kv - ] - mV_cur = seqlen.offset_batch_K(mV, batch_idx, dim=3)[ - None, None, head_idx_kv - ] + mK_cur = seqlen.offset_batch_K(mK, batch_idx, dim=3)[None, None, head_idx_kv] + mV_cur = seqlen.offset_batch_K(mV, batch_idx, dim=3)[None, None, head_idx_kv] mKC_cur = mKC[None, None, head_idx_kv, batch_idx] mVC_cur = mVC[None, None, head_idx_kv, batch_idx] + mQScale_cur = None + mKScale_cur = None + mVScale_cur = None + mKCScale_cur = None + if const_expr(self.fp8_inputs): + mQScale_cur = mQScale[None, head_idx, batch_idx] + mKScale_cur = mKScale[None, head_idx_kv, batch_idx] + mVScale_cur = mVScale[None, head_idx_kv, batch_idx] + mKCScale_cur = mKCScale[None, head_idx_kv, batch_idx] gQ = cute.local_tile(mQ_cur, (self.tile_m, self.tile_hdim), (m_block, 0)) gK = cute.local_tile(mK_cur, (self.tile_n, self.tile_hdim), (None, 0)) @@ -1599,38 +1698,22 @@ def mma_one_warpgroup_sol_attn_route_tma( gKC = cute.local_tile(mKC_cur, (self.tile_n, self.tile_hdim), (None, 0)) gVC = cute.local_tile(mVC_cur, (self.tile_n, self.tile_hdimv), (None, 0)) - load_Q, _, _ = copy_utils.tma_get_copy_fn( - tma_atom_Q, 0, cute.make_layout(1), gQ, sQ, single_stage=True - ) - tma_load_K_fn, _, _ = copy_utils.tma_get_copy_fn( - tma_atom_K, 0, cute.make_layout(1), gK, sK - ) + load_Q, _, _ = copy_utils.tma_get_copy_fn(tma_atom_Q, 0, cute.make_layout(1), gQ, sQ, single_stage=True) + tma_load_K_fn, _, _ = copy_utils.tma_get_copy_fn(tma_atom_K, 0, cute.make_layout(1), gK, sK) tma_load_K_fn = copy_utils.tma_producer_copy_fn(tma_load_K_fn, pipeline_k) - tma_load_V_fn, _, _ = copy_utils.tma_get_copy_fn( - tma_atom_V, 0, cute.make_layout(1), gV, sV - ) + tma_load_V_fn, _, _ = copy_utils.tma_get_copy_fn(tma_atom_V, 0, cute.make_layout(1), gV, sV) tma_load_V_fn = copy_utils.tma_producer_copy_fn(tma_load_V_fn, pipeline_v) - tma_load_KC_fn, _, _ = copy_utils.tma_get_copy_fn( - tma_atom_KC, 0, cute.make_layout(1), gKC, sK - ) - tma_load_KC_fn = copy_utils.tma_producer_copy_fn( - tma_load_KC_fn, pipeline_k - ) - tma_load_VC_fn, _, _ = copy_utils.tma_get_copy_fn( - tma_atom_VC, 0, cute.make_layout(1), gVC, sV - ) - tma_load_VC_fn = copy_utils.tma_producer_copy_fn( - tma_load_VC_fn, pipeline_v - ) + tma_load_KC_fn, _, _ = copy_utils.tma_get_copy_fn(tma_atom_KC, 0, cute.make_layout(1), gKC, sK) + tma_load_KC_fn = copy_utils.tma_producer_copy_fn(tma_load_KC_fn, pipeline_k) + tma_load_VC_fn, _, _ = copy_utils.tma_get_copy_fn(tma_atom_VC, 0, cute.make_layout(1), gVC, sV) + tma_load_VC_fn = copy_utils.tma_producer_copy_fn(tma_load_VC_fn, pipeline_v) if warp_idx == Int32(0): pipeline_q.producer_acquire_w_index_phase(0, q_producer_phase) load_Q(tma_bar_ptr=pipeline_q.sync_object_full.get_barrier(0)) pipeline_q.consumer_wait_w_index_phase(0, q_consumer_phase) - warp_group_thread_layout = cute.make_layout( - 1, stride=self.num_threads_per_warp_group - ) + warp_group_thread_layout = cute.make_layout(1, stride=self.num_threads_per_warp_group) thr_mma_qk = tiled_mma_qk.get_slice(tidx) wg_mma_qk = tiled_mma_qk.get_slice(warp_group_thread_layout(Int32(0))) wg_mma_pv = tiled_mma_pv.get_slice(warp_group_thread_layout(Int32(0))) @@ -1647,23 +1730,32 @@ def mma_one_warpgroup_sol_attn_route_tma( acc_O, tOrP, tOrVt = sm90_utils.partition_fragment_ABC( wg_mma_pv, (self.tile_m, self.tile_hdimv, self.tile_n), sP, sVt ) - mma_pv_fn = partial(sm90_utils.gemm_w_idx, tiled_mma_pv, acc_O, tOrP, tOrVt) + tile_acc_O = ( + cute.make_rmem_tensor_like(acc_O, Float32) + if const_expr(self.fp8_inputs and self.sol_attn_num_splits == 1) + else None + ) + mma_pv_fn = partial( + self.sol_attn_pv_gemm, + tiled_mma_pv, + acc_O, + tile_acc_O, + tOrP, + tOrVt, + ) smem_copy_atom_P = utils.get_smem_store_atom( - self.arch.major * 10 + self.arch.minor, self.dtype + self.arch.major * 10 + self.arch.minor, + self.p_dtype, ) - smem_thr_copy_P = cute.make_tiled_copy_C( - smem_copy_atom_P, tiled_mma_qk - ).get_slice(tidx) + smem_thr_copy_P = cute.make_tiled_copy_C(smem_copy_atom_P, tiled_mma_qk).get_slice(tidx) tPsP = smem_thr_copy_P.partition_D(sP) if const_expr(sP is not None) else None + cS_route = cute.make_identity_tensor((self.tile_m, self.tile_n)) smem_copy_params = SimpleNamespace( smem_thr_copy_P=smem_thr_copy_P, tPsP=tPsP, ) acc_O.fill(0.0) - cS_route = cute.make_identity_tensor((self.tile_m, self.tile_n)) - tScS_route_mn = layout_utils.reshape_acc_to_mn( - thr_mma_qk.partition_C(cS_route) - ) + tScS_route_mn = layout_utils.reshape_acc_to_mn(thr_mma_qk.partition_C(cS_route)) mask = AttentionMaskCls(seqlen) mask_fn = partial( mask.apply_mask, @@ -1696,6 +1788,14 @@ def mma_one_warpgroup_sol_attn_route_tma( if const_expr(self.sol_attn_neutral_softmax_state): softmax.row_max.fill(-Float32.inf) softmax.row_sum.fill(0.0) + exact_score_scale_fn = None + if const_expr(self.fp8_inputs): + exact_score_scale_fn = partial( + self.sol_attn_scale_exact_scores, + q_block=m_block, + q_scale=mQScale_cur, + k_scale=mKScale_cur, + ) exact_mma_one_n_block = partial( self.mma_one_n_block, mma_qk_fn=mma_qk_fn, @@ -1706,7 +1806,7 @@ def mma_one_warpgroup_sol_attn_route_tma( smem_copy_params=smem_copy_params, softmax=softmax, score_mod_fn=score_mod_fn, - score_scale_fn=None, + score_scale_fn=exact_score_scale_fn, check_inf=not self.sol_attn_assume_nonempty_rows, ) n_block_min, n_block_max = block_info.get_n_block_min_max(seqlen, m_block) @@ -1718,16 +1818,11 @@ def mma_one_warpgroup_sol_attn_route_tma( else: tail_valid_count = Int32(0) elif const_expr(self.sol_attn_assume_full_route_groups): - num_full_route_groups = cute.ceil_div( - route_block_count, self.sol_attn_group_size - ) + num_full_route_groups = cute.ceil_div(route_block_count, self.sol_attn_group_size) tail_valid_count = Int32(0) else: num_full_route_groups = route_block_count // Int32(self.sol_attn_group_size) - tail_valid_count = ( - route_block_count - - num_full_route_groups * Int32(self.sol_attn_group_size) - ) + tail_valid_count = route_block_count - num_full_route_groups * Int32(self.sol_attn_group_size) num_route_groups = num_full_route_groups if tail_valid_count > Int32(0): num_route_groups += Int32(1) @@ -1735,34 +1830,22 @@ def mma_one_warpgroup_sol_attn_route_tma( split_group_begin = Int32(0) split_num_route_groups = num_route_groups else: - groups_per_split = ( - num_route_groups + self.sol_attn_num_splits - 1 - ) // self.sol_attn_num_splits + groups_per_split = (num_route_groups + self.sol_attn_num_splits - 1) // self.sol_attn_num_splits split_group_begin = split_idx * groups_per_split - split_group_end = cutlass.min( - split_group_begin + groups_per_split, num_route_groups - ) - split_num_route_groups = cutlass.max( - split_group_end - split_group_begin, Int32(0) - ) + split_group_end = cutlass.min(split_group_begin + groups_per_split, num_route_groups) + split_num_route_groups = cutlass.max(split_group_end - split_group_begin, Int32(0)) O_should_accumulate = self.sol_attn_neutral_softmax_state - for local_group_iter in cutlass.range( - split_num_route_groups, unroll=1 - ): + for local_group_iter in cutlass.range(split_num_route_groups, unroll=1): group_iter = split_group_begin + local_group_iter group_start = n_block_min + group_iter * Int32(self.sol_attn_group_size) route_valid_count = Int32(self.sol_attn_group_size) if const_expr(not self.sol_attn_assume_full_route_groups): if group_iter == num_full_route_groups and tail_valid_count > Int32(0): route_valid_count = tail_valid_count - route_col_offset = group_start - ( - group_start // Int32(self.tile_n) - ) * Int32(self.tile_n) + route_col_offset = group_start - (group_start // Int32(self.tile_n)) * Int32(self.tile_n) route_n_block = group_start - route_col_offset route_tile = route_n_block // Int32(self.tile_n) - has_next_route_group = ( - local_group_iter + Int32(1) < split_num_route_groups - ) + has_next_route_group = local_group_iter + Int32(1) < split_num_route_groups next_route_tile = Int32(-1) if has_next_route_group: next_group_start = group_start + Int32(self.sol_attn_group_size) @@ -1802,6 +1885,15 @@ def mma_one_warpgroup_sol_attn_route_tma( acc_S = mma_qk_fn(B_idx=kv_consumer_state.index, wg_wait=-1) warpgroup.wait_group(0) pipeline_k.consumer_release(kv_consumer_state) + if const_expr(self.fp8_inputs): + self.sol_attn_scale_route_scores( + acc_S, + tScS_route_mn, + m_block, + route_n_block, + mQScale_cur, + mKCScale_cur, + ) mask0, mask1, mask2, mask3 = self.sol_attn_build_route_mask_from_acc( acc_S, route_sums, @@ -1846,34 +1938,23 @@ def mma_one_warpgroup_sol_attn_route_tma( exact_mask3 = mask3 first_exact_exists = ( - (mask0 != Int32(0)) - or (mask1 != Int32(0)) - or (mask2 != Int32(0)) - or (mask3 != Int32(0)) + (mask0 != Int32(0)) or (mask1 != Int32(0)) or (mask2 != Int32(0)) or (mask3 != Int32(0)) ) if mask0 != Int32(0): first_lowbit = mask0 & (Int32(0) - mask0) - first_exact_n_block += sol_attn_selector.sol_attn_bfind_b32( - first_lowbit - ) + first_exact_n_block += sol_attn_selector.sol_attn_bfind_b32(first_lowbit) exact_mask0 = mask0 & (mask0 - Int32(1)) elif mask1 != Int32(0): first_lowbit = mask1 & (Int32(0) - mask1) - first_exact_n_block += Int32(32) + ( - sol_attn_selector.sol_attn_bfind_b32(first_lowbit) - ) + first_exact_n_block += Int32(32) + (sol_attn_selector.sol_attn_bfind_b32(first_lowbit)) exact_mask1 = mask1 & (mask1 - Int32(1)) elif mask2 != Int32(0): first_lowbit = mask2 & (Int32(0) - mask2) - first_exact_n_block += Int32(64) + ( - sol_attn_selector.sol_attn_bfind_b32(first_lowbit) - ) + first_exact_n_block += Int32(64) + (sol_attn_selector.sol_attn_bfind_b32(first_lowbit)) exact_mask2 = mask2 & (mask2 - Int32(1)) elif mask3 != Int32(0): first_lowbit = mask3 & (Int32(0) - mask3) - first_exact_n_block += Int32(96) + ( - sol_attn_selector.sol_attn_bfind_b32(first_lowbit) - ) + first_exact_n_block += Int32(96) + (sol_attn_selector.sol_attn_bfind_b32(first_lowbit)) exact_mask3 = mask3 & (mask3 - Int32(1)) if first_exact_exists and warp_idx == Int32(0): pipeline_k.producer_acquire(kv_producer_state) @@ -1911,9 +1992,8 @@ def mma_one_warpgroup_sol_attn_route_tma( valid_bits1 = Int32(-1) if valid1 < Int32(32): valid_bits1 = (Int32(1) << valid1) - Int32(1) - route_has_approx = ( - ((mask0 & valid_bits0) != valid_bits0) - or ((mask1 & valid_bits1) != valid_bits1) + route_has_approx = ((mask0 & valid_bits0) != valid_bits0) or ( + (mask1 & valid_bits1) != valid_bits1 ) self.sol_attn_mask_route_approx_columns( acc_S, @@ -1934,17 +2014,13 @@ def mma_one_warpgroup_sol_attn_route_tma( ) if route_has_approx: row_sum_prev = None - if const_expr( - self.sol_attn_fast_route_lens - and not self.sol_attn_full_block_row_sum_prescale - ): + if const_expr(self.sol_attn_fast_route_lens and not self.sol_attn_full_block_row_sum_prescale): row_sum_prev = cute.make_fragment_like(softmax.row_sum, Float32) row_sum_prev.store(softmax.row_sum.load()) + row_scale = cute.make_fragment_like(softmax.row_sum, Float32) if O_should_accumulate: if const_expr(self.sol_attn_full_block_row_sum_prescale): - for r in cutlass.range( - cute.size(softmax.row_sum), unroll_full=True - ): + for r in cutlass.range(cute.size(softmax.row_sum), unroll_full=True): softmax.row_sum[r] *= Float32(1.0 / self.tile_n) row_scale = softmax.online_softmax( acc_S, @@ -1953,9 +2029,7 @@ def mma_one_warpgroup_sol_attn_route_tma( ) softmax.rescale_O(acc_O, row_scale) if const_expr(self.sol_attn_full_block_row_sum_prescale): - for r in cutlass.range( - cute.size(softmax.row_sum), unroll_full=True - ): + for r in cutlass.range(cute.size(softmax.row_sum), unroll_full=True): softmax.row_sum[r] *= Float32(self.tile_n) elif const_expr(self.sol_attn_fast_route_lens): self.sol_attn_apply_route_current_lens_to_row_sum_fast( @@ -1987,9 +2061,7 @@ def mma_one_warpgroup_sol_attn_route_tma( check_inf=not self.sol_attn_assume_nonempty_rows, ) if const_expr(self.sol_attn_full_block_row_sum_prescale): - for r in cutlass.range( - cute.size(softmax.row_sum), unroll_full=True - ): + for r in cutlass.range(cute.size(softmax.row_sum), unroll_full=True): softmax.row_sum[r] *= Float32(self.tile_n) elif const_expr(self.sol_attn_fast_route_lens): self.sol_attn_apply_route_current_lens_to_row_sum_fast( @@ -2014,13 +2086,20 @@ def mma_one_warpgroup_sol_attn_route_tma( seqlen, softmax, ) + if const_expr(self.fp8_inputs): + self.sol_attn_scale_route_probabilities( + acc_S, + tScS_route_mn, + route_n_block, + seqlen, + ) tOrP_acc = layout_utils.reshape_acc_to_frgA(acc_S) tOrP_cur = ( tOrP if const_expr(self.mma_pv_is_rs) - else cute.make_rmem_tensor_like(tOrP_acc, self.dtype) + else cute.make_rmem_tensor_like(tOrP_acc, self.p_dtype) ) - utils.cvt_f16(tOrP_acc, tOrP_cur) + self.sol_attn_convert_probability(tOrP_acc, tOrP_cur) if const_expr(not self.mma_pv_is_rs): tPrP = smem_copy_params.smem_thr_copy_P.retile(tOrP_cur) cute.copy( @@ -2030,39 +2109,48 @@ def mma_one_warpgroup_sol_attn_route_tma( ) cute.arch.fence_view_async_shared() cute.arch.sync_warp() - if O_should_accumulate: + if const_expr(self.fp8_inputs and self.sol_attn_num_splits == 1): + sm90_utils.gemm_w_idx( + tiled_mma_pv, + tile_acc_O, + tOrP_cur if const_expr(self.mma_pv_is_rs) else tOrP, + tOrVt, + zero_init=True, + B_idx=kv_consumer_state.index, + wg_wait=0, + ) + acc_O.store(acc_O.load() + tile_acc_O.load()) + elif O_should_accumulate: sm90_utils.gemm_w_idx( tiled_mma_pv, acc_O, - tOrP_cur, + tOrP_cur if const_expr(self.mma_pv_is_rs) else tOrP, tOrVt, zero_init=False, B_idx=kv_consumer_state.index, wg_wait=-1, ) + warpgroup.wait_group(0) else: sm90_utils.gemm_w_idx( tiled_mma_pv, acc_O, - tOrP_cur, + tOrP_cur if const_expr(self.mma_pv_is_rs) else tOrP, tOrVt, zero_init=True, B_idx=kv_consumer_state.index, wg_wait=-1, ) - warpgroup.wait_group(0) + warpgroup.wait_group(0) O_should_accumulate = True pipeline_v.consumer_release(kv_consumer_state) kv_consumer_state.advance() last_n_block = Int32(-1) if const_expr( - (not self.sol_attn_assume_full_k_exact_blocks) - or self.sol_attn_exact_mask_seqlen_last_only + (not self.sol_attn_assume_full_k_exact_blocks) or self.sol_attn_exact_mask_seqlen_last_only ): - last_n_block = ( - (seqlen.seqlen_k + Int32(self.tile_n - 1)) // Int32(self.tile_n) - ) - Int32(1) + last_n_block = ((seqlen.seqlen_k + Int32(self.tile_n - 1)) // Int32(self.tile_n)) - Int32(1) if O_should_accumulate: ( kv_producer_state, @@ -2141,6 +2229,13 @@ def mma_one_warpgroup_sol_attn_route_tma( pipeline_q.consumer_release_w_index(0) final_scale = softmax.finalize(sink_val=None) softmax.rescale_O(acc_O, final_scale) + if const_expr(self.fp8_inputs): + self.sol_attn_apply_v_scale( + acc_O, + tiled_mma_pv, + tidx, + mVScale_cur, + ) if const_expr(self.use_tma_O): self.epilogue_one_warpgroup_tma_o( acc_O, @@ -2222,12 +2317,8 @@ def mma_one_n_block( row_scale = softmax.online_softmax(acc_S, is_first=is_first_n_block, check_inf=check_inf) tOrP_acc = layout_utils.reshape_acc_to_frgA(acc_S) - tOrP_cur = ( - tOrP - if const_expr(self.mma_pv_is_rs) - else cute.make_rmem_tensor_like(tOrP_acc, self.dtype) - ) - utils.cvt_f16(tOrP_acc, tOrP_cur) + tOrP_cur = tOrP if const_expr(self.mma_pv_is_rs) else cute.make_rmem_tensor_like(tOrP_acc, self.p_dtype) + self.sol_attn_convert_probability(tOrP_acc, tOrP_cur) if const_expr(not self.mma_pv_is_rs): tPrP = smem_copy_params.smem_thr_copy_P.retile(tOrP_cur) cute.copy(smem_copy_params.smem_thr_copy_P, tPrP, smem_copy_params.tPsP) @@ -2291,9 +2382,7 @@ def apply_score_mod( def warp_scheduler_barrier_sync(self): if const_expr(self.use_scheduler_barrier): cute.arch.barrier( - barrier_id=int(NamedBarrierFwd.WarpSchedulerWG1) - - 1 - + utils.canonical_warp_group_idx(sync=False), + barrier_id=int(NamedBarrierFwd.WarpSchedulerWG1) - 1 + utils.canonical_warp_group_idx(sync=False), number_of_threads=2 * self.num_threads_per_warp_group, ) diff --git a/telefuser/kernel/sol_attn/triton_ref/fwd.py b/telefuser/kernel/sol_attn/triton_ref/fwd.py index 67edcc5..5327284 100644 --- a/telefuser/kernel/sol_attn/triton_ref/fwd.py +++ b/telefuser/kernel/sol_attn/triton_ref/fwd.py @@ -16,7 +16,6 @@ from .preprocess import prepare as prepare_ptr - BLOCK = 64 GROUP = 32 @@ -29,11 +28,7 @@ def _use_tma(device) -> bool: @triton.autotune( - configs=[ - triton.Config({}, num_warps=warps, num_stages=stages) - for warps in (4, 8) - for stages in (1, 2, 3, 4) - ], + configs=[triton.Config({}, num_warps=warps, num_stages=stages) for warps in (4, 8) for stages in (1, 2, 3, 4)], key=["T"], ) @triton.jit @@ -80,39 +75,23 @@ def _forward_tma( row_max = tl.full((BLOCK_SIZE,), -float("inf"), tl.float32) scale_log2 = scale * 1.4426950408889634 tail_length = T - (NT - 1) * BLOCK_SIZE - route_threshold = tl.load( - threshold + (batch * NT + q_block) * H + head - ) + route_threshold = tl.load(threshold + (batch * NT + q_block) * H + head) for group_start in range(0, NT, GROUP_SIZE): block_indices = group_start + group_offsets valid = block_indices < NT - kc = kc_desc.load( - [batch, group_start, head, 0] - ).reshape([GROUP_SIZE, D]) - vc = vc_desc.load( - [batch, group_start, head, v_tile * BV] - ).reshape([GROUP_SIZE, BV]) + kc = kc_desc.load([batch, group_start, head, 0]).reshape([GROUP_SIZE, D]) + vc = vc_desc.load([batch, group_start, head, v_tile * BV]).reshape([GROUP_SIZE, BV]) scores = tl.dot(q, kc.T).to(tl.float32) * scale_log2 - exact = ( - (tl.sum(scores, axis=0) / q_len > route_threshold) - | (tl.abs(q_block - block_indices) <= 1) - ) + exact = (tl.sum(scores, axis=0) / q_len > route_threshold) | (tl.abs(q_block - block_indices) <= 1) if HAS_SINK: - exact = exact | ( - (block_indices >= sink_start_block) - & (block_indices < sink_end_block) - ) + exact = exact | ((block_indices >= sink_start_block) & (block_indices < sink_end_block)) exact = exact & valid approximate = valid & ~exact - approximate_scores = tl.where( - approximate[None, :], scores, -float("inf") - ) + approximate_scores = tl.where(approximate[None, :], scores, -float("inf")) new_max = tl.maximum(row_max, tl.max(approximate_scores, axis=1)) - alpha = tl.math.exp2( - tl.where(row_max == new_max, 0.0, row_max - new_max) - ) + alpha = tl.math.exp2(tl.where(row_max == new_max, 0.0, row_max - new_max)) approximate_probability = tl.where( approximate[None, :], tl.math.exp2(approximate_scores - new_max[:, None]), @@ -122,9 +101,7 @@ def _forward_tma( approximate_probability.to(vc.dtype), vc, ) - lengths = tl.where( - block_indices == NT - 1, tail_length, BLOCK_SIZE - ).to(tl.float32) + lengths = tl.where(block_indices == NT - 1, tail_length, BLOCK_SIZE).to(tl.float32) row_sum = row_sum * alpha + tl.sum( approximate_probability * lengths[None, :], axis=1, @@ -141,9 +118,7 @@ def _forward_tma( exact_offsets, ) kv_start = block * BLOCK_SIZE - k = k_desc.load( - [batch, kv_start, head, 0] - ).reshape([BLOCK_SIZE, D]) + k = k_desc.load([batch, kv_start, head, 0]).reshape([BLOCK_SIZE, D]) exact_scores = tl.dot(q, k.T).to(tl.float32) * scale_log2 exact_scores += tl.where( (kv_start + token_offsets)[None, :] < T, @@ -152,16 +127,12 @@ def _forward_tma( ) new_max = tl.maximum(row_max, tl.max(exact_scores, axis=1)) alpha = tl.math.exp2(row_max - new_max) - exact_probability = tl.math.exp2( - exact_scores - new_max[:, None] - ) + exact_probability = tl.math.exp2(exact_scores - new_max[:, None]) row_sum = row_sum * alpha + tl.sum( exact_probability, axis=1, ) - v = v_desc.load( - [batch, kv_start, head, v_tile * BV] - ).reshape([BLOCK_SIZE, BV]) + v = v_desc.load([batch, kv_start, head, v_tile * BV]).reshape([BLOCK_SIZE, BV]) output = output * alpha[:, None] + tl.dot( exact_probability.to(v.dtype), v, @@ -188,6 +159,9 @@ def _forward_ptr( q_ptr, k_ptr, v_ptr, + q_scale_ptr, + k_scale_ptr, + v_scale_ptr, kc_ptr, vc_ptr, threshold_ptr, @@ -204,6 +178,7 @@ def _forward_ptr( BV: tl.constexpr, BLOCK_SIZE: tl.constexpr, GROUP_SIZE: tl.constexpr, + FP8: tl.constexpr, ): v_tile, q_block, batch_head = ( tl.program_id(0), @@ -223,44 +198,35 @@ def _forward_ptr( value_dims = v_tile * BV + tl.arange(0, BV) q_tokens = q_block * BLOCK_SIZE + token_offsets q_valid = q_tokens < T - q_offsets = ( - ((batch * T + q_tokens[:, None]).to(tl.int64) * H + head) * D - + dims[None, :] - ) + q_offsets = ((batch * T + q_tokens[:, None]).to(tl.int64) * H + head) * D + dims[None, :] q = tl.load(q_ptr + q_offsets, mask=q_valid[:, None], other=0.0) + if FP8: + q_scale = tl.load(q_scale_ptr + (batch * NT + q_block) * H + head) + q_route = q.to(tl.float32) * q_scale + else: + q_route = q q_len = tl.minimum(BLOCK_SIZE, T - q_block * BLOCK_SIZE).to(tl.float32) output = tl.zeros([BLOCK_SIZE, BV], dtype=tl.float32) row_sum = tl.zeros((BLOCK_SIZE,), dtype=tl.float32) row_max = tl.full((BLOCK_SIZE,), -float("inf"), tl.float32) scale_log2 = scale * 1.4426950408889634 - route_threshold = tl.load( - threshold_ptr + (batch * NT + q_block) * H + head - ) + route_threshold = tl.load(threshold_ptr + (batch * NT + q_block) * H + head) for group_start in range(0, NT, GROUP_SIZE): block_indices = group_start + group_offsets valid = block_indices < NT - kc_offsets = ( - ((batch * NPAD + block_indices[:, None]) * H + head) * D - + dims[None, :] - ) - vc_offsets = ( - ((batch * NPAD + block_indices[:, None]) * H + head) * D - + value_dims[None, :] - ) + kc_offsets = ((batch * NPAD + block_indices[:, None]) * H + head) * D + dims[None, :] + vc_offsets = ((batch * NPAD + block_indices[:, None]) * H + head) * D + value_dims[None, :] kc = tl.load(kc_ptr + kc_offsets) vc = tl.load(vc_ptr + vc_offsets) - scores = tl.dot(q, kc.T).to(tl.float32) * scale_log2 - exact = ( - (tl.sum(scores, axis=0) / q_len > route_threshold) - | (tl.abs(q_block - block_indices) <= 1) - ) + if FP8: + kc = kc.to(tl.float32) + vc = vc.to(tl.float32) + scores = tl.dot(q_route, kc.T).to(tl.float32) * scale_log2 + exact = (tl.sum(scores, axis=0) / q_len > route_threshold) | (tl.abs(q_block - block_indices) <= 1) if HAS_SINK: - exact = exact | ( - (block_indices >= sink_start_block) - & (block_indices < sink_end_block) - ) + exact = exact | ((block_indices >= sink_start_block) & (block_indices < sink_end_block)) exact = exact & valid approximate = valid & ~exact @@ -273,13 +239,8 @@ def _forward_ptr( safe_scores = tl.where(has_approximate, approximate_scores, 0.0) candidate_max = tl.maximum(row_max, tl.max(safe_scores, axis=1)) new_max = tl.where(has_approximate, candidate_max, row_max) - alpha = tl.math.exp2( - tl.where(has_approximate, row_max - new_max, 0.0) - ) - probability = tl.math.exp2( - safe_scores - - tl.where(has_approximate, new_max, 0.0)[:, None] - ) + alpha = tl.math.exp2(tl.where(has_approximate, row_max - new_max, 0.0)) + probability = tl.math.exp2(safe_scores - tl.where(has_approximate, new_max, 0.0)[:, None]) probability = tl.where( has_approximate & approximate[None, :], probability, @@ -311,17 +272,17 @@ def _forward_ptr( ) kv_tokens = block * BLOCK_SIZE + token_offsets kv_valid = kv_tokens < T - k_offsets = ( - ((batch * T + kv_tokens[:, None]).to(tl.int64) * H + head) - * D - + dims[None, :] - ) + k_offsets = ((batch * T + kv_tokens[:, None]).to(tl.int64) * H + head) * D + dims[None, :] k = tl.load( k_ptr + k_offsets, mask=kv_valid[:, None], other=0.0, ) - exact_scores = tl.dot(q, k.T).to(tl.float32) * scale_log2 + if FP8: + k_scale = tl.load(k_scale_ptr + (batch * NT + block) * H + head) + exact_scores = tl.dot(q, k.T, out_dtype=tl.float32) * (q_scale * k_scale * scale_log2) + else: + exact_scores = tl.dot(q, k.T).to(tl.float32) * scale_log2 exact_scores += tl.where( kv_valid[None, :], 0.0, @@ -329,33 +290,38 @@ def _forward_ptr( ) new_max = tl.maximum(row_max, tl.max(exact_scores, axis=1)) alpha = tl.math.exp2(row_max - new_max) - exact_probability = tl.math.exp2( - exact_scores - new_max[:, None] - ) + exact_probability = tl.math.exp2(exact_scores - new_max[:, None]) row_sum = row_sum * alpha + tl.sum( exact_probability, axis=1, ) - v_offsets = ( - ((batch * T + kv_tokens[:, None]).to(tl.int64) * H + head) - * D - + value_dims[None, :] - ) + v_offsets = ((batch * T + kv_tokens[:, None]).to(tl.int64) * H + head) * D + value_dims[None, :] v = tl.load( v_ptr + v_offsets, mask=kv_valid[:, None], other=0.0, ) - output = output * alpha[:, None] + tl.dot( - exact_probability.to(v.dtype), - v, - ) + if FP8: + v_scale = tl.load(v_scale_ptr + (batch * NT + block) * H + head) + probability_max = tl.maximum( + 1.0e-6, + tl.minimum(1.0, tl.max(exact_probability, axis=1)), + ) + probability_scale = probability_max / 448.0 + probability_fp8 = (exact_probability / probability_scale[:, None]).to(tl.float8e4nv) + output = output * alpha[:, None] + tl.dot( + probability_fp8, + v, + out_dtype=tl.float32, + ) * (probability_scale[:, None] * v_scale) + else: + output = output * alpha[:, None] + tl.dot( + exact_probability.to(v.dtype), + v, + ) row_max = new_max - output_offsets = ( - ((batch * T + q_tokens[:, None]).to(tl.int64) * H + head) * D - + value_dims[None, :] - ) + output_offsets = ((batch * T + q_tokens[:, None]).to(tl.int64) * H + head) * D + value_dims[None, :] tl.store( o_ptr + output_offsets, (output / row_sum[:, None]).to(tl.bfloat16), @@ -373,9 +339,18 @@ def sol_attn( thresh_type: str = "diag", sink_tokens: int = 0, sink_start: int | None = None, + q_scale: torch.Tensor | None = None, + k_scale: torch.Tensor | None = None, + v_scale: torch.Tensor | None = None, ) -> torch.Tensor: - """Run Triton Sol-Attn on contiguous BF16 BTHD inputs.""" + """Run Triton Sol-Attn on contiguous BF16 or block-scaled FP8 BTHD inputs.""" + fp8_inputs = q.dtype == torch.float8_e4m3fn + if fp8_inputs: + if k.dtype != q.dtype or v.dtype != q.dtype: + raise TypeError("FP8 Sol-Attn requires q, k, and v to share dtype") + if any(scale is None for scale in (q_scale, k_scale, v_scale)): + raise ValueError("FP8 Sol-Attn requires q_scale, k_scale, and v_scale") arch = _validate_inputs( q, k, @@ -386,8 +361,7 @@ def sol_attn( ) if arch[0] < 8: raise RuntimeError( - "Triton Sol-Attn requires an NVIDIA GPU with compute " - f"capability >= 8.0; got SM{arch[0]}{arch[1]}" + f"Triton Sol-Attn requires an NVIDIA GPU with compute capability >= 8.0; got SM{arch[0]}{arch[1]}" ) scale = q.shape[-1] ** -0.5 if scale is None else float(scale) tau = float(tau) @@ -400,7 +374,7 @@ def sol_attn( ) use_tma = _use_tma(q.device) - if use_tma: + if use_tma and not fp8_inputs: # Keep the original descriptor-backed preprocessing on TMA devices. # The pointer preprocessing below exists only for older architectures. from ..preprocess import prepare as prepare_tma @@ -446,13 +420,22 @@ def sol_attn( tau=tau, thresh_type=thresh_type, tokens=tokens, + q_scale=q_scale, + k_scale=k_scale, + v_scale=v_scale, ) - output = torch.empty_like(v) - grid = lambda meta: (head_dim // meta["BV"], blocks, batch * heads) + output = torch.empty(v.shape, device=v.device, dtype=torch.bfloat16) + + def grid(meta): + return (head_dim // meta["BV"], blocks, batch * heads) + _forward_ptr[grid]( q, k, v, + q_scale if fp8_inputs else q.new_ones((1,), dtype=torch.float32), + k_scale if fp8_inputs else q.new_ones((1,), dtype=torch.float32), + v_scale if fp8_inputs else q.new_ones((1,), dtype=torch.float32), kc, vc, threshold, @@ -468,6 +451,7 @@ def sol_attn( NT=blocks, BLOCK_SIZE=BLOCK, GROUP_SIZE=GROUP, + FP8=fp8_inputs, ) return output diff --git a/telefuser/kernel/sol_attn/triton_ref/preprocess.py b/telefuser/kernel/sol_attn/triton_ref/preprocess.py index 6341736..bec3bb8 100644 --- a/telefuser/kernel/sol_attn/triton_ref/preprocess.py +++ b/telefuser/kernel/sol_attn/triton_ref/preprocess.py @@ -6,7 +6,6 @@ import triton import triton.language as tl - BLOCK_SIZE = 64 HEAD_DIM = 128 THRESHOLD_GROUP_SIZE = 64 @@ -14,43 +13,76 @@ @triton.autotune( - configs=[ - triton.Config({}, num_warps=warps, num_stages=stages) - for warps in (4, 8) - for stages in (1, 2) - ], + configs=[triton.Config({}, num_warps=warps, num_stages=stages) for warps in (4, 8) for stages in (1, 2)], key=["T"], ) @triton.jit def _reduce_kv_kernel( k, v, + k_scale, + v_scale, kc, vc, + kc_fp8, + vc_fp8, + kc_out_scale, T, TP, NPAD, H: tl.constexpr, + N: tl.constexpr, D: tl.constexpr, BLOCK: tl.constexpr, + FP8: tl.constexpr, + TOKEN_SCALES: tl.constexpr, + V_CHANNEL_SCALE: tl.constexpr, + V_TOKEN_CONTIGUOUS: tl.constexpr, + SM90_FP8_OUTPUTS: tl.constexpr, ): block, batch_head = tl.program_id(0), tl.program_id(1) batch, head = batch_head // H, batch_head % H tokens = block * BLOCK + tl.arange(0, BLOCK) dims = tl.arange(0, D) valid = tokens < T - offsets = ( - ((batch * TP + tokens[:, None]).to(tl.int64) * H + head) * D - + dims[None, :] - ) + offsets = ((batch * TP + tokens[:, None]).to(tl.int64) * H + head) * D + dims[None, :] + v_offsets = offsets + if V_TOKEN_CONTIGUOUS: + v_offsets = ((batch * H + head) * D + dims[None, :]) * TP + tokens[:, None] k_values = tl.load(k + offsets, mask=valid[:, None], other=0.0) - v_values = tl.load(v + offsets, mask=valid[:, None], other=0.0) + v_raw = tl.load(v + v_offsets, mask=valid[:, None], other=0.0) + v_values = v_raw + if FP8: + scale_offset = (batch * N + block) * H + head + if TOKEN_SCALES: + token_scale_offsets = (batch * TP + tokens) * H + head + k_values = ( + k_values.to(tl.float32) + * tl.load( + k_scale + token_scale_offsets, + mask=valid, + other=0.0, + )[:, None] + ) + else: + k_values = k_values.to(tl.float32) * tl.load(k_scale + scale_offset) + if V_CHANNEL_SCALE: + channel_offsets = (batch * H + head) * D + dims + v_values = v_values.to(tl.float32) * tl.load(v_scale + channel_offsets)[None, :] + else: + v_values = v_values.to(tl.float32) * tl.load(v_scale + scale_offset) block_len = tl.minimum(BLOCK, T - block * BLOCK).to(tl.float32) - summary_offsets = ( - ((batch * NPAD + block) * H + head) * D + dims - ) - tl.store(kc + summary_offsets, tl.sum(k_values, axis=0) / block_len) - tl.store(vc + summary_offsets, tl.sum(v_values, axis=0)) + summary_offsets = ((batch * NPAD + block) * H + head) * D + dims + k_summary = tl.sum(k_values, axis=0) / block_len + tl.store(kc + summary_offsets, k_summary) + if SM90_FP8_OUTPUTS: + kc_s = tl.maximum(tl.max(tl.abs(k_summary), axis=0), 1.0e-6) / 448.0 + tl.store(kc_out_scale + (batch * NPAD + block) * H + head, kc_s) + tl.store(kc_fp8 + summary_offsets, k_summary / kc_s) + vc_offsets = ((batch * H + head) * D + dims) * NPAD + block + tl.store(vc_fp8 + vc_offsets, tl.sum(v_raw.to(tl.float32), axis=0) / block_len) + else: + tl.store(vc + summary_offsets, tl.sum(v_values, axis=0)) @triton.jit @@ -74,10 +106,7 @@ def _reduce_kc_stats_kernel( for start in range(0, N, GROUP): block_indices = start + blocks valid = block_indices < N - offsets = ( - ((batch * NPAD + block_indices[:, None]) * H + head) * D - + dims[None, :] - ) + offsets = ((batch * NPAD + block_indices[:, None]) * H + head) * D + dims[None, :] values = tl.load( kc + offsets, mask=valid[:, None], @@ -95,6 +124,7 @@ def _reduce_kc_stats_kernel( @triton.jit def _diag_threshold_kernel( q, + q_scale, kc_mean, kc_var_diag, threshold, @@ -106,17 +136,29 @@ def _diag_threshold_kernel( D: tl.constexpr, BLOCK: tl.constexpr, TAU: tl.constexpr, + FP8: tl.constexpr, + TOKEN_SCALES: tl.constexpr, ): q_block, batch_head = tl.program_id(0), tl.program_id(1) batch, head = batch_head // H, batch_head % H tokens = q_block * BLOCK + tl.arange(0, BLOCK) dims = tl.arange(0, D) valid = tokens < T - offsets = ( - ((batch * TP + tokens[:, None]).to(tl.int64) * H + head) * D - + dims[None, :] - ) + offsets = ((batch * TP + tokens[:, None]).to(tl.int64) * H + head) * D + dims[None, :] q_values = tl.load(q + offsets, mask=valid[:, None], other=0.0) + if FP8: + if TOKEN_SCALES: + token_scale_offsets = (batch * TP + tokens) * H + head + q_values = ( + q_values.to(tl.float32) + * tl.load( + q_scale + token_scale_offsets, + mask=valid, + other=0.0, + )[:, None] + ) + else: + q_values = q_values.to(tl.float32) * tl.load(q_scale + (batch * N + q_block) * H + head) q_len = tl.minimum(BLOCK, T - q_block * BLOCK).to(tl.float32) q_centroid = tl.sum(q_values.to(tl.float32), axis=0) / q_len mean_kc = tl.load(kc_mean + batch_head * D + dims) @@ -137,6 +179,7 @@ def _diag_threshold_kernel( @triton.jit def _pool_query_kernel( q, + q_scale, q_bar, T, TP, @@ -144,17 +187,29 @@ def _pool_query_kernel( N: tl.constexpr, D: tl.constexpr, BLOCK: tl.constexpr, + FP8: tl.constexpr, + TOKEN_SCALES: tl.constexpr, ): q_block, batch_head = tl.program_id(0), tl.program_id(1) batch, head = batch_head // H, batch_head % H tokens = q_block * BLOCK + tl.arange(0, BLOCK) dims = tl.arange(0, D) valid = tokens < T - offsets = ( - ((batch * TP + tokens[:, None]).to(tl.int64) * H + head) * D - + dims[None, :] - ) + offsets = ((batch * TP + tokens[:, None]).to(tl.int64) * H + head) * D + dims[None, :] values = tl.load(q + offsets, mask=valid[:, None], other=0.0) + if FP8: + if TOKEN_SCALES: + token_scale_offsets = (batch * TP + tokens) * H + head + values = ( + values.to(tl.float32) + * tl.load( + q_scale + token_scale_offsets, + mask=valid, + other=0.0, + )[:, None] + ) + else: + values = values.to(tl.float32) * tl.load(q_scale + (batch * N + q_block) * H + head) q_len = tl.minimum(BLOCK, T - q_block * BLOCK).to(tl.float32) centroid = tl.sum(values.to(tl.float32), axis=0) / q_len tl.store(q_bar + (batch_head * N + q_block) * D + dims, centroid) @@ -183,12 +238,7 @@ def _exact_fused_threshold_kernel( other=0.0, ) mean_kc = tl.load(kc_mean + batch_head * D + dims) - second_moment = tl.load( - kc_second_moment - + batch_head * D * D - + dims[:, None] * D - + dims[None, :] - ) + second_moment = tl.load(kc_second_moment + batch_head * D * D + dims[:, None] * D + dims[None, :]) raw_mean = tl.sum(q_centroid.to(tl.float32) * mean_kc[None, :], axis=1) projected = tl.dot(q_centroid, second_moment, out_dtype=tl.float32) raw_second_moment = tl.sum( @@ -215,6 +265,8 @@ def _reduce_kv( v: torch.Tensor, *, tokens: int | None = None, + k_scale: torch.Tensor | None = None, + v_scale: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: batch, padded_tokens, heads, head_dim = k.shape tokens = padded_tokens if tokens is None else int(tokens) @@ -226,21 +278,115 @@ def _reduce_kv( dtype=torch.bfloat16, ) vc = torch.zeros_like(kc) + fp8_inputs = k.dtype == torch.float8_e4m3fn + v_channel_scale = fp8_inputs and v_scale is not None and v_scale.shape == (batch, heads, head_dim) + v_token_contiguous = v.stride(1) == 1 + dummy_scale = torch.ones((1,), device=k.device, dtype=torch.float32) _reduce_kv_kernel[(blocks, batch * heads)]( k, v, + k_scale if fp8_inputs else dummy_scale, + v_scale if fp8_inputs else dummy_scale, + kc, + vc, kc, vc, + dummy_scale, tokens, padded_tokens, padded_blocks, heads, + blocks, head_dim, BLOCK_SIZE, + FP8=fp8_inputs, + TOKEN_SCALES=False, + V_CHANNEL_SCALE=v_channel_scale, + V_TOKEN_CONTIGUOUS=v_token_contiguous, + SM90_FP8_OUTPUTS=False, ) return kc, vc +def prepare_sm90_fp8( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + *, + tau: float, + scale: float, + thresh_type: str = "diag", + tokens: int | None = None, + q_scale: torch.Tensor, + k_scale: torch.Tensor, + v_scale: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Build BF16 threshold stats and TMA-ready FP8 summaries in one reduction.""" + + batch, padded_tokens, heads, head_dim = k.shape + tokens = padded_tokens if tokens is None else int(tokens) + blocks = triton.cdiv(tokens, BLOCK_SIZE) + padded_blocks = triton.cdiv(blocks, SUMMARY_PAD) * SUMMARY_PAD + kc_stats = torch.zeros( + (batch, padded_blocks, heads, head_dim), + device=k.device, + dtype=torch.bfloat16, + ) + kc_fp8 = torch.zeros_like(kc_stats, dtype=torch.float8_e4m3fn) + vc_storage = torch.zeros( + (batch, heads, head_dim, padded_blocks), + device=k.device, + dtype=torch.float8_e4m3fn, + ) + kc_out_scale = torch.ones( + (batch, padded_blocks, heads), + device=k.device, + dtype=torch.float32, + ) + _reduce_kv_kernel[(blocks, batch * heads)]( + k, + v, + k_scale, + v_scale, + kc_stats, + kc_stats, + kc_fp8, + vc_storage, + kc_out_scale, + tokens, + padded_tokens, + padded_blocks, + heads, + blocks, + head_dim, + BLOCK_SIZE, + FP8=True, + TOKEN_SCALES=False, + V_CHANNEL_SCALE=True, + V_TOKEN_CONTIGUOUS=True, + SM90_FP8_OUTPUTS=True, + ) + if thresh_type == "exact": + threshold = _compute_exact_threshold( + q, + kc_stats, + tau=tau, + scale=scale, + tokens=tokens, + q_scale=q_scale, + ) + else: + threshold = _compute_diag_threshold( + q, + kc_stats, + tau=tau, + scale=scale, + tokens=tokens, + q_scale=q_scale, + ) + return kc_fp8, vc_storage.permute(0, 3, 1, 2), threshold, kc_out_scale + + def _compute_diag_threshold( q: torch.Tensor, kc: torch.Tensor, @@ -248,6 +394,7 @@ def _compute_diag_threshold( tau: float, scale: float, tokens: int | None = None, + q_scale: torch.Tensor | None = None, ) -> torch.Tensor: batch, padded_tokens, heads, head_dim = q.shape tokens = padded_tokens if tokens is None else int(tokens) @@ -278,6 +425,7 @@ def _compute_diag_threshold( ) _diag_threshold_kernel[(blocks, batch_heads)]( q, + q_scale if q_scale is not None else torch.ones((1,), device=q.device, dtype=torch.float32), kc_mean, kc_var_diag, threshold, @@ -289,6 +437,8 @@ def _compute_diag_threshold( head_dim, BLOCK_SIZE, tau, + FP8=q.dtype == torch.float8_e4m3fn, + TOKEN_SCALES=False, num_warps=4, num_stages=2, ) @@ -302,6 +452,7 @@ def _compute_exact_threshold( tau: float, scale: float, tokens: int | None = None, + q_scale: torch.Tensor | None = None, ) -> torch.Tensor: batch, padded_tokens, heads, head_dim = q.shape tokens = padded_tokens if tokens is None else int(tokens) @@ -326,6 +477,7 @@ def _compute_exact_threshold( ) _pool_query_kernel[(blocks, batch_heads)]( q, + q_scale if q_scale is not None else torch.ones((1,), device=q.device, dtype=torch.float32), q_bar, tokens, padded_tokens, @@ -333,13 +485,13 @@ def _compute_exact_threshold( blocks, head_dim, BLOCK_SIZE, + FP8=q.dtype == torch.float8_e4m3fn, + TOKEN_SCALES=False, num_warps=4, num_stages=1, ) block_m = 64 - _exact_fused_threshold_kernel[ - (triton.cdiv(blocks, block_m), batch_heads) - ]( + _exact_fused_threshold_kernel[(triton.cdiv(blocks, block_m), batch_heads)]( q_bar, kc_mean, kc_second_moment, @@ -365,8 +517,11 @@ def prepare( scale: float, thresh_type: str = "diag", tokens: int | None = None, + q_scale: torch.Tensor | None = None, + k_scale: torch.Tensor | None = None, + v_scale: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - kc, vc = _reduce_kv(k, v, tokens=tokens) + kc, vc = _reduce_kv(k, v, tokens=tokens, k_scale=k_scale, v_scale=v_scale) if thresh_type == "exact": threshold = _compute_exact_threshold( q, @@ -374,6 +529,7 @@ def prepare( tau=tau, scale=scale, tokens=tokens, + q_scale=q_scale, ) else: threshold = _compute_diag_threshold( @@ -382,8 +538,9 @@ def prepare( tau=tau, scale=scale, tokens=tokens, + q_scale=q_scale, ) return kc, vc, threshold -__all__ = ["prepare"] +__all__ = ["prepare", "prepare_sm90_fp8"] diff --git a/telefuser/models/minimax_h3_dit.py b/telefuser/models/minimax_h3_dit.py index 98522b9..fad880d 100644 --- a/telefuser/models/minimax_h3_dit.py +++ b/telefuser/models/minimax_h3_dit.py @@ -14,6 +14,7 @@ import torch import torch.distributed as dist import torch.nn as nn +import torch.nn.functional as F from telefuser.core.base_model import BaseModel from telefuser.core.config import AttentionConfig, AttnImplType, QuantConfig, QuantKernelBackend, QuantType @@ -29,7 +30,8 @@ from telefuser.distributed.ulysses_comm import ulysses_gather_heads_destination_major, ulysses_scatter_heads from telefuser.feature_cache import AdaTaylorCacheCalibrator, NoOpCache from telefuser.ops import RMSNorm, apply_qk_norm_rope_neox, indexed_gate, indexed_scale_shift, silu_and_mul_reuse_input -from telefuser.ops.attention import attention +from telefuser.ops.attention import SparseAttentionState, attention +from telefuser.ops.fp8_attention import quantize_fp8_per_block, quantize_fp8_qkv from telefuser.ops.rotary import apply_rotary_emb_neox from telefuser.utils.logging import logger @@ -446,7 +448,7 @@ def enable_tp(self, group: dist.ProcessGroup, *, rank: int, world_size: int) -> self.tp_group = group @staticmethod - def _sage_live_tokens(sequence_lengths: list[int], total_tokens: int) -> int: + def _live_tokens(sequence_lengths: list[int], total_tokens: int) -> int: if len(sequence_lengths) == 1 and sequence_lengths[0] == total_tokens: return total_tokens if ( @@ -457,7 +459,42 @@ def _sage_live_tokens(sequence_lengths: list[int], total_tokens: int) -> int: and total_tokens % 64 == 0 ): return sequence_lengths[0] - raise ValueError("MiniMax H3 SageAttention requires one live sequence with optional trailing alignment padding") + raise ValueError("MiniMax H3 optimized attention requires one live sequence with optional trailing padding") + + @staticmethod + def _is_sol_active(sparse_state: SparseAttentionState | None) -> bool: + return sparse_state is not None and not sparse_state.should_use_dense() + + @classmethod + def _prepare_sol_qkv( + cls, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + sparse_state: SparseAttentionState, + ) -> tuple[ + torch.Tensor, + torch.Tensor, + torch.Tensor, + tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None, + ]: + config = sparse_state.config + layer_end = config.sol_fp8_layer_end + fp8_layer_active = ( + cls._is_sol_active(sparse_state) + and config.sol_fp8 + and sparse_state.layer_idx >= config.sol_fp8_layer_start + and (layer_end is None or sparse_state.layer_idx < layer_end) + ) + if not fp8_layer_active: + return query, key, value, None + if query.is_cuda and torch.cuda.get_device_capability(query.device) == (9, 0): + query, key, value, q_scale, k_scale, v_scale = quantize_fp8_qkv(query, key, value) + else: + query, q_scale = quantize_fp8_per_block(query) + key, k_scale = quantize_fp8_per_block(key) + value, v_scale = quantize_fp8_per_block(value) + return query, key, value, (q_scale, k_scale, v_scale) def forward( self, @@ -467,6 +504,8 @@ def forward( rope_cos_sin_cache: torch.Tensor | None, attention_config: AttentionConfig | None, cu_seqlens: torch.Tensor | None = None, + sparse_state: SparseAttentionState | None = None, + prefix_tokens: int = 0, ) -> torch.Tensor: sequence, _ = hidden.shape qkv = self.qkv_proj(hidden).reshape(sequence, 3, self.num_heads, self.head_dim) @@ -500,16 +539,53 @@ def forward( query = query_wait() key = key_wait() value = value_wait() - if attention_config is not None and attention_config.attn_impl == AttnImplType.SAGE_ATTN_2_8_8_SM90: + optimized_impls = {AttnImplType.SAGE_ATTN_2_8_8_SM90, AttnImplType.SOL_ATTN} + if attention_config is not None and attention_config.attn_impl in optimized_impls: total_tokens = query.shape[1] - live_tokens = self._sage_live_tokens(sequence_lengths, total_tokens) + live_tokens = self._live_tokens(sequence_lengths, total_tokens) + live_query = query[:, :live_tokens].contiguous() + live_key = key[:, :live_tokens].contiguous() + live_value = value[:, :live_tokens].contiguous() + scales = None + runtime_attention_config = attention_config + runtime_sparse_state = sparse_state + if attention_config.attn_impl == AttnImplType.SOL_ATTN: + if sparse_state is None: + raise RuntimeError("MiniMax H3 Sol-Attn requires sparse runtime state") + if not 0 <= prefix_tokens <= live_tokens: + raise ValueError("MiniMax H3 Sol-Attn prefix must be within the live packed sequence") + sol_query, sol_key, sol_value, scales = self._prepare_sol_qkv( + live_query, + live_key, + live_value, + sparse_state, + ) + if sparse_state.should_use_dense(): + runtime_attention_config = AttentionConfig.dense_attention(AttnImplType.FLASH_ATTN_4) + runtime_sparse_state = None + else: + sol_query, sol_key, sol_value = live_query, live_key, live_value live_output = attention( - query[:, :live_tokens].contiguous(), - key[:, :live_tokens].contiguous(), - value[:, :live_tokens].contiguous(), - attention_config=attention_config, + sol_query, + sol_key, + sol_value, + attention_config=runtime_attention_config, + sparse_state=runtime_sparse_state, scale=self.head_dim**-0.5, + q_scale=None if scales is None else scales[0], + k_scale=None if scales is None else scales[1], + v_scale=None if scales is None else scales[2], + sink_start=0, + sink_tokens=prefix_tokens, ) + if self._is_sol_active(sparse_state) and prefix_tokens: + dense_prefix = F.scaled_dot_product_attention( + live_query[:, :prefix_tokens].transpose(1, 2), + live_key.transpose(1, 2), + live_value.transpose(1, 2), + scale=self.head_dim**-0.5, + ).transpose(1, 2) + live_output = torch.cat((dense_prefix, live_output[:, prefix_tokens:]), dim=1) if live_tokens == total_tokens: output = live_output else: @@ -683,6 +759,8 @@ def forward( rope_cos_sin_cache: torch.Tensor, attention_config: AttentionConfig | None, cu_seqlens: torch.Tensor | None = None, + sparse_state: SparseAttentionState | None = None, + prefix_tokens: int = 0, adaln_params: tuple[torch.Tensor, ...] | None = None, ) -> torch.Tensor: if adaln_params is None: @@ -696,6 +774,8 @@ def forward( rope_cos_sin_cache=rope_cos_sin_cache, attention_config=attention_config, cu_seqlens=cu_seqlens, + sparse_state=sparse_state, + prefix_tokens=prefix_tokens, ) hidden = indexed_gate(residual, gate_msa, value, combined_indices) residual = hidden @@ -759,6 +839,25 @@ def __init__(self, config: MiniMaxH3DiTConfig | None = None) -> None: self._online_adaln_rows: dict[str, tuple[float, tuple[torch.Tensor, ...], torch.Tensor]] = {} self._online_adaln_batches: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = [] self._online_adaln_copy_device: torch.device | None = None + self.sparse_attention_state: SparseAttentionState | None = None + + def set_attention_config(self, attention_config: AttentionConfig) -> None: + super().set_attention_config(attention_config) + if attention_config.attn_impl == AttnImplType.SOL_ATTN: + if attention_config.sparse_config is None: + raise ValueError("MiniMax H3 Sol-Attn requires sparse attention configuration") + self.sparse_attention_state = SparseAttentionState( + config=attention_config.sparse_config, + mask_map=None, + model_type="minimax_h3", + ) + else: + self.sparse_attention_state = None + + def _token_refiner_attention_config(self) -> AttentionConfig: + if self.attention_config.is_sparse(): + return AttentionConfig.dense_attention(AttnImplType.FLASH_ATTN_4) + return self.attention_config def adaln_fingerprint(self) -> str: if self.time_embedder is None: @@ -977,7 +1076,7 @@ def _static_inputs( prompt = kwargs["prompt_embeds"].to(device=device, dtype=torch.bfloat16) prompt = self.condition_proj(prompt[: text_positions.numel()]) - prompt = self.token_refiner(prompt, attention_config=self.attention_config) + prompt = self.token_refiner(prompt, attention_config=self._token_refiner_attention_config()) rope_position_ids = kwargs["img_position_ids"].to(device) rope_position_ids = rope_position_ids[:, rope_row_start:rope_row_stop] rope_frequencies = self.rope(rope_position_ids) @@ -1026,6 +1125,13 @@ def forward(self, **kwargs: Any) -> tuple[torch.Tensor, torch.Tensor]: output_positions = self._position_ids( kwargs["img_pos_for_infer_output_info"], "img_pos_for_infer_output_info" ).to(device) + sparse_state = self.sparse_attention_state + prefix_tokens = 0 + if self.attention_config.attn_impl == AttnImplType.SOL_ATTN: + if sparse_state is None: + raise RuntimeError("MiniMax H3 Sol-Attn was not initialized through set_attention_config") + sparse_state.update(numeral_timestep=int(kwargs.get("sparse_step_index", 0))) + prefix_tokens = int(kwargs.get("sol_prefix_tokens", output_positions.min().item())) local_embedding_layout = kwargs.get("local_embedding_layout") use_local_embedding = self.usp_flag and local_embedding_layout is not None @@ -1146,6 +1252,8 @@ def layout_tensor(name: str) -> torch.Tensor: block.adaln_proj.split_output(output) for block, output in zip(self.blocks, gathered_adaln) ) for index, block in enumerate(self.blocks): + if sparse_state is not None: + sparse_state.update(layer_idx=index) hidden = block( hidden, adaln_input=adaln_input, @@ -1154,6 +1262,8 @@ def layout_tensor(name: str) -> torch.Tensor: rope_cos_sin_cache=rope_cos_sin_cache, attention_config=self.attention_config, cu_seqlens=cu_seqlens, + sparse_state=sparse_state, + prefix_tokens=prefix_tokens, adaln_params=None if block_adaln_params is None else block_adaln_params[index], ) if isinstance(feature_cache, AdaTaylorCacheCalibrator): diff --git a/telefuser/models/wan_video_dit.py b/telefuser/models/wan_video_dit.py index c38419d..9748408 100755 --- a/telefuser/models/wan_video_dit.py +++ b/telefuser/models/wan_video_dit.py @@ -37,6 +37,8 @@ from telefuser.offload.async_offload import AsyncOffloadManager from telefuser.ops.attention import MaskMap, SparseAttentionState from telefuser.ops.attention import attention as attn_func +from telefuser.ops.fp8_attention import quantize_fp8_per_block, quantize_fp8_qkv +from telefuser.ops.fp8_gemm import FP8Linear, fp8_linear_forward_many from telefuser.ops.normalization import LayerNorm, RMSNorm, fused_scale_shift, modulate from telefuser.ops.rotary import apply_rotary_emb from telefuser.utils.logging import logger @@ -121,6 +123,56 @@ def _resolve_attention_config(self, sparse_state: SparseAttentionState | None) - return AttentionConfig.dense_attention(AttnImplType.FLASH_ATTN_2) return self.attention_config + @staticmethod + def _is_sol_active(sparse_state: SparseAttentionState | None) -> bool: + return ( + sparse_state is not None + and sparse_state.config.sparse_impl == "sol" + and not sparse_state.should_use_dense() + ) + + def _prepare_sol_projection_input( + self, + x: torch.Tensor, + sparse_state: SparseAttentionState | None, + ) -> torch.Tensor: + # Cast once before q/k/v instead of inside all three FP8 projections. + # Keeping the outer projection dtype in sync also avoids a redundant + # BF16 -> FP32 -> BF16 round trip before the output FP8 Linear. + shared_fp8_qkv = all(isinstance(projection, FP8Linear) for projection in (self.q, self.k, self.v)) + if ( + (self._is_sol_active(sparse_state) or shared_fp8_qkv) + and x.dtype != torch.bfloat16 + and torch.is_autocast_enabled(x.device.type) + ): + return x.to(torch.bfloat16) + return x + + def _prepare_sol_qkv( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + sparse_state: SparseAttentionState | None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None]: + fp8_layer_active = False + if self._is_sol_active(sparse_state) and sparse_state is not None and sparse_state.config.sol_fp8: + layer_end = sparse_state.config.sol_fp8_layer_end + fp8_layer_active = sparse_state.layer_idx >= sparse_state.config.sol_fp8_layer_start and ( + layer_end is None or sparse_state.layer_idx < layer_end + ) + if fp8_layer_active: + if q.is_cuda and torch.cuda.get_device_capability(q.device) == (9, 0): + q, k, v, q_scale, k_scale, v_scale = quantize_fp8_qkv(q, k, v) + return q, k, v, (q_scale, k_scale, v_scale) + q, q_scale = quantize_fp8_per_block(q) + k, k_scale = quantize_fp8_per_block(k) + v, v_scale = quantize_fp8_per_block(v) + return q, k, v, (q_scale, k_scale, v_scale) + if self._is_sol_active(sparse_state) and q.dtype != torch.bfloat16: + return q.to(torch.bfloat16), k.to(torch.bfloat16), v.to(torch.bfloat16), None + return q, k, v, None + def async_usp_forward( self, x: torch.Tensor, @@ -185,14 +237,23 @@ def default_forward( sparse_state: SparseAttentionState | None = None, device_mesh: DeviceMesh | None = None, ) -> torch.Tensor: - q = self.norm_q(self.q(x)) - k = self.norm_k(self.k(x)) - v = self.v(x) + input_dtype = x.dtype + x = self._prepare_sol_projection_input(x, sparse_state) + projection_dtype = x.dtype + if all(isinstance(projection, FP8Linear) for projection in (self.q, self.k, self.v)): + q, k, v = fp8_linear_forward_many((self.q, self.k, self.v), x) + q = self.norm_q(q) + k = self.norm_k(k) + else: + q = self.norm_q(self.q(x)) + k = self.norm_k(self.k(x)) + v = self.v(x) q = rope_apply(q, freqs_cos, freqs_sin, self.num_heads) k = rope_apply(k, freqs_cos, freqs_sin, self.num_heads) q = rearrange(q, "b s (n d) -> b s n d", n=self.num_heads) k = rearrange(k, "b s (n d) -> b s n d", n=self.num_heads) v = rearrange(v, "b s (n d) -> b s n d", n=self.num_heads) + q, k, v, scales = self._prepare_sol_qkv(q, k, v, sparse_state) if sparse_state is not None and sparse_state.config.sparse_impl == "radial": seqlen = q.shape[2] q = rearrange(q, "b s n d -> (b s) n d", s=seqlen, n=self.num_heads) @@ -206,8 +267,16 @@ def default_forward( sparse_state=sparse_state, input_layout="BSND", output_layout="BSND", + q_scale=None if scales is None else scales[0], + k_scale=None if scales is None else scales[1], + v_scale=None if scales is None else scales[2], ) x = rearrange(x, "b s n d -> b s (n d)", n=self.num_heads) + if projection_dtype != input_dtype: + x = self.o(x) + return x.to(input_dtype) + if x.dtype != projection_dtype: + x = x.to(input_dtype) return self.o(x) @@ -439,7 +508,7 @@ def reset_y_camera_status(self): def enable_quant(self, quant_type: str | torch.dtype): """Enable quantization for transformer blocks.""" - from telefuser.core.config import QuantConfig, QuantType + from telefuser.core.config import QuantConfig, QuantKernelBackend, QuantType if isinstance(quant_type, QuantConfig): if quant_type.quant_type == QuantType.BNB_NF4: @@ -470,6 +539,38 @@ def enable_quant(self, quant_type: str | torch.dtype): logger.info(f"TorchAO FP8 converted {replaced} Linear layers") self.quant_type = quant_type.quant_type return + if quant_type.quant_type == QuantType.FP8: + if quant_type.kernel_backend not in (QuantKernelBackend.AUTO, QuantKernelBackend.TF_KERNEL): + raise ValueError( + "Wan FP8 online quantization requires the tf-kernel backend; " + f"got {quant_type.kernel_backend.name}" + ) + logger.info("loading weights with tf-kernel FP8, start quantize linear layers") + from telefuser.ops.fp8_gemm import FP8GemmOptions, count_linear_layers, enable_fp8_gemm + + include_names = quant_type.quantize_modules or ("blocks.",) + + def module_filter(name: str, _module: nn.Module) -> bool: + return any(token in name for token in include_names) and not any( + token and token in name for token in quant_type.skip_modules + ) + + replaced = count_linear_layers(self, module_filter=module_filter) + enable_fp8_gemm( + self, + options=FP8GemmOptions( + cast_output_back=False, + fp16_weight_storage="keep" if quant_type.keep_fp16_weight else "discard", + materialize_fp8_on_wrap=True, + ), + module_filter=module_filter, + ) + if replaced == 0: + raise RuntimeError("Wan FP8 online quantization did not select any Linear layers") + self.tf_kernel_fp8_replaced_linear = replaced + self.quant_type = quant_type.quant_type + logger.info(f"Wan tf-kernel FP8 converted {replaced} transformer Linear layers") + return quant_type = torch.float8_e4m3fn if quant_type.quant_type == QuantType.FP8 else quant_type.quant_type if quant_type in [torch.float8_e4m3fn]: diff --git a/telefuser/ops/attention/attention_impl.py b/telefuser/ops/attention/attention_impl.py index ca322c4..f03427e 100755 --- a/telefuser/ops/attention/attention_impl.py +++ b/telefuser/ops/attention/attention_impl.py @@ -166,8 +166,11 @@ def _resolve_sol_kv_splits(q: Tensor, kv_splits: int | str) -> int: """Match the official Sol-Engine automatic split policy.""" if kv_splits != "auto": return int(kv_splits) - if torch.cuda.get_device_capability(q.device) == (9, 0) and q.shape[1] >= 65536: - return 4 + if torch.cuda.get_device_capability(q.device) == (9, 0): + if q.dtype == torch.float8_e4m3fn and q.shape[1] >= 16384: + return 2 + if q.shape[1] >= 65536: + return 4 return 1 @@ -187,6 +190,11 @@ def attention( return_lse: bool = False, sequence_lengths: list[int] | None = None, cu_seqlens: Tensor | None = None, + q_scale: Tensor | None = None, + k_scale: Tensor | None = None, + v_scale: Tensor | None = None, + sink_start: int | None = None, + sink_tokens: int = 0, **kwargs: Any, ) -> Tensor | tuple[Tensor, Tensor]: """Unified attention function. @@ -206,6 +214,8 @@ def attention( return_lse: Return log-sum-exp values. sequence_lengths: Length of each sequence packed along the sequence axis. cu_seqlens: Optional precomputed cumulative sequence lengths for varlen kernels. + sink_start: Start of the exact KV sink used by Sol-Attn. + sink_tokens: Number of exact KV sink tokens used by Sol-Attn. **kwargs: Implementation-specific arguments. Returns: @@ -237,6 +247,11 @@ def attention( elif attn_impl == AttnImplType.SOL_ATTN: if sparse_state.should_use_dense(): attn_impl = AttnImplType.FLASH_ATTN_2 + elif sparse_state.config.sol_fp8 and sparse_state.config.sol_tau < 0.0 and q.dtype == torch.bfloat16: + # FP8 Dense only needs the CuTe exact mainloop in quantized + # layers. Unquantized layers use the faster dense backend and + # avoid compiling a second BF16 CuTe specialization. + attn_impl = AttnImplType.TORCH_SDPA else: if sparse_state.mask_map is None: raise RuntimeError("Radial attention requires a mask map") @@ -403,6 +418,9 @@ def attention( # Sol-Attn elif attn_impl == AttnImplType.SOL_ATTN and SOL_ATTN_AVAILABLE and sol_attn is not None: + sparse_config = attention_config.sparse_config + if sparse_config is None: + raise RuntimeError("Sol-Attn requires sparse attention configuration") eligible = ( attn_mask is None and not is_causal @@ -411,28 +429,58 @@ def attention( and q.shape == k.shape == v.shape and q.ndim == 4 and q.shape[-1] == 128 - and q.dtype == torch.bfloat16 + and (q.dtype == torch.bfloat16 or (sparse_config.sol_fp8 and q.dtype == torch.float8_e4m3fn)) and q.is_cuda ) if eligible: - sparse_config = attention_config.sparse_config - if sparse_config is None: - raise RuntimeError("Sol-Attn requires sparse attention configuration") + if q.dtype == torch.float8_e4m3fn and any(scale is None for scale in (q_scale, k_scale, v_scale)): + raise ValueError("FP8 Sol-Attn requires q_scale, k_scale, and v_scale") try: output = sol_attn( q.contiguous(), k.contiguous(), - v.contiguous(), + v if q.dtype == torch.float8_e4m3fn else v.contiguous(), scale=scale, tau=sparse_config.sol_tau, thresh_type=sparse_config.sol_threshold_type, kv_splits=_resolve_sol_kv_splits(q, sparse_config.sol_kv_splits), + q_scale=q_scale, + k_scale=k_scale, + v_scale=v_scale, + sink_start=sink_start, + sink_tokens=sink_tokens, + # A partial FP8 layer range otherwise compiles both BF16 + # and FP8 CuTe specializations on the first sparse step. + # Triton is a better cold-start tradeoff for the remaining + # sparse BF16 layers; exact FP8 Dense keeps CuTe throughout. + force_triton=(sparse_config.sol_fp8 and q.dtype == torch.bfloat16 and sparse_config.sol_tau >= 0.0), ) except (RuntimeError, TypeError, ValueError) as error: msg = "Sol-Attn execution failed, falling back to TORCH_SDPA" if msg not in _warned_attn_fallback: _warned_attn_fallback.add(msg) logger.warning("%s: %s", msg, error) + if q.dtype == torch.float8_e4m3fn: + from telefuser.ops.fp8_attention import ( + dequantize_fp8_per_block, + dequantize_fp8_per_channel, + dequantize_fp8_per_token, + ) + + if ( + q_scale.shape[0] == q.shape[0] + and q_scale.shape[1] >= q.shape[1] + and q_scale.shape[2] == q.shape[2] + ): + q = dequantize_fp8_per_token(q, q_scale, torch.bfloat16) + k = dequantize_fp8_per_token(k, k_scale, torch.bfloat16) + else: + q = dequantize_fp8_per_block(q, q_scale, torch.bfloat16) + k = dequantize_fp8_per_block(k, k_scale, torch.bfloat16) + if v_scale.shape == (v.shape[0], v.shape[2], v.shape[3]): + v = dequantize_fp8_per_channel(v, v_scale, torch.bfloat16) + else: + v = dequantize_fp8_per_block(v, v_scale, torch.bfloat16) # Fallback to SDPA if output is None: diff --git a/telefuser/ops/fp8_attention.py b/telefuser/ops/fp8_attention.py new file mode 100644 index 0000000..7a2f36d --- /dev/null +++ b/telefuser/ops/fp8_attention.py @@ -0,0 +1,216 @@ +"""Block-scaled FP8 activation helpers for attention boundaries.""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F +import triton +import triton.language as tl + +FP8_ATTENTION_BLOCK_SIZE = 64 + + +@triton.jit +def _quantize_qkv_fp8_stage1( + q, + k, + v, + q_out, + k_out, + q_scale, + k_scale, + v_scale, + tokens: tl.constexpr, + heads: tl.constexpr, + head_dim: tl.constexpr, + block: tl.constexpr, +): + block_idx = tl.program_id(0) + batch_head = tl.program_id(1) + batch = batch_head // heads + head = batch_head % heads + token_offsets = block_idx * block + tl.arange(0, block) + dim_offsets = tl.arange(0, head_dim) + valid = token_offsets < tokens + offsets = ((batch * tokens + token_offsets[:, None]) * heads + head) * head_dim + dim_offsets[None, :] + q_values = tl.load(q + offsets, mask=valid[:, None], other=0.0).to(tl.float32) + k_values = tl.load(k + offsets, mask=valid[:, None], other=0.0).to(tl.float32) + v_values = tl.load(v + offsets, mask=valid[:, None], other=0.0).to(tl.float32) + + q_s = tl.maximum(tl.max(tl.max(tl.abs(q_values), axis=1), axis=0), 1.0e-6) / 448.0 + k_s = tl.maximum(tl.max(tl.max(tl.abs(k_values), axis=1), axis=0), 1.0e-6) / 448.0 + scale_offset = (batch * tl.cdiv(tokens, block) + block_idx) * heads + head + tl.store(q_scale + scale_offset, q_s) + tl.store(k_scale + scale_offset, k_s) + tl.store(q_out + offsets, q_values / q_s, mask=valid[:, None]) + tl.store(k_out + offsets, k_values / k_s, mask=valid[:, None]) + + v_s = tl.max(tl.abs(v_values), axis=0) / 448.0 + v_scale_offsets = (batch * heads + head) * head_dim + dim_offsets + tl.atomic_max(v_scale + v_scale_offsets, v_s) + + +@triton.jit +def _quantize_qkv_fp8_stage2_v( + v, + v_out, + v_scale, + tokens: tl.constexpr, + heads: tl.constexpr, + head_dim: tl.constexpr, + block: tl.constexpr, +): + block_idx = tl.program_id(0) + batch_head = tl.program_id(1) + batch = batch_head // heads + head = batch_head % heads + token_offsets = block_idx * block + tl.arange(0, block) + dim_offsets = tl.arange(0, head_dim) + valid = token_offsets < tokens + input_offsets = ((batch * tokens + token_offsets[:, None]) * heads + head) * head_dim + dim_offsets[None, :] + output_offsets = ((batch * heads + head) * head_dim + dim_offsets[None, :]) * tokens + token_offsets[:, None] + scale_offsets = (batch * heads + head) * head_dim + dim_offsets + scale = tl.maximum(tl.load(v_scale + scale_offsets), 1.0e-6 / 448.0) + values = tl.load(v + input_offsets, mask=valid[:, None], other=0.0).to(tl.float32) + tl.store(v_out + output_offsets, values / scale[None, :], mask=valid[:, None]) + + +def quantize_fp8_qkv( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Fused SM90 block-scaled Q/K and layout-aware V-channel E4M3 quantization.""" + + if q.shape != k.shape or q.shape != v.shape or q.ndim != 4: + raise ValueError("q, k, and v must share shape [B, T, H, D]") + if not (q.is_cuda and q.is_contiguous() and k.is_contiguous() and v.is_contiguous()): + raise ValueError("fused FP8 QKV quantization requires contiguous CUDA tensors") + batch, tokens, heads, head_dim = q.shape + if head_dim != 128: + raise ValueError("fused FP8 QKV quantization requires head dimension 128") + blocks = triton.cdiv(tokens, FP8_ATTENTION_BLOCK_SIZE) + q_out = torch.empty(q.shape, device=q.device, dtype=torch.float8_e4m3fn) + k_out = torch.empty_like(q_out) + v_storage = torch.empty((batch, heads, head_dim, tokens), device=q.device, dtype=torch.float8_e4m3fn) + q_scale = torch.empty((batch, blocks, heads), device=q.device, dtype=torch.float32) + k_scale = torch.ones_like(q_scale) + v_scale = torch.zeros((batch, heads, head_dim), device=q.device, dtype=torch.float32) + grid = (blocks, batch * heads) + _quantize_qkv_fp8_stage1[grid]( + q, + k, + v, + q_out, + k_out, + q_scale, + k_scale, + v_scale, + tokens, + heads, + head_dim, + FP8_ATTENTION_BLOCK_SIZE, + num_warps=8, + num_stages=1, + ) + _quantize_qkv_fp8_stage2_v[grid]( + v, + v_storage, + v_scale, + tokens, + heads, + head_dim, + FP8_ATTENTION_BLOCK_SIZE, + num_warps=8, + num_stages=1, + ) + v_out = v_storage.permute(0, 3, 1, 2) + return q_out, k_out, v_out, q_scale, k_scale, v_scale + + +def quantize_fp8_per_block( + x: torch.Tensor, + block_size: int = FP8_ATTENTION_BLOCK_SIZE, +) -> tuple[torch.Tensor, torch.Tensor]: + """Quantize a [B, T, H, D] tensor with one E4M3 scale per block/head.""" + if x.ndim != 4 or not x.is_floating_point(): + raise ValueError("FP8 attention quantization expects a floating-point [B, T, H, D] tensor") + batch, tokens, heads, head_dim = x.shape + blocks = (tokens + block_size - 1) // block_size + padded_tokens = blocks * block_size + padded = F.pad(x, (0, 0, 0, 0, 0, padded_tokens - tokens)) + blocked = padded.reshape(batch, blocks, block_size, heads, head_dim) + scale = blocked.detach().abs().amax(dim=(2, 4)).float().clamp_min(1e-6) / 448.0 + quantized = (blocked / scale.to(x.dtype)[:, :, None, :, None]).to(torch.float8_e4m3fn) + return quantized.reshape(batch, padded_tokens, heads, head_dim)[:, :tokens].contiguous(), scale + + +def dequantize_fp8_per_block( + x: torch.Tensor, + scale: torch.Tensor, + dtype: torch.dtype, + block_size: int = FP8_ATTENTION_BLOCK_SIZE, +) -> torch.Tensor: + """Restore block-scaled FP8 activations to ``dtype``.""" + if x.dtype != torch.float8_e4m3fn: + raise TypeError("expected torch.float8_e4m3fn activations") + token_scale = scale.repeat_interleave(block_size, dim=1)[:, : x.shape[1]] + return x.to(dtype) * token_scale.to(dtype).unsqueeze(-1) + + +def dequantize_fp8_per_token( + x: torch.Tensor, + scale: torch.Tensor, + dtype: torch.dtype, +) -> torch.Tensor: + """Restore token-scaled FP8 [B, T, H, D] activations.""" + if x.dtype != torch.float8_e4m3fn: + raise TypeError("expected torch.float8_e4m3fn activations") + if scale.shape[0] != x.shape[0] or scale.shape[1] < x.shape[1] or scale.shape[2] != x.shape[2]: + raise ValueError("scale must have shape [B, padded_T, H] with padded_T >= T") + return x.to(dtype) * scale[:, : x.shape[1]].to(dtype).unsqueeze(-1) + + +def quantize_fp8_per_channel( + x: torch.Tensor, + *, + token_contiguous: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: + """Quantize BTHD with one E4M3 scale per head/channel. + + ``token_contiguous`` stores the same BTHD view over B,H,D,T-contiguous + backing memory, matching the SM90 K-major PV WGMMA operand. + """ + if x.ndim != 4 or not x.is_floating_point(): + raise ValueError("FP8 attention quantization expects a floating-point [B, T, H, D] tensor") + scale = x.detach().abs().amax(dim=1).float().clamp_min(1e-6) / 448.0 + quantized = (x / scale.to(x.dtype).unsqueeze(1)).to(torch.float8_e4m3fn) + if token_contiguous: + quantized = quantized.permute(0, 2, 3, 1).contiguous().permute(0, 3, 1, 2) + else: + quantized = quantized.contiguous() + return quantized, scale.contiguous() + + +def dequantize_fp8_per_channel( + x: torch.Tensor, + scale: torch.Tensor, + dtype: torch.dtype, +) -> torch.Tensor: + """Restore channel-scaled FP8 [B, T, H, D] activations.""" + if x.dtype != torch.float8_e4m3fn: + raise TypeError("expected torch.float8_e4m3fn activations") + if scale.shape != (x.shape[0], x.shape[2], x.shape[3]): + raise ValueError("scale must have shape [B, H, D]") + return x.to(dtype) * scale.to(dtype).unsqueeze(1) + + +__all__ = [ + "FP8_ATTENTION_BLOCK_SIZE", + "dequantize_fp8_per_block", + "dequantize_fp8_per_channel", + "dequantize_fp8_per_token", + "quantize_fp8_qkv", + "quantize_fp8_per_block", + "quantize_fp8_per_channel", +] diff --git a/telefuser/ops/fp8_gemm.py b/telefuser/ops/fp8_gemm.py index 0cd04e3..5603161 100644 --- a/telefuser/ops/fp8_gemm.py +++ b/telefuser/ops/fp8_gemm.py @@ -239,45 +239,56 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: "Use fp16_weight_storage='cpu_offload' (or 'keep') for CPU fallback." ) - # tf-kernel FP8 GEMM only supports fp16/bf16 outputs. + if x.dtype not in (torch.float16, torch.bfloat16) and not self.options.cast_inputs: + if self.linear is not None: + return self.linear(x) + if self._fp16_weight_cpu is not None: + weight = self._fp16_weight_cpu.to(device=x.device, dtype=x.dtype) + bias = self._fp16_bias_cpu + bias = bias.to(device=x.device, dtype=x.dtype) if bias is not None else None + return torch.nn.functional.linear(x, weight, bias) + raise RuntimeError("cast_inputs=False requires FP16 weights for fallback, but they were discarded.") + + x_fp, in_dtype, out_dtype = self._prepare_cuda_input(x) + x_shape = x_fp.shape + x_2d = x_fp.reshape(-1, x_shape[-1]).contiguous() + qinput = torch.empty_like(x_2d, dtype=torch.float8_e4m3fn) + input_scale = torch.empty((x_2d.shape[0], 1), dtype=torch.float32, device=x_fp.device) + self._tf_kernel.tf_per_token_quant_fp8(x_2d, qinput, input_scale) + return self._forward_quantized(qinput, input_scale, x_shape, in_dtype, out_dtype) + + def _prepare_cuda_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.dtype, torch.dtype]: + """Cast an activation to a dtype supported by tf-kernel FP8 GEMM.""" in_dtype = x.dtype if in_dtype not in (torch.float16, torch.bfloat16): - if not self.options.cast_inputs: - # Fall back if we still have FP16 weights. - if self.linear is not None: - return self.linear(x) - if self._fp16_weight_cpu is not None: - w = self._fp16_weight_cpu.to(device=x.device, dtype=in_dtype) - b = self._fp16_bias_cpu - b = b.to(device=x.device, dtype=in_dtype) if b is not None else None - return torch.nn.functional.linear(x, w, b) - raise RuntimeError("cast_inputs=False requires FP16 weights for fallback, but they were discarded.") - # import nvtx - # nvtx.push_range(f"cast_input") x_fp = x.to(torch.bfloat16) - # nvtx.pop_range() out_dtype = torch.bfloat16 else: x_fp = x out_dtype = in_dtype - - self._maybe_requantize_weight(x_fp.device) + return x_fp, in_dtype, out_dtype + + def _forward_quantized( + self, + qinput: torch.Tensor, + input_scale: torch.Tensor, + x_shape: torch.Size, + in_dtype: torch.dtype, + out_dtype: torch.dtype, + ) -> torch.Tensor: + """Run this Linear using an already quantized shared activation.""" + self._maybe_requantize_weight(qinput.device) if self.linear is not None: bias = self.linear.bias else: bias = self.bias if bias is not None: - if bias.device != x_fp.device: - bias = bias.to(device=x_fp.device, non_blocking=True) + if bias.device != qinput.device: + bias = bias.to(device=qinput.device, non_blocking=True) if bias.dtype != out_dtype: bias = bias.to(dtype=out_dtype) - x_shape = x_fp.shape - x_2d = x_fp.reshape(-1, x_shape[-1]).contiguous() - qinput = torch.empty_like(x_2d, dtype=torch.float8_e4m3fn) - input_scale = torch.empty((x_2d.shape[0], 1), dtype=torch.float32, device=x_fp.device) - self._tf_kernel.tf_per_token_quant_fp8(x_2d, qinput, input_scale) y = self._tf_kernel.fp8_scaled_mm( qinput, self._fp8_weight, @@ -293,6 +304,24 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: return y +def fp8_linear_forward_many(linears: tuple[FP8Linear, ...], x: torch.Tensor) -> tuple[torch.Tensor, ...]: + """Reuse one dynamic activation quantization across compatible FP8 Linears.""" + if not linears: + return () + first = linears[0] + unsupported_no_cast = x.dtype not in (torch.float16, torch.bfloat16) and not first.options.cast_inputs + if not x.is_cuda or unsupported_no_cast or any(linear.options != first.options for linear in linears[1:]): + return tuple(linear(x) for linear in linears) + + x_fp, in_dtype, out_dtype = first._prepare_cuda_input(x) + x_shape = x_fp.shape + x_2d = x_fp.reshape(-1, x_shape[-1]).contiguous() + qinput = torch.empty_like(x_2d, dtype=torch.float8_e4m3fn) + input_scale = torch.empty((x_2d.shape[0], 1), dtype=torch.float32, device=x_fp.device) + first._tf_kernel.tf_per_token_quant_fp8(x_2d, qinput, input_scale) + return tuple(linear._forward_quantized(qinput, input_scale, x_shape, in_dtype, out_dtype) for linear in linears) + + def enable_fp8_gemm( model: nn.Module, *, diff --git a/telefuser/pipelines/minimax_h3/denoising.py b/telefuser/pipelines/minimax_h3/denoising.py index 89d02ee..e2fa27d 100644 --- a/telefuser/pipelines/minimax_h3/denoising.py +++ b/telefuser/pipelines/minimax_h3/denoising.py @@ -126,6 +126,9 @@ def __init__( self.transformer = module_manager.fetch_module("minimax_h3_transformer") if self.transformer is None: raise ValueError("ModuleManager must contain 'minimax_h3_transformer'") + set_attention_config = getattr(self.transformer, "set_attention_config", None) + if callable(set_attention_config): + set_attention_config(model_runtime_config.attention_config) if model_runtime_config.lora_configs: MiniMaxH3LoraAdapter.apply(self.transformer, model_runtime_config.lora_configs) step_update = "training_euler" if model_runtime_config.lora_configs else "reference_blend" @@ -159,7 +162,6 @@ def parallel_models(self) -> None: raise NotImplementedError(f"MiniMax H3 does not support these parallel degrees yet: {invalid}") device_mesh = create_device_mesh_from_config(parallel_config) self.transformer.device_mesh = device_mesh - self.transformer.set_attention_config(self.model_runtime_config.attention_config) if parallel_config.tp_degree > 1: if parallel_config.enable_fsdp: raise ValueError("MiniMax H3 DiT tensor parallelism cannot be combined with FSDP") @@ -374,6 +376,7 @@ def denoise( text_pos_cpu = packed["text_pos"] text_pos = text_pos_cpu.to(device) target_img_pos = img_pos[video_update] + sol_prefix_tokens = int(img_pos_cpu[video_update_cpu][0]) target_video_row_start = int((~video_update_cpu).sum()) target_audio_row_start = int((~audio_update_cpu).sum()) condition_img_pos = img_pos_cpu[~video_update_cpu] @@ -496,6 +499,8 @@ def denoise( block_combined_indices=block_combined_indices, local_embedding_layout=local_embedding_layout, static_cache_key=static_cache_key, + sparse_step_index=step, + sol_prefix_tokens=sol_prefix_tokens, skip_mask_out_condition=True, ) audio_target_velocity = audio_velocity[audio_target_slice] diff --git a/tests/unit/models/test_minimax_h3_dit.py b/tests/unit/models/test_minimax_h3_dit.py index 35147de..8bc228a 100644 --- a/tests/unit/models/test_minimax_h3_dit.py +++ b/tests/unit/models/test_minimax_h3_dit.py @@ -14,6 +14,7 @@ MiniMaxH3DiTConfig, _reorder_grouped_qkv_to_qkv, ) +from telefuser.ops.attention import SparseAttentionState from telefuser.ops.rotary import apply_qk_norm_rope_neox, apply_rotary_emb_neox @@ -159,6 +160,106 @@ def sage_output(query: torch.Tensor, *_: torch.Tensor, **__: object) -> torch.Te assert torch.count_nonzero(output[61:]) == 0 +def test_sol_attention_preserves_prefix_sink_and_dense_prefix_queries() -> None: + module = MiniMaxH3Attention(_small_config()).eval() + hidden = torch.randn(64, 32, dtype=torch.bfloat16) + config = AttentionConfig.sol_attention(dense_timesteps=0, dense_layers=0, threshold_type="exact") + state = SparseAttentionState(config.sparse_config, mask_map=None, model_type="minimax_h3") + + with ( + patch("telefuser.models.minimax_h3_dit.attention", side_effect=lambda query, *_args, **_kwargs: query) as sol, + patch( + "telefuser.models.minimax_h3_dit.F.scaled_dot_product_attention", + side_effect=lambda query, *_args, **_kwargs: query, + ) as dense_prefix, + ): + output = module( + hidden, + sequence_lengths=[61, 3], + rope_cos_sin_cache=None, + attention_config=config, + sparse_state=state, + prefix_tokens=13, + ) + + assert sol.call_args.args[0].shape == (1, 61, 4, 8) + assert sol.call_args.kwargs["sparse_state"] is state + assert sol.call_args.kwargs["sink_start"] == 0 + assert sol.call_args.kwargs["sink_tokens"] == 13 + assert dense_prefix.call_args.args[0].shape == (1, 4, 13, 8) + assert dense_prefix.call_args.args[1].shape == (1, 4, 61, 8) + assert torch.count_nonzero(output[61:]) == 0 + + +def test_sol_dense_guard_uses_packed_flash_attention_4() -> None: + module = MiniMaxH3Attention(_small_config()).eval() + hidden = torch.randn(64, 32, dtype=torch.bfloat16) + config = AttentionConfig.sol_attention(dense_timesteps=10, dense_layers=2, threshold_type="exact") + state = SparseAttentionState(config.sparse_config, mask_map=None, model_type="minimax_h3") + + with patch("telefuser.models.minimax_h3_dit.attention", side_effect=lambda query, *_args, **_kwargs: query) as call: + module( + hidden, + sequence_lengths=[61, 3], + rope_cos_sin_cache=None, + attention_config=config, + sparse_state=state, + prefix_tokens=13, + ) + + runtime_config = call.call_args.kwargs["attention_config"] + assert runtime_config.attn_impl == AttnImplType.FLASH_ATTN_4 + assert call.call_args.kwargs["sparse_state"] is None + + +def test_sol_fp8_passes_quantized_qkv_scales_to_attention() -> None: + module = MiniMaxH3Attention(_small_config()).eval() + hidden = torch.randn(64, 32, dtype=torch.bfloat16) + config = AttentionConfig.sol_attention( + dense_timesteps=0, + dense_layers=0, + threshold_type="exact", + sol_fp8=True, + ) + state = SparseAttentionState(config.sparse_config, mask_map=None, model_type="minimax_h3") + scales = [torch.ones(1, 1, 4) * value for value in (1, 2, 3)] + + def quantize(value: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + return value.to(torch.float8_e4m3fn), scales.pop(0) + + with ( + patch("telefuser.models.minimax_h3_dit.quantize_fp8_per_block", side_effect=quantize), + patch( + "telefuser.models.minimax_h3_dit.attention", + side_effect=lambda query, *_args, **_kwargs: query.float(), + ) as attention_call, + ): + module( + hidden, + sequence_lengths=[61, 3], + rope_cos_sin_cache=None, + attention_config=config, + sparse_state=state, + ) + + assert attention_call.call_args.args[0].dtype == torch.float8_e4m3fn + assert attention_call.call_args.kwargs["q_scale"].flatten()[0].item() == 1 + assert attention_call.call_args.kwargs["k_scale"].flatten()[0].item() == 2 + assert attention_call.call_args.kwargs["v_scale"].flatten()[0].item() == 3 + + +def test_minimax_h3_initializes_sol_runtime_state() -> None: + model = MiniMaxH3DiT(_small_config()) + config = AttentionConfig.sol_attention(dense_timesteps=10, dense_layers=2, threshold_type="exact") + + model.set_attention_config(config) + + assert model.sparse_attention_state is not None + assert model.sparse_attention_state.config is config.sparse_config + assert model.sparse_attention_state.model_type == "minimax_h3" + assert model._token_refiner_attention_config().attn_impl == AttnImplType.FLASH_ATTN_4 + + def test_ulysses_overlaps_strided_value_scatter_with_qk_preprocessing() -> None: module = MiniMaxH3Attention(_small_config()).eval() module.ulysses_group = MagicMock() diff --git a/tests/unit/models/test_wan_video_sol_attention.py b/tests/unit/models/test_wan_video_sol_attention.py index d0c8afd..08532f6 100644 --- a/tests/unit/models/test_wan_video_sol_attention.py +++ b/tests/unit/models/test_wan_video_sol_attention.py @@ -6,6 +6,13 @@ from telefuser.core.config import AttentionConfig, AttnImplType, SparseAttentionConfig from telefuser.models.wan_video_dit import SelfAttention, WanModel, precompute_freqs_cis_3d from telefuser.ops.attention import SparseAttentionState, attention_impl +from telefuser.ops.fp8_attention import ( + dequantize_fp8_per_block, + dequantize_fp8_per_channel, + quantize_fp8_per_block, + quantize_fp8_qkv, +) +from telefuser.ops.fp8_gemm import FP8Linear def test_wan_model_enables_sol_attention_state() -> None: @@ -99,6 +106,96 @@ def fake_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, **kwargs) assert captured["sparse_state"] is state +def test_wan_self_attention_casts_fp32_qkv_only_for_active_sol() -> None: + module = SelfAttention(dim=128, num_heads=1) + config = SparseAttentionConfig(sparse_impl="sol", dense_timesteps=1, dense_layers=0) + state = SparseAttentionState(config, mask_map=None) + q = torch.randn(1, 4, 1, 128) + + dense_qkv = module._prepare_sol_qkv(q, q, q, state) + assert all(tensor.dtype is torch.float32 for tensor in dense_qkv[:3]) + assert dense_qkv[3] is None + + state.update(numeral_timestep=1) + sol_qkv = module._prepare_sol_qkv(q, q, q, state) + assert all(tensor.dtype is torch.bfloat16 for tensor in sol_qkv[:3]) + assert sol_qkv[3] is None + + +def test_wan_self_attention_casts_projection_input_once_under_autocast() -> None: + module = SelfAttention(dim=128, num_heads=1) + config = SparseAttentionConfig(sparse_impl="sol", dense_timesteps=1, dense_layers=0) + state = SparseAttentionState(config, mask_map=None) + x = torch.randn(1, 4, 128) + + with patch("telefuser.models.wan_video_dit.torch.is_autocast_enabled", return_value=True): + dense_x = module._prepare_sol_projection_input(x, state) + state.update(numeral_timestep=1) + sol_x = module._prepare_sol_projection_input(x, state) + + assert dense_x is x + assert sol_x.dtype is torch.bfloat16 + + +def test_wan_self_attention_casts_shared_fp8_qkv_input_during_dense_warmup() -> None: + module = SelfAttention(dim=128, num_heads=1) + for name in ("q", "k", "v"): + projection = FP8Linear.__new__(FP8Linear) + torch.nn.Module.__init__(projection) + setattr(module, name, projection) + config = SparseAttentionConfig(sparse_impl="sol", dense_timesteps=1, dense_layers=0) + state = SparseAttentionState(config, mask_map=None) + x = torch.randn(1, 4, 128) + + with patch("telefuser.models.wan_video_dit.torch.is_autocast_enabled", return_value=True): + prepared = module._prepare_sol_projection_input(x, state) + + assert state.should_use_dense() + assert prepared.dtype is torch.bfloat16 + + +def test_wan_self_attention_quantizes_qkv_for_fp8_sol() -> None: + module = SelfAttention(dim=128, num_heads=1).to(torch.bfloat16) + config = SparseAttentionConfig(sparse_impl="sol", dense_timesteps=0, sol_fp8=True) + state = SparseAttentionState(config, mask_map=None) + captured = {} + + def fake_attention(q, k, v, **kwargs): + captured.update({"q": q, "k": k, "v": v, **kwargs}) + return q.to(torch.bfloat16) + + x = torch.randn(1, 65, 128, dtype=torch.bfloat16) + freqs = torch.zeros(65, 64, dtype=torch.bfloat16) + with patch("telefuser.models.wan_video_dit.attn_func", side_effect=fake_attention): + module.default_forward(x, freqs, freqs, sparse_state=state) + + assert captured["q"].dtype is torch.float8_e4m3fn + assert captured["q_scale"].shape == (1, 2, 1) + restored = dequantize_fp8_per_block(captured["q"], captured["q_scale"], torch.bfloat16) + assert torch.isfinite(restored).all() + + +def test_wan_self_attention_limits_fp8_sol_to_configured_layers() -> None: + module = SelfAttention(dim=128, num_heads=1) + config = SparseAttentionConfig( + sparse_impl="sol", + dense_timesteps=0, + sol_fp8=True, + sol_fp8_layer_start=1, + sol_fp8_layer_end=2, + ) + state = SparseAttentionState(config, mask_map=None) + q = torch.randn(1, 64, 1, 128) + + bf16_qkv = module._prepare_sol_qkv(q, q, q, state) + assert bf16_qkv[3] is None + + state.update(layer_idx=1) + fp8_qkv = module._prepare_sol_qkv(q, q, q, state) + assert fp8_qkv[0].dtype is torch.float8_e4m3fn + assert fp8_qkv[3] is not None + + @pytest.mark.gpu def test_wan_self_attention_executes_sol_on_h100(monkeypatch: pytest.MonkeyPatch) -> None: if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (9, 0): @@ -130,3 +227,154 @@ def tracked_sol_attn(*args, **kwargs): assert output.shape == x.shape assert torch.isfinite(output).all() assert kernel_calls == 1 + + +@pytest.mark.gpu +def test_wan_self_attention_executes_fp8_sol_on_h100(monkeypatch: pytest.MonkeyPatch) -> None: + if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (9, 0): + pytest.skip("Wan FP8 Sol-Attn execution test requires H100") + + assert attention_impl.SOL_ATTN_AVAILABLE + assert attention_impl.sol_attn is not None + captured = {} + sol_attn = attention_impl.sol_attn + + def tracked_sol_attn(q, k, v, **kwargs): + captured.update({"q": q, "k": k, "v": v, **kwargs}) + return sol_attn(q, k, v, **kwargs) + + monkeypatch.setattr(attention_impl, "sol_attn", tracked_sol_attn) + + module = SelfAttention(dim=128, num_heads=1).eval().cuda().to(torch.bfloat16) + x = torch.randn(1, 256, 128, device="cuda", dtype=torch.bfloat16) + freqs = precompute_freqs_cis_3d(128) + freqs_cos = torch.cat([freq.real for freq in freqs], dim=-1)[:256].cuda() + freqs_sin = torch.cat([freq.imag for freq in freqs], dim=-1)[:256].cuda() + sparse_config = SparseAttentionConfig( + sparse_impl="sol", + dense_timesteps=0, + sol_tau=-1000.0, + sol_fp8=True, + ) + module.attention_config = AttentionConfig(attn_impl=AttnImplType.SOL_ATTN, sparse_config=sparse_config) + state = SparseAttentionState(sparse_config, mask_map=None) + + output = module(x, freqs_cos, freqs_sin, sparse_state=state) + + assert output.shape == x.shape + assert torch.isfinite(output).all() + assert captured["q"].dtype is torch.float8_e4m3fn + assert captured["k"].dtype is torch.float8_e4m3fn + assert captured["v"].dtype is torch.float8_e4m3fn + assert captured["v"].stride(1) == 1 + assert captured["q_scale"].shape == (1, 4, 1) + assert captured["k_scale"].shape == (1, 4, 1) + assert captured["v_scale"].shape == (1, 1, 128) + + +@pytest.mark.gpu +def test_fused_fp8_qkv_quantization_on_h100() -> None: + if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (9, 0): + pytest.skip("fused FP8 QKV quantization test requires H100") + + q = torch.randn(1, 130, 2, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn_like(q) + v = torch.randn_like(q) + q_fp8, k_fp8, v_fp8, q_scale, k_scale, v_scale = quantize_fp8_qkv(q, k, v) + + assert q_fp8.shape == q.shape + assert k_fp8.shape == k.shape + assert v_fp8.shape == v.shape + assert q_scale.shape == (1, 3, 2) + assert k_scale.shape == (1, 3, 2) + assert v_scale.shape == (1, 2, 128) + assert v_fp8.stride(1) == 1 + torch.testing.assert_close( + dequantize_fp8_per_block(q_fp8, q_scale, torch.bfloat16), + q, + rtol=0.15, + atol=0.05, + ) + torch.testing.assert_close( + dequantize_fp8_per_channel(v_fp8, v_scale, torch.bfloat16), + v, + rtol=0.15, + atol=0.05, + ) + + +@pytest.mark.gpu +def test_fp8_sol_handles_partial_tail_on_h100() -> None: + if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (9, 0): + pytest.skip("partial-tail FP8 Sol-Attn test requires H100") + + q = torch.randn(1, 130, 1, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn_like(q) + v = torch.randn_like(q) + q_fp8, k_fp8, v_fp8, q_scale, k_scale, v_scale = quantize_fp8_qkv(q, k, v) + output = attention_impl.sol_attn( + q_fp8, + k_fp8, + v_fp8, + tau=-1000.0, + q_scale=q_scale, + k_scale=k_scale, + v_scale=v_scale, + ) + reference = attention_impl.sol_attn( + dequantize_fp8_per_block(q_fp8, q_scale, torch.bfloat16).contiguous(), + dequantize_fp8_per_block(k_fp8, k_scale, torch.bfloat16).contiguous(), + dequantize_fp8_per_channel(v_fp8, v_scale, torch.bfloat16).contiguous(), + tau=-1000.0, + ) + + cosine = torch.nn.functional.cosine_similarity(output.float().flatten(), reference.float().flatten(), dim=0) + assert cosine > 0.99 + + +@pytest.mark.gpu +def test_fp8_sol_preserves_constant_values_at_long_sequence_on_h100() -> None: + if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (9, 0): + pytest.skip("long-sequence FP8 Sol-Attn test requires H100") + + q = torch.randn(1, 2048, 1, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn_like(q) + v = torch.ones_like(q) + q, k, v, q_scale, k_scale, v_scale = quantize_fp8_qkv(q, k, v) + output = attention_impl.sol_attn( + q, + k, + v, + tau=-1000.0, + q_scale=q_scale, + k_scale=k_scale, + v_scale=v_scale, + ) + + torch.testing.assert_close(output.float(), torch.ones_like(output, dtype=torch.float32), rtol=0.02, atol=0.02) + + +@pytest.mark.gpu +def test_fp8_sol_split_preserves_sparse_route_weights_on_h100() -> None: + if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (9, 0): + pytest.skip("split FP8 Sol-Attn test requires H100") + + torch.manual_seed(7) + q = torch.randn(1, 4160, 1, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn_like(q) + v = torch.randn_like(q) + reference = attention_impl.sol_attn(q, k, v, tau=1.0, kv_splits=1) + q, k, v, q_scale, k_scale, v_scale = quantize_fp8_qkv(q, k, v) + output = attention_impl.sol_attn( + q, + k, + v, + tau=1.0, + kv_splits=2, + q_scale=q_scale, + k_scale=k_scale, + v_scale=v_scale, + ) + + cosine = torch.nn.functional.cosine_similarity(output.float().flatten(), reference.float().flatten(), dim=0) + assert cosine > 0.98 diff --git a/tests/unit/ops/test_fp8_gemm.py b/tests/unit/ops/test_fp8_gemm.py index 256fb3f..d6324f1 100644 --- a/tests/unit/ops/test_fp8_gemm.py +++ b/tests/unit/ops/test_fp8_gemm.py @@ -51,3 +51,51 @@ def test_fp8_linear_tf_kernel_forward() -> None: assert actual.dtype == expected.dtype assert torch.isfinite(actual).all() torch.testing.assert_close(actual.float(), expected.float(), atol=0.1, rtol=0.1) + + +@pytest.mark.gpu +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.skipif(fp8_gemm.tf_kernel is None, reason="tf-kernel is required") +def test_fp8_linear_preserves_cuda_fallback_when_casting_is_disabled() -> None: + linear = nn.Linear(8, 4, device="cuda", dtype=torch.float32) + inputs = torch.randn(2, 8, device="cuda", dtype=torch.float32) + wrapped = fp8_gemm.FP8Linear( + linear, + options=fp8_gemm.FP8GemmOptions( + cast_inputs=False, + fp16_weight_storage="keep", + materialize_fp8_on_wrap=False, + ), + ) + + torch.testing.assert_close(wrapped(inputs), linear(inputs)) + + +@pytest.mark.gpu +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.skipif(fp8_gemm.tf_kernel is None, reason="tf-kernel is required") +def test_fp8_linear_forward_many_reuses_activation_quantization(monkeypatch: pytest.MonkeyPatch) -> None: + torch.manual_seed(0) + modules = tuple( + fp8_gemm.FP8Linear( + nn.Linear(64, 128, device="cuda", dtype=torch.bfloat16), + options=fp8_gemm.FP8GemmOptions(fp16_weight_storage="keep"), + ) + for _ in range(3) + ) + inputs = torch.randn(2, 3, 64, device="cuda", dtype=torch.bfloat16) + expected = tuple(module(inputs) for module in modules) + quantization_calls = 0 + quantize = modules[0]._tf_kernel.tf_per_token_quant_fp8 + + def tracked_quantize(*args, **kwargs): + nonlocal quantization_calls + quantization_calls += 1 + return quantize(*args, **kwargs) + + monkeypatch.setattr(modules[0]._tf_kernel, "tf_per_token_quant_fp8", tracked_quantize) + actual = fp8_gemm.fp8_linear_forward_many(modules, inputs) + + assert quantization_calls == 1 + for result, reference in zip(actual, expected): + torch.testing.assert_close(result.float(), reference.float(), atol=0.1, rtol=0.1) diff --git a/tests/unit/ops/test_sol_attention.py b/tests/unit/ops/test_sol_attention.py index ab663fd..f87414b 100644 --- a/tests/unit/ops/test_sol_attention.py +++ b/tests/unit/ops/test_sol_attention.py @@ -94,6 +94,60 @@ def test_sol_attention_dense_guard_does_not_call_kernel() -> None: kernel.assert_not_called() +@pytest.mark.gpu +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_fp8_sol_uses_triton_for_unquantized_bf16_layers() -> None: + q = torch.randn(1, 64, 1, 128, device="cuda", dtype=torch.bfloat16) + kernel = MagicMock(side_effect=lambda q, _k, _v, **_kwargs: q) + config = AttentionConfig.sol_attention( + dense_timesteps=0, + dense_layers=0, + sol_fp8=True, + sol_fp8_layer_start=10, + sol_fp8_layer_end=20, + ) + assert config.sparse_config is not None + state = SparseAttentionState(config.sparse_config, mask_map=None) + + with ( + patch.object(attention_impl, "SOL_ATTN_AVAILABLE", True), + patch.object(attention_impl, "sol_attn", kernel), + ): + output = attention_impl.attention(q, q, q, attention_config=config, sparse_state=state) + + assert output.shape == q.shape + assert kernel.call_args.kwargs["force_triton"] is True + + +def test_fp8_dense_uses_sdpa_for_unquantized_bf16_layers() -> None: + q = torch.randn(1, 64, 1, 128, dtype=torch.bfloat16) + kernel = MagicMock() + config = AttentionConfig.sol_attention( + dense_timesteps=0, + dense_layers=0, + tau=-1000.0, + sol_fp8=True, + sol_fp8_layer_start=10, + sol_fp8_layer_end=20, + ) + assert config.sparse_config is not None + state = SparseAttentionState(config.sparse_config, mask_map=None) + + with ( + patch.object(attention_impl, "SOL_ATTN_AVAILABLE", True), + patch.object(attention_impl, "sol_attn", kernel), + ): + output = attention_impl.attention(q, q, q, attention_config=config, sparse_state=state) + + expected = torch.nn.functional.scaled_dot_product_attention( + q.transpose(1, 2), + q.transpose(1, 2), + q.transpose(1, 2), + ).transpose(1, 2) + torch.testing.assert_close(output, expected) + kernel.assert_not_called() + + @pytest.mark.gpu def test_sol_attention_public_ops_matches_sdpa_on_h100(monkeypatch: pytest.MonkeyPatch) -> None: if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (9, 0): diff --git a/tests/unit/pipelines/minimax_h3/test_examples.py b/tests/unit/pipelines/minimax_h3/test_examples.py index 4655763..336113d 100644 --- a/tests/unit/pipelines/minimax_h3/test_examples.py +++ b/tests/unit/pipelines/minimax_h3/test_examples.py @@ -134,6 +134,13 @@ def fake_loader(model_root: str, **kwargs: object) -> object: "enable_fsdp": True, "online_adaln_cache": True, "attn_impl": AttnImplType.FLASH_ATTN_4, + "sol_fp8": False, + "sol_dense_steps": 10, + "sol_dense_layers": 2, + "sol_tau": 1.0, + "sol_threshold_type": "exact", + "sol_fp8_layer_start": 0, + "sol_fp8_layer_end": None, "feature_cache_config": FeatureCacheConfig( enabled=True, model_type="MiniMax-H3-Base", @@ -166,6 +173,7 @@ def test_cache_calibration_applies_validated_h3_profile(tmp_path: Path) -> None: [ ("torchao-fp8", QuantType.TORCHAO_FP8, QuantKernelBackend.TORCHAO), ("torchao_fp8", QuantType.TORCHAO_FP8, QuantKernelBackend.TORCHAO), + ("fp8", QuantType.FP8, QuantKernelBackend.TF_KERNEL), ("tf-kernel-fp8", QuantType.FP8, QuantKernelBackend.TF_KERNEL), ("bnb-nf4", QuantType.BNB_NF4, QuantKernelBackend.BITSANDBYTES), ], @@ -221,6 +229,52 @@ def fake_get_pipeline(*args: object, **kwargs: object) -> object: assert fl2va_example.PIPELINE_MANIFEST["pipeline_name"] == fl2va_example.PPL_CONFIG["name"] +def test_standard_example_forwards_fp8_sol_configuration(monkeypatch: pytest.MonkeyPatch) -> None: + calls = [] + sentinel = object() + + def fake_get_pipeline(*args: object, **kwargs: object) -> object: + calls.append((args, kwargs)) + return sentinel + + monkeypatch.setattr(fl2va_example, "load_minimax_h3_pipeline", fake_get_pipeline) + result = fl2va_example.get_pipeline( + 1, + "/models/h3", + attn_impl="SOL_ATTN", + sol_fp8=True, + sol_dense_steps=10, + sol_dense_layers=2, + sol_tau=0.9, + sol_threshold_type="diag", + sol_fp8_layer_start=2, + sol_fp8_layer_end=40, + quantization="tf-kernel-fp8", + ) + + assert result is sentinel + options = calls[0][1] + assert options["attn_impl"] == "SOL_ATTN" + assert options["sol_fp8"] is True + assert options["sol_dense_steps"] == 10 + assert options["sol_dense_layers"] == 2 + assert options["sol_tau"] == 0.9 + assert options["sol_threshold_type"] == "diag" + assert options["sol_fp8_layer_start"] == 2 + assert options["sol_fp8_layer_end"] == 40 + assert options["quantization"] == "tf-kernel-fp8" + + +def test_sol_fp8_rejects_dense_attention(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="sol_fp8 requires"): + load_minimax_h3_pipeline( + tmp_path, + partition="FL2VA", + attn_impl=AttnImplType.FLASH_ATTN_4, + sol_fp8=True, + ) + + def test_fl2va_run_maps_standard_service_tasks_to_model_conditions() -> None: calls = [] marker = object() diff --git a/tests/unit/pipelines/minimax_h3/test_parallelism.py b/tests/unit/pipelines/minimax_h3/test_parallelism.py index 379366a..da9537e 100644 --- a/tests/unit/pipelines/minimax_h3/test_parallelism.py +++ b/tests/unit/pipelines/minimax_h3/test_parallelism.py @@ -4,6 +4,8 @@ import torch from telefuser.core.config import ( + AttentionConfig, + AttnImplType, ModelRuntimeConfig, OffloadConfig, ParallelConfig, @@ -19,7 +21,10 @@ from telefuser.pipelines.minimax_h3.vae import MiniMaxH3VideoVAEStage -def _stage(parallel_config: ParallelConfig) -> tuple[MiniMaxH3DenoisingStage, MagicMock]: +def _stage( + parallel_config: ParallelConfig, + attention_config: AttentionConfig | None = None, +) -> tuple[MiniMaxH3DenoisingStage, MagicMock]: transformer = MagicMock() transformer.parameters.return_value = [torch.nn.Parameter(torch.zeros(1, dtype=torch.float32))] transformer.get_fsdp_module_names.return_value = ["blocks"] @@ -31,10 +36,19 @@ def _stage(parallel_config: ParallelConfig) -> tuple[MiniMaxH3DenoisingStage, Ma torch_dtype=torch.bfloat16, parallel_config=parallel_config, offload_config=OffloadConfig(offload_type=WeightOffloadType.NO_CPU_OFFLOAD), + attention_config=attention_config or AttentionConfig.dense_attention(), ) return MiniMaxH3DenoisingStage(manager, runtime), transformer +def test_single_gpu_stage_applies_sol_attention_config() -> None: + config = AttentionConfig.sol_attention(dense_timesteps=10, dense_layers=2, threshold_type="exact", sol_fp8=True) + + _, transformer = _stage(ParallelConfig(device_ids=[0]), config) + + transformer.set_attention_config.assert_called_once_with(config) + + def test_local_embedding_layout_selects_only_rank_owned_rows() -> None: layout = _build_local_embedding_layout( seq_len=12, diff --git a/tests/unit/pipelines/wan_video/test_optimized_example.py b/tests/unit/pipelines/wan_video/test_optimized_example.py new file mode 100644 index 0000000..38b3eb0 --- /dev/null +++ b/tests/unit/pipelines/wan_video/test_optimized_example.py @@ -0,0 +1,178 @@ +from __future__ import annotations + +import pytest +import torch + +from examples.wan_video.wan21_1_3b_text_to_video_optimized_h100 import ( + make_attention_config, + make_quant_config, + resolve_fp8_linear_scope, + run, +) +from telefuser.core.config import AttnImplType, QuantConfig, QuantKernelBackend, QuantType +from telefuser.models.wan_video_dit import WanModel + + +def test_wan_optimized_example_builds_compatible_configs() -> None: + attention = make_attention_config("sol") + quant = make_quant_config("torchao-fp8") + + assert attention.attn_impl is AttnImplType.SOL_ATTN + assert attention.sparse_config is not None + assert attention.sparse_config.sol_tau == 1.0 + assert not attention.sparse_config.sol_fp8 + assert quant.enabled + assert quant.quant_type is QuantType.TORCHAO_FP8 + assert quant.kernel_backend is QuantKernelBackend.TORCHAO + + +def test_wan_optimized_example_builds_dense_attention_config() -> None: + attention = make_attention_config("dense") + + assert attention.attn_impl is AttnImplType.TORCH_SDPA + assert attention.sparse_config is None + + +def test_wan_optimized_example_builds_fp8_sol_config() -> None: + attention = make_attention_config("fp8-sol", fp8_layer_start=10, fp8_layer_end=20) + assert attention.attn_impl is AttnImplType.SOL_ATTN + assert attention.sparse_config is not None + assert attention.sparse_config.sol_fp8 + assert attention.sparse_config.sol_tau == 1.0 + assert attention.sparse_config.sol_fp8_layer_start == 10 + assert attention.sparse_config.sol_fp8_layer_end == 20 + + +def test_wan_optimized_example_builds_fp8_dense_config() -> None: + attention = make_attention_config("fp8-dense", fp8_layer_start=10, fp8_layer_end=20) + + assert attention.attn_impl is AttnImplType.SOL_ATTN + assert attention.sparse_config is not None + assert attention.sparse_config.sol_fp8 + assert attention.sparse_config.dense_timesteps == 0 + assert attention.sparse_config.dense_layers == 0 + assert attention.sparse_config.sol_tau == -1000.0 + assert attention.sparse_config.sol_fp8_layer_start == 10 + assert attention.sparse_config.sol_fp8_layer_end == 20 + + +def test_wan_optimized_example_rejects_unknown_attention() -> None: + with pytest.raises(ValueError, match="attention must be"): + make_attention_config("radial") + + +@pytest.mark.parametrize( + ("name", "quant_type", "backend"), + [ + ("none", QuantType.FP8, QuantKernelBackend.AUTO), + ("tf-kernel-fp8", QuantType.FP8, QuantKernelBackend.TF_KERNEL), + ("torchao-fp8", QuantType.TORCHAO_FP8, QuantKernelBackend.TORCHAO), + ("bnb-nf4", QuantType.BNB_NF4, QuantKernelBackend.BITSANDBYTES), + ], +) +def test_wan_optimized_example_quantization_choices(name, quant_type, backend) -> None: + config = make_quant_config(name) + if name == "none": + assert not config.enabled + else: + assert config.enabled + assert config.quant_type is quant_type + assert config.kernel_backend is backend + + +def test_wan_optimized_example_rejects_unknown_quantization() -> None: + with pytest.raises(ValueError, match="quantization must be"): + make_quant_config("int8") + + +def test_wan_optimized_example_builds_ffn_only_fp8_config() -> None: + config = make_quant_config("tf-kernel-fp8", fp8_linear_scope="ffn") + + assert config.quantize_modules == (".ffn.",) + + +def test_wan_optimized_example_uses_all_linear_layers_for_auto_fp8_scope() -> None: + assert resolve_fp8_linear_scope("dense", "auto") == "all" + assert resolve_fp8_linear_scope("sol", "auto") == "all" + assert resolve_fp8_linear_scope("fp8-dense", "auto") == "all" + assert resolve_fp8_linear_scope("fp8-sol", "auto") == "all" + + +def test_wan_optimized_example_rejects_unknown_fp8_linear_scope() -> None: + with pytest.raises(ValueError, match="fp8_linear_scope must be"): + resolve_fp8_linear_scope("fp8-sol", "attention") + + +def test_wan_model_enables_tf_kernel_fp8_on_transformer_blocks(monkeypatch: pytest.MonkeyPatch) -> None: + model = WanModel.__new__(WanModel) + torch.nn.Module.__init__(model) + model.blocks = torch.nn.ModuleList([torch.nn.Sequential(torch.nn.Linear(8, 16), torch.nn.Linear(16, 8))]) + calls = [] + + def fake_count(module, *, module_filter=None): + calls.append(("count", module, module_filter)) + return 2 + + def fake_enable(module, *, options, module_filter=None): + calls.append(("enable", module, options, module_filter)) + return module + + monkeypatch.setattr("telefuser.ops.fp8_gemm.tf_kernel", None) + monkeypatch.setattr("telefuser.ops.fp8_gemm.count_linear_layers", fake_count) + monkeypatch.setattr("telefuser.ops.fp8_gemm.enable_fp8_gemm", fake_enable) + + model.enable_quant( + QuantConfig( + enabled=True, + quant_type=QuantType.FP8, + kernel_backend=QuantKernelBackend.TF_KERNEL, + ) + ) + + assert model.tf_kernel_fp8_replaced_linear == 2 + assert model.quant_type is QuantType.FP8 + assert calls[0][0] == "count" + assert calls[0][1] is model + assert calls[1][0] == "enable" + assert calls[1][1] is model + options = calls[1][2] + assert not options.cast_output_back + assert options.fp16_weight_storage == "discard" + assert options.materialize_fp8_on_wrap + module_filter = calls[1][3] + assert module_filter("blocks.0.0", model.blocks[0][0]) + assert not module_filter("head", model.blocks[0][0]) + + +def test_wan_optimized_run_forwards_explicit_benchmark_parameters() -> None: + captured = {} + + def fake_pipeline(**kwargs): + captured.update(kwargs) + return object() + + output = run( + fake_pipeline, + "prompt", + seed=7, + width=832, + height=480, + num_inference_steps=50, + num_frames=81, + cfg_scale=5.0, + sigma_shift=5.0, + ) + + assert output is not None + assert captured["seed"] == 7 + assert captured["width"] == 832 + assert captured["height"] == 480 + assert captured["num_inference_steps"] == 50 + assert captured["num_frames"] == 81 + assert captured["cfg_scale"] == 5.0 + assert captured["sigma_shift"] == 5.0 + + +def test_wan_optimized_run_requires_width_and_height_together() -> None: + with pytest.raises(ValueError, match="width and height must be provided together"): + run(lambda **_kwargs: object(), "prompt", width=832) diff --git a/tests/unit/service/test_example_service_parity.py b/tests/unit/service/test_example_service_parity.py index 8de6f59..587d008 100644 --- a/tests/unit/service/test_example_service_parity.py +++ b/tests/unit/service/test_example_service_parity.py @@ -30,6 +30,7 @@ "wan21_i2v_service": (Path("examples/wan_video/wan21_14b_image_to_video_480p_service.py"), "i2v", True), "minimax_h3_fl2va": (Path("examples/minimax_h3/minimax_h3_fl2va_h100.py"), "t2v", True), "minimax_h3_ref2va": (Path("examples/minimax_h3/minimax_h3_ref2va_h100.py"), "s2v", True), + "minimax_h3_turbo_lora": (Path("examples/minimax_h3/minimax_h3_turbo_lora_h100.py"), "i2v", True), "wan22_i2v_distill": (Path("examples/wan_video/wan22_14b_image_to_video_distill_h100.py"), "i2v", True), "lingbot_video_dense": (Path("examples/lingbot_video/lingbot_video_dense_1_3b.py"), "t2i", True), "lingbot_video_moe": (Path("examples/lingbot_video/lingbot_video_moe_30b.py"), "t2i", True), diff --git a/tools/validation/benchmark_minimax_h3_quantization.py b/tools/validation/benchmark_minimax_h3_quantization.py index 9ec5ea1..50d43a7 100644 --- a/tools/validation/benchmark_minimax_h3_quantization.py +++ b/tools/validation/benchmark_minimax_h3_quantization.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -"""Benchmark MiniMax H3 BF16 and online-quantized single-GPU profiles.""" +"""Benchmark MiniMax H3 dense/Sol and BF16/FP8 single-GPU profiles.""" from __future__ import annotations @@ -23,8 +23,13 @@ def _package_version(name: str) -> str | None: def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--model-root", default="/hhb-data/aigc/model_zoo/MiniMaxAI_MiniMax-H3") - parser.add_argument("--backend", choices=("bf16", "torchao-fp8", "tf-kernel-fp8", "bnb-nf4"), required=True) + parser.add_argument( + "--backend", + choices=("bf16", "bf16-sol", "fp8", "fp8-sol", "torchao-fp8", "bnb-nf4"), + required=True, + ) parser.add_argument("--prompt", default="Steam rises from the ramen while the family talks in the background.") + parser.add_argument("--prompt-file", type=Path, help="JSON file containing a top-level prompt string.") parser.add_argument("--duration", type=float, default=5.0) parser.add_argument("--steps", type=int, default=50) parser.add_argument("--seed", type=int, default=0) @@ -33,15 +38,22 @@ def main() -> None: parser.add_argument("--output", type=Path, required=True) parser.add_argument("--metrics-json", type=Path) args = parser.parse_args() + prompt = args.prompt + if args.prompt_file is not None: + prompt = json.loads(args.prompt_file.read_text(encoding="utf-8"))["prompt"] - quantization = None if args.backend == "bf16" else args.backend + uses_sol = args.backend.endswith("-sol") + quantization = "tf-kernel-fp8" if args.backend in {"fp8", "fp8-sol"} else args.backend + if args.backend in {"bf16", "bf16-sol"}: + quantization = None load_started = time.perf_counter() pipeline = load_minimax_h3_pipeline( args.model_root, partition="FL2VA", device=args.device, num_inference_steps=args.steps, - attn_impl=AttnImplType.FLASH_ATTN_4, + attn_impl=AttnImplType.SOL_ATTN if uses_sol else AttnImplType.FLASH_ATTN_4, + sol_fp8=args.backend == "fp8-sol", quantization=quantization, ) load_seconds = time.perf_counter() - load_started @@ -49,7 +61,7 @@ def main() -> None: generation_started = time.perf_counter() result = pipeline( task="t2va", - prompt=args.prompt, + prompt=prompt, conditions=[], target={ "short_edge": 768, @@ -65,17 +77,20 @@ def main() -> None: finally: pipeline.stop() + denoising_seconds = float(result.runtime_metrics["denoising_seconds"]) report = { "backend": args.backend, "model_root": str(Path(args.model_root)), "output": str(args.output), - "prompt": args.prompt, + "prompt": prompt, "duration_seconds": args.duration, "num_inference_steps": args.steps, "seed": args.seed, "aspect_ratio": args.aspect_ratio, "load_seconds": load_seconds, "generation_seconds": generation_seconds, + "denoising_steps_per_second": args.steps / denoising_seconds, + "generated_video_seconds_per_second": args.duration / generation_seconds, "save_seconds": save_seconds, "runtime_metrics": result.runtime_metrics, "versions": {