Merge remote-tracking branch 'staging/pr-5053-tests' into pr-5053-head
This commit is contained in:
commit
78e2ece54c
3 changed files with 259 additions and 0 deletions
108
test_attach_logger_error_no_misleading_warning.py
Normal file
108
test_attach_logger_error_no_misleading_warning.py
Normal 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]})"
|
||||
)
|
||||
81
test_infer_buffer_only_multi_device_recurses.py
Normal file
81
test_infer_buffer_only_multi_device_recurses.py
Normal 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
|
||||
70
test_llama_seq_class_hook_call_passes_false.py
Normal file
70
test_llama_seq_class_hook_call_passes_false.py
Normal 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)})"
|
||||
)
|
||||
Loading…
Add table
Add a link
Reference in a new issue