feat: [Qwen3.5] Support block-wise FP8 quantization and model adaptation (#18926)

This commit is contained in:
Zheng Li
2026-02-18 11:44:25 +08:00
committed by GitHub
parent 83e24e2eb4
commit fa5698d791
4 changed files with 57 additions and 12 deletions
+48
View File
@@ -728,6 +728,51 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
)
self.weight_loader_v2(param, loaded_weight_shard, shard_id)
def _load_merged_block_scale(
self, param: BasevLLMParameter, loaded_weight: torch.Tensor
):
"""
Handle block-wise scale loading for MergedColumnParallelLinear.
Similar to QKVParallelLinear._load_qkv_block_scale, but for merged column layers.
"""
weight_block_size = self.quant_method.quant_config.weight_block_size
block_n, _ = weight_block_size[0], weight_block_size[1]
block_n = 1 if getattr(param, "format_ue8m0", False) else block_n
# Calculate block sizes for each shard
shard_block_sizes = []
shard_block_offsets = []
current_block_offset = 0
for output_size in self.output_sizes:
shard_block_size = (output_size + block_n - 1) // block_n
shard_block_sizes.append(shard_block_size)
shard_block_offsets.append(current_block_offset)
current_block_offset += shard_block_size
# Load each shard
for shard_id, (shard_block_offset, shard_block_size) in enumerate(
zip(shard_block_offsets, shard_block_sizes)
):
# Extract the shard from loaded_weight
loaded_weight_shard = loaded_weight.narrow(
param.output_dim, shard_block_offset, shard_block_size
)
# Calculate per-rank offset and size (considering TP)
rank_shard_offset = shard_block_offset // self.tp_size
rank_shard_size = shard_block_size // self.tp_size
# Load into the parameter
param.load_merged_column_weight(
loaded_weight=loaded_weight_shard,
shard_id=shard_id,
shard_offset=rank_shard_offset,
shard_size=rank_shard_size,
tp_rank=self.tp_rank,
tp_size=self.tp_size,
use_presharded_weights=self.use_presharded_weights,
)
def weight_loader_v2(
self,
param: BasevLLMParameter,
@@ -743,6 +788,9 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
tp_size=self.tp_size,
)
return
elif isinstance(param, BlockQuantScaleParameter):
self._load_merged_block_scale(param, loaded_weight)
return
elif type(param) in (RowvLLMParameter, BasevLLMParameter):
param.load_merged_column_weight(
loaded_weight=loaded_weight,