Fix nested distillation forward contexts - #2653
Michael-RDev wants to merge 1 commit into
Conversation
|
Navigate logical layers of code changes, visualize relationships, and explore their blast radius. No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (2)
Included review availability: This review used your included allowance. Your plan provides up to 12 included reviews per hour; 10 remain after this review. 📝 WalkthroughWalkthroughThe student-only and teacher-only forward contexts now restore their prior mode flags when they exit. Tests cover nested contexts, disabled contexts, and exceptions. ChangesNested forward contexts
Priority: ⬇️ Low Estimated code review effort: 2 (Simple) | ~10 minutes Change: Bug fix Merge Risk: ⚪ Minimal · up to Nested distillation contexts now restore the enclosing forward mode, preventing inner contexts from affecting later forwards. The implementation and regression coverage align; no actionable merge-blocking risk remains. 🚥 Pre-merge checks | ✅ 5 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (5 passed)
✨ Finishing Touches🧪 Generate unit tests (beta)
Comment |
- Restore the previous teacher/student forward flags on context exit so nested and disabled contexts preserve the enclosing mode - Add regression tests for nesting, exceptions, and overlapping contexts Signed-off-by: Michael Rusu <mi660386@ucf.edu>
c25fd72 to
5c33d27
Compare
What does this PR do?
Type of change: Bug fix
I found that nesting
only_student_forward()oronly_teacher_forward()causes the outer context to stop working when the inner one exits. This also happens when the inner context usesenable=False.Both context managers reset their flags to
Falseon exit. This change saves the previous value and restores it infinally, so the enclosing context keeps its execution mode, including after an exception.Public APIs, checkpoint handling, and loss calculations are unchanged.
Testing
Added eight regression cases using small local models and real forward passes. They cover:
All eight cases fail before the fix and pass afterward.
Verified with CPU PyTorch:
GPU and distributed integration tests were not run.
Before this PR is ready for review
Summary by CodeRabbit