fix directml backend

This commit is contained in:
Seunghoon Lee
2025-03-12 00:50:27 +09:00
parent 9b7bb5f213
commit 2be555dc82
2 changed files with 12 additions and 1 deletions
+6
View File
@@ -162,6 +162,12 @@ class OffloadHook(accelerate.hooks.ModelHook):
if device_map is None or max_memory != getattr(module, "balanced_offload_max_memory", None):
device_map = accelerate.infer_auto_device_map(module, max_memory=max_memory)
offload_dir = getattr(module, "offload_dir", os.path.join(shared.opts.accelerate_offload_path, module.__class__.__name__))
keys = device_map.keys()
for v in keys:
if isinstance(device_map[v], int):
# int implies CUDA device, but it will break DirectML backend.
# Therefore, the type of device should be added.
device_map[v] = f"{devices.device.type}:{device_map[v]}"
module = accelerate.dispatch_model(module, device_map=device_map, offload_dir=offload_dir)
module._hf_hook.execution_device = torch.device(devices.device) # pylint: disable=protected-access
module.balanced_offload_device_map = device_map