chore: import upstream snapshot with attribution
This commit is contained in:
@@ -0,0 +1,11 @@
|
||||
from .modeling import CrossDecoderModel
|
||||
from .runner import DecoderOnlyRerankerRunner
|
||||
from .arguments import RerankerModelArguments
|
||||
from .trainer import DecoderOnlyRerankerTrainer
|
||||
|
||||
__all__ = [
|
||||
"CrossDecoderModel",
|
||||
"DecoderOnlyRerankerRunner",
|
||||
"DecoderOnlyRerankerTrainer",
|
||||
"RerankerModelArguments",
|
||||
]
|
||||
@@ -0,0 +1,30 @@
|
||||
from transformers import HfArgumentParser
|
||||
|
||||
from FlagEmbedding.abc.finetune.reranker import (
|
||||
AbsRerankerDataArguments,
|
||||
AbsRerankerTrainingArguments
|
||||
)
|
||||
|
||||
from FlagEmbedding.finetune.reranker.decoder_only.base import (
|
||||
DecoderOnlyRerankerRunner,
|
||||
RerankerModelArguments
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
parser = HfArgumentParser((RerankerModelArguments, AbsRerankerDataArguments, AbsRerankerTrainingArguments))
|
||||
model_args, data_args, training_args = parser.parse_args_into_dataclasses()
|
||||
model_args: RerankerModelArguments
|
||||
data_args: AbsRerankerDataArguments
|
||||
training_args: AbsRerankerTrainingArguments
|
||||
|
||||
runner = DecoderOnlyRerankerRunner(
|
||||
model_args=model_args,
|
||||
data_args=data_args,
|
||||
training_args=training_args
|
||||
)
|
||||
runner.run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,58 @@
|
||||
from typing import List
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from FlagEmbedding.abc.finetune.reranker import AbsRerankerModelArguments
|
||||
|
||||
|
||||
def default_target_modules() -> List[int]:
|
||||
return ['v_proj', 'q_proj', 'k_proj', 'gate_proj', 'down_proj', 'o_proj', 'up_proj']
|
||||
|
||||
|
||||
@dataclass
|
||||
class RerankerModelArguments(AbsRerankerModelArguments):
|
||||
"""
|
||||
Model argument class for decoder only reranker.
|
||||
"""
|
||||
use_lora: bool = field(
|
||||
default=True,
|
||||
metadata={"help": "If passed, will use LORA (low-rank parameter-efficient training) to train the model."}
|
||||
)
|
||||
lora_rank: int = field(
|
||||
default=64,
|
||||
metadata={"help": "The rank of lora."}
|
||||
)
|
||||
lora_alpha: float = field(
|
||||
default=16,
|
||||
metadata={"help": "The alpha parameter of lora."}
|
||||
)
|
||||
lora_dropout: float = field(
|
||||
default=0.1,
|
||||
metadata={"help": "The dropout rate of lora modules."}
|
||||
)
|
||||
target_modules: List[str] = field(
|
||||
default_factory=default_target_modules,
|
||||
metadata={"help": "The target modules to apply LORA."}
|
||||
)
|
||||
modules_to_save: List[str] = field(
|
||||
default=None,
|
||||
metadata={"help": "List of modules that should be saved in the final checkpoint."}
|
||||
)
|
||||
use_flash_attn: bool = field(
|
||||
default=False,
|
||||
metadata={"help": "If passed, will use flash attention to train the model."}
|
||||
)
|
||||
# use_slow_tokenizer: bool = field(
|
||||
# default=False,
|
||||
# metadata={"help": "If passed, will use a slow tokenizer (not backed by the 🤗 Tokenizers library)."}
|
||||
# )
|
||||
from_peft: str = field(
|
||||
default=None
|
||||
)
|
||||
raw_peft: List[str] = field(
|
||||
default=None
|
||||
)
|
||||
|
||||
save_merged_lora_model: bool = field(
|
||||
default=False,
|
||||
metadata={"help": "If passed, will merge the lora modules and save the entire model."}
|
||||
)
|
||||
@@ -0,0 +1,168 @@
|
||||
import os
|
||||
import re
|
||||
import logging
|
||||
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
|
||||
from peft import LoraConfig, TaskType, get_peft_model, PeftModel
|
||||
|
||||
from FlagEmbedding.finetune.reranker.decoder_only.base.arguments import RerankerModelArguments
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def find_largest_checkpoint(checkpoint_dir):
|
||||
"""Find the largest checkpoint from directory.
|
||||
|
||||
Args:
|
||||
checkpoint_dir (str): Directory to the checkpoint.
|
||||
|
||||
Returns:
|
||||
str: Directory to the checkpoint, None no matching found.
|
||||
"""
|
||||
checkpoint_pattern = re.compile(r'checkpoint-(\d+)')
|
||||
max_number = -1
|
||||
max_checkpoint_file = None
|
||||
for file in os.listdir(checkpoint_dir):
|
||||
match = checkpoint_pattern.search(file)
|
||||
if match:
|
||||
number = int(match.group(1))
|
||||
if number > max_number:
|
||||
max_number = number
|
||||
max_checkpoint_file = file
|
||||
if max_checkpoint_file:
|
||||
return os.path.join(checkpoint_dir, max_checkpoint_file)
|
||||
else:
|
||||
return None
|
||||
|
||||
|
||||
def get_model(model_args: RerankerModelArguments):
|
||||
"""Get the model.
|
||||
|
||||
Args:
|
||||
model_args (RerankerModelArguments): Model arguments instance.
|
||||
|
||||
Returns:
|
||||
transformers.PreTrainedModel or PeftModel: The loaded model.
|
||||
"""
|
||||
if model_args.config_name:
|
||||
config = AutoConfig.from_pretrained(
|
||||
model_args.config_name,
|
||||
trust_remote_code=model_args.trust_remote_code,
|
||||
token=model_args.token,
|
||||
cache_dir=model_args.cache_dir
|
||||
)
|
||||
elif model_args.model_name_or_path:
|
||||
config = AutoConfig.from_pretrained(
|
||||
model_args.model_name_or_path,
|
||||
trust_remote_code=model_args.trust_remote_code,
|
||||
token=model_args.token,
|
||||
cache_dir=model_args.cache_dir
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"You are instantiating a new config instance from scratch. This is not supported by this script."
|
||||
)
|
||||
config.use_cache = False
|
||||
|
||||
if model_args.model_name_or_path:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
model_args.model_name_or_path,
|
||||
# torch_dtype=torch.bfloat16,
|
||||
attn_implementation = "flash_attention_2" if model_args.use_flash_attn else None,
|
||||
token=model_args.token,
|
||||
cache_dir=model_args.cache_dir,
|
||||
from_tf=bool(".ckpt" in model_args.model_name_or_path),
|
||||
config=config,
|
||||
trust_remote_code=model_args.trust_remote_code,
|
||||
)
|
||||
else:
|
||||
logger.info("Training new model from scratch")
|
||||
model = model_args.from_config(config)
|
||||
|
||||
if model_args.raw_peft is not None:
|
||||
for peft_path in model_args.raw_peft:
|
||||
model = PeftModel.from_pretrained(model, peft_path)
|
||||
model = model.merge_and_unload()
|
||||
|
||||
if model_args.from_peft is not None:
|
||||
model = PeftModel.from_pretrained(model, model_args.from_peft, is_trainable=True)
|
||||
model.print_trainable_parameters()
|
||||
else:
|
||||
if model_args.use_lora:
|
||||
peft_config = LoraConfig(
|
||||
task_type=TaskType.CAUSAL_LM,
|
||||
inference_mode=False,
|
||||
r=model_args.lora_rank,
|
||||
target_modules=model_args.target_modules,
|
||||
modules_to_save=model_args.modules_to_save,
|
||||
lora_alpha=model_args.lora_alpha,
|
||||
lora_dropout=model_args.lora_dropout
|
||||
)
|
||||
model = get_peft_model(model, peft_config)
|
||||
model.print_trainable_parameters()
|
||||
|
||||
return model
|
||||
|
||||
|
||||
def save_merged_model(model_args: RerankerModelArguments, output_dir: str):
|
||||
"""
|
||||
Loads and save a model with specified configurations, merges it with PEFT layers if available.
|
||||
|
||||
Args:
|
||||
model_args (RerankerModelArguments): Model arguments instance.
|
||||
output_dir (str): Directory to save the model.
|
||||
"""
|
||||
if model_args.config_name:
|
||||
config = AutoConfig.from_pretrained(
|
||||
model_args.config_name,
|
||||
token=model_args.token,
|
||||
trust_remote_code=model_args.trust_remote_code,
|
||||
cache_dir=model_args.cache_dir
|
||||
)
|
||||
elif model_args.model_name_or_path:
|
||||
config = AutoConfig.from_pretrained(
|
||||
model_args.model_name_or_path,
|
||||
token=model_args.token,
|
||||
trust_remote_code=model_args.trust_remote_code,
|
||||
cache_dir=model_args.cache_dir
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"You are instantiating a new config instance from scratch. This is not supported by this script."
|
||||
)
|
||||
config.use_cache = False
|
||||
|
||||
if model_args.model_name_or_path:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
model_args.model_name_or_path,
|
||||
# torch_dtype=torch.bfloat16,
|
||||
attn_implementation = "flash_attention_2" if model_args.use_flash_attn else None,
|
||||
token=model_args.token,
|
||||
cache_dir=model_args.cache_dir,
|
||||
from_tf=bool(".ckpt" in model_args.model_name_or_path),
|
||||
trust_remote_code=model_args.trust_remote_code,
|
||||
config=config,
|
||||
)
|
||||
else:
|
||||
logger.info("Training new model from scratch")
|
||||
model = model_args.from_config(config)
|
||||
|
||||
if model_args.raw_peft is not None:
|
||||
for peft_path in model_args.raw_peft:
|
||||
model = PeftModel.from_pretrained(model, peft_path)
|
||||
model = model.merge_and_unload()
|
||||
|
||||
try:
|
||||
model = PeftModel.from_pretrained(model, output_dir)
|
||||
model = model.merge_and_unload()
|
||||
except:
|
||||
model = PeftModel.from_pretrained(model, find_largest_checkpoint(output_dir))
|
||||
model = model.merge_and_unload()
|
||||
|
||||
model.save_pretrained(os.path.join(output_dir, 'merged_model'))
|
||||
|
||||
try:
|
||||
tokenizer = AutoTokenizer.from_pretrained(output_dir)
|
||||
except:
|
||||
tokenizer = AutoTokenizer.from_pretrained(find_largest_checkpoint(output_dir))
|
||||
|
||||
tokenizer.save_pretrained(os.path.join(output_dir, 'merged_model'))
|
||||
@@ -0,0 +1,51 @@
|
||||
import torch
|
||||
from transformers import PreTrainedModel, AutoTokenizer
|
||||
import logging
|
||||
|
||||
from FlagEmbedding.abc.finetune.reranker import AbsRerankerModel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CrossDecoderModel(AbsRerankerModel):
|
||||
"""
|
||||
Model class for decoder only reranker.
|
||||
|
||||
Args:
|
||||
base_model (PreTrainedModel): The underlying pre-trained model used for encoding and scoring input pairs.
|
||||
tokenizer (AutoTokenizer, optional): The tokenizer for encoding input text. Defaults to ``None``.
|
||||
train_batch_size (int, optional): The batch size to use. Defaults to ``4``.
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
base_model: PreTrainedModel,
|
||||
tokenizer: AutoTokenizer = None,
|
||||
train_batch_size: int = 4,
|
||||
):
|
||||
super().__init__(
|
||||
base_model,
|
||||
tokenizer=tokenizer,
|
||||
train_batch_size=train_batch_size,
|
||||
)
|
||||
|
||||
def encode(self, features):
|
||||
"""Encodes input features to logits.
|
||||
|
||||
Args:
|
||||
features (dict): Dictionary with input features.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The logits output from the model.
|
||||
"""
|
||||
if features is None:
|
||||
return None
|
||||
outputs = self.model(input_ids=features['input_ids'],
|
||||
attention_mask=features['attention_mask'],
|
||||
position_ids=features['position_ids'] if 'position_ids' in features.keys() else None,
|
||||
output_hidden_states=True)
|
||||
# _, max_indices = torch.max(features['labels'], dim=1)
|
||||
# predict_indices = max_indices
|
||||
# logits = [outputs.logits[i, predict_indices[i], :] for i in range(outputs.logits.shape[0])]
|
||||
# logits = torch.stack(logits, dim=0)
|
||||
scores = outputs.logits[:, -1, self.yes_loc]
|
||||
return scores.contiguous()
|
||||
@@ -0,0 +1,108 @@
|
||||
import logging
|
||||
from typing import Tuple
|
||||
from pathlib import Path
|
||||
from FlagEmbedding.abc.finetune.reranker.AbsArguments import AbsRerankerDataArguments, AbsRerankerTrainingArguments
|
||||
from transformers import (
|
||||
AutoTokenizer, PreTrainedTokenizer
|
||||
)
|
||||
|
||||
from FlagEmbedding.abc.finetune.reranker import AbsRerankerRunner, AbsRerankerModel
|
||||
|
||||
from .modeling import CrossDecoderModel
|
||||
from .arguments import RerankerModelArguments
|
||||
from .trainer import DecoderOnlyRerankerTrainer
|
||||
from .load_model import get_model, save_merged_model
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DecoderOnlyRerankerRunner(AbsRerankerRunner):
|
||||
"""
|
||||
Decoder only reranker runner for finetuning.
|
||||
|
||||
Args:
|
||||
model_args (RerankerModelArguments): Model arguments instance.
|
||||
data_args (AbsRerankerDataArguments): Data arguments instance.
|
||||
training_args (AbsRerankerTrainingArguments): Trainer arguments.
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
model_args: RerankerModelArguments,
|
||||
data_args: AbsRerankerDataArguments,
|
||||
training_args: AbsRerankerTrainingArguments
|
||||
):
|
||||
super().__init__(model_args, data_args, training_args)
|
||||
|
||||
def load_tokenizer_and_model(self) -> Tuple[PreTrainedTokenizer, AbsRerankerModel]:
|
||||
"""Load the tokenizer and model.
|
||||
|
||||
Returns:
|
||||
Tuple[PreTrainedTokenizer, AbsEmbedderModel]: Tokenizer and model instances.
|
||||
"""
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
self.model_args.tokenizer_name if self.model_args.tokenizer_name else self.model_args.model_name_or_path,
|
||||
token=self.model_args.token,
|
||||
cache_dir=self.model_args.cache_dir,
|
||||
use_fast=self.model_args.use_fast_tokenizer,
|
||||
add_eos_token=False,
|
||||
trust_remote_code=self.model_args.trust_remote_code,
|
||||
)
|
||||
|
||||
if tokenizer.pad_token is None:
|
||||
if tokenizer.unk_token is not None:
|
||||
tokenizer.pad_token = tokenizer.unk_token
|
||||
tokenizer.pad_token_id = tokenizer.unk_token_id
|
||||
elif tokenizer.eod_id is not None:
|
||||
tokenizer.pad_token = tokenizer.eod
|
||||
tokenizer.pad_token_id = tokenizer.eod_id
|
||||
tokenizer.bos_token = tokenizer.im_start
|
||||
tokenizer.bos_token_id = tokenizer.im_start_id
|
||||
tokenizer.eos_token = tokenizer.im_end
|
||||
tokenizer.eos_token_id = tokenizer.im_end_id
|
||||
else:
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
tokenizer.pad_token_id = tokenizer.eos_token_id
|
||||
# if 'mistral' in self.model_args.model_name_or_path.lower():
|
||||
tokenizer.padding_side = 'left'
|
||||
|
||||
base_model = get_model(self.model_args)
|
||||
|
||||
model = CrossDecoderModel(
|
||||
base_model,
|
||||
tokenizer=tokenizer,
|
||||
train_batch_size=self.training_args.per_device_train_batch_size,
|
||||
)
|
||||
|
||||
if self.training_args.gradient_checkpointing:
|
||||
model.enable_input_require_grads()
|
||||
|
||||
return tokenizer, model
|
||||
|
||||
def load_trainer(self) -> DecoderOnlyRerankerTrainer:
|
||||
"""Load the trainer.
|
||||
|
||||
Returns:
|
||||
DecoderOnlyRerankerTrainer: Loaded trainer instance.
|
||||
"""
|
||||
trainer = DecoderOnlyRerankerTrainer(
|
||||
model=self.model,
|
||||
args=self.training_args,
|
||||
train_dataset=self.train_dataset,
|
||||
data_collator=self.data_collator,
|
||||
tokenizer=self.tokenizer
|
||||
)
|
||||
return trainer
|
||||
|
||||
def run(self):
|
||||
"""
|
||||
Run the finetuning.
|
||||
"""
|
||||
Path(self.training_args.output_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Training
|
||||
self.trainer.train(resume_from_checkpoint=self.training_args.resume_from_checkpoint)
|
||||
self.trainer.save_model()
|
||||
|
||||
# save merged model
|
||||
if self.model_args.save_merged_lora_model and self.training_args.process_index == 0:
|
||||
save_merged_model(self.model_args, self.training_args.output_dir)
|
||||
@@ -0,0 +1,52 @@
|
||||
import os
|
||||
import torch
|
||||
import logging
|
||||
from typing import Optional
|
||||
# from transformers.deepspeed import is_deepspeed_zero3_enabled
|
||||
|
||||
from FlagEmbedding.abc.finetune.reranker import AbsRerankerTrainer
|
||||
from peft import get_peft_model_state_dict
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DecoderOnlyRerankerTrainer(AbsRerankerTrainer):
|
||||
"""
|
||||
Trainer class for encoder only base reranker models.
|
||||
"""
|
||||
def _save(self, output_dir: Optional[str] = None, state_dict=None):
|
||||
"""Save the model to directory.
|
||||
|
||||
Args:
|
||||
output_dir (Optional[str], optional): Output directory to save the model. Defaults to ``None``.
|
||||
|
||||
Raises:
|
||||
NotImplementedError
|
||||
"""
|
||||
output_dir = output_dir if output_dir is not None else self.args.output_dir
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
logger.info("Saving model checkpoint to %s", output_dir)
|
||||
# Save a trained model and configuration using `save_pretrained()`.
|
||||
# They can then be reloaded using `from_pretrained()`
|
||||
if not hasattr(self.model, 'save'):
|
||||
raise NotImplementedError(
|
||||
f'MODEL {self.model.__class__.__name__} '
|
||||
f'does not support save interface')
|
||||
else:
|
||||
self.model.save(output_dir)
|
||||
|
||||
if self.tokenizer is not None and self.is_world_process_zero():
|
||||
self.tokenizer.save_pretrained(output_dir)
|
||||
|
||||
torch.save(self.args, os.path.join(output_dir, "training_args.bin"))
|
||||
|
||||
# if is_deepspeed_zero3_enabled():
|
||||
# if state_dict is None:
|
||||
# state_dict = self.model.state_dict()
|
||||
# prefix = 'model.'
|
||||
# assert all(k.startswith(prefix) for k in state_dict.keys()), list(state_dict.keys())
|
||||
# state_dict = {k[len(prefix):]: v for k, v in state_dict.items()}
|
||||
# lora_state_dict = get_peft_model_state_dict(self.model.model, state_dict)
|
||||
# if self.args.process_index <= 0:
|
||||
# torch.save(lora_state_dict, os.path.join(output_dir, "adapter_model.bin"))
|
||||
# print(f"Save adapter model at {output_dir}")
|
||||
Reference in New Issue
Block a user