diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 6383e14ac9..fbdfb84602 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -196,18 +196,15 @@ from unsloth_zoo.temporary_patches import ( def _unsloth_install_pretrain_detector(model): - """Attach a one-shot forward pre-hook that records whether a forward ran - before trainer.train(). Used by prepare_for_training_mode to drop a - torch.compile graph cache poisoned by a stray manual forward/backward. - Idempotent and a no-op if the model cannot take hooks.""" + """Attach a one-shot forward pre-hook recording whether a forward ran before + trainer.train(), so prepare_for_training_mode can drop a torch.compile graph cache poisoned + by a stray manual forward/backward. Idempotent; no-op if the model cannot take hooks.""" if model is None or not hasattr(model, "register_forward_pre_hook"): return model marker = getattr(model, "_unsloth_pretrain_marker", None) if isinstance(marker, dict): marker["seen"] = False - # Re-register only if the previous hook was torn down (e.g. by an earlier - # train()); if the hook is still live this is a strict no-op so we never - # stack duplicate hooks. + # Re-register only if the previous hook was torn down; a live hook stays (no duplicates). if "hook" in marker: return model else: @@ -218,10 +215,8 @@ def _unsloth_install_pretrain_detector(model): return model def _mark(_module, _inp): - # Only a GRAD-ENABLED forward can poison the AOTAutograd/torch.compile - # backward-graph cache. A no-grad probe (`with torch.no_grad(): model(...)`, - # the sanity check the warning recommends) builds no backward graph, so - # treat it as clean and avoid a needless dynamo reset + recompile + warning. + # Only a grad-enabled forward poisons the AOTAutograd backward-graph cache; a no-grad + # probe builds no backward graph, so treat it as clean (avoids a needless dynamo reset). if torch.is_grad_enabled(): marker["seen"] = True diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index f399a7b3e8..1a00f21416 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -391,21 +391,17 @@ try: except: def reset_unsloth_gradient_checkpointing_buffers(): pass def _unsloth_reset_stray_compile_cache(self): - # A manual forward / forward+backward run under torch.compile BEFORE - # trainer.train() (e.g. a pre-train grad-norm probe) compiles and caches the - # model's forward and (via AOTAutograd) its backward graph in a one-off - # context that does not match the training loop. Reusing that cached graph - # poisons training with NaN/zero gradients (loss never moves). If a pre-train - # forward was seen and torch.compile is enabled, drop the compiled-graph cache - # so training recompiles cleanly. No-op on the normal path. + # 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. import os model = getattr(self, "model", None) if model is None: return - # The detector hook may sit on any wrapper in the chain (PeftModel / DDP / - # the base model), and a pre-train probe could have run on a different - # wrapper than self.model. Walk the chain so a "seen" marker anywhere is - # detected, and collect every marker so all hooks are torn down below. + # 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. markers = [] seen = False _curr = model @@ -417,9 +413,7 @@ def _unsloth_reset_stray_compile_cache(self): markers.append(_m) if _m.get("seen"): seen = True - # Follow the wrapper chain: Unsloth/HF (.model), PEFT (.base_model) and - # DDP / FSDP (.module). A pre-train probe can fire on the model below a - # DDP wrapper, so .module must be walked too or the marker is missed. + # 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)