mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Add support for postprocess exported model for block scale swizzling and support for different padding strategy (#1195)
### What does this PR do? Type of change: ? new feature <!-- Details about the change. --> Adds post-processing support for exported diffusion model checkpoints to enable NVFP4 block scale swizzling and configurable padding strategies. This allows exported quantized checkpoints to be directly consumed by inference runtimes (e.g., ComfyUI with comfy_kitchen) that require cuBLAS 2-D block-scaling-factors layout. Changes: 1) Unified post-processing step (_postprocess_safetensors): Loads saved safetensors files and applies merge, padding, swizzle, and quantization metadata injection in a single pass. 2) NVFP4 scale swizzle (swizzle_nvfp4_scales): Rearranges block scales from ModelOpt's flat [rows, cols // 16] layout to cuBLAS 2-D tiled layout per the cuBLAS specification. 3) Configurable padding (pad_nvfp4_weights): Pads NVFP4 weight and scale tensors to multiples of 16, with "row" (rows only) or "row_col" (both dimensions) strategies. 4) Standalone quantization metadata (build_layerwise_quant_metadata): Extracted from merge_diffusion_checkpoint so _quantization_metadata can be injected independently of merging — works for both merged (LTX-2) and standalone (Flux2) exports. 5) Bug fix (conversion.py): Wrapped yield in try/finally in set_quantizer_by_cfg_context so quantizer states are always restored, fixing an issue when yield fails. ### Usage ```python # LTX-2 export with merge + swizzle + padding export_hf_checkpoint( pipeline, export_dir="./output", merged_base_safetensor_path="./ltx-2-22b-dev.safetensors", enable_swizzle_layout=True, padding_strategy="row_col", enable_layerwise_quant_metadata=True, ) # Flux2 standalone export with swizzle + padding (no merge needed) export_hf_checkpoint( transformer, export_dir="./output", enable_swizzle_layout=True, padding_strategy="row_col", ) # Via quantize.py CLI python quantize.py \ --model ltx-2 --format fp4 \ --extra-param merged_base_safetensor_path=./ltx-2-22b-dev.safetensors \ --extra-param enable_swizzle_layout=true \ --extra-param padding_strategy=row_col \ --hf-ckpt-dir ./output ``` ### Testing 1) Exported LTX-2.3 NVFP4 with swizzle + padding + merged base checkpoint. Verified checkpoint has correct uint8 weights, float8_e4m3fn scales in swizzled layout, and _quantization_metadata . Ran the checkpoint with ComfyUI 2) Exported Flux2 NVFP4 with swizzle + padding. Verified checkpoint has correct uint8 weights, float8_e4m3fn scales in swizzled layout, and _quantization_metadata . Ran the checkpoint with ComfyUI ### Before your PR is "*Ready for review*" Make sure you read and follow [Contributor guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md) and your commits are signed (`git commit -s -S`). Make sure you read and follow the [Security Best Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors) (e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(..., weights_only=False)`, `pickle`, etc.). - Is this change backward compatible?: ✅ - If you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: N/A - Did you write any new necessary tests?: N/A - Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?: N/A <!--- Only for new features, API changes, critical bug fixes or backward incompatible changes. --> ### Additional Information <!-- E.g. related issue. --> <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Diffusers export: optional NVFP4 support — swizzle layout, row/row_col padding, and optional per-layer quantization metadata; exports are now post-processed to apply these options. * Export flow accepts new flags to enable swizzle, padding strategy, and layerwise metadata. * **Bug Fixes** * Quantizer context manager now always restores state, including on exceptions. * **Tests** * Added unit tests for NVFP4 padding, swizzling, metadata injection, and post-processing. * **Documentation** * README example updated to show swizzle and padding flags. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: ynankani <ynankani@nvidia.com> Signed-off-by: YASH Nankani <ynankani@2u1g-x570-0073.ipp2a1.colossus.nvidia.com> Signed-off-by: ynankani-nv <ynankani@nvidia.com> Signed-off-by: YASH Nankani <ynankani@dl325g11-1979.ipp2a2.colossus.nvidia.com> Signed-off-by: YASH Nankani <ynankani@dl325g11-0771.ipp4a1.colossus.nvidia.com> Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com> Co-authored-by: YASH Nankani <ynankani@2u1g-x570-0073.ipp2a1.colossus.nvidia.com> Co-authored-by: YASH Nankani <ynankani@dl325g11-1979.ipp2a2.colossus.nvidia.com> Co-authored-by: YASH Nankani <ynankani@dl325g11-0771.ipp4a1.colossus.nvidia.com>
This commit is contained in:
co-authored by
coderabbitai[bot]
YASH Nankani
YASH Nankani
YASH Nankani
parent
b02e888550
commit
3ff15ccef3
@@ -162,6 +162,13 @@ python quantize.py \
|
||||
--extra-param merged_base_safetensor_path=./ltx-2-19b-dev-fp8.safetensors
|
||||
```
|
||||
|
||||
To additionally apply NVFP4 scale swizzle and padding , add:
|
||||
|
||||
```sh
|
||||
--extra-param enable_swizzle_layout=true \
|
||||
--extra-param padding_strategy=row_col
|
||||
```
|
||||
|
||||
#### Important Parameters
|
||||
|
||||
- `percentile`: Control quantization scaling factors (amax) collecting range, meaning that we will collect the chosen amax in the range of `(n_steps * percentile)` steps. Recommendation: 1.0
|
||||
|
||||
@@ -216,6 +216,9 @@ class PipelineManager:
|
||||
"fp8transformer", False
|
||||
)
|
||||
params.pop("merged_base_safetensor_path", None)
|
||||
params.pop("enable_swizzle_layout", None)
|
||||
params.pop("padding_strategy", None)
|
||||
params.pop("enable_layerwise_quant_metadata", None)
|
||||
|
||||
if not checkpoint_path:
|
||||
raise ValueError("Missing required extra_param: checkpoint_path.")
|
||||
|
||||
@@ -357,6 +357,28 @@ class ExportManager:
|
||||
if merged_path:
|
||||
self.logger.info(f"Merging base safetensors from {merged_path} for LTX2 export")
|
||||
kwargs["merged_base_safetensor_path"] = merged_path
|
||||
if model_config:
|
||||
for key in ("enable_swizzle_layout", "enable_layerwise_quant_metadata"):
|
||||
val = model_config.extra_params.get(key)
|
||||
if val is not None:
|
||||
normalized = str(val).strip().lower()
|
||||
if normalized in ("true", "1", "yes"):
|
||||
kwargs[key] = True
|
||||
elif normalized in ("false", "0", "no"):
|
||||
kwargs[key] = False
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid value for {key}: {val!r}. "
|
||||
"Expected true/false, 1/0, or yes/no."
|
||||
)
|
||||
padding = model_config.extra_params.get("padding_strategy")
|
||||
if padding is not None:
|
||||
padding = str(padding).strip().lower()
|
||||
if padding not in ("row", "row_col"):
|
||||
raise ValueError(
|
||||
f"Invalid padding_strategy: {padding!r}. Expected 'row' or 'row_col'."
|
||||
)
|
||||
kwargs["padding_strategy"] = padding
|
||||
export_hf_checkpoint(pipe, export_dir=self.config.hf_ckpt_dir, **kwargs)
|
||||
self.logger.info("HuggingFace checkpoint export completed successfully")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user