diff --git a/test_attach_combined_load_flags_detection.py b/test_attach_combined_load_flags_detection.py index bf298f13fc..5872fa3c1a 100644 --- a/test_attach_combined_load_flags_detection.py +++ b/test_attach_combined_load_flags_detection.py @@ -9,7 +9,9 @@ def _find_vision(): for p in [ Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py", Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py", - Path("/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py"), + Path( + "/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py" + ), ]: if p.exists(): return p @@ -24,8 +26,17 @@ def _load_fns(): "_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"] + 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: @@ -34,30 +45,32 @@ class _P: class _FakeMod: - def __init__(self, params=None, buffers=None, children=None, hf_device_map=None): + def __init__(self, params = None, buffers = None, children = None, hf_device_map = None): self._p = list(params or []) self._b = list(buffers or []) self._c = list(children or []) self.hf_device_map = hf_device_map - def named_parameters(self, recurse=True, remove_duplicate=False): + def named_parameters(self, recurse = True, remove_duplicate = False): for n, d in self._p: yield n, _P(d) if recurse: for cn, cm in self._c: - for pn, pp in cm.named_parameters(recurse=True, remove_duplicate=remove_duplicate): + for pn, pp in cm.named_parameters( + recurse = True, remove_duplicate = remove_duplicate + ): yield f"{cn}.{pn}", pp - def parameters(self, recurse=True): - for _, p in self.named_parameters(recurse=recurse): + def parameters(self, recurse = True): + for _, p in self.named_parameters(recurse = recurse): yield p - def named_buffers(self, recurse=True): + 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): + for bn, bb in cm.named_buffers(recurse = True): yield f"{cn}.{bn}", bb def named_children(self): @@ -68,9 +81,20 @@ def test_attach_load_in_4bit_bool_alone_activates(monkeypatch): """Minimal contract: when only load_in_4bit=True is passed, the helper activates without requiring model.is_loaded_in_4bit to also be True.""" import accelerate + called = {"n": 0} - monkeypatch.setattr(accelerate, "dispatch_model", lambda *a, **kw: called.__setitem__("n", called["n"] + 1)) + monkeypatch.setattr( + accelerate, + "dispatch_model", + lambda *a, **kw: called.__setitem__("n", called["n"] + 1), + ) _, attach = _load_fns() - m = _FakeMod(params=[("w", "cuda:1")]) # no is_loaded_in_* attrs set - attach(m, load_in_4bit=True, load_in_8bit=False, offload_embedding=False, fast_inference=False) + m = _FakeMod(params = [("w", "cuda:1")]) # no is_loaded_in_* attrs set + attach( + m, + load_in_4bit = True, + load_in_8bit = False, + offload_embedding = False, + fast_inference = False, + ) assert called["n"] == 1 diff --git a/test_attach_detects_is_loaded_in_8bit.py b/test_attach_detects_is_loaded_in_8bit.py index 0d2b86a987..6e51640968 100644 --- a/test_attach_detects_is_loaded_in_8bit.py +++ b/test_attach_detects_is_loaded_in_8bit.py @@ -9,7 +9,9 @@ def _find_vision(): for p in [ Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py", Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py", - Path("/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py"), + Path( + "/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py" + ), ]: if p.exists(): return p @@ -24,8 +26,17 @@ def _load_fns(): "_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"] + 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: @@ -34,30 +45,32 @@ class _P: class _FakeMod: - def __init__(self, params=None, buffers=None, children=None, hf_device_map=None): + def __init__(self, params = None, buffers = None, children = None, hf_device_map = None): self._p = list(params or []) self._b = list(buffers or []) self._c = list(children or []) self.hf_device_map = hf_device_map - def named_parameters(self, recurse=True, remove_duplicate=False): + def named_parameters(self, recurse = True, remove_duplicate = False): for n, d in self._p: yield n, _P(d) if recurse: for cn, cm in self._c: - for pn, pp in cm.named_parameters(recurse=True, remove_duplicate=remove_duplicate): + for pn, pp in cm.named_parameters( + recurse = True, remove_duplicate = remove_duplicate + ): yield f"{cn}.{pn}", pp - def parameters(self, recurse=True): - for _, p in self.named_parameters(recurse=recurse): + def parameters(self, recurse = True): + for _, p in self.named_parameters(recurse = recurse): yield p - def named_buffers(self, recurse=True): + 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): + for bn, bb in cm.named_buffers(recurse = True): yield f"{cn}.{bn}", bb def named_children(self): @@ -68,10 +81,21 @@ def test_attach_detects_bnb_via_is_loaded_in_8bit(monkeypatch): """8-bit loads via quantization_config zero out load_in_*bit booleans; detection must still fire via model.is_loaded_in_8bit attribute.""" import accelerate + called = {"n": 0} - monkeypatch.setattr(accelerate, "dispatch_model", lambda *a, **kw: called.__setitem__("n", called["n"] + 1)) + monkeypatch.setattr( + accelerate, + "dispatch_model", + lambda *a, **kw: called.__setitem__("n", called["n"] + 1), + ) _, attach = _load_fns() - m = _FakeMod(params=[("w", "cuda:1")]) + m = _FakeMod(params = [("w", "cuda:1")]) m.is_loaded_in_8bit = True - attach(m, load_in_4bit=False, load_in_8bit=False, offload_embedding=False, fast_inference=False) + attach( + m, + load_in_4bit = False, + load_in_8bit = False, + offload_embedding = False, + fast_inference = False, + ) assert called["n"] == 1 diff --git a/test_attach_device_map_disk_entries.py b/test_attach_device_map_disk_entries.py index 0b9ec56b72..b06276c79f 100644 --- a/test_attach_device_map_disk_entries.py +++ b/test_attach_device_map_disk_entries.py @@ -9,7 +9,9 @@ def _find_vision(): for p in [ Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py", Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py", - Path("/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py"), + Path( + "/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py" + ), ]: if p.exists(): return p @@ -24,8 +26,17 @@ def _load_fns(): "_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"] + 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: @@ -34,30 +45,32 @@ class _P: class _FakeMod: - def __init__(self, params=None, buffers=None, children=None, hf_device_map=None): + def __init__(self, params = None, buffers = None, children = None, hf_device_map = None): self._p = list(params or []) self._b = list(buffers or []) self._c = list(children or []) self.hf_device_map = hf_device_map - def named_parameters(self, recurse=True, remove_duplicate=False): + def named_parameters(self, recurse = True, remove_duplicate = False): for n, d in self._p: yield n, _P(d) if recurse: for cn, cm in self._c: - for pn, pp in cm.named_parameters(recurse=True, remove_duplicate=remove_duplicate): + for pn, pp in cm.named_parameters( + recurse = True, remove_duplicate = remove_duplicate + ): yield f"{cn}.{pn}", pp - def parameters(self, recurse=True): - for _, p in self.named_parameters(recurse=recurse): + def parameters(self, recurse = True): + for _, p in self.named_parameters(recurse = recurse): yield p - def named_buffers(self, recurse=True): + 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): + for bn, bb in cm.named_buffers(recurse = True): yield f"{cn}.{bn}", bb def named_children(self): @@ -69,13 +82,22 @@ def test_attach_main_device_skips_cpu_and_disk_candidates(monkeypatch): non-device entries. Verifies the iter-4 `d not in ("cpu", "disk")` filter handles both string constants.""" import accelerate + rec = {} - monkeypatch.setattr(accelerate, "dispatch_model", lambda model, **kw: rec.update(kw)) + monkeypatch.setattr( + accelerate, "dispatch_model", lambda model, **kw: rec.update(kw) + ) _, attach = _load_fns() # First entry is cpu; fallback must find the cuda:1 entry instead. - a = _FakeMod(params=[("w", "cpu")]) - b = _FakeMod(params=[("w", "cuda:1")]) - m = _FakeMod(children=[("a", a), ("b", b)]) - attach(m, load_in_4bit=True, load_in_8bit=False, offload_embedding=False, fast_inference=False) + a = _FakeMod(params = [("w", "cpu")]) + b = _FakeMod(params = [("w", "cuda:1")]) + m = _FakeMod(children = [("a", a), ("b", b)]) + attach( + m, + load_in_4bit = True, + load_in_8bit = False, + offload_embedding = False, + fast_inference = False, + ) md = rec.get("main_device") assert md == 1, f"main_device must skip cpu/disk strings, got {md!r}" diff --git a/test_attach_no_side_effects_on_early_exit.py b/test_attach_no_side_effects_on_early_exit.py index c38c2c4186..04f9c389a9 100644 --- a/test_attach_no_side_effects_on_early_exit.py +++ b/test_attach_no_side_effects_on_early_exit.py @@ -9,7 +9,9 @@ def _find_vision(): for p in [ Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py", Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py", - Path("/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py"), + Path( + "/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py" + ), ]: if p.exists(): return p @@ -24,8 +26,17 @@ def _load_fns(): "_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"] + 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: @@ -35,19 +46,20 @@ class _P: class _ObservableMod: """Tracks whether any mutation happened on its parameters.""" + def __init__(self, params): self._p = [(n, _P(d)) for n, d in params] self.hf_device_map = None - def named_parameters(self, recurse=True, remove_duplicate=False): + def named_parameters(self, recurse = True, remove_duplicate = False): for n, p in self._p: yield n, p - def parameters(self, recurse=True): + def parameters(self, recurse = True): for _, p in self._p: yield p - def named_buffers(self, recurse=True): + def named_buffers(self, recurse = True): return iter([]) def named_children(self): @@ -58,12 +70,19 @@ def test_early_exit_does_not_strip_params(monkeypatch): """When a guard triggers (fast_inference=True), the helper must return before touching any parameter attributes (strip loop is never entered).""" import accelerate + monkeypatch.setattr(accelerate, "dispatch_model", lambda *a, **kw: None) _, attach = _load_fns() m = _ObservableMod([("w", "cuda:1")]) p = m._p[0][1] p._is_hf_initialized = "original" p._other_attr = "keep" - attach(m, load_in_4bit=True, load_in_8bit=False, offload_embedding=False, fast_inference=True) + attach( + m, + load_in_4bit = True, + load_in_8bit = False, + offload_embedding = False, + fast_inference = True, + ) assert p._is_hf_initialized == "original" assert p._other_attr == "keep" diff --git a/test_attach_restore_on_dispatch_exception.py b/test_attach_restore_on_dispatch_exception.py index 2a184fbb5e..7e320e5e91 100644 --- a/test_attach_restore_on_dispatch_exception.py +++ b/test_attach_restore_on_dispatch_exception.py @@ -9,7 +9,9 @@ def _find_vision(): for p in [ Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py", Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py", - Path("/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py"), + Path( + "/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py" + ), ]: if p.exists(): return p @@ -24,8 +26,17 @@ def _load_fns(): "_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"] + 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: @@ -38,15 +49,15 @@ class _TrackMod: self._p = [(n, _P(d)) for n, d in params] self.hf_device_map = None - def named_parameters(self, recurse=True, remove_duplicate=False): + def named_parameters(self, recurse = True, remove_duplicate = False): for n, p in self._p: yield n, p - def parameters(self, recurse=True): + def parameters(self, recurse = True): for _, p in self._p: yield p - def named_buffers(self, recurse=True): + def named_buffers(self, recurse = True): return iter([]) def named_children(self): @@ -57,8 +68,10 @@ def test_attach_restores_is_hf_initialized_after_dispatch_raises(monkeypatch): """If dispatch_model raises, the inner finally must still restore the stripped _is_hf_initialized attribute on every param.""" import accelerate + def boom(*a, **kw): raise RuntimeError("dispatch blew up") + monkeypatch.setattr(accelerate, "dispatch_model", boom) _, attach = _load_fns() m = _TrackMod([("w", "cuda:1")]) @@ -66,5 +79,11 @@ def test_attach_restores_is_hf_initialized_after_dispatch_raises(monkeypatch): p._is_hf_initialized = True with _warnings.catch_warnings(): _warnings.simplefilter("ignore") - attach(m, load_in_4bit=True, load_in_8bit=False, offload_embedding=False, fast_inference=False) + attach( + m, + load_in_4bit = True, + load_in_8bit = False, + offload_embedding = False, + fast_inference = False, + ) assert p.__dict__.get("_is_hf_initialized") is True diff --git a/test_infer_params_and_buffers_same_module.py b/test_infer_params_and_buffers_same_module.py index 8876eea7c9..bb9dd5bf07 100644 --- a/test_infer_params_and_buffers_same_module.py +++ b/test_infer_params_and_buffers_same_module.py @@ -9,7 +9,9 @@ def _find_vision(): for p in [ Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py", Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py", - Path("/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py"), + Path( + "/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py" + ), ]: if p.exists(): return p @@ -24,8 +26,17 @@ def _load_fns(): "_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"] + 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: @@ -34,30 +45,32 @@ class _P: class _FakeMod: - def __init__(self, params=None, buffers=None, children=None, hf_device_map=None): + def __init__(self, params = None, buffers = None, children = None, hf_device_map = None): self._p = list(params or []) self._b = list(buffers or []) self._c = list(children or []) self.hf_device_map = hf_device_map - def named_parameters(self, recurse=True, remove_duplicate=False): + def named_parameters(self, recurse = True, remove_duplicate = False): for n, d in self._p: yield n, _P(d) if recurse: for cn, cm in self._c: - for pn, pp in cm.named_parameters(recurse=True, remove_duplicate=remove_duplicate): + for pn, pp in cm.named_parameters( + recurse = True, remove_duplicate = remove_duplicate + ): yield f"{cn}.{pn}", pp - def parameters(self, recurse=True): - for _, p in self.named_parameters(recurse=recurse): + def parameters(self, recurse = True): + for _, p in self.named_parameters(recurse = recurse): yield p - def named_buffers(self, recurse=True): + 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): + for bn, bb in cm.named_buffers(recurse = True): yield f"{cn}.{bn}", bb def named_children(self): @@ -69,8 +82,8 @@ def test_infer_module_with_both_params_and_buffers(): to a single entry; the buffer must not confuse the single-device path.""" infer, _ = _load_fns() m = _FakeMod( - params=[("weight", "cuda:1")], - buffers=[("running_mean", "cuda:1"), ("running_var", "cuda:1")], + params = [("weight", "cuda:1")], + buffers = [("running_mean", "cuda:1"), ("running_var", "cuda:1")], ) dm = infer(m) assert dm == {"": torch.device("cuda", 1)} diff --git a/test_infer_three_level_deep_split.py b/test_infer_three_level_deep_split.py index 855674a484..f336afc190 100644 --- a/test_infer_three_level_deep_split.py +++ b/test_infer_three_level_deep_split.py @@ -9,7 +9,9 @@ def _find_vision(): for p in [ Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py", Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py", - Path("/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py"), + Path( + "/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py" + ), ]: if p.exists(): return p @@ -24,8 +26,17 @@ def _load_fns(): "_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"] + 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: @@ -34,30 +45,32 @@ class _P: class _FakeMod: - def __init__(self, params=None, buffers=None, children=None, hf_device_map=None): + def __init__(self, params = None, buffers = None, children = None, hf_device_map = None): self._p = list(params or []) self._b = list(buffers or []) self._c = list(children or []) self.hf_device_map = hf_device_map - def named_parameters(self, recurse=True, remove_duplicate=False): + def named_parameters(self, recurse = True, remove_duplicate = False): for n, d in self._p: yield n, _P(d) if recurse: for cn, cm in self._c: - for pn, pp in cm.named_parameters(recurse=True, remove_duplicate=remove_duplicate): + for pn, pp in cm.named_parameters( + recurse = True, remove_duplicate = remove_duplicate + ): yield f"{cn}.{pn}", pp - def parameters(self, recurse=True): - for _, p in self.named_parameters(recurse=recurse): + def parameters(self, recurse = True): + for _, p in self.named_parameters(recurse = recurse): yield p - def named_buffers(self, recurse=True): + 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): + for bn, bb in cm.named_buffers(recurse = True): yield f"{cn}.{bn}", bb def named_children(self): @@ -68,11 +81,11 @@ def test_infer_three_level_deep_mixed(): """Split at the third level of nesting: the algorithm must recurse deep enough to distinguish grandchildren on different devices.""" infer, _ = _load_fns() - g1 = _FakeMod(params=[("w", "cuda:0")]) - g2 = _FakeMod(params=[("w", "cuda:1")]) - level2 = _FakeMod(children=[("g1", g1), ("g2", g2)]) - level1 = _FakeMod(children=[("l2", level2)]) - root = _FakeMod(children=[("l1", level1)]) + g1 = _FakeMod(params = [("w", "cuda:0")]) + g2 = _FakeMod(params = [("w", "cuda:1")]) + level2 = _FakeMod(children = [("g1", g1), ("g2", g2)]) + level1 = _FakeMod(children = [("l2", level2)]) + root = _FakeMod(children = [("l1", level1)]) dm = infer(root) assert dm.get("l1.l2.g1") == torch.device("cuda", 0) assert dm.get("l1.l2.g2") == torch.device("cuda", 1) diff --git a/test_infer_tied_params.py b/test_infer_tied_params.py index 7ae0dbdbee..dbf93055b4 100644 --- a/test_infer_tied_params.py +++ b/test_infer_tied_params.py @@ -9,7 +9,9 @@ def _find_vision(): for p in [ Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py", Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py", - Path("/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py"), + Path( + "/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py" + ), ]: if p.exists(): return p @@ -24,8 +26,17 @@ def _load_fns(): "_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"] + 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: @@ -37,19 +48,20 @@ class _TiedMod: """Emits the same parameter object under two different names to simulate tied weights (lm_head.weight == embed.weight). With remove_duplicate=False we yield both names; devices unioned must still be a single device.""" + def __init__(self, dev): self._shared = _P(dev) self.hf_device_map = None - def named_parameters(self, recurse=True, remove_duplicate=False): + def named_parameters(self, recurse = True, remove_duplicate = False): yield "embed.weight", self._shared if not remove_duplicate: yield "lm_head.weight", self._shared - def parameters(self, recurse=True): + def parameters(self, recurse = True): yield self._shared - def named_buffers(self, recurse=True): + def named_buffers(self, recurse = True): return iter([]) def named_children(self): diff --git a/test_infer_xpu_device_type_recursion.py b/test_infer_xpu_device_type_recursion.py index 1e319b97e8..63e4d17ad9 100644 --- a/test_infer_xpu_device_type_recursion.py +++ b/test_infer_xpu_device_type_recursion.py @@ -9,7 +9,9 @@ def _find_vision(): for p in [ Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py", Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py", - Path("/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py"), + Path( + "/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py" + ), ]: if p.exists(): return p @@ -24,8 +26,17 @@ def _load_fns(): "_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"] + 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: @@ -34,30 +45,32 @@ class _P: class _FakeMod: - def __init__(self, params=None, buffers=None, children=None, hf_device_map=None): + def __init__(self, params = None, buffers = None, children = None, hf_device_map = None): self._p = list(params or []) self._b = list(buffers or []) self._c = list(children or []) self.hf_device_map = hf_device_map - def named_parameters(self, recurse=True, remove_duplicate=False): + def named_parameters(self, recurse = True, remove_duplicate = False): for n, d in self._p: yield n, _P(d) if recurse: for cn, cm in self._c: - for pn, pp in cm.named_parameters(recurse=True, remove_duplicate=remove_duplicate): + for pn, pp in cm.named_parameters( + recurse = True, remove_duplicate = remove_duplicate + ): yield f"{cn}.{pn}", pp - def parameters(self, recurse=True): - for _, p in self.named_parameters(recurse=recurse): + def parameters(self, recurse = True): + for _, p in self.named_parameters(recurse = recurse): yield p - def named_buffers(self, recurse=True): + 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): + for bn, bb in cm.named_buffers(recurse = True): yield f"{cn}.{bn}", bb def named_children(self): @@ -70,9 +83,9 @@ def test_infer_handles_xpu_device_recursion(): indices as distinct devices for map-building purposes.""" infer, _ = _load_fns() # xpu:0 and xpu:1 — avoids any cuda-specific codepath in infer - a = _FakeMod(params=[("w", torch.device("xpu", 0))]) - b = _FakeMod(params=[("w", torch.device("xpu", 1))]) - root = _FakeMod(children=[("a", a), ("b", b)]) + a = _FakeMod(params = [("w", torch.device("xpu", 0))]) + b = _FakeMod(params = [("w", torch.device("xpu", 1))]) + root = _FakeMod(children = [("a", a), ("b", b)]) dm = infer(root) assert dm.get("a") == torch.device("xpu", 0) assert dm.get("b") == torch.device("xpu", 1) diff --git a/test_vision_fastbasemodel_calls_helper.py b/test_vision_fastbasemodel_calls_helper.py index c089c5383b..e23c75e482 100644 --- a/test_vision_fastbasemodel_calls_helper.py +++ b/test_vision_fastbasemodel_calls_helper.py @@ -6,7 +6,9 @@ def _find_vision(): for p in [ Path(__file__).resolve().parent / "unsloth" / "models" / "vision.py", Path(__file__).resolve().parents[1] / "unsloth" / "models" / "vision.py", - Path("/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py"), + Path( + "/mnt/disks/unslothai/ubuntu/workspace_25/github_review/unsloth-pr-5053-staging-3/unsloth/models/vision.py" + ), ]: if p.exists(): return p @@ -29,8 +31,11 @@ def test_vision_fastbasemodel_from_pretrained_calls_helper(): for node in ast.walk(fn): if ( isinstance(node, ast.Call) - and getattr(node.func, "id", None) == "_attach_bnb_multidevice_hooks" + and getattr(node.func, "id", None) + == "_attach_bnb_multidevice_hooks" ): found = True break - assert found, "FastBaseModel.from_pretrained must call _attach_bnb_multidevice_hooks" + assert ( + found + ), "FastBaseModel.from_pretrained must call _attach_bnb_multidevice_hooks" diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index 030db49168..61e0822aa5 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -115,7 +115,9 @@ def _infer_device_map_from_loaded_model(model): device_map[prefix] = next(iter(buf_devs)) else: for child_name, child in module.named_children(): - child_prefix = f"{prefix}.{child_name}" if prefix else child_name + child_prefix = ( + f"{prefix}.{child_name}" if prefix else child_name + ) _assign(child, child_prefix) return if len(subtree_devs) == 1: @@ -125,8 +127,10 @@ def _infer_device_map_from_loaded_model(model): child_prefix = f"{prefix}.{child_name}" if prefix else child_name _assign(child, child_prefix) local_devs = { - p.device for _, p in module.named_parameters( - recurse = False, remove_duplicate = False, + p.device + for _, p in module.named_parameters( + recurse = False, + remove_duplicate = False, ) } if local_devs and len(local_devs) == 1: