diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 1aa2c6e820..25c83af852 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -63,6 +63,7 @@ __all__ = [ "unpatch_unsloth_smart_gradient_checkpointing", "apply_unsloth_gradient_checkpointing", "_unsloth_install_pretrain_detector", + "_unsloth_flag_pretrain_generation", "_unsloth_reset_stray_compile_cache", "patch_compiled_autograd", "process_vision_info", @@ -314,23 +315,58 @@ def _unsloth_install_pretrain_detector(model): return model +def _unsloth_flag_pretrain_generation(model): + """Record that a torch.compile generation ran before trainer.train(). + + model.generate() does not trip the pre-train forward hook: generation drives the inner + model, not the wrapped forward the hook sits on, so the hook never fires (and generation + is no-grad, which _mark ignores anyway). But a no-grad generate still specializes the + fullgraph MoE kernels (e.g. gpt-oss); reusing those inference graphs hard-fails the + training recompile (FailOnRecompileLimitHit under fullgraph=True). Set a persistent marker + so the next prepare_for_training_mode drops the cache. Walk the wrapper chain since the + marker can sit on a different node than the one for_inference is called on. No-op if no + detector was installed (e.g. an untrained model).""" + _curr = model + _visited = set() + while _curr is not None and id(_curr) not in _visited: + _visited.add(id(_curr)) + _m = getattr(_curr, "_unsloth_pretrain_marker", None) + if isinstance(_m, dict): + _m["generation"] = True + # Follow the wrapper chain: Unsloth/HF (.model), PEFT (.base_model), DDP/FSDP (.module). + _nxt = getattr(_curr, "model", None) + if _nxt is None: + _nxt = getattr(_curr, "base_model", None) + if _nxt is None: + _nxt = getattr(_curr, "module", None) + _curr = _nxt + return model + + def _unsloth_reset_stray_compile_cache(self): - # A manual forward/backward under torch.compile BEFORE trainer.train() (e.g. a grad-norm - # probe) caches a forward + AOTAutograd backward graph in a one-off context; reusing it - # poisons training with NaN/zero gradients. If such a forward was seen and compile is on, - # drop the compiled-graph cache so training recompiles cleanly. No-op on the normal path. - # Module-level (not just inside the RL trainer template) so the SFT auto-packing wrapper and - # the plain-Trainer loop can import and run it too. + # A forward under torch.compile BEFORE trainer.train() specializes graphs that poison + # training when reused. Two flavours, both reset here: + # * "seen" - a grad-enabled manual forward/backward (e.g. a grad-norm probe) caches + # a forward + AOTAutograd backward graph, giving NaN/zero gradients. + # * "generation" - a model.generate() (even no-grad) specializes the fullgraph MoE kernels + # (e.g. gpt-oss); the training recompile then hard-fails with + # FailOnRecompileLimitHit. The forward hook cannot catch this (generate + # drives the inner model), so for_inference flags it via + # _unsloth_flag_pretrain_generation. + # If either was seen and compile is on, drop the compiled-graph cache so training recompiles + # cleanly. No-op on the normal path. Module-level (not just inside the RL trainer template) + # so the SFT auto-packing wrapper and the plain-Trainer loop can import and run it too. import os model = getattr(self, "model", None) if model is None: return - # The detector hook can sit on any wrapper in the chain, and the probe may have run on a - # different one than self.model, so walk the chain: detect a "seen" marker anywhere and - # collect every marker to tear down below. + # The detector marker can sit on any wrapper in the chain, and the probe/generation may have + # run on a different node than self.model, so walk the chain: detect a flagged marker + # anywhere and collect every marker to tear down below. markers = [] seen = False + generation = False _curr = model _visited = set() while _curr is not None and id(_curr) not in _visited: @@ -340,6 +376,8 @@ def _unsloth_reset_stray_compile_cache(self): markers.append(_m) if _m.get("seen"): seen = True + if _m.get("generation"): + generation = True # Follow the wrapper chain: Unsloth/HF (.model), PEFT (.base_model), DDP/FSDP (.module). _nxt = getattr(_curr, "model", None) if _nxt is None: @@ -347,7 +385,7 @@ def _unsloth_reset_stray_compile_cache(self): if _nxt is None: _nxt = getattr(_curr, "module", None) _curr = _nxt - if seen and os.environ.get("UNSLOTH_COMPILE_DISABLE", "0") != "1": + if (seen or generation) and os.environ.get("UNSLOTH_COMPILE_DISABLE", "0") != "1": try: import torch._dynamo as _dynamo _dynamo.reset() @@ -364,13 +402,15 @@ def _unsloth_reset_stray_compile_cache(self): model.zero_grad(set_to_none = True) except Exception: pass - import warnings - - warnings.warn( - "Unsloth: detected a manual forward/backward run before trainer.train(); " - "reset the torch.compile graph cache it poisoned so training starts clean. " - "To avoid this, run any pre-train probe under `with torch.no_grad():`." - ) + # Only the grad-enabled probe is a likely mistake worth flagging; a pre-train generate is + # a normal pattern, so reset its cache silently. + if seen: + import warnings + warnings.warn( + "Unsloth: detected a manual forward/backward run before trainer.train(); " + "reset the torch.compile graph cache it poisoned so training starts clean. " + "To avoid this, run any pre-train probe under `with torch.no_grad():`." + ) # Tear down every one-shot detector hook in the chain so none adds per-step cost. for _m in markers: hook = _m.pop("hook", None) @@ -380,6 +420,7 @@ def _unsloth_reset_stray_compile_cache(self): except Exception: pass _m["seen"] = False + _m["generation"] = False def apply_unsloth_gradient_checkpointing(use_gradient_checkpointing, max_seq_length, dtype): diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index a1da099758..e7e871805b 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -3770,6 +3770,9 @@ class FastLlamaModel: restore_use_cache(model) except ImportError: pass + # Flag that a compiled generation ran, so the next trainer.train() drops the + # torch.compile cache it specialized before it hard-fails the training recompile. + _unsloth_flag_pretrain_generation(model) return model @staticmethod diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index 0a68a49fee..1e703ec86f 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -2113,6 +2113,9 @@ class FastBaseModel: # Turn off skip guards and set stance to default if torch_compiler_set_stance is not None: torch_compiler_set_stance(stance = "default", skip_guard_eval_unsafe = False) + # Flag that a compiled generation ran, so the next trainer.train() drops the + # torch.compile cache it specialized before it hard-fails the training recompile. + _unsloth_flag_pretrain_generation(model) return model @staticmethod