Add review tests

This commit is contained in:
Daniel Han 2026-04-21 00:34:27 +00:00
commit a38ceafc88
2 changed files with 162 additions and 0 deletions

View file

@ -0,0 +1,92 @@
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
def test_buffer_only_subtree_with_single_device_collapses():
"""Param-less subtree with buffers all on one device still collapses to a
single-entry prefix; no accidental recursion when not needed."""
infer, _ = _load_fns()
a = _BufMod(buffers = [("cache", "cuda:1")])
b = _BufMod(buffers = [("cache", "cuda:1")])
root = _BufMod(children = [("a", a), ("b", b)])
dm = infer(root)
assert dm == {"": torch.device("cuda", 1)}, 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)})"
)