Gemini code review suggestions

This commit is contained in:
electroglyph 2025-12-12 03:45:13 -08:00
commit 855f719e76

View file

@ -13,39 +13,42 @@
# limitations under the License.
from .loader import FastModel
import torch
import inspect
import json
import os
from huggingface_hub import hf_hub_download
class FastSentenceTransformer(FastModel):
@staticmethod
def from_pretrained(
model_name,
max_seq_length = None,
dtype = None,
load_in_4bit = True,
load_in_8bit = False,
load_in_16bit = False,
full_finetuning = False,
token = None,
device_map = "sequential",
rope_scaling = None,
fix_tokenizer = True,
trust_remote_code = False,
use_gradient_checkpointing = "unsloth",
resize_model_vocab = None,
revision = None,
use_exact_model_name = False,
offload_embedding = False,
random_state = 3407,
max_lora_rank = 64,
disable_log_stats = True,
qat_scheme = None,
load_in_fp8 = False,
unsloth_tiled_mlp = False,
pooling_mode = "mean",
max_seq_length=None,
dtype=None,
load_in_4bit=True,
load_in_8bit=False,
load_in_16bit=False,
full_finetuning=False,
token=None,
device_map="sequential",
rope_scaling=None,
fix_tokenizer=True,
trust_remote_code=False,
use_gradient_checkpointing="unsloth",
resize_model_vocab=None,
revision=None,
use_exact_model_name=False,
offload_embedding=False,
random_state=3407,
max_lora_rank=64,
disable_log_stats=True,
qat_scheme=None,
load_in_fp8=False,
unsloth_tiled_mlp=False,
pooling_mode="mean",
**kwargs,
):
try:
import sentence_transformers
from sentence_transformers import SentenceTransformer
from sentence_transformers.models import Transformer, Pooling, Normalize
from transformers import AutoModel
@ -62,46 +65,41 @@ class FastSentenceTransformer(FastModel):
kwargs["add_pooling_layer"] = False
model, tokenizer = FastModel.from_pretrained(
model_name = model_name,
max_seq_length = max_seq_length,
dtype = dtype,
load_in_4bit = load_in_4bit,
load_in_8bit = load_in_8bit,
load_in_16bit = load_in_16bit,
full_finetuning = full_finetuning,
token = token,
device_map = device_map,
rope_scaling = rope_scaling,
fix_tokenizer = fix_tokenizer,
trust_remote_code = trust_remote_code,
use_gradient_checkpointing = use_gradient_checkpointing,
resize_model_vocab = resize_model_vocab,
revision = revision,
return_logits = False,
use_exact_model_name = use_exact_model_name,
offload_embedding = offload_embedding,
random_state = random_state,
max_lora_rank = max_lora_rank,
disable_log_stats = disable_log_stats,
qat_scheme = qat_scheme,
load_in_fp8 = load_in_fp8,
unsloth_tiled_mlp = unsloth_tiled_mlp,
model_name=model_name,
max_seq_length=max_seq_length,
dtype=dtype,
load_in_4bit=load_in_4bit,
load_in_8bit=load_in_8bit,
load_in_16bit=load_in_16bit,
full_finetuning=full_finetuning,
token=token,
device_map=device_map,
rope_scaling=rope_scaling,
fix_tokenizer=fix_tokenizer,
trust_remote_code=trust_remote_code,
use_gradient_checkpointing=use_gradient_checkpointing,
resize_model_vocab=resize_model_vocab,
revision=revision,
return_logits=False,
use_exact_model_name=use_exact_model_name,
offload_embedding=offload_embedding,
random_state=random_state,
max_lora_rank=max_lora_rank,
disable_log_stats=disable_log_stats,
qat_scheme=qat_scheme,
load_in_fp8=load_in_fp8,
unsloth_tiled_mlp=unsloth_tiled_mlp,
**kwargs,
)
transformer_module = Transformer.__new__(Transformer)
import torch
torch.nn.Module.__init__(transformer_module)
transformer_module.auto_model = model
transformer_module.tokenizer = tokenizer
transformer_module.do_lower_case = False
if hasattr(tokenizer, "do_lower_case"):
transformer_module.do_lower_case = tokenizer.do_lower_case
import inspect
model_forward_params = list(inspect.signature(model.forward).parameters)
transformer_module.model_forward_params = set(model_forward_params) | {
"input_ids",
@ -137,17 +135,13 @@ class FastSentenceTransformer(FastModel):
# detect pooling mode if not specified/default
if pooling_mode == "mean":
try:
from huggingface_hub import hf_hub_download
import json
import os
if os.path.exists(model_name) and os.path.exists(
os.path.join(model_name, "modules.json")
):
modules_json_path = os.path.join(model_name, "modules.json")
else:
modules_json_path = hf_hub_download(
model_name, "modules.json", token = token
model_name, "modules.json", token=token
)
with open(modules_json_path, "r") as f:
@ -169,37 +163,24 @@ class FastSentenceTransformer(FastModel):
pooling_config_path = hf_hub_download(
model_name,
os.path.join(pooling_path, "config.json"),
token = token,
token=token,
)
break
if pooling_config_path:
with open(pooling_config_path, "r") as f:
pooling_config = json.load(f)
if (
"pooling_mode_cls_token" in pooling_config
and pooling_config["pooling_mode_cls_token"]
):
print("Pooling mode detected as cls, updating...")
pooling_mode = "cls"
elif (
"pooling_mode_mean_tokens" in pooling_config
and pooling_config["pooling_mode_mean_tokens"]
):
print("Pooling mode detected as mean, updating...")
pooling_mode = "mean"
elif (
"pooling_mode_max_tokens" in pooling_config
and pooling_config["pooling_mode_max_tokens"]
):
print("Pooling mode detected as max, updating...")
pooling_mode = "max"
elif (
"pooling_mode_mean_sqrt_len_tokens" in pooling_config
and pooling_config["pooling_mode_mean_sqrt_len_tokens"]
):
print("Pooling mode detected as mean_sqrt_len, updating...")
pooling_mode = "mean_sqrt_len"
pooling_map = {
"pooling_mode_cls_token": "cls",
"pooling_mode_mean_tokens": "mean",
"pooling_mode_max_tokens": "max",
"pooling_mode_mean_sqrt_len_tokens": "mean_sqrt_len",
}
for config_key, mode in pooling_map.items():
if pooling_config.get(config_key):
print(f"Pooling mode detected as {mode}, updating...")
pooling_mode = mode
break
except Exception as e:
print(
@ -207,36 +188,36 @@ class FastSentenceTransformer(FastModel):
)
pooling_module = Pooling(
word_embedding_dimension = hidden_size,
pooling_mode = pooling_mode,
word_embedding_dimension=hidden_size,
pooling_mode=pooling_mode,
)
normalize_module = Normalize()
modules = [transformer_module, pooling_module, normalize_module]
st_model = SentenceTransformer(modules = modules)
st_model = SentenceTransformer(modules=modules)
return st_model
@staticmethod
def get_peft_model(
model,
r = 16,
target_modules = [
r=16,
target_modules=[
"query",
"key",
"value",
"dense",
],
lora_alpha = 16,
lora_dropout = 0.0,
bias = "none",
layers_to_transform = None,
layers_pattern = None,
use_gradient_checkpointing = "unsloth",
random_state = 3407,
max_seq_length = 2048,
use_rslora = False,
modules_to_save = None,
init_lora_weights = True,
loftq_config = {},
lora_alpha=16,
lora_dropout=0.0,
bias="none",
layers_to_transform=None,
layers_pattern=None,
use_gradient_checkpointing="unsloth",
random_state=3407,
max_seq_length=2048,
use_rslora=False,
modules_to_save=None,
init_lora_weights=True,
loftq_config={},
**kwargs,
):
from sentence_transformers import SentenceTransformer
@ -251,21 +232,21 @@ class FastSentenceTransformer(FastModel):
inner_model = transformer_module.auto_model
peft_model = FastModel.get_peft_model(
model = inner_model,
r = r,
target_modules = target_modules,
lora_alpha = lora_alpha,
lora_dropout = lora_dropout,
bias = bias,
layers_to_transform = layers_to_transform,
layers_pattern = layers_pattern,
use_gradient_checkpointing = use_gradient_checkpointing,
random_state = random_state,
max_seq_length = max_seq_length,
use_rslora = use_rslora,
modules_to_save = modules_to_save,
init_lora_weights = init_lora_weights,
loftq_config = loftq_config,
model=inner_model,
r=r,
target_modules=target_modules,
lora_alpha=lora_alpha,
lora_dropout=lora_dropout,
bias=bias,
layers_to_transform=layers_to_transform,
layers_pattern=layers_pattern,
use_gradient_checkpointing=use_gradient_checkpointing,
random_state=random_state,
max_seq_length=max_seq_length,
use_rslora=use_rslora,
modules_to_save=modules_to_save,
init_lora_weights=init_lora_weights,
loftq_config=loftq_config,
**kwargs,
)
@ -274,20 +255,20 @@ class FastSentenceTransformer(FastModel):
return model
else:
return FastModel.get_peft_model(
model = model,
r = r,
target_modules = target_modules,
lora_alpha = lora_alpha,
lora_dropout = lora_dropout,
bias = bias,
layers_to_transform = layers_to_transform,
layers_pattern = layers_pattern,
use_gradient_checkpointing = use_gradient_checkpointing,
random_state = random_state,
max_seq_length = max_seq_length,
use_rslora = use_rslora,
modules_to_save = modules_to_save,
init_lora_weights = init_lora_weights,
loftq_config = loftq_config,
model=model,
r=r,
target_modules=target_modules,
lora_alpha=lora_alpha,
lora_dropout=lora_dropout,
bias=bias,
layers_to_transform=layers_to_transform,
layers_pattern=layers_pattern,
use_gradient_checkpointing=use_gradient_checkpointing,
random_state=random_state,
max_seq_length=max_seq_length,
use_rslora=use_rslora,
modules_to_save=modules_to_save,
init_lora_weights=init_lora_weights,
loftq_config=loftq_config,
**kwargs,
)