Merge remote-tracking branch 'staging/pr-5053-tests' into pr-5053-head

This commit is contained in:
Daniel Han 2026-04-21 00:36:39 +00:00
commit 78e2ece54c
3 changed files with 259 additions and 0 deletions

View file

@ -0,0 +1,108 @@
import ast
import warnings as _warnings
from pathlib import Path
import torch
def _find_vision():
for p in [
Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py",
Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py",
]:
if p.exists():
return p
raise FileNotFoundError("vision.py not found")
class _RaisingLogger:
def info(self, *a, **kw):
raise RuntimeError("simulated broken logging handler")
def _load_fns(logger):
tree = ast.parse(_find_vision().read_text())
ns = {"torch": torch, "warnings": _warnings, "logger": logger}
for node in tree.body:
if isinstance(node, ast.FunctionDef) and node.name in {
"_infer_device_map_from_loaded_model",
"_attach_bnb_multidevice_hooks",
}:
exec(
compile(
ast.Module(body = [node], type_ignores = []),
str(_find_vision()),
"exec",
),
ns,
)
return ns["_infer_device_map_from_loaded_model"], ns[
"_attach_bnb_multidevice_hooks"
]
class _P:
def __init__(self, dev):
self.device = torch.device(dev) if isinstance(dev, str) else dev
class _FakeMod:
def __init__(self, params = None):
self._p = list(params or [])
self.hf_device_map = None
def named_parameters(self, recurse = True, remove_duplicate = False):
for n, d in self._p:
yield n, _P(d)
def parameters(self, recurse = True):
for _, p in self.named_parameters(recurse = recurse):
yield p
def named_buffers(self, recurse = True):
return iter([])
def named_children(self):
return iter([])
def test_successful_dispatch_does_not_emit_misleading_warning_when_logger_raises(monkeypatch):
"""When dispatch_model succeeds but the user's logger.info raises
(broken handler / strict test harness), the helper must NOT emit the
'Could not attach multi-device dispatch hooks automatically' warning,
because the hooks *are* installed."""
import accelerate
dispatch_called = {"n": 0}
monkeypatch.setattr(
accelerate,
"dispatch_model",
lambda *a, **kw: dispatch_called.__setitem__("n", dispatch_called["n"] + 1),
)
_, attach = _load_fns(_RaisingLogger())
m = _FakeMod(params = [("w", "cuda:1")])
with _warnings.catch_warnings(record = True) as caught:
_warnings.simplefilter("always")
try:
attach(
m,
load_in_4bit = True,
load_in_8bit = False,
offload_embedding = False,
fast_inference = False,
)
except RuntimeError:
# Post-fix: logger error may propagate (hooks are installed;
# surfacing a real logging misconfiguration is acceptable).
pass
assert dispatch_called["n"] == 1, "dispatch_model must have been called"
misleading = [
w
for w in caught
if "Could not attach multi-device dispatch hooks" in str(w.message)
]
assert not misleading, (
f"Must not emit the 'Could not attach' warning when dispatch "
f"actually succeeded (got {[str(w.message) for w in misleading]})"
)

View file

@ -0,0 +1,81 @@
import ast
from pathlib import Path
import logging
import warnings as _warnings
import torch
def _find_vision():
for p in [
Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py",
Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py",
]:
if p.exists():
return p
raise FileNotFoundError("vision.py not found")
def _load_fns():
tree = ast.parse(_find_vision().read_text())
ns = {"torch": torch, "warnings": _warnings, "logger": logging.getLogger("test")}
for node in tree.body:
if isinstance(node, ast.FunctionDef) and node.name in {
"_infer_device_map_from_loaded_model",
"_attach_bnb_multidevice_hooks",
}:
exec(
compile(
ast.Module(body = [node], type_ignores = []),
str(_find_vision()),
"exec",
),
ns,
)
return ns["_infer_device_map_from_loaded_model"], ns[
"_attach_bnb_multidevice_hooks"
]
class _P:
def __init__(self, dev):
self.device = torch.device(dev) if isinstance(dev, str) else dev
class _BufMod:
def __init__(self, buffers = None, children = None):
self._b = list(buffers or [])
self._c = list(children or [])
self.hf_device_map = None
def named_parameters(self, recurse = True, remove_duplicate = False):
if False:
yield
def parameters(self, recurse = True):
if False:
yield
def named_buffers(self, recurse = True):
for n, d in self._b:
yield n, _P(d)
if recurse:
for cn, cm in self._c:
for bn, bb in cm.named_buffers(recurse = True):
yield f"{cn}.{bn}", bb
def named_children(self):
yield from self._c
def test_buffer_only_subtree_with_multi_device_recurses():
"""Param-less subtree whose buffers span multiple devices must NOT collapse
to the first buffer's device. It must recurse into children so each child
gets its own map entry keyed by its own prefix."""
infer, _ = _load_fns()
a = _BufMod(buffers = [("cache", "cuda:0")])
b = _BufMod(buffers = [("cache", "cuda:1")])
root = _BufMod(children = [("a", a), ("b", b)])
dm = infer(root)
assert dm.get("a") == torch.device("cuda", 0), dm
assert dm.get("b") == torch.device("cuda", 1), dm
assert "" not in dm, dm

View file

@ -0,0 +1,70 @@
import ast
from pathlib import Path
def _find_llama():
for p in [
Path(__file__).resolve().parent / "unsloth" / "models" / "llama.py",
Path(__file__).resolve().parents[1] / "unsloth" / "models" / "llama.py",
]:
if p.exists():
return p
raise FileNotFoundError("llama.py not found")
def _collect_hook_calls_under(node):
"""Yield every Call node invoking _attach_bnb_multidevice_hooks within the
AST subtree rooted at `node`."""
for sub in ast.walk(node):
if (
isinstance(sub, ast.Call)
and getattr(sub.func, "id", None) == "_attach_bnb_multidevice_hooks"
):
yield sub
def _kwarg_literal(call, name):
for kw in call.keywords:
if kw.arg == name and isinstance(kw.value, ast.Constant):
return kw.value.value
if kw.arg == name:
return kw.value
return None
def test_seq_class_branch_passes_fast_inference_false():
"""The AutoModelForSequenceClassification branch (gated by
`if num_labels is not None:`) must pass fast_inference=False to
_attach_bnb_multidevice_hooks. That branch never reaches vLLM, so
forwarding a truthy fast_inference would short-circuit hook install
and regress multi-GPU bnb seq-class inference."""
tree = ast.parse(_find_llama().read_text())
hook_calls = []
for if_node in ast.walk(tree):
if not isinstance(if_node, ast.If):
continue
test = if_node.test
is_num_labels_not_none = (
isinstance(test, ast.Compare)
and isinstance(test.left, ast.Name)
and test.left.id == "num_labels"
and len(test.ops) == 1
and isinstance(test.ops[0], ast.IsNot)
and len(test.comparators) == 1
and isinstance(test.comparators[0], ast.Constant)
and test.comparators[0].value is None
)
if not is_num_labels_not_none:
continue
for body_node in if_node.body:
hook_calls.extend(_collect_hook_calls_under(body_node))
assert hook_calls, (
"No _attach_bnb_multidevice_hooks call found under `if num_labels is not None:`"
)
for call in hook_calls:
v = _kwarg_literal(call, "fast_inference")
assert v is False, (
f"seq-class hook call must use fast_inference=False (got {ast.dump(call)})"
)