Add regression tests for the stray-forward compile-cache reset - #6569
Conversation
Follow-up to #6511, which fixed the bug but whose squash merge did not include the tests. These cover the two issues that fix addressed, under the GPU-free tests/conftest.py harness: - _unsloth_reset_stray_compile_cache is an exported module-level symbol in unsloth.models._utils (it previously lived only inside the RL trainer template string, so every non-RL import silently no-op'd) - _unsloth_install_pretrain_detector keeps a recorded "seen" forward on an idempotent reinstall with a live hook, and only resets it after teardown - only a grad-enabled pre-train forward marks the cache poisoned - the reset warns and clears seen when a stray forward was seen, tears the hook down even on the clean path, and walks the .model/.base_model/.module wrapper chain to reach a nested marker
for more information, see https://pre-commit.ci
There was a problem hiding this comment.
Code Review
This pull request adds a new test file tests/test_pretrain_compile_reset.py to verify the stray-pre-train-forward detector and its torch.compile cache reset behavior. The review feedback points out that test_reset_clears_seen_and_warns_when_a_stray_forward_was_seen could fail if UNSLOTH_COMPILE_DISABLE is set to "1" in the test environment, and suggests using monkeypatch to explicitly set it to "0" during the test.
Important
The consumer version of Gemini Code Assist on GitHub is being sunset. Starting June 18, 2026, new organization installations will be blocked, and all code review activity will officially cease on July 17, 2026.
For more details on the timeline and next steps, please review the Help Documentation.
| m = torch.nn.Linear(2, 2) | ||
| _unsloth_install_pretrain_detector(m) | ||
| m._unsloth_pretrain_marker["seen"] = True # a stray pre-train forward | ||
| trainer = _Trainer() | ||
| trainer.model = m | ||
|
|
||
| with warnings.catch_warnings(record = True) as caught: | ||
| warnings.simplefilter("always") | ||
| _unsloth_reset_stray_compile_cache(trainer) | ||
|
|
There was a problem hiding this comment.
If UNSLOTH_COMPILE_DISABLE is set to "1" in the test environment (which is common for GPU-free test suites to avoid compilation overhead or errors), _unsloth_reset_stray_compile_cache will skip raising the warning. This will cause test_reset_clears_seen_and_warns_when_a_stray_forward_was_seen to fail because caught will be empty.
To make this test robust against environment settings, use pytest's monkeypatch fixture to temporarily set UNSLOTH_COMPILE_DISABLE to "0" during the test.
| m = torch.nn.Linear(2, 2) | |
| _unsloth_install_pretrain_detector(m) | |
| m._unsloth_pretrain_marker["seen"] = True # a stray pre-train forward | |
| trainer = _Trainer() | |
| trainer.model = m | |
| with warnings.catch_warnings(record = True) as caught: | |
| warnings.simplefilter("always") | |
| _unsloth_reset_stray_compile_cache(trainer) | |
| def test_reset_clears_seen_and_warns_when_a_stray_forward_was_seen(monkeypatch): | |
| monkeypatch.setenv("UNSLOTH_COMPILE_DISABLE", "0") | |
| m = torch.nn.Linear(2, 2) | |
| _unsloth_install_pretrain_detector(m) | |
| m._unsloth_pretrain_marker["seen"] = True # a stray pre-train forward | |
| trainer = _Trainer() | |
| trainer.model = m | |
| with warnings.catch_warnings(record=True) as caught: | |
| warnings.simplefilter("always") | |
| _unsloth_reset_stray_compile_cache(trainer) |
The reset only warns and resets Dynamo when UNSLOTH_COMPILE_DISABLE != "1". A GPU-free CI env that sets it to "1" would make the warn assertion in test_reset_clears_seen_and_warns_when_a_stray_forward_was_seen flaky. monkeypatch it to "0" in both warn-path tests so the warn / no-warn assertions are deterministic and test the seen flag, not the env.
|
Good catch, addressed. Both warn-path tests now
Verified by running the suite with @codex review |
|
Codex Review: Didn't find any major issues. Delightful! Reviewed commit: ℹ️ About Codex in GitHubYour team has set up Codex to review pull requests in this repo. Reviews are triggered when you
If Codex has suggestions, it will comment; otherwise it will react with 👍. Codex can also answer questions or update the PR. Try commenting "@codex address that feedback". |
Follow-up to #6511. That PR fixed the stray-forward torch.compile cache poisoning, but the squash merge captured the branch before the test commit, so the regression tests never reached main. This adds them back, test-only, no source changes.
tests/test_pretrain_compile_reset.pyruns under the GPU-freetests/conftest.pyharness and guards the two issues #6511 addressed:_unsloth_reset_stray_compile_cacheis an exported module-level symbol inunsloth.models._utils. It previously existed only inside theRLTrainer_replacementtemplate string inrl.py, sofrom unsloth.models.rl import _unsloth_reset_stray_compile_cacheraised ImportError and every non-RL consumer silently no-op'd. This pins it as importable and in__all__._unsloth_install_pretrain_detectorkeeps a recorded "seen" forward on an idempotent reinstall with a live hook, and only resets it after teardown.seenwhen a stray forward was seen, tears the detector hook down even on the clean path, and walks the.model/.base_model/.modulewrapper chain to reach a nested marker.All 8 tests pass against current main.
@codex review