[diffusion] fix: per-shard FP8 scale shape for single-GPU fused linears (#32157)
This commit is contained in:
@@ -490,8 +490,9 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
|||||||
tp_group: dist.ProcessGroup = None,
|
tp_group: dist.ProcessGroup = None,
|
||||||
):
|
):
|
||||||
tp_group = tp_group or get_tp_group()
|
tp_group = tp_group or get_tp_group()
|
||||||
if get_group_size(tp_group) > 1:
|
# Set output_sizes BEFORE super().__init__() so ColumnParallelLinear derives
|
||||||
self.output_sizes = output_sizes
|
# per-shard output_partition_sizes.
|
||||||
|
self.output_sizes = output_sizes
|
||||||
super().__init__(
|
super().__init__(
|
||||||
input_size=input_size,
|
input_size=input_size,
|
||||||
output_size=sum(output_sizes),
|
output_size=sum(output_sizes),
|
||||||
@@ -503,7 +504,6 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
|||||||
prefix=prefix,
|
prefix=prefix,
|
||||||
tp_group=tp_group,
|
tp_group=tp_group,
|
||||||
)
|
)
|
||||||
self.output_sizes = output_sizes
|
|
||||||
assert all(output_size % self.tp_size == 0 for output_size in output_sizes)
|
assert all(output_size % self.tp_size == 0 for output_size in output_sizes)
|
||||||
|
|
||||||
def weight_loader(
|
def weight_loader(
|
||||||
|
|||||||
Reference in New Issue
Block a user