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:
co-authored by
Claude Opus 4.6
parent
b5c4df2700
commit
000058aece
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user