fix: fall back to non-lazy mx.load when lazy param unavailable

MLX < 0.31 doesn't support mx.load(lazy=True). Try lazy first,
fall back to eager loading on TypeError.

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 773bfb6dc4
commit 2e814f8a6f
+6 -3
View File
@@ -331,17 +331,20 @@ def broadcast_model_weights(
"""
all_sum = partial(_all_sum_cpu, group=group)
# Source lazily loads weights (data stays on disk until mx.eval per tensor)
# Source loads weights (lazy if supported, so only one tensor in memory at a time)
weights: dict[str, mx.array] = {}
if is_source:
weight_files = sorted(model_path.glob("*.safetensors"))
if not weight_files:
weight_files = sorted(model_path.glob("**/*.safetensors"))
for wf in weight_files:
loaded = cast(dict[str, mx.array], mx.load(str(wf), lazy=True)) # pyright: ignore[reportCallIssue]
try:
loaded = cast(dict[str, mx.array], mx.load(str(wf), lazy=True)) # pyright: ignore[reportCallIssue]
except TypeError:
loaded = cast(dict[str, mx.array], mx.load(str(wf)))
weights.update(loaded)
logger.info(
f"Source mapped {len(weights)} weight tensors from {len(weight_files)} files"
f"Source loaded {len(weights)} weight tensors from {len(weight_files)} files"
)
# Broadcast weight metadata: {name: {shape, dtype}}