From 9ea964a535c42d41a8feda9991edc274b55f81fd Mon Sep 17 00:00:00 2001 From: Yihao Wang <42559837+AgainstEntropy@users.noreply.github.com> Date: Tue, 28 Jul 2026 20:28:01 -0700 Subject: [PATCH] [diffusion] fix: per-shard FP8 scale shape for single-GPU fused linears (#32157) --- python/sglang/multimodal_gen/runtime/layers/linear.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/layers/linear.py b/python/sglang/multimodal_gen/runtime/layers/linear.py index aff7acc25..4e6657388 100644 --- a/python/sglang/multimodal_gen/runtime/layers/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/linear.py @@ -490,8 +490,9 @@ class MergedColumnParallelLinear(ColumnParallelLinear): tp_group: dist.ProcessGroup = None, ): tp_group = tp_group or get_tp_group() - if get_group_size(tp_group) > 1: - self.output_sizes = output_sizes + # Set output_sizes BEFORE super().__init__() so ColumnParallelLinear derives + # per-shard output_partition_sizes. + self.output_sizes = output_sizes super().__init__( input_size=input_size, output_size=sum(output_sizes), @@ -503,7 +504,6 @@ class MergedColumnParallelLinear(ColumnParallelLinear): prefix=prefix, tp_group=tp_group, ) - self.output_sizes = output_sizes assert all(output_size % self.tp_size == 0 for output_size in output_sizes) def weight_loader(