fix: call weight_loader in Step35ShardingStrategy for distributed transfer

Step35ShardingStrategy.shard_model() accepted the weight_loader parameter
but never called it, unlike all other sharding strategies. This meant
receiver nodes using distributed weight broadcast with tensor parallelism
on Step35 models would get zero/garbage weights after sharding.

Co-Authored-By: Claude Opus 4.6 <[email protected]>
This commit is contained in:
Alex Cheema
2026-02-17 10:05:43 -08:00
co-authored by Claude Opus 4.6
parent b5c4df2700
commit 000058aece
+3 -1
View File
@@ -1041,7 +1041,9 @@ class Step35ShardingStrategy(TensorParallelShardingStrategy):
) -> nn.Module:
model = cast(Step35Model, model)
for layer in model.layers:
for i, layer in enumerate(model.layers):
if weight_loader is not None:
weight_loader(model, i)
eval_with_timeout(
layer.parameters(), timeout_seconds / len(model.layers), on_timeout
)