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.
This commit is contained in:
parent
14ab6fbfae
commit
45572dedd8
1 changed files with 186 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue