| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| from typing import Dict, List, Optional |
|
|
| from .quant import * |
| from .quant import DynamicDiTQuantizer |
|
|
| DEFAULT_FP8_INCLUDE_PATTERNS = ["blocks"] |
| DEFAULT_FP8_EXCLUDE_PATTERNS = [] |
|
|
|
|
| def apply_fp8_quantization( |
| model, |
| quant_type: str = "fp8-per-token", |
| include_patterns: Optional[List[str]] = None, |
| exclude_patterns: Optional[List[str]] = None, |
| ) -> Dict[str, object]: |
| """Apply DynamicDiTQuantizer to the provided DiT model.""" |
| final_include_patterns = ( |
| include_patterns if include_patterns is not None else DEFAULT_FP8_INCLUDE_PATTERNS |
| ) |
| final_exclude_patterns = ( |
| exclude_patterns if exclude_patterns is not None else DEFAULT_FP8_EXCLUDE_PATTERNS |
| ) |
|
|
| quantizer = DynamicDiTQuantizer( |
| quant_type=quant_type, |
| include_patterns=final_include_patterns, |
| exclude_patterns=final_exclude_patterns, |
| ) |
| quantizer.convert_linear(model) |
| return { |
| "quantizer": quantizer, |
| "quant_type": quant_type, |
| "include_patterns": final_include_patterns, |
| "exclude_patterns": final_exclude_patterns, |
| } |
|
|