From a38ceafc888482a2ac3efad7fcb45a15266bc665 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Tue, 21 Apr 2026 00:34:27 +0000 Subject: [PATCH] Add review tests --- ...infer_buffer_only_multi_device_recurses.py | 92 +++++++++++++++++++ ..._llama_seq_class_hook_call_passes_false.py | 70 ++++++++++++++ 2 files changed, 162 insertions(+) create mode 100644 test_infer_buffer_only_multi_device_recurses.py create mode 100644 test_llama_seq_class_hook_call_passes_false.py diff --git a/test_infer_buffer_only_multi_device_recurses.py b/test_infer_buffer_only_multi_device_recurses.py new file mode 100644 index 0000000000..773efb51b6 --- /dev/null +++ b/test_infer_buffer_only_multi_device_recurses.py @@ -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 diff --git a/test_llama_seq_class_hook_call_passes_false.py b/test_llama_seq_class_hook_call_passes_false.py new file mode 100644 index 0000000000..a4252b46b9 --- /dev/null +++ b/test_llama_seq_class_hook_call_passes_false.py @@ -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)})" + )