mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 09:38:23 +02:00
fix directml backend
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user