[codex] Optimize Z-Image packed QKV (#24117)

This commit is contained in:
Xiaoyu Zhang
2026-05-07 07:51:22 +08:00
committed by GitHub
parent 9e1336d406
commit a9a8b20a90
3 changed files with 89 additions and 20 deletions
@@ -39,6 +39,39 @@ class ZImageArchConfig(DiTArchConfig):
param_names_mapping: dict = field(
default_factory=lambda: {
r"(.*)\.attention\.to_q\.weight$": (r"\1.attention.to_qkv.weight", 0, 3),
r"(.*)\.attention\.to_k\.weight$": (r"\1.attention.to_qkv.weight", 1, 3),
r"(.*)\.attention\.to_v\.weight$": (r"\1.attention.to_qkv.weight", 2, 3),
r"(.*)\.attention\.to_q\.weight_scale_inv$": (
r"\1.attention.to_qkv.weight_scale_inv",
0,
3,
),
r"(.*)\.attention\.to_k\.weight_scale_inv$": (
r"\1.attention.to_qkv.weight_scale_inv",
1,
3,
),
r"(.*)\.attention\.to_v\.weight_scale_inv$": (
r"\1.attention.to_qkv.weight_scale_inv",
2,
3,
),
r"(.*)\.attention\.to_q\.(lora_A|lora_B)$": (
r"\1.attention.to_qkv.\2",
0,
3,
),
r"(.*)\.attention\.to_k\.(lora_A|lora_B)$": (
r"\1.attention.to_qkv.\2",
1,
3,
),
r"(.*)\.attention\.to_v\.(lora_A|lora_B)$": (
r"\1.attention.to_qkv.\2",
2,
3,
),
r"(.*)\.feed_forward\.w1\.weight$": (r"\1.feed_forward.w13.weight", 0, 2),
r"(.*)\.feed_forward\.w3\.weight$": (r"\1.feed_forward.w13.weight", 1, 2),
r"(.*)\.feed_forward\.w1\.(lora_A|lora_B)$": (
@@ -613,6 +613,12 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
loaded_weight: torch.Tensor,
loaded_shard_id: int | None = None,
) -> None:
if isinstance(param, BlockQuantScaleParameter):
self._weight_loader_v2_block_quant_scale(
param, loaded_weight, loaded_shard_id
)
return
if loaded_shard_id is None:
if isinstance(param, PerTensorScaleParameter):
param.load_merged_column_weight(loaded_weight=loaded_weight, shard_id=0)
@@ -628,25 +634,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
tp_size = self.tp_size
if isinstance(param, BlockQuantScaleParameter):
raise NotImplementedError("FP8 is not implemented yet")
# FIXME(will): add fp8 support
# from vllm.model_executor.layers.quantization.fp8 import (
# Fp8LinearMethod, Fp8MoEMethod)
# assert self.quant_method is not None
# assert isinstance(self.quant_method,
# (Fp8LinearMethod, Fp8MoEMethod))
# weight_block_size = self.quant_method.quant_config.weight_block_size
# assert weight_block_size is not None
# block_n, _ = weight_block_size[0], weight_block_size[1]
# shard_offset = (
# (sum(self.output_sizes[:loaded_shard_id]) + block_n - 1) //
# block_n) // tp_size
# shard_size = ((self.output_sizes[loaded_shard_id] + block_n - 1) //
# block_n // tp_size)
else:
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
shard_size = self.output_sizes[loaded_shard_id] // tp_size
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
shard_size = self.output_sizes[loaded_shard_id] // tp_size
param.load_merged_column_weight(
loaded_weight=loaded_weight,
@@ -655,6 +644,53 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
shard_size=shard_size,
)
def _weight_loader_v2_block_quant_scale(
self,
param: BlockQuantScaleParameter,
loaded_weight: torch.Tensor,
loaded_shard_id: int | None = None,
) -> None:
assert self.quant_method is not None
weight_block_size = getattr(
self.quant_method.quant_config, "weight_block_size", None
)
if weight_block_size is None:
raise ValueError(
"MergedColumnParallelLinear block-scale loading requires "
"quant_config.weight_block_size."
)
block_n, _ = weight_block_size
output_dim = param.output_dim
if loaded_shard_id is None:
if param.data.shape == loaded_weight.shape:
param.data.copy_(loaded_weight)
return
block_offset = 0
for shard_id, output_size in enumerate(self.output_sizes):
block_size = divide(output_size, block_n)
loaded_weight_shard = loaded_weight.narrow(
output_dim, block_offset, block_size
)
self._weight_loader_v2_block_quant_scale(
param, loaded_weight_shard, shard_id
)
block_offset += block_size
return
assert loaded_shard_id < len(self.output_sizes)
shard_offset = divide(sum(self.output_sizes[:loaded_shard_id]), self.tp_size)
shard_size = divide(self.output_sizes[loaded_shard_id], self.tp_size)
block_shard_offset = divide(shard_offset, block_n)
block_shard_size = divide(shard_size, block_n)
param_data = param.data.narrow(output_dim, block_shard_offset, block_shard_size)
start_idx = self.tp_rank * block_shard_size
loaded_weight = loaded_weight.narrow(output_dim, start_idx, block_shard_size)
assert param_data.shape == loaded_weight.shape
param_data.copy_(loaded_weight)
class QKVParallelLinear(ColumnParallelLinear):
"""Linear layers for the attention's QKV transformation.
@@ -159,7 +159,7 @@ class ZImageAttention(nn.Module):
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)
self.use_fused_qkv = True
if self.use_fused_qkv:
self.to_qkv = MergedColumnParallelLinear(