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

for more information, see https://pre-commit.ci
This commit is contained in:
pre-commit-ci[bot] 2026-01-05 06:56:57 +00:00
commit 739aa923fb

View file

@ -659,7 +659,15 @@ class FastSentenceTransformer(FastModel):
return modules, True
# Encoder model types that benefit from native torch.compile instead of Unsloth patching
ENCODER_MODEL_TYPES = {"mpnet", "bert", "distilbert", "roberta", "xlm-roberta", "albert", "electra"}
ENCODER_MODEL_TYPES = {
"mpnet",
"bert",
"distilbert",
"roberta",
"xlm-roberta",
"albert",
"electra",
}
@staticmethod
def from_pretrained(
@ -754,10 +762,11 @@ class FastSentenceTransformer(FastModel):
# NOTE: The old Unsloth path is BROKEN for encoder models with torch 2.9+ due to
# conflicting @torch.compile and @torch.compiler.disable decorators.
# Set UNSLOTH_COMPILE_DISABLE=1 to disable torch.compile and use the old path.
is_encoder_model = model_type.lower() in FastSentenceTransformer.ENCODER_MODEL_TYPES
is_encoder_model = (
model_type.lower() in FastSentenceTransformer.ENCODER_MODEL_TYPES
)
use_fast_encoder = os.environ.get("UNSLOTH_COMPILE_DISABLE", "0") != "1"
if use_fast_encoder and is_encoder_model:
# torch.compile mode: "reduce-overhead" is optimal for training
compile_mode = "reduce-overhead"
@ -776,7 +785,9 @@ class FastSentenceTransformer(FastModel):
supports_sdpa = False
if config is not None:
try:
model_class = _get_model_class(config, kwargs.get("auto_model", AutoModel)._model_mapping)
model_class = _get_model_class(
config, kwargs.get("auto_model", AutoModel)._model_mapping
)
supports_sdpa = getattr(model_class, "_supports_sdpa", False)
except:
pass
@ -791,13 +802,18 @@ class FastSentenceTransformer(FastModel):
# Print optimization status
sdpa_str = " + SDPA" if supports_sdpa else ""
if load_in_4bit:
print(f"Unsloth: Using fast encoder path for {model_type} with 4-bit quantization{sdpa_str}")
print(
f"Unsloth: Using fast encoder path for {model_type} with 4-bit quantization{sdpa_str}"
)
else:
print(f"Unsloth: Using fast encoder path for {model_type} (torch.compile{sdpa_str})")
print(
f"Unsloth: Using fast encoder path for {model_type} (torch.compile{sdpa_str})"
)
# Handle 4-bit quantization via BitsAndBytesConfig
if load_in_4bit:
from transformers import BitsAndBytesConfig
bnb_config = BitsAndBytesConfig(
load_in_4bit = True,
bnb_4bit_compute_dtype = dtype,
@ -811,7 +827,9 @@ class FastSentenceTransformer(FastModel):
# Handle gradient checkpointing - warn user it conflicts with torch.compile
_use_gc = use_gradient_checkpointing
if _use_gc and _use_gc != False:
print("Unsloth Warning: Gradient checkpointing is incompatible with torch.compile.")
print(
"Unsloth Warning: Gradient checkpointing is incompatible with torch.compile."
)
print("Disabling torch.compile to enable gradient checkpointing.")
compile_mode = None # Disable compilation
@ -850,7 +868,9 @@ class FastSentenceTransformer(FastModel):
tokenizer.save_pretrained(save_directory)
FastSentenceTransformer._add_unsloth_branding(save_directory)
st_model.save_pretrained_merged = types.MethodType(_save_pretrained_merged, st_model)
st_model.save_pretrained_merged = types.MethodType(
_save_pretrained_merged, st_model
)
def _push_to_hub_merged(self, repo_id, **push_kwargs):
hub_token = push_kwargs.get("token", None) or get_token()
@ -858,22 +878,37 @@ class FastSentenceTransformer(FastModel):
raise ValueError("No HF token provided")
api = HfApi(token = hub_token)
try:
api.create_repo(repo_id = repo_id, private = push_kwargs.get("private"), exist_ok = True, repo_type = "model")
api.create_repo(
repo_id = repo_id,
private = push_kwargs.get("private"),
exist_ok = True,
repo_type = "model",
)
except:
pass
FastSentenceTransformer._add_unsloth_tags(repo_id, hub_token)
with tempfile.TemporaryDirectory() as temp_dir:
self.save_pretrained_merged(temp_dir, **push_kwargs)
api.upload_folder(folder_path = temp_dir, repo_id = repo_id, commit_message = push_kwargs.get("commit_message", "Upload model"))
api.upload_folder(
folder_path = temp_dir,
repo_id = repo_id,
commit_message = push_kwargs.get(
"commit_message", "Upload model"
),
)
print(f"Unsloth: Pushed to https://huggingface.co/{repo_id}")
st_model.push_to_hub_merged = types.MethodType(_push_to_hub_merged, st_model)
st_model.push_to_hub_merged = types.MethodType(
_push_to_hub_merged, st_model
)
return st_model
# Warn if using 4-bit with encoder (slow due to dequantization overhead)
if is_encoder_model and load_in_4bit:
print("Unsloth Warning: 4-bit quantization adds ~2.3x overhead for encoder models.")
print(
"Unsloth Warning: 4-bit quantization adds ~2.3x overhead for encoder models."
)
print("Consider using load_in_16bit=True for better performance.")
# check if the model supports add_pooling_layer
@ -1111,8 +1146,11 @@ class FastSentenceTransformer(FastModel):
inner_model = transformer_module.auto_model
# Check if model is quantized (4-bit/8-bit)
is_quantized = getattr(inner_model, "is_quantized", False) or \
getattr(inner_model.config, "quantization_config", None) is not None
is_quantized = (
getattr(inner_model, "is_quantized", False)
or getattr(inner_model.config, "quantization_config", None)
is not None
)
# Track if gradient checkpointing was actually enabled
gc_enabled = False
@ -1120,7 +1158,12 @@ class FastSentenceTransformer(FastModel):
# Prepare for k-bit training if quantized
if is_quantized:
from peft import prepare_model_for_kbit_training
_gc_for_kbit = use_gradient_checkpointing if use_gradient_checkpointing else False
_gc_for_kbit = (
use_gradient_checkpointing
if use_gradient_checkpointing
else False
)
try:
inner_model = prepare_model_for_kbit_training(
inner_model,
@ -1131,12 +1174,16 @@ class FastSentenceTransformer(FastModel):
except ValueError as e:
if "does not support gradient checkpointing" in str(e):
# Model doesn't support gradient checkpointing, disable it
print(f"Unsloth Warning: {inner_model.__class__.__name__} does not support gradient checkpointing. Skipping.")
print(
f"Unsloth Warning: {inner_model.__class__.__name__} does not support gradient checkpointing. Skipping."
)
inner_model = prepare_model_for_kbit_training(
inner_model,
use_gradient_checkpointing = False,
)
print("Unsloth: Prepared quantized model for k-bit training (without gradient checkpointing)")
print(
"Unsloth: Prepared quantized model for k-bit training (without gradient checkpointing)"
)
else:
raise
@ -1149,7 +1196,9 @@ class FastSentenceTransformer(FastModel):
gc_enabled = True
except ValueError as e:
if "does not support gradient checkpointing" in str(e):
print(f"Unsloth Warning: {inner_model.__class__.__name__} does not support gradient checkpointing. Skipping.")
print(
f"Unsloth Warning: {inner_model.__class__.__name__} does not support gradient checkpointing. Skipping."
)
# Create LoRA config
lora_config = LoraConfig(
@ -1169,13 +1218,19 @@ class FastSentenceTransformer(FastModel):
# Re-enable torch.compile if gradient checkpointing was requested but couldn't be enabled
if compile_mode is None and not gc_enabled:
compile_mode = "reduce-overhead"
print("Unsloth: Re-enabling torch.compile since gradient checkpointing is not supported")
print(
"Unsloth: Re-enabling torch.compile since gradient checkpointing is not supported"
)
if compile_mode is not None:
print(f"Unsloth: Applying torch.compile with mode='{compile_mode}' for 6x speedup")
print(
f"Unsloth: Applying torch.compile with mode='{compile_mode}' for 6x speedup"
)
peft_model = torch.compile(peft_model, mode = compile_mode)
else:
print("Unsloth: torch.compile disabled (gradient checkpointing enabled)")
print(
"Unsloth: torch.compile disabled (gradient checkpointing enabled)"
)
# Re-assign the peft model back to the transformer module
transformer_module.auto_model = peft_model