Revert "Fix: MPS allocator error in model.to() transfer (#305)"

This reverts commit 9424020687.
This commit is contained in:
Adrien Toupet
2025-11-27 14:52:12 -05:00
parent 9424020687
commit 79e7f41216
+1 -25
View File
@@ -875,31 +875,7 @@ def _standard_model_movement(model: torch.nn.Module, current_device: torch.devic
debug.start_timer(timer_name)
# Move model and clear gradients
try:
model.to(target_device)
except RuntimeError as e:
# MPS allocator workaround: Some PyTorch/macOS versions fail with bulk transfers
error_msg = str(e).lower()
is_mps_alloc_error = (
torch.device(target_device).type == "mps" and
any(keyword in error_msg for keyword in ["watermark", "allocat", "memory"])
)
if is_mps_alloc_error:
if debug:
debug.log("MPS bulk transfer failed, moving parameters individually",
category="memory", indent_level=1)
# Move parameters and buffers individually to avoid allocator bug
with torch.no_grad():
for param in model.parameters():
param.data = param.data.to(target_device)
for buffer in model.buffers():
buffer.data = buffer.data.to(target_device)
else:
# Re-raise if it's a different error
raise
model.to(target_device)
model.zero_grad(set_to_none=True)
# Clear VAE memory buffers when moving to CPU