From 45572dedd86bde7ebed5e6b25307adf6de6b6d2c Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 16 Apr 2026 11:47:29 +0000 Subject: [PATCH] BUG: fix multi-GPU inference crash for bnb 4-bit/8-bit models When load_in_4bit=True is used with device_map="sequential" and the model is placed across multiple GPUs (or entirely on a non-default GPU like cuda:1), the bitsandbytes loading path in transformers places weights on the target CUDA devices but never calls dispatch_model. This leaves zero AlignDevicesHook instances installed, so the first cross-device forward call crashes with: RuntimeError: Expected all tensors to be on the same device, but got index is on cuda:0, different from other tensors on cuda:1 This commonly happens on Kaggle 2xT4 setups where Gemma-4-31B in 4-bit (~17 GB) does not fit on a single 15 GB T4, forcing accelerate to shard across both devices. Fix: after from_pretrained returns a bnb-quantized model, infer a device map from post-load parameter placement and call accelerate's attach_align_device_hook_on_blocks to install input-routing hooks. Two scenarios are handled: 1. Multi-device (weights span multiple GPUs): per-block hooks route each module's inputs to the device holding its weights. 2. Single non-default device (all weights on e.g. cuda:1 because cuda:0 was too small): a root-level hook with io_same_device=True auto-moves inputs from any device to the model's device. Guards ensure zero overhead for the common single-GPU-on-cuda:0 path: - Only activates for load_in_4bit or load_in_8bit - Skips fast_inference (vLLM), offload_embedding, already-dispatched models - Skips when all weights are already on cuda:0 Tested on 2x B200 with Gemma-4-31B-it-unsloth-bnb-4bit across 6 different max_memory configurations (1 GiB to 15 GiB per device). All pass. --- unsloth/models/vision.py | 186 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 186 insertions(+) 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