diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index e31617f89a..9ba0bb0e50 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -97,6 +97,180 @@ __all__ = [ "FastBaseModel", ] + +def _infer_device_map_from_loaded_model(model): + """ + Build a compact device_map dict by inspecting where each parameter of + *model* actually lives after a bitsandbytes multi-device load. + + The resulting map satisfies accelerate's check_device_map invariant: + every parameter in the state_dict is covered by exactly one prefix key. + + Algorithm: recurse over the module tree top-down. When all parameters in + a subtree share the same device, emit one entry for the whole subtree. + Otherwise recurse into children and emit entries for direct-parameter + leaves that are not already covered by a child entry. + """ + device_map = {} + + def _assign(module, prefix): + params = list(module.named_parameters(remove_duplicate=False)) + if not params: + bufs = list(module.named_buffers()) + if bufs: + device_map[prefix] = bufs[0][1].device + return + devices = {p.device for _, p in params} + if len(devices) == 1: + device_map[prefix] = next(iter(devices)) + else: + for child_name, child in module.named_children(): + child_prefix = f"{prefix}.{child_name}" if prefix else child_name + _assign(child, child_prefix) + for pname, param in module.named_parameters(remove_duplicate=False): + if "." not in pname: + full = f"{prefix}.{pname}" if prefix else pname + if not any( + full == k or full.startswith(k + ".") + for k in device_map + ): + device_map[full] = param.device + + _assign(model, "") + # Remove the root key only when child entries already cover all params. + # For single-device models, "" may be the sole entry and must be kept. + if "" in device_map and len(device_map) > 1: + device_map.pop("") + return device_map + + +def _attach_bnb_multidevice_hooks( + model, load_in_4bit, load_in_8bit, offload_embedding, fast_inference +): + """ + Retroactively attach accelerate AlignDevicesHook on a bitsandbytes- + quantised model that was loaded with a multi-device (or non-default-device) + device_map. + + When load_in_4bit or load_in_8bit is used together with an accelerate + device_map, AutoModel.from_pretrained places weights on the target CUDA + devices but does NOT call dispatch_model, so no AlignDevicesHook is + installed. Any cross-device forward pass then crashes with: + RuntimeError: Expected all tensors to be on the same device + + This function fixes that by installing hooks after loading, without + re-moving any quantised weight tensors. Two scenarios are handled: + + 1. Multi-device: weights span multiple CUDA devices. Per-block hooks + route each module's inputs to the correct device. + 2. Single non-default device: all weights land on e.g. cuda:1 but the + caller may pass inputs on cuda:0. A root-level hook fixes this. + + Guards + ------ + - Only runs when load_in_4bit or load_in_8bit is True. + - Skips the vLLM path (fast_inference=True). + - Skips the offload_embedding CPU-offload path (handled separately). + - Skips models that already have hf_device_map set (already dispatched). + - Skips models with no CUDA parameters (CPU-only or meta-device paths). + """ + if fast_inference: + return + if not (load_in_4bit or load_in_8bit): + return + if offload_embedding: + return + if getattr(model, "hf_device_map", None) is not None: + return # already dispatched + + try: + cuda_devs = { + p.device + for p in model.parameters() + if hasattr(p, "device") and p.device.type == "cuda" + } + except Exception: + return + + if not cuda_devs: + return # no CUDA parameters -- nothing to do + + # All weights on the default device -- the common single-GPU case. + default_cuda = torch.device("cuda", 0) + if cuda_devs == {default_cuda}: + return + + try: + from accelerate.hooks import attach_align_device_hook_on_blocks + from accelerate.utils import find_tied_parameters, retie_parameters + except ImportError: + return # accelerate not available + + try: + inferred_map = _infer_device_map_from_loaded_model(model) + if not inferred_map: + return + + # Determine the "main" device (first CUDA device encountered). + cuda_device_vals = [ + v for v in inferred_map.values() + if isinstance(v, torch.device) and v.type == "cuda" + ] + main_device = cuda_device_vals[0] if cuda_device_vals else next(iter(cuda_devs)) + + # Preserve tied-parameter references before installing hooks. + tied_params = find_tied_parameters(model) + + if len(cuda_devs) > 1: + # Multi-device: each block's hook routes inputs to its own device. + execution_device = dict(inferred_map) + execution_device[""] = main_device + offload_dict = {k: False for k in execution_device} + attach_align_device_hook_on_blocks( + model, + execution_device=execution_device, + offload=offload_dict, + weights_map=None, + offload_buffers=False, + ) + desc = f"{len(inferred_map)} block(s) across {len(cuda_devs)} device(s)" + else: + # Single non-default device: one root hook sends all inputs to it. + attach_align_device_hook_on_blocks( + model, + execution_device=main_device, + offload=False, + weights_map=None, + offload_buffers=False, + ) + desc = f"root hook -> {main_device} (single non-default device)" + + # Retie parameters that hook installation may have unlinked. + retie_parameters(model, tied_params) + + # Expose hf_device_map so downstream generation helpers work. + # Use the device index (int) as the value, matching HF's convention. + model.hf_device_map = { + k: v.index if isinstance(v, torch.device) else v + for k, v in inferred_map.items() + } + + print( + f"Unsloth: Attached accelerate AlignDevicesHook ({desc}) " + f"for bnb multi-GPU inference." + ) + except Exception as exc: + import warnings + warnings.warn( + f"Unsloth: Could not attach multi-device dispatch hooks automatically " + f"({type(exc).__name__}: {exc}). " + "Cross-device inference may fail. Consider using a single GPU or " + "calling accelerate.dispatch_model() manually.", + RuntimeWarning, + stacklevel=2, + ) + + global NUM_LOGITS_TO_KEEP NUM_LOGITS_TO_KEEP = dict() @@ -835,6 +1009,18 @@ class FastBaseModel: # attn_implementation = attn_implementation, **kwargs, ) + # Attach AlignDevicesHook for bnb multi-device / non-default-device + # loads. The bnb loading path places weights on CUDA devices but + # never calls dispatch_model, so no hooks are installed. Without + # hooks, any cross-device module call crashes at forward time with + # "Expected all tensors to be on the same device". + _attach_bnb_multidevice_hooks( + model, + load_in_4bit = load_in_4bit, + load_in_8bit = load_in_8bit, + offload_embedding = offload_embedding, + fast_inference = fast_inference, + ) if hasattr(model, "generate"): model.fast_generate = make_fast_generate_wrapper(model.generate) model.fast_generate_batches = error_out_no_vllm