[diffusion] feat: support nunchaku for Z-Image-Turbo and flux.1 (int4) (#18959)
This commit is contained in:
@@ -23,9 +23,36 @@ class FluxArchConfig(DiTArchConfig):
|
|||||||
|
|
||||||
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
|
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
|
||||||
|
|
||||||
|
# nunchaku checkpoint uses different weight names; map to sglang flux layout
|
||||||
param_names_mapping: dict = field(
|
param_names_mapping: dict = field(
|
||||||
default_factory=lambda: {
|
default_factory=lambda: {
|
||||||
r"transformer\.(\w*)\.(.*)$": r"\1.\2",
|
# HF diffusers format
|
||||||
|
r"^transformer\.(\w*)\.(.*)$": r"\1.\2",
|
||||||
|
# transformer_blocks nunchaku format (raw export - before internal conversion)
|
||||||
|
r"^transformer_blocks\.(\d+)\.mlp_fc1\.(.*)$": r"transformer_blocks.\1.ff.net.0.proj.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.mlp_fc2\.(.*)$": r"transformer_blocks.\1.ff.net.2.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.mlp_context_fc1\.(.*)$": r"transformer_blocks.\1.ff_context.net.0.proj.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.mlp_context_fc2\.(.*)$": r"transformer_blocks.\1.ff_context.net.2.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.qkv_proj\.(.*)$": r"transformer_blocks.\1.attn.to_qkv.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.qkv_proj_context\.(.*)$": r"transformer_blocks.\1.attn.to_added_qkv.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.out_proj\.(.*)$": r"transformer_blocks.\1.attn.to_out.0.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.out_proj_context\.(.*)$": r"transformer_blocks.\1.attn.to_add_out.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.norm_q\.(.*)$": r"transformer_blocks.\1.attn.norm_q.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.norm_k\.(.*)$": r"transformer_blocks.\1.attn.norm_k.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.norm_added_q\.(.*)$": r"transformer_blocks.\1.attn.norm_added_q.\2",
|
||||||
|
r"^transformer_blocks\.(\d+)\.norm_added_k\.(.*)$": r"transformer_blocks.\1.attn.norm_added_k.\2",
|
||||||
|
# transformer_blocks nunchaku format (already converted with convert_flux_state_dict)
|
||||||
|
r"^transformer_blocks\.(\d+)\.attn\.add_qkv_proj\.(.*)$": r"transformer_blocks.\1.attn.to_added_qkv.\2",
|
||||||
|
# single_transformer_blocks nunchaku format (raw export - before internal conversion)
|
||||||
|
r"^single_transformer_blocks\.(\d+)\.qkv_proj\.(.*)$": r"single_transformer_blocks.\1.attn.to_qkv.\2",
|
||||||
|
r"^single_transformer_blocks\.(\d+)\.out_proj\.(.*)$": r"single_transformer_blocks.\1.attn.to_out.0.\2",
|
||||||
|
r"^single_transformer_blocks\.(\d+)\.norm_q\.(.*)$": r"single_transformer_blocks.\1.attn.norm_q.\2",
|
||||||
|
r"^single_transformer_blocks\.(\d+)\.norm_k\.(.*)$": r"single_transformer_blocks.\1.attn.norm_k.\2",
|
||||||
|
# nunchaku quantization parameter name conversions (apply to all blocks)
|
||||||
|
r"^(.*)\.smooth_orig$": r"\1.smooth_factor_orig",
|
||||||
|
r"^(.*)\.smooth$": r"\1.smooth_factor",
|
||||||
|
r"^(.*)\.lora_down$": r"\1.proj_down",
|
||||||
|
r"^(.*)\.lora_up$": r"\1.proj_up",
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
+31
-38
@@ -16,27 +16,6 @@ from .base_config import QuantizationConfig, QuantizeMethodBase
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
SVDQ_W4A4_LAYER_PATTERNS = [
|
|
||||||
"attn.to_qkv",
|
|
||||||
"attn.to_out",
|
|
||||||
"attn.add_qkv_proj",
|
|
||||||
"attn.to_add_out",
|
|
||||||
"img_mlp",
|
|
||||||
"txt_mlp",
|
|
||||||
]
|
|
||||||
|
|
||||||
AWQ_W4A16_LAYER_PATTERNS = [
|
|
||||||
"img_mod",
|
|
||||||
"txt_mod",
|
|
||||||
]
|
|
||||||
|
|
||||||
SKIP_QUANTIZATION_PATTERNS = [
|
|
||||||
"norm",
|
|
||||||
"embed",
|
|
||||||
"rotary",
|
|
||||||
"pos_embed",
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
@lru_cache(maxsize=1)
|
@lru_cache(maxsize=1)
|
||||||
def is_nunchaku_available() -> bool:
|
def is_nunchaku_available() -> bool:
|
||||||
@@ -61,13 +40,15 @@ class NunchakuConfig(QuantizationConfig):
|
|||||||
group_size: Quantization group size (automatically set based on precision)
|
group_size: Quantization group size (automatically set based on precision)
|
||||||
act_unsigned: Use unsigned activation quantization
|
act_unsigned: Use unsigned activation quantization
|
||||||
quantized_model_path: Path to pre-quantized model weights (.safetensors)
|
quantized_model_path: Path to pre-quantized model weights (.safetensors)
|
||||||
|
model_cls: DiT model class that provides quantization rules via get_nunchaku_quant_rules()
|
||||||
"""
|
"""
|
||||||
|
|
||||||
precision: str = "int4" # "int4" or "nvfp4"
|
precision: str = "int4"
|
||||||
rank: int = 32
|
rank: int = 32
|
||||||
group_size: Optional[int] = None
|
group_size: Optional[int] = None
|
||||||
act_unsigned: bool = False
|
act_unsigned: bool = False
|
||||||
quantized_model_path: Optional[str] = None
|
quantized_model_path: Optional[str] = None
|
||||||
|
model_cls: Optional[type] = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_name(cls) -> str:
|
def get_name(cls) -> str:
|
||||||
@@ -99,32 +80,28 @@ class NunchakuConfig(QuantizationConfig):
|
|||||||
def get_quant_method(
|
def get_quant_method(
|
||||||
self, layer: torch.nn.Module, prefix: str
|
self, layer: torch.nn.Module, prefix: str
|
||||||
) -> Optional[QuantizeMethodBase]:
|
) -> Optional[QuantizeMethodBase]:
|
||||||
|
|
||||||
if not isinstance(layer, LinearBase):
|
if not isinstance(layer, LinearBase):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
for pattern in SKIP_QUANTIZATION_PATTERNS:
|
# get quantization rules from model class
|
||||||
|
quant_rules = self._get_quant_rules()
|
||||||
|
|
||||||
|
# priority: skip > awq_w4a16 > svdq_w4a4 > default
|
||||||
|
skip_patterns = quant_rules.get("skip", [])
|
||||||
|
for pattern in skip_patterns:
|
||||||
if pattern in prefix.lower():
|
if pattern in prefix.lower():
|
||||||
return None
|
return None
|
||||||
|
|
||||||
for pattern in SVDQ_W4A4_LAYER_PATTERNS:
|
awq_patterns = quant_rules.get("awq_w4a16", [])
|
||||||
if pattern in prefix:
|
for pattern in awq_patterns:
|
||||||
from ..nunchaku_linear import NunchakuSVDQLinearMethod
|
|
||||||
|
|
||||||
return NunchakuSVDQLinearMethod(
|
|
||||||
precision=self.precision,
|
|
||||||
rank=self.rank,
|
|
||||||
act_unsigned=self.act_unsigned,
|
|
||||||
)
|
|
||||||
|
|
||||||
for pattern in AWQ_W4A16_LAYER_PATTERNS:
|
|
||||||
if pattern in prefix:
|
if pattern in prefix:
|
||||||
from ..nunchaku_linear import NunchakuAWQLinearMethod
|
from ..nunchaku_linear import NunchakuAWQLinearMethod
|
||||||
|
|
||||||
return NunchakuAWQLinearMethod(
|
return NunchakuAWQLinearMethod(group_size=64)
|
||||||
group_size=64,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
svdq_patterns = quant_rules.get("svdq_w4a4", [])
|
||||||
|
for pattern in svdq_patterns:
|
||||||
|
if pattern in prefix:
|
||||||
from ..nunchaku_linear import NunchakuSVDQLinearMethod
|
from ..nunchaku_linear import NunchakuSVDQLinearMethod
|
||||||
|
|
||||||
return NunchakuSVDQLinearMethod(
|
return NunchakuSVDQLinearMethod(
|
||||||
@@ -133,6 +110,22 @@ class NunchakuConfig(QuantizationConfig):
|
|||||||
act_unsigned=self.act_unsigned,
|
act_unsigned=self.act_unsigned,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# default: apply svdq_w4a4 to all remaining linear layers
|
||||||
|
from ..nunchaku_linear import NunchakuSVDQLinearMethod
|
||||||
|
|
||||||
|
return NunchakuSVDQLinearMethod(
|
||||||
|
precision=self.precision,
|
||||||
|
rank=self.rank,
|
||||||
|
act_unsigned=self.act_unsigned,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _get_quant_rules(self) -> dict[str, list[str]]:
|
||||||
|
if self.model_cls is not None and hasattr(
|
||||||
|
self.model_cls, "get_nunchaku_quant_rules"
|
||||||
|
):
|
||||||
|
return self.model_cls.get_nunchaku_quant_rules()
|
||||||
|
return {}
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
if self.group_size is None:
|
if self.group_size is None:
|
||||||
if self.precision == "nvfp4":
|
if self.precision == "nvfp4":
|
||||||
|
|||||||
@@ -94,8 +94,9 @@ class TransformerLoader(ComponentLoader):
|
|||||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||||
|
|
||||||
nunchaku_config = server_args.nunchaku_config
|
nunchaku_config = server_args.nunchaku_config
|
||||||
|
|
||||||
if nunchaku_config is not None:
|
if nunchaku_config is not None:
|
||||||
|
nunchaku_config.model_cls = model_cls
|
||||||
|
|
||||||
# respect dtype from checkpoint
|
# respect dtype from checkpoint
|
||||||
# TODO: improve the condition
|
# TODO: improve the condition
|
||||||
param_dtype = None
|
param_dtype = None
|
||||||
@@ -158,7 +159,7 @@ class TransformerLoader(ComponentLoader):
|
|||||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||||
|
|
||||||
# considering the existent of mixed-precision models (e.g., nunchaku)
|
# considering the existent of mixed-precision models (e.g., nunchaku)
|
||||||
if next(model.parameters()).dtype != param_dtype:
|
if next(model.parameters()).dtype != param_dtype and param_dtype:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Model dtype does not match expected param dtype, {next(model.parameters()).dtype} vs {param_dtype}"
|
f"Model dtype does not match expected param dtype, {next(model.parameters()).dtype} vs {param_dtype}"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -238,7 +238,8 @@ def load_model_from_full_model_state_dict(
|
|||||||
"""
|
"""
|
||||||
meta_sd = model.state_dict()
|
meta_sd = model.state_dict()
|
||||||
param_dict = dict(model.named_parameters())
|
param_dict = dict(model.named_parameters())
|
||||||
sharded_sd = {}
|
|
||||||
|
# map names from checkpoint to customized names
|
||||||
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
|
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
|
||||||
full_sd_iterator, param_names_mapping
|
full_sd_iterator, param_names_mapping
|
||||||
) # type: ignore
|
) # type: ignore
|
||||||
@@ -250,7 +251,7 @@ def load_model_from_full_model_state_dict(
|
|||||||
# sort parameter names to ensure all ranks process parameters in the same order
|
# sort parameter names to ensure all ranks process parameters in the same order
|
||||||
sorted_param_names = sorted(custom_param_sd.keys())
|
sorted_param_names = sorted(custom_param_sd.keys())
|
||||||
|
|
||||||
requires_grad = False
|
sharded_sd = {}
|
||||||
|
|
||||||
# shard from loaded state_dict, custom_param_sd -> sharded_sd
|
# shard from loaded state_dict, custom_param_sd -> sharded_sd
|
||||||
for target_param_name in sorted_param_names:
|
for target_param_name in sorted_param_names:
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ def set_default_torch_dtype(dtype: torch.dtype):
|
|||||||
|
|
||||||
|
|
||||||
def get_param_names_mapping(
|
def get_param_names_mapping(
|
||||||
mapping_dict: dict[str, str],
|
mapping_dict: dict[str, str | tuple[str, int, int]],
|
||||||
) -> Callable[[str], tuple[str, Any, Any]]:
|
) -> Callable[[str], tuple[str, Any, Any]]:
|
||||||
"""
|
"""
|
||||||
Creates a mapping function that transforms parameter names using regex patterns.
|
Creates a mapping function that transforms parameter names using regex patterns.
|
||||||
@@ -44,21 +44,50 @@ def get_param_names_mapping(
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def mapping_fn(name: str) -> tuple[str, Any, Any]:
|
def mapping_fn(name: str) -> tuple[str, Any, Any]:
|
||||||
# Try to match and transform the name using the regex patterns in mapping_dict
|
# support chained conversions, e.g.:
|
||||||
for pattern, replacement in mapping_dict.items():
|
# transformer.xxx.lora_down -> xxx.lora_down -> xxx.proj_down
|
||||||
match = re.match(pattern, name)
|
|
||||||
if match:
|
|
||||||
merge_index = None
|
merge_index = None
|
||||||
total_split_params = None
|
total_split_params = None
|
||||||
if isinstance(replacement, tuple):
|
max_steps = max(8, len(mapping_dict) * 2)
|
||||||
merge_index = replacement[1]
|
applied_patterns: set[str] = set()
|
||||||
total_split_params = replacement[2]
|
visited_names: set[str] = {name}
|
||||||
replacement = replacement[0]
|
|
||||||
name = re.sub(pattern, replacement, name)
|
|
||||||
return name, merge_index, total_split_params
|
|
||||||
|
|
||||||
# If no pattern matches, return the original name
|
for _ in range(max_steps):
|
||||||
return name, None, None
|
transformed = False
|
||||||
|
for pattern, replacement in mapping_dict.items():
|
||||||
|
# avoid re-applying the same rule on its own output
|
||||||
|
if pattern in applied_patterns:
|
||||||
|
continue
|
||||||
|
if re.match(pattern, name) is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
curr_merge_index = None
|
||||||
|
curr_total_split_params = None
|
||||||
|
if isinstance(replacement, tuple):
|
||||||
|
curr_merge_index = replacement[1]
|
||||||
|
curr_total_split_params = replacement[2]
|
||||||
|
replacement = replacement[0]
|
||||||
|
|
||||||
|
new_name = re.sub(pattern, replacement, name)
|
||||||
|
|
||||||
|
if new_name != name:
|
||||||
|
if curr_merge_index is not None:
|
||||||
|
merge_index = curr_merge_index
|
||||||
|
total_split_params = curr_total_split_params
|
||||||
|
|
||||||
|
name = new_name
|
||||||
|
applied_patterns.add(pattern)
|
||||||
|
if name in visited_names:
|
||||||
|
transformed = False
|
||||||
|
break
|
||||||
|
visited_names.add(name)
|
||||||
|
transformed = True
|
||||||
|
break
|
||||||
|
|
||||||
|
if not transformed:
|
||||||
|
break
|
||||||
|
|
||||||
|
return name, merge_index, total_split_params
|
||||||
|
|
||||||
return mapping_fn
|
return mapping_fn
|
||||||
|
|
||||||
@@ -150,25 +179,5 @@ def _list_safetensors_files(model_path: str) -> list[str]:
|
|||||||
|
|
||||||
BYTES_PER_GB = 1024**3
|
BYTES_PER_GB = 1024**3
|
||||||
|
|
||||||
|
|
||||||
def get_memory_usage_of_component(module) -> float | None:
|
|
||||||
"""
|
|
||||||
returned value is in GB, rounded to 2 decimal digits
|
|
||||||
"""
|
|
||||||
if not isinstance(module, nn.Module):
|
|
||||||
return None
|
|
||||||
if hasattr(module, "get_memory_footprint"):
|
|
||||||
usage = module.get_memory_footprint() / BYTES_PER_GB
|
|
||||||
else:
|
|
||||||
# manually
|
|
||||||
param_size = sum(p.numel() * p.element_size() for p in module.parameters())
|
|
||||||
buffer_size = sum(b.numel() * b.element_size() for b in module.buffers())
|
|
||||||
|
|
||||||
total_size_bytes = param_size + buffer_size
|
|
||||||
usage = total_size_bytes / (1024**3)
|
|
||||||
|
|
||||||
return round(usage, 2)
|
|
||||||
|
|
||||||
|
|
||||||
# component name -> ComponentLoader class
|
# component name -> ComponentLoader class
|
||||||
component_name_to_loader_cls: Dict[str, Type[Any]] = {}
|
component_name_to_loader_cls: Dict[str, Type[Any]] = {}
|
||||||
|
|||||||
@@ -107,3 +107,17 @@ class CachableDiT(TeaCacheMixin, BaseDiT):
|
|||||||
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
||||||
super().__init__(config, **kwargs)
|
super().__init__(config, **kwargs)
|
||||||
self._init_teacache_state()
|
self._init_teacache_state()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_nunchaku_quant_rules(cls) -> dict[str, dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
Get quantization rules for Nunchaku quantization.
|
||||||
|
|
||||||
|
Returns a dict mapping layer name patterns to quantization configs:
|
||||||
|
{
|
||||||
|
"skip": [list of patterns to skip quantization],
|
||||||
|
"svdq_w4a4": [list of patterns for SVDQ W4A4],
|
||||||
|
"awq_w4a16": [list of patterns for AWQ W4A16],
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
return {}
|
||||||
|
|||||||
@@ -32,8 +32,17 @@ from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
|||||||
|
|
||||||
# from sglang.multimodal_gen.runtime.layers.layernorm import LayerNorm as LayerNorm
|
# from sglang.multimodal_gen.runtime.layers.layernorm import LayerNorm as LayerNorm
|
||||||
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, apply_qk_norm
|
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, apply_qk_norm
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import ColumnParallelLinear
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
ColumnParallelLinear,
|
||||||
|
MergedColumnParallelLinear,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
|
QuantizationConfig,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
|
||||||
|
NunchakuConfig,
|
||||||
|
is_nunchaku_available,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
NDRotaryEmbedding,
|
NDRotaryEmbedding,
|
||||||
apply_flashinfer_rope_qk_inplace,
|
apply_flashinfer_rope_qk_inplace,
|
||||||
@@ -49,16 +58,147 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
|||||||
|
|
||||||
logger = init_logger(__name__) # pylint: disable=invalid-name
|
logger = init_logger(__name__) # pylint: disable=invalid-name
|
||||||
|
|
||||||
|
try:
|
||||||
|
from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import]
|
||||||
|
from nunchaku.models.normalization import ( # type: ignore[import]
|
||||||
|
NunchakuAdaLayerNormZero,
|
||||||
|
NunchakuAdaLayerNormZeroSingle,
|
||||||
|
)
|
||||||
|
from nunchaku.ops.gemm import (
|
||||||
|
svdq_gemm_w4a4_cuda as _svdq_gemm_w4a4, # type: ignore[import]
|
||||||
|
)
|
||||||
|
from nunchaku.ops.quantize import (
|
||||||
|
svdq_quantize_w4a4_act_fuse_lora_cuda as _svdq_quantize_w4a4, # type: ignore[import]
|
||||||
|
)
|
||||||
|
|
||||||
|
_nunchaku_fused_ops_available = True
|
||||||
|
except Exception:
|
||||||
|
NunchakuFeedForward = None
|
||||||
|
NunchakuAdaLayerNormZero = None
|
||||||
|
NunchakuAdaLayerNormZeroSingle = None
|
||||||
|
_svdq_gemm_w4a4 = None
|
||||||
|
_svdq_quantize_w4a4 = None
|
||||||
|
_nunchaku_fused_ops_available = False
|
||||||
|
|
||||||
|
|
||||||
|
def _fused_gelu_mlp(
|
||||||
|
x: torch.Tensor,
|
||||||
|
fc1,
|
||||||
|
fc2,
|
||||||
|
pad_size: int = 256,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
Fused GELU MLP matching nunchaku's fused_gelu_mlp kernel path.
|
||||||
|
|
||||||
|
nunchaku's single-block MLP checkpoint is calibrated for the fused path where:
|
||||||
|
1. fc1 GEMM + GELU + 0.171875 shift + unsigned re-quantization + fc2.lora_down
|
||||||
|
are all done in a single fused kernel call
|
||||||
|
2. fc2 GEMM then receives unsigned INT4 activations (act_unsigned=True)
|
||||||
|
|
||||||
|
Using the sequential path (fc1 → GELU → fc2 with symmetric quantization) is
|
||||||
|
fundamentally incompatible with these wscales, causing visually wrong outputs.
|
||||||
|
"""
|
||||||
|
batch_size, seq_len, channels = x.shape
|
||||||
|
x_2d = x.view(batch_size * seq_len, channels)
|
||||||
|
|
||||||
|
quantized_x, ascales, lora_act = _svdq_quantize_w4a4(
|
||||||
|
x_2d,
|
||||||
|
lora_down=fc1.proj_down,
|
||||||
|
smooth=fc1.smooth_factor,
|
||||||
|
fp4=fc1.precision == "nvfp4",
|
||||||
|
pad_size=pad_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
batch_size_pad = (batch_size * seq_len + pad_size - 1) // pad_size * pad_size
|
||||||
|
is_fp4 = fc2.precision == "nvfp4"
|
||||||
|
|
||||||
|
qout_act = torch.empty(
|
||||||
|
batch_size_pad,
|
||||||
|
fc1.output_size_per_partition // 2,
|
||||||
|
dtype=torch.uint8,
|
||||||
|
device=x_2d.device,
|
||||||
|
)
|
||||||
|
if is_fp4:
|
||||||
|
qout_ascales = torch.empty(
|
||||||
|
fc1.output_size_per_partition // 16,
|
||||||
|
batch_size_pad,
|
||||||
|
dtype=torch.float8_e4m3fn,
|
||||||
|
device=x_2d.device,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
qout_ascales = torch.empty(
|
||||||
|
fc1.output_size_per_partition // 64,
|
||||||
|
batch_size_pad,
|
||||||
|
dtype=x_2d.dtype,
|
||||||
|
device=x_2d.device,
|
||||||
|
)
|
||||||
|
qout_lora_act = torch.empty(
|
||||||
|
batch_size_pad, fc2.proj_down.shape[1], dtype=torch.float32, device=x_2d.device
|
||||||
|
)
|
||||||
|
|
||||||
|
# fused: fc1 GEMM + GELU + shift + unsigned quantize + fc2.lora_down
|
||||||
|
_svdq_gemm_w4a4(
|
||||||
|
act=quantized_x,
|
||||||
|
wgt=fc1.qweight,
|
||||||
|
qout=qout_act,
|
||||||
|
ascales=ascales,
|
||||||
|
wscales=fc1.wscales,
|
||||||
|
oscales=qout_ascales,
|
||||||
|
lora_act_in=lora_act,
|
||||||
|
lora_up=fc1.proj_up,
|
||||||
|
lora_down=fc2.proj_down,
|
||||||
|
lora_act_out=qout_lora_act,
|
||||||
|
bias=fc1.bias,
|
||||||
|
smooth_factor=fc2.smooth_factor,
|
||||||
|
fp4=is_fp4,
|
||||||
|
alpha=getattr(fc1, "_nunchaku_alpha", None),
|
||||||
|
wcscales=getattr(fc1, "wcscales", None),
|
||||||
|
)
|
||||||
|
|
||||||
|
output = torch.empty(
|
||||||
|
batch_size * seq_len,
|
||||||
|
fc2.output_size_per_partition,
|
||||||
|
dtype=x_2d.dtype,
|
||||||
|
device=x_2d.device,
|
||||||
|
)
|
||||||
|
# fc2 GEMM with unsigned INT4 activations (fused kernel shifted by 0.171875)
|
||||||
|
_svdq_gemm_w4a4(
|
||||||
|
act=qout_act,
|
||||||
|
wgt=fc2.qweight,
|
||||||
|
out=output,
|
||||||
|
ascales=qout_ascales,
|
||||||
|
wscales=fc2.wscales,
|
||||||
|
lora_act_in=qout_lora_act,
|
||||||
|
lora_up=fc2.proj_up,
|
||||||
|
bias=fc2.bias,
|
||||||
|
fp4=is_fp4,
|
||||||
|
alpha=getattr(fc2, "_nunchaku_alpha", None),
|
||||||
|
wcscales=getattr(fc2, "wcscales", None),
|
||||||
|
act_unsigned=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
return output.view(batch_size, seq_len, -1)
|
||||||
|
|
||||||
|
|
||||||
def _get_qkv_projections(
|
def _get_qkv_projections(
|
||||||
attn: "FluxAttention", hidden_states, encoder_hidden_states=None
|
attn: "FluxAttention", hidden_states, encoder_hidden_states=None
|
||||||
):
|
):
|
||||||
|
if getattr(attn, "use_fused_qkv", False):
|
||||||
|
qkv, _ = attn.to_qkv(hidden_states)
|
||||||
|
query, key, value = [x.contiguous() for x in qkv.chunk(3, dim=-1)]
|
||||||
|
else:
|
||||||
query, _ = attn.to_q(hidden_states)
|
query, _ = attn.to_q(hidden_states)
|
||||||
key, _ = attn.to_k(hidden_states)
|
key, _ = attn.to_k(hidden_states)
|
||||||
value, _ = attn.to_v(hidden_states)
|
value, _ = attn.to_v(hidden_states)
|
||||||
|
|
||||||
encoder_query = encoder_key = encoder_value = None
|
encoder_query = encoder_key = encoder_value = None
|
||||||
if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None:
|
if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None:
|
||||||
|
if getattr(attn, "use_fused_added_qkv", False):
|
||||||
|
added_qkv, _ = attn.to_added_qkv(encoder_hidden_states)
|
||||||
|
encoder_query, encoder_key, encoder_value = [
|
||||||
|
x.contiguous() for x in added_qkv.chunk(3, dim=-1)
|
||||||
|
]
|
||||||
|
else:
|
||||||
encoder_query, _ = attn.add_q_proj(encoder_hidden_states)
|
encoder_query, _ = attn.add_q_proj(encoder_hidden_states)
|
||||||
encoder_key, _ = attn.add_k_proj(encoder_hidden_states)
|
encoder_key, _ = attn.add_k_proj(encoder_hidden_states)
|
||||||
encoder_value, _ = attn.add_v_proj(encoder_hidden_states)
|
encoder_value, _ = attn.add_v_proj(encoder_hidden_states)
|
||||||
@@ -81,6 +221,8 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
out_dim: int = None,
|
out_dim: int = None,
|
||||||
context_pre_only: Optional[bool] = None,
|
context_pre_only: Optional[bool] = None,
|
||||||
pre_only: bool = False,
|
pre_only: bool = False,
|
||||||
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
|
prefix: str = "",
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@@ -96,9 +238,22 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
self.added_kv_proj_dim = added_kv_proj_dim
|
self.added_kv_proj_dim = added_kv_proj_dim
|
||||||
self.added_proj_bias = added_proj_bias
|
self.added_proj_bias = added_proj_bias
|
||||||
|
|
||||||
|
self.use_fused_qkv = isinstance(quant_config, NunchakuConfig)
|
||||||
|
self.use_fused_added_qkv = isinstance(quant_config, NunchakuConfig)
|
||||||
|
|
||||||
self.norm_q = RMSNorm(dim_head, eps=eps)
|
self.norm_q = RMSNorm(dim_head, eps=eps)
|
||||||
self.norm_k = RMSNorm(dim_head, eps=eps)
|
self.norm_k = RMSNorm(dim_head, eps=eps)
|
||||||
|
|
||||||
|
if self.use_fused_qkv:
|
||||||
|
self.to_qkv = MergedColumnParallelLinear(
|
||||||
|
query_dim,
|
||||||
|
[self.inner_dim] * 3,
|
||||||
|
bias=bias,
|
||||||
|
gather_output=True,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.to_qkv" if prefix else "to_qkv",
|
||||||
|
)
|
||||||
|
else:
|
||||||
self.to_q = ColumnParallelLinear(
|
self.to_q = ColumnParallelLinear(
|
||||||
query_dim, self.inner_dim, bias=bias, gather_output=True
|
query_dim, self.inner_dim, bias=bias, gather_output=True
|
||||||
)
|
)
|
||||||
@@ -108,12 +263,16 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
self.to_v = ColumnParallelLinear(
|
self.to_v = ColumnParallelLinear(
|
||||||
query_dim, self.inner_dim, bias=bias, gather_output=True
|
query_dim, self.inner_dim, bias=bias, gather_output=True
|
||||||
)
|
)
|
||||||
|
|
||||||
if not self.pre_only:
|
if not self.pre_only:
|
||||||
self.to_out = torch.nn.ModuleList([])
|
self.to_out = torch.nn.ModuleList([])
|
||||||
self.to_out.append(
|
self.to_out.append(
|
||||||
ColumnParallelLinear(
|
ColumnParallelLinear(
|
||||||
self.inner_dim, self.out_dim, bias=out_bias, gather_output=True
|
self.inner_dim,
|
||||||
|
self.out_dim,
|
||||||
|
bias=out_bias,
|
||||||
|
gather_output=True,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.to_out.0" if prefix else "",
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
if dropout != 0.0:
|
if dropout != 0.0:
|
||||||
@@ -122,6 +281,16 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
if added_kv_proj_dim is not None:
|
if added_kv_proj_dim is not None:
|
||||||
self.norm_added_q = RMSNorm(dim_head, eps=eps)
|
self.norm_added_q = RMSNorm(dim_head, eps=eps)
|
||||||
self.norm_added_k = RMSNorm(dim_head, eps=eps)
|
self.norm_added_k = RMSNorm(dim_head, eps=eps)
|
||||||
|
if self.use_fused_added_qkv:
|
||||||
|
self.to_added_qkv = MergedColumnParallelLinear(
|
||||||
|
added_kv_proj_dim,
|
||||||
|
[self.inner_dim] * 3,
|
||||||
|
bias=added_proj_bias,
|
||||||
|
gather_output=True,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.to_added_qkv" if prefix else "to_added_qkv",
|
||||||
|
)
|
||||||
|
else:
|
||||||
self.add_q_proj = ColumnParallelLinear(
|
self.add_q_proj = ColumnParallelLinear(
|
||||||
added_kv_proj_dim,
|
added_kv_proj_dim,
|
||||||
self.inner_dim,
|
self.inner_dim,
|
||||||
@@ -141,7 +310,12 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
gather_output=True,
|
gather_output=True,
|
||||||
)
|
)
|
||||||
self.to_add_out = ColumnParallelLinear(
|
self.to_add_out = ColumnParallelLinear(
|
||||||
self.inner_dim, query_dim, bias=out_bias, gather_output=True
|
self.inner_dim,
|
||||||
|
query_dim,
|
||||||
|
bias=out_bias,
|
||||||
|
gather_output=True,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.to_add_out" if prefix else "",
|
||||||
)
|
)
|
||||||
|
|
||||||
self.attn = USPAttention(
|
self.attn = USPAttention(
|
||||||
@@ -218,6 +392,7 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
],
|
],
|
||||||
dim=1,
|
dim=1,
|
||||||
)
|
)
|
||||||
|
if not self.pre_only:
|
||||||
x, _ = self.to_out[0](x)
|
x, _ = self.to_out[0](x)
|
||||||
if len(self.to_out) == 2:
|
if len(self.to_out) == 2:
|
||||||
x = self.to_out[1](x)
|
x = self.to_out[1](x)
|
||||||
@@ -225,6 +400,10 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
|
|
||||||
return x, encoder_hidden_states
|
return x, encoder_hidden_states
|
||||||
else:
|
else:
|
||||||
|
if not self.pre_only:
|
||||||
|
x, _ = self.to_out[0](x)
|
||||||
|
if len(self.to_out) == 2:
|
||||||
|
x = self.to_out[1](x)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
@@ -235,19 +414,61 @@ class FluxSingleTransformerBlock(nn.Module):
|
|||||||
num_attention_heads: int,
|
num_attention_heads: int,
|
||||||
attention_head_dim: int,
|
attention_head_dim: int,
|
||||||
mlp_ratio: float = 4.0,
|
mlp_ratio: float = 4.0,
|
||||||
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
|
prefix: str = "",
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.mlp_hidden_dim = int(dim * mlp_ratio)
|
self.mlp_hidden_dim = int(dim * mlp_ratio)
|
||||||
|
self.use_nunchaku_structure = isinstance(quant_config, NunchakuConfig)
|
||||||
|
|
||||||
self.norm = AdaLayerNormZeroSingle(dim)
|
self.norm = AdaLayerNormZeroSingle(dim)
|
||||||
|
|
||||||
|
if self.use_nunchaku_structure:
|
||||||
|
self.mlp_fc1 = ColumnParallelLinear(
|
||||||
|
dim,
|
||||||
|
self.mlp_hidden_dim,
|
||||||
|
bias=True,
|
||||||
|
gather_output=True,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.mlp_fc1" if prefix else "mlp_fc1",
|
||||||
|
)
|
||||||
|
self.act_mlp = nn.GELU(approximate="tanh")
|
||||||
|
self.mlp_fc2 = ColumnParallelLinear(
|
||||||
|
self.mlp_hidden_dim,
|
||||||
|
dim,
|
||||||
|
bias=True,
|
||||||
|
gather_output=True,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.mlp_fc2" if prefix else "mlp_fc2",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.attn = FluxAttention(
|
||||||
|
query_dim=dim,
|
||||||
|
dim_head=attention_head_dim,
|
||||||
|
num_heads=num_attention_heads,
|
||||||
|
out_dim=dim,
|
||||||
|
bias=True,
|
||||||
|
eps=1e-6,
|
||||||
|
pre_only=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.attn" if prefix else "attn",
|
||||||
|
)
|
||||||
|
if NunchakuAdaLayerNormZeroSingle is not None:
|
||||||
|
self.norm = NunchakuAdaLayerNormZeroSingle(self.norm, scale_shift=0)
|
||||||
|
else:
|
||||||
self.proj_mlp = ColumnParallelLinear(
|
self.proj_mlp = ColumnParallelLinear(
|
||||||
dim, self.mlp_hidden_dim, bias=True, gather_output=True
|
dim,
|
||||||
|
self.mlp_hidden_dim,
|
||||||
|
bias=True,
|
||||||
|
gather_output=True,
|
||||||
)
|
)
|
||||||
self.act_mlp = nn.GELU(approximate="tanh")
|
self.act_mlp = nn.GELU(approximate="tanh")
|
||||||
self.proj_out = ColumnParallelLinear(
|
self.proj_out = ColumnParallelLinear(
|
||||||
dim + self.mlp_hidden_dim, dim, bias=True, gather_output=True
|
dim + self.mlp_hidden_dim,
|
||||||
|
dim,
|
||||||
|
bias=True,
|
||||||
|
gather_output=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.attn = FluxAttention(
|
self.attn = FluxAttention(
|
||||||
query_dim=dim,
|
query_dim=dim,
|
||||||
dim_head=attention_head_dim,
|
dim_head=attention_head_dim,
|
||||||
@@ -271,9 +492,34 @@ class FluxSingleTransformerBlock(nn.Module):
|
|||||||
|
|
||||||
residual = hidden_states
|
residual = hidden_states
|
||||||
norm_hidden_states, gate = self.norm(hidden_states, emb=temb)
|
norm_hidden_states, gate = self.norm(hidden_states, emb=temb)
|
||||||
|
joint_attention_kwargs = joint_attention_kwargs or {}
|
||||||
|
|
||||||
|
if self.use_nunchaku_structure:
|
||||||
|
if _nunchaku_fused_ops_available:
|
||||||
|
mlp_hidden_states = _fused_gelu_mlp(
|
||||||
|
norm_hidden_states, self.mlp_fc1, self.mlp_fc2
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
mlp_out, _ = self.mlp_fc1(norm_hidden_states)
|
||||||
|
mlp_hidden_states = self.act_mlp(mlp_out)
|
||||||
|
mlp_hidden_states, _ = self.mlp_fc2(mlp_hidden_states)
|
||||||
|
|
||||||
|
attn_output = self.attn(
|
||||||
|
x=norm_hidden_states,
|
||||||
|
freqs_cis=freqs_cis,
|
||||||
|
**joint_attention_kwargs,
|
||||||
|
)
|
||||||
|
if isinstance(attn_output, tuple):
|
||||||
|
attn_output = attn_output[0]
|
||||||
|
|
||||||
|
hidden_states = attn_output + mlp_hidden_states
|
||||||
|
gate = gate.unsqueeze(1)
|
||||||
|
hidden_states = gate * hidden_states
|
||||||
|
hidden_states = residual + hidden_states
|
||||||
|
else:
|
||||||
proj_hidden_states, _ = self.proj_mlp(norm_hidden_states)
|
proj_hidden_states, _ = self.proj_mlp(norm_hidden_states)
|
||||||
mlp_hidden_states = self.act_mlp(proj_hidden_states)
|
mlp_hidden_states = self.act_mlp(proj_hidden_states)
|
||||||
joint_attention_kwargs = joint_attention_kwargs or {}
|
|
||||||
attn_output = self.attn(
|
attn_output = self.attn(
|
||||||
x=norm_hidden_states,
|
x=norm_hidden_states,
|
||||||
freqs_cis=freqs_cis,
|
freqs_cis=freqs_cis,
|
||||||
@@ -285,6 +531,7 @@ class FluxSingleTransformerBlock(nn.Module):
|
|||||||
proj_out, _ = self.proj_out(hidden_states)
|
proj_out, _ = self.proj_out(hidden_states)
|
||||||
hidden_states = gate * proj_out
|
hidden_states = gate * proj_out
|
||||||
hidden_states = residual + hidden_states
|
hidden_states = residual + hidden_states
|
||||||
|
|
||||||
if hidden_states.dtype == torch.float16:
|
if hidden_states.dtype == torch.float16:
|
||||||
hidden_states = hidden_states.clip(-65504, 65504)
|
hidden_states = hidden_states.clip(-65504, 65504)
|
||||||
|
|
||||||
@@ -303,6 +550,8 @@ class FluxTransformerBlock(nn.Module):
|
|||||||
attention_head_dim: int,
|
attention_head_dim: int,
|
||||||
qk_norm: str = "rms_norm",
|
qk_norm: str = "rms_norm",
|
||||||
eps: float = 1e-6,
|
eps: float = 1e-6,
|
||||||
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
|
prefix: str = "",
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@@ -318,22 +567,38 @@ class FluxTransformerBlock(nn.Module):
|
|||||||
context_pre_only=False,
|
context_pre_only=False,
|
||||||
bias=True,
|
bias=True,
|
||||||
eps=eps,
|
eps=eps,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.attn" if prefix else "attn",
|
||||||
)
|
)
|
||||||
|
|
||||||
self.norm2 = LayerNorm(dim, eps=1e-6, elementwise_affine=False)
|
self.norm2 = LayerNorm(dim, eps=1e-6, elementwise_affine=False)
|
||||||
self.ff = MLP(
|
|
||||||
input_dim=dim, mlp_hidden_dim=dim * 4, output_dim=dim, act_type="gelu"
|
|
||||||
)
|
|
||||||
self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate")
|
|
||||||
|
|
||||||
self.norm2_context = LayerNorm(dim, eps=1e-6, elementwise_affine=False)
|
self.norm2_context = LayerNorm(dim, eps=1e-6, elementwise_affine=False)
|
||||||
self.ff_context = MLP(
|
|
||||||
input_dim=dim, mlp_hidden_dim=dim * 4, output_dim=dim, act_type="gelu"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
nunchaku_enabled = (
|
||||||
|
quant_config is not None
|
||||||
|
and hasattr(quant_config, "get_name")
|
||||||
|
and quant_config.get_name() == "svdquant"
|
||||||
|
and is_nunchaku_available()
|
||||||
|
and NunchakuFeedForward is not None
|
||||||
|
)
|
||||||
|
self.use_nunchaku_structure = nunchaku_enabled
|
||||||
|
self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate")
|
||||||
self.ff_context = FeedForward(
|
self.ff_context = FeedForward(
|
||||||
dim=dim, dim_out=dim, activation_fn="gelu-approximate"
|
dim=dim, dim_out=dim, activation_fn="gelu-approximate"
|
||||||
)
|
)
|
||||||
|
if nunchaku_enabled:
|
||||||
|
nunchaku_kwargs = {
|
||||||
|
"precision": quant_config.precision,
|
||||||
|
"rank": quant_config.rank,
|
||||||
|
"act_unsigned": quant_config.act_unsigned,
|
||||||
|
}
|
||||||
|
self.ff = NunchakuFeedForward(self.ff, **nunchaku_kwargs)
|
||||||
|
self.ff_context = NunchakuFeedForward(self.ff_context, **nunchaku_kwargs)
|
||||||
|
if NunchakuAdaLayerNormZero is not None:
|
||||||
|
self.norm1 = NunchakuAdaLayerNormZero(self.norm1, scale_shift=0)
|
||||||
|
self.norm1_context = NunchakuAdaLayerNormZero(
|
||||||
|
self.norm1_context, scale_shift=0
|
||||||
|
)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -369,6 +634,11 @@ class FluxTransformerBlock(nn.Module):
|
|||||||
attn_output = gate_msa.unsqueeze(1) * attn_output
|
attn_output = gate_msa.unsqueeze(1) * attn_output
|
||||||
hidden_states = hidden_states + attn_output
|
hidden_states = hidden_states + attn_output
|
||||||
norm_hidden_states = self.norm2(hidden_states)
|
norm_hidden_states = self.norm2(hidden_states)
|
||||||
|
if self.use_nunchaku_structure:
|
||||||
|
norm_hidden_states = (
|
||||||
|
norm_hidden_states * scale_mlp[:, None] + shift_mlp[:, None]
|
||||||
|
)
|
||||||
|
else:
|
||||||
norm_hidden_states = (
|
norm_hidden_states = (
|
||||||
norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
|
norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
|
||||||
)
|
)
|
||||||
@@ -385,6 +655,11 @@ class FluxTransformerBlock(nn.Module):
|
|||||||
encoder_hidden_states = encoder_hidden_states + context_attn_output
|
encoder_hidden_states = encoder_hidden_states + context_attn_output
|
||||||
|
|
||||||
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
|
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
|
||||||
|
if self.use_nunchaku_structure:
|
||||||
|
norm_encoder_hidden_states = (
|
||||||
|
norm_encoder_hidden_states * c_scale_mlp[:, None] + c_shift_mlp[:, None]
|
||||||
|
)
|
||||||
|
else:
|
||||||
norm_encoder_hidden_states = (
|
norm_encoder_hidden_states = (
|
||||||
norm_encoder_hidden_states * (1 + c_scale_mlp[:, None])
|
norm_encoder_hidden_states * (1 + c_scale_mlp[:, None])
|
||||||
+ c_shift_mlp[:, None]
|
+ c_shift_mlp[:, None]
|
||||||
@@ -433,7 +708,44 @@ class FluxTransformer2DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
|
|
||||||
param_names_mapping = FluxConfig().arch_config.param_names_mapping
|
param_names_mapping = FluxConfig().arch_config.param_names_mapping
|
||||||
|
|
||||||
def __init__(self, config: FluxConfig, hf_config: dict[str, Any]) -> None:
|
@classmethod
|
||||||
|
def get_nunchaku_quant_rules(cls) -> dict[str, list[str]]:
|
||||||
|
return {
|
||||||
|
"skip": [
|
||||||
|
"norm",
|
||||||
|
"embed",
|
||||||
|
"rotary",
|
||||||
|
"pos_embed",
|
||||||
|
],
|
||||||
|
"svdq_w4a4": [
|
||||||
|
"attn.to_qkv",
|
||||||
|
"attn.to_out",
|
||||||
|
"attn.add_qkv_proj",
|
||||||
|
"attn.to_added_qkv",
|
||||||
|
"attn.to_add_out",
|
||||||
|
"img_mlp",
|
||||||
|
"txt_mlp",
|
||||||
|
"attention.to_qkv",
|
||||||
|
"attention.to_out",
|
||||||
|
"proj_mlp",
|
||||||
|
"proj_out",
|
||||||
|
"mlp_fc1",
|
||||||
|
"mlp_fc2",
|
||||||
|
"ff.net",
|
||||||
|
"ff_context.net",
|
||||||
|
],
|
||||||
|
"awq_w4a16": [
|
||||||
|
"img_mod",
|
||||||
|
"txt_mod",
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: FluxConfig,
|
||||||
|
hf_config: dict[str, Any],
|
||||||
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
|
) -> None:
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
self.config = config.arch_config
|
self.config = config.arch_config
|
||||||
|
|
||||||
@@ -471,8 +783,10 @@ class FluxTransformer2DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
dim=self.inner_dim,
|
dim=self.inner_dim,
|
||||||
num_attention_heads=self.config.num_attention_heads,
|
num_attention_heads=self.config.num_attention_heads,
|
||||||
attention_head_dim=self.config.attention_head_dim,
|
attention_head_dim=self.config.attention_head_dim,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"transformer_blocks.{i}",
|
||||||
)
|
)
|
||||||
for _ in range(self.config.num_layers)
|
for i in range(self.config.num_layers)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -482,8 +796,10 @@ class FluxTransformer2DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
dim=self.inner_dim,
|
dim=self.inner_dim,
|
||||||
num_attention_heads=self.config.num_attention_heads,
|
num_attention_heads=self.config.num_attention_heads,
|
||||||
attention_head_dim=self.config.attention_head_dim,
|
attention_head_dim=self.config.attention_head_dim,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"single_transformer_blocks.{i}",
|
||||||
)
|
)
|
||||||
for _ in range(self.config.num_single_layers)
|
for i in range(self.config.num_single_layers)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1123,6 +1123,29 @@ class QwenImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
|
|
||||||
param_names_mapping = QwenImageDitConfig().arch_config.param_names_mapping
|
param_names_mapping = QwenImageDitConfig().arch_config.param_names_mapping
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_nunchaku_quant_rules(cls) -> dict[str, list[str]]:
|
||||||
|
return {
|
||||||
|
"skip": [
|
||||||
|
"norm",
|
||||||
|
"embed",
|
||||||
|
"rotary",
|
||||||
|
"pos_embed",
|
||||||
|
],
|
||||||
|
"svdq_w4a4": [
|
||||||
|
"attn.to_qkv",
|
||||||
|
"attn.to_out",
|
||||||
|
"attn.add_qkv_proj",
|
||||||
|
"attn.to_add_out",
|
||||||
|
"img_mlp",
|
||||||
|
"txt_mlp",
|
||||||
|
],
|
||||||
|
"awq_w4a16": [
|
||||||
|
"img_mod",
|
||||||
|
"txt_mod",
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: QwenImageDitConfig,
|
config: QwenImageDitConfig,
|
||||||
|
|||||||
@@ -15,6 +15,13 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
ReplicatedLinear,
|
ReplicatedLinear,
|
||||||
RowParallelLinear,
|
RowParallelLinear,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
|
||||||
|
QuantizationConfig,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
|
||||||
|
NunchakuConfig,
|
||||||
|
is_nunchaku_available,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
_apply_rotary_emb,
|
_apply_rotary_emb,
|
||||||
apply_flashinfer_rope_qk_inplace,
|
apply_flashinfer_rope_qk_inplace,
|
||||||
@@ -24,6 +31,11 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
|
|||||||
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
|
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
try:
|
||||||
|
from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import]
|
||||||
|
except Exception:
|
||||||
|
NunchakuFeedForward = None
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
_is_cuda = current_platform.is_cuda()
|
_is_cuda = current_platform.is_cuda()
|
||||||
|
|
||||||
@@ -111,6 +123,8 @@ class ZImageAttention(nn.Module):
|
|||||||
num_kv_heads: int,
|
num_kv_heads: int,
|
||||||
qk_norm: bool = True,
|
qk_norm: bool = True,
|
||||||
eps: float = 1e-6,
|
eps: float = 1e-6,
|
||||||
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
|
prefix: str = "",
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dim = dim
|
self.dim = dim
|
||||||
@@ -129,12 +143,25 @@ class ZImageAttention(nn.Module):
|
|||||||
self.local_num_heads = num_heads // tp_size
|
self.local_num_heads = num_heads // tp_size
|
||||||
self.local_num_kv_heads = num_kv_heads // tp_size
|
self.local_num_kv_heads = num_kv_heads // tp_size
|
||||||
|
|
||||||
|
kv_dim = self.head_dim * num_kv_heads
|
||||||
|
self.use_fused_qkv = isinstance(quant_config, NunchakuConfig)
|
||||||
|
|
||||||
|
if self.use_fused_qkv:
|
||||||
|
self.to_qkv = MergedColumnParallelLinear(
|
||||||
|
dim,
|
||||||
|
[dim, kv_dim, kv_dim],
|
||||||
|
bias=False,
|
||||||
|
gather_output=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.to_qkv",
|
||||||
|
)
|
||||||
|
else:
|
||||||
self.to_q = ColumnParallelLinear(dim, dim, bias=False, gather_output=False)
|
self.to_q = ColumnParallelLinear(dim, dim, bias=False, gather_output=False)
|
||||||
self.to_k = ColumnParallelLinear(
|
self.to_k = ColumnParallelLinear(
|
||||||
dim, self.head_dim * num_kv_heads, bias=False, gather_output=False
|
dim, kv_dim, bias=False, gather_output=False
|
||||||
)
|
)
|
||||||
self.to_v = ColumnParallelLinear(
|
self.to_v = ColumnParallelLinear(
|
||||||
dim, self.head_dim * num_kv_heads, bias=False, gather_output=False
|
dim, kv_dim, bias=False, gather_output=False
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.qk_norm:
|
if self.qk_norm:
|
||||||
@@ -145,7 +172,16 @@ class ZImageAttention(nn.Module):
|
|||||||
self.norm_k = None
|
self.norm_k = None
|
||||||
|
|
||||||
self.to_out = nn.ModuleList(
|
self.to_out = nn.ModuleList(
|
||||||
[RowParallelLinear(dim, dim, bias=False, input_is_parallel=True)]
|
[
|
||||||
|
RowParallelLinear(
|
||||||
|
dim,
|
||||||
|
dim,
|
||||||
|
bias=False,
|
||||||
|
input_is_parallel=True,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.to_out.0",
|
||||||
|
)
|
||||||
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
self.attn = USPAttention(
|
self.attn = USPAttention(
|
||||||
@@ -162,6 +198,20 @@ class ZImageAttention(nn.Module):
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||||
):
|
):
|
||||||
|
if self.use_fused_qkv:
|
||||||
|
qkv, _ = self.to_qkv(hidden_states)
|
||||||
|
q, k, v = qkv.split(
|
||||||
|
[
|
||||||
|
self.local_num_heads * self.head_dim,
|
||||||
|
self.local_num_kv_heads * self.head_dim,
|
||||||
|
self.local_num_kv_heads * self.head_dim,
|
||||||
|
],
|
||||||
|
dim=-1,
|
||||||
|
)
|
||||||
|
q = q.contiguous()
|
||||||
|
k = k.contiguous()
|
||||||
|
v = v.contiguous()
|
||||||
|
else:
|
||||||
q, _ = self.to_q(hidden_states)
|
q, _ = self.to_q(hidden_states)
|
||||||
k, _ = self.to_k(hidden_states)
|
k, _ = self.to_k(hidden_states)
|
||||||
v, _ = self.to_v(hidden_states)
|
v, _ = self.to_v(hidden_states)
|
||||||
@@ -214,6 +264,8 @@ class ZImageTransformerBlock(nn.Module):
|
|||||||
norm_eps: float,
|
norm_eps: float,
|
||||||
qk_norm: bool,
|
qk_norm: bool,
|
||||||
modulation=True,
|
modulation=True,
|
||||||
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
|
prefix: str = "",
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dim = dim
|
self.dim = dim
|
||||||
@@ -227,9 +279,37 @@ class ZImageTransformerBlock(nn.Module):
|
|||||||
num_kv_heads=n_kv_heads,
|
num_kv_heads=n_kv_heads,
|
||||||
qk_norm=qk_norm,
|
qk_norm=qk_norm,
|
||||||
eps=1e-5,
|
eps=1e-5,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"{prefix}.attention",
|
||||||
)
|
)
|
||||||
|
|
||||||
self.feed_forward = FeedForward(dim=dim, hidden_dim=int(dim / 3 * 8))
|
hidden_dim = int(dim / 3 * 8)
|
||||||
|
nunchaku_enabled = (
|
||||||
|
isinstance(quant_config, NunchakuConfig) and is_nunchaku_available()
|
||||||
|
)
|
||||||
|
if nunchaku_enabled:
|
||||||
|
import diffusers
|
||||||
|
|
||||||
|
ff = diffusers.models.attention.FeedForward(
|
||||||
|
dim=dim,
|
||||||
|
dim_out=dim,
|
||||||
|
activation_fn="swiglu",
|
||||||
|
inner_dim=hidden_dim,
|
||||||
|
bias=False,
|
||||||
|
)
|
||||||
|
nunchaku_kwargs = {
|
||||||
|
"precision": quant_config.precision,
|
||||||
|
"rank": quant_config.rank,
|
||||||
|
"act_unsigned": quant_config.act_unsigned,
|
||||||
|
}
|
||||||
|
self.feed_forward = NunchakuFeedForward(ff, **nunchaku_kwargs)
|
||||||
|
# NunchakuFeedForward overrides net[2].act_unsigned=True for int4 (GELU-specific
|
||||||
|
# optimization for non-negative activations). Z-Image uses SwiGLU whose output
|
||||||
|
# can be negative, so we must restore the original act_unsigned value.
|
||||||
|
if hasattr(self.feed_forward, "net") and len(self.feed_forward.net) > 2:
|
||||||
|
self.feed_forward.net[2].act_unsigned = quant_config.act_unsigned
|
||||||
|
else:
|
||||||
|
self.feed_forward = FeedForward(dim=dim, hidden_dim=hidden_dim)
|
||||||
|
|
||||||
self.attention_norm1 = RMSNorm(dim, eps=norm_eps)
|
self.attention_norm1 = RMSNorm(dim, eps=norm_eps)
|
||||||
self.ffn_norm1 = RMSNorm(dim, eps=norm_eps)
|
self.ffn_norm1 = RMSNorm(dim, eps=norm_eps)
|
||||||
@@ -389,10 +469,32 @@ class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
ZImageDitConfig().arch_config.reverse_param_names_mapping
|
ZImageDitConfig().arch_config.reverse_param_names_mapping
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_nunchaku_quant_rules(cls) -> dict[str, list[str]]:
|
||||||
|
return {
|
||||||
|
"skip": [
|
||||||
|
"norm",
|
||||||
|
"embed",
|
||||||
|
"rotary",
|
||||||
|
"pos_embed",
|
||||||
|
],
|
||||||
|
"svdq_w4a4": [
|
||||||
|
"attention.to_qkv",
|
||||||
|
"attention.to_out",
|
||||||
|
"img_mlp",
|
||||||
|
"txt_mlp",
|
||||||
|
],
|
||||||
|
"awq_w4a16": [
|
||||||
|
"img_mod",
|
||||||
|
"txt_mod",
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: ZImageDitConfig,
|
config: ZImageDitConfig,
|
||||||
hf_config: dict[str, Any],
|
hf_config: dict[str, Any],
|
||||||
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(config=config, hf_config=hf_config)
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
|
|
||||||
@@ -443,6 +545,8 @@ class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
arch_config.norm_eps,
|
arch_config.norm_eps,
|
||||||
arch_config.qk_norm,
|
arch_config.qk_norm,
|
||||||
modulation=True,
|
modulation=True,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"noise_refiner.{layer_id}",
|
||||||
)
|
)
|
||||||
for layer_id in range(arch_config.n_refiner_layers)
|
for layer_id in range(arch_config.n_refiner_layers)
|
||||||
]
|
]
|
||||||
@@ -457,6 +561,8 @@ class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
arch_config.norm_eps,
|
arch_config.norm_eps,
|
||||||
arch_config.qk_norm,
|
arch_config.qk_norm,
|
||||||
modulation=False,
|
modulation=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"context_refiner.{layer_id}",
|
||||||
)
|
)
|
||||||
for layer_id in range(arch_config.n_refiner_layers)
|
for layer_id in range(arch_config.n_refiner_layers)
|
||||||
]
|
]
|
||||||
@@ -482,6 +588,8 @@ class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
arch_config.n_kv_heads,
|
arch_config.n_kv_heads,
|
||||||
arch_config.norm_eps,
|
arch_config.norm_eps,
|
||||||
arch_config.qk_norm,
|
arch_config.qk_norm,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=f"layers.{layer_id}",
|
||||||
)
|
)
|
||||||
for layer_id in range(arch_config.num_layers)
|
for layer_id in range(arch_config.num_layers)
|
||||||
]
|
]
|
||||||
|
|||||||
Reference in New Issue
Block a user