[FIX] [Transformers] VLM input embeds fix for gradients (#3715)

* Fix get_input_embeds call for VLMs

* patch input_require_grads instead

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* cleanup old patch

* cleanup old patch

* cleanup

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Apply suggestion from @danielhanchen

* use logger instead of prints

* Move unsloth present set

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Daniel Han <danielhanchen@gmail.com>
This commit is contained in:
Datta Nimmaturi 2025-12-12 17:03:39 +05:30 committed by GitHub
commit 10e8518ae3
2 changed files with 85 additions and 28 deletions

View file

@ -17,6 +17,9 @@ from packaging.version import Version
import os, re, subprocess, inspect, functools
import numpy as np
# Log Unsloth is being used
# We want logger in import_fixes and hence setting it here for zoo to be importable
os.environ["UNSLOTH_IS_PRESENT"] = "1"
# Fix some issues before importing other packages
from .import_fixes import (
fix_message_factory_issue,
@ -63,8 +66,6 @@ os.environ["PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION"] = "python"
# "pinned_use_cuda_host_register:True,"\
# "pinned_num_register_threads:8"
# Log Unsloth is being used
os.environ["UNSLOTH_IS_PRESENT"] = "1"
from importlib.metadata import version as importlib_version
from importlib.metadata import PackageNotFoundError
@ -123,6 +124,7 @@ from .import_fixes import (
patch_ipykernel_hf_xet,
patch_trackio,
patch_datasets,
patch_enable_input_require_grads,
)
fix_xformers_performance_issue()
@ -132,6 +134,7 @@ ignore_logger_messages()
patch_ipykernel_hf_xet()
patch_trackio()
patch_datasets()
patch_enable_input_require_grads()
del fix_xformers_performance_issue
del fix_vllm_aimv2_issue
@ -140,6 +143,7 @@ del ignore_logger_messages
del patch_ipykernel_hf_xet
del patch_trackio
del patch_datasets
del patch_enable_input_require_grads
# Torch 2.4 has including_emulation
if DEVICE_TYPE == "cuda":

View file

@ -19,8 +19,7 @@ from importlib.metadata import version as importlib_version
from packaging.version import Version as TrueVersion
import re
import logging
UNSLOTH_ENABLE_LOGGING = os.environ.get("UNSLOTH_ENABLE_LOGGING", "0") == "1"
from unsloth_zoo.log import logger
def Version(version):
@ -71,8 +70,7 @@ def fix_message_factory_issue():
return
if not hasattr(google.protobuf.message_factory, "MessageFactory"):
if UNSLOTH_ENABLE_LOGGING:
print("Unsloth: Patching protobuf.MessageFactory as it doesn't exist")
logger.info("Unsloth: Patching protobuf.MessageFactory as it doesn't exist")
google.protobuf.message_factory.MessageFactory = MessageFactory
elif (
hasattr(google.protobuf.message_factory, "MessageFactory")
@ -82,8 +80,7 @@ def fix_message_factory_issue():
and not hasattr(google.protobuf.message_factory, "GetMessageClass")
):
google.protobuf.message_factory.MessageFactory = MessageFactory
if UNSLOTH_ENABLE_LOGGING:
print("Unsloth: Patching protobuf.MessageFactory as it doesn't exist")
logger.info("Unsloth: Patching protobuf.MessageFactory as it doesn't exist")
elif (
hasattr(google.protobuf.message_factory, "MessageFactory")
and not hasattr(
@ -97,8 +94,7 @@ def fix_message_factory_issue():
return GetMessageClass(descriptor)
google.protobuf.message_factory.MessageFactory.GetPrototype = GetPrototype
if UNSLOTH_ENABLE_LOGGING:
print("Unsloth: Patching protobuf.MessageFactory.GetPrototype")
logger.info("Unsloth: Patching protobuf.MessageFactory.GetPrototype")
pass
except:
pass
@ -126,13 +122,11 @@ def fix_xformers_performance_issue():
f.seek(0)
f.write(text)
f.truncate()
if UNSLOTH_ENABLE_LOGGING:
print(
"Unsloth: Patching Xformers to fix some performance issues."
)
logger.info(
"Unsloth: Patching Xformers to fix some performance issues."
)
except Exception as e:
if UNSLOTH_ENABLE_LOGGING:
print(f"Unsloth: Failed patching Xformers with error = {str(e)}")
logger.info(f"Unsloth: Failed patching Xformers with error = {str(e)}")
# ValueError: 'aimv2' is already used by a Transformers config, pick another name.
@ -167,13 +161,11 @@ def fix_vllm_aimv2_issue():
f.seek(0)
f.write(text)
f.truncate()
if UNSLOTH_ENABLE_LOGGING:
print(
"Unsloth: Patching vLLM to fix `'aimv2' is already used by a Transformers config, pick another name.`"
)
logger.info(
"Unsloth: Patching vLLM to fix `'aimv2' is already used by a Transformers config, pick another name.`"
)
except Exception as e:
if UNSLOTH_ENABLE_LOGGING:
print(f"Unsloth: Failed patching vLLM with error = {str(e)}")
logger.info(f"Unsloth: Failed patching vLLM with error = {str(e)}")
def fix_vllm_guided_decoding_params():
@ -274,8 +266,70 @@ def check_fbgemm_gpu_version():
raise ImportError(
f"Unsloth: fbgemm_gpu_genai=={fbgemm_gpu_version} detected. It might cause unexpected issues like segmentation faults. Please uninstall the current one by doing `pip uninstall fbgemm-gpu` && `pip install fbgemm-gpu` to install fbgemm-gpu 1.4.0 or newer!"
)
elif UNSLOTH_ENABLE_LOGGING:
print(f"Unsloth: fbgemm_gpu_genai=={fbgemm_gpu_version} detected.")
logger.info(f"Unsloth: fbgemm_gpu_genai=={fbgemm_gpu_version} detected.")
def patch_enable_input_require_grads():
"""
Patch transformers PreTrainedModel.enable_input_require_grads to handle vision models
that raise NotImplementedError from get_input_embeddings().
"""
import inspect
from transformers import PreTrainedModel
# Check if the original function iterates over self.modules() instead of just returning the enable_input_require_grads
# Ref: https://github.com/huggingface/transformers/pull/41993/files#diff-6b72b98c4c2dcfc6cc606843917733f5d858374fbc22a735ff483bbc0c1e63eaL1979-R1996
try:
original_source = inspect.getsource(PreTrainedModel.enable_input_require_grads)
except (OSError, TypeError):
return
# Only patch if the new pattern exists (iterating over self.modules())
if "for module in self.modules()" not in original_source:
return
def _patched_enable_input_require_grads(self):
def make_inputs_require_grads(module, input, output):
output.requires_grad_(True)
hooks = []
seen_modules = set()
for module in self.modules():
if not (
isinstance(module, PreTrainedModel)
and hasattr(module, "get_input_embeddings")
):
continue
try:
input_embeddings = module.get_input_embeddings()
except NotImplementedError:
# Vision models may not implement get_input_embeddings - skip them
# For GLM V4.6 for example, this skips only `self.visual`
continue
if input_embeddings is None:
continue
embedding_id = id(input_embeddings)
if embedding_id in seen_modules:
continue
seen_modules.add(embedding_id)
hooks.append(
input_embeddings.register_forward_hook(make_inputs_require_grads)
)
self._require_grads_hooks = hooks
if hooks:
self._require_grads_hook = hooks[0]
PreTrainedModel.enable_input_require_grads = _patched_enable_input_require_grads
logger.info(
"Unsloth: Patched enable_input_require_grads for vision model compatibility"
)
def torchvision_compatibility_check():
@ -313,7 +367,6 @@ def torchvision_compatibility_check():
f"but found torchvision=={torchvision_version}. "
f"Please refer to https://pytorch.org/get-started/previous-versions/ for more information."
)
elif UNSLOTH_ENABLE_LOGGING:
print(
f"Unsloth: torch=={torch_version} and torchvision=={torchvision_version} are compatible."
)
logger.info(
f"Unsloth: torch=={torch_version} and torchvision=={torchvision_version} are compatible."
)