From 432864cb25f128be6bd234f97d6eb82d1a497791 Mon Sep 17 00:00:00 2001 From: Duc-Viet Hoang Date: Mon, 12 Jan 2026 10:03:54 +0700 Subject: [PATCH] Complete disable `gradient_checkpointing` for vision when `use_gradient_checkpointing=False` --- unsloth/models/rl.py | 6 ++++-- unsloth/models/vision.py | 2 +- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index e945a80354..9ec4d76ed3 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -264,17 +264,19 @@ def prepare_for_training_mode(f): def wrapper(self, *args, **kwargs): # Enable training mode _was_training = None + # Get gradient checkpointing setting from training arguments + use_gc = getattr(self.args, 'gradient_checkpointing', True) if hasattr(self, 'model') and hasattr(self.model, "training"): _was_training = self.model.training if hasattr(self, 'model') and hasattr(self.model, "for_training"): - self.model.for_training() + self.model.for_training(use_gradient_checkpointing=use_gc) output = f(self, *args, **kwargs) # Restore previous mode when possible if hasattr(self, 'model') and hasattr(self.model, "for_inference"): if _was_training is False: self.model.for_inference() elif _was_training is True and hasattr(self.model, "for_training"): - self.model.for_training() + self.model.for_training(use_gradient_checkpointing=use_gc) # Reset gradient checkpointing buffers to free memory while staying ready for next run try: reset_unsloth_gradient_checkpointing_buffers() diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index 6de942d7d2..4e03e0a168 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -1273,7 +1273,7 @@ class FastBaseModel: # Since transformers 4.53, must turn on explicitly for module in model.modules(): if hasattr(module, "gradient_checkpointing"): - module.gradient_checkpointing = True + module.gradient_checkpointing = use_gradient_checkpointing # Also re-enable training for embeddings for NEFTune if hasattr(model, "get_input_embeddings"):