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}")
|
||||
@@ -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.layerwise 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,78 @@
|
||||
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."}
|
||||
)
|
||||
|
||||
model_type: str = field(
|
||||
default='from_raw_model' # should be one of ['from_raw_model', 'from_finetuned_model']
|
||||
# from_raw_model -- openbmb/MiniCPM-2B-dpo-bf16
|
||||
# from_finetuned_model -- BAAI/bge-reranker-v2-minicpm-layerwise
|
||||
)
|
||||
|
||||
start_layer: int = field(
|
||||
default=8,
|
||||
metadata={"help": "which layer to start to compute score"}
|
||||
)
|
||||
|
||||
head_multi: bool = field(
|
||||
default=False,
|
||||
metadata={"help": "use one / multi classifier"}
|
||||
)
|
||||
head_type: str = field(
|
||||
default='simple',
|
||||
metadata={"help": "the type of the classifier"}
|
||||
)
|
||||
+208
@@ -0,0 +1,208 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
|
||||
# and OPT implementations in this library. It has been modified from its
|
||||
# original forms to accommodate minor architectural differences compared
|
||||
# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
""" MiniCPM model configuration"""
|
||||
|
||||
from transformers.configuration_utils import PretrainedConfig
|
||||
from transformers.utils import logging
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
MINICPM_PRETRAINED_CONFIG_ARCHIVE_MAP = {}
|
||||
|
||||
class LayerWiseMiniCPMConfig(PretrainedConfig):
|
||||
r"""
|
||||
This is the configuration class to store the configuration of a [`MiniCPMModel`]. It is used to instantiate an MiniCPM
|
||||
model according to the specified arguments, defining the model architecture. Instantiating a configuration with the
|
||||
defaults will yield a similar configuration to that of the MiniCPM-7B.
|
||||
|
||||
Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
|
||||
documentation from [`PretrainedConfig`] for more information.
|
||||
|
||||
|
||||
Args:
|
||||
vocab_size (`int`, *optional*, defaults to 32000):
|
||||
Vocabulary size of the MiniCPM model. Defines the number of different tokens that can be represented by the
|
||||
`inputs_ids` passed when calling [`MiniCPMModel`]
|
||||
hidden_size (`int`, *optional*, defaults to 4096):
|
||||
Dimension of the hidden representations.
|
||||
intermediate_size (`int`, *optional*, defaults to 11008):
|
||||
Dimension of the MLP representations.
|
||||
num_hidden_layers (`int`, *optional*, defaults to 32):
|
||||
Number of hidden layers in the Transformer decoder.
|
||||
num_attention_heads (`int`, *optional*, defaults to 32):
|
||||
Number of attention heads for each attention layer in the Transformer decoder.
|
||||
num_key_value_heads (`int`, *optional*):
|
||||
This is the number of key_value heads that should be used to implement Grouped Query Attention. If
|
||||
`num_key_value_heads=num_attention_heads`, the model will use Multi Head Attention (MHA), if
|
||||
`num_key_value_heads=1 the model will use Multi Query Attention (MQA) otherwise GQA is used. When
|
||||
converting a multi-head checkpoint to a GQA checkpoint, each group key and value head should be constructed
|
||||
by meanpooling all the original heads within that group. For more details checkout [this
|
||||
paper](https://arxiv.org/pdf/2305.13245.pdf). If it is not specified, will default to
|
||||
`num_attention_heads`.
|
||||
hidden_act (`str` or `function`, *optional*, defaults to `"silu"`):
|
||||
The non-linear activation function (function or string) in the decoder.
|
||||
max_position_embeddings (`int`, *optional*, defaults to 2048):
|
||||
The maximum sequence length that this model might ever be used with. MiniCPM 1 supports up to 2048 tokens,
|
||||
MiniCPM 2 up to 4096, CodeMiniCPM up to 16384.
|
||||
initializer_range (`float`, *optional*, defaults to 0.02):
|
||||
The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
|
||||
rms_norm_eps (`float`, *optional*, defaults to 1e-06):
|
||||
The epsilon used by the rms normalization layers.
|
||||
use_cache (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not the model should return the last key/values attentions (not used by all models). Only
|
||||
relevant if `config.is_decoder=True`.
|
||||
pad_token_id (`int`, *optional*):
|
||||
Padding token id.
|
||||
bos_token_id (`int`, *optional*, defaults to 1):
|
||||
Beginning of stream token id.
|
||||
eos_token_id (`int`, *optional*, defaults to 2):
|
||||
End of stream token id.
|
||||
pretraining_tp (`int`, *optional*, defaults to 1):
|
||||
Experimental feature. Tensor parallelism rank used during pretraining. Please refer to [this
|
||||
document](https://huggingface.co/docs/transformers/parallelism) to understand more about it. This value is
|
||||
necessary to ensure exact reproducibility of the pretraining results. Please refer to [this
|
||||
issue](https://github.com/pytorch/pytorch/issues/76232).
|
||||
tie_word_embeddings (`bool`, *optional*, defaults to `False`):
|
||||
Whether to tie weight embeddings
|
||||
rope_theta (`float`, *optional*, defaults to 10000.0):
|
||||
The base period of the RoPE embeddings.
|
||||
rope_scaling (`Dict`, *optional*):
|
||||
Dictionary containing the scaling configuration for the RoPE embeddings. Currently supports two scaling
|
||||
strategies: linear and dynamic. Their scaling factor must be a float greater than 1. The expected format is
|
||||
`{"type": strategy name, "factor": scaling factor}`. When using this flag, don't update
|
||||
`max_position_embeddings` to the expected new maximum. See the following thread for more information on how
|
||||
these scaling strategies behave:
|
||||
https://www.reddit.com/r/LocalMiniCPM/comments/14mrgpr/dynamically_scaled_rope_further_increases/. This is an
|
||||
experimental feature, subject to breaking API changes in future versions.
|
||||
attention_bias (`bool`, defaults to `False`, *optional*, defaults to `False`):
|
||||
Whether to use a bias in the query, key, value and output projection layers during self-attention.
|
||||
attention_dropout (`float`, *optional*, defaults to 0.0):
|
||||
The dropout ratio for the attention probabilities.
|
||||
|
||||
```python
|
||||
>>> from transformers import MiniCPMModel, MiniCPMConfig
|
||||
|
||||
>>> # Initializing a MiniCPM minicpm-7b style configuration
|
||||
>>> configuration = MiniCPMConfig()
|
||||
|
||||
>>> # Initializing a model from the minicpm-7b style configuration
|
||||
>>> model = MiniCPMModel(configuration)
|
||||
|
||||
>>> # Accessing the model configuration
|
||||
>>> configuration = model.config
|
||||
```"""
|
||||
|
||||
model_type = "minicpm"
|
||||
keys_to_ignore_at_inference = ["past_key_values"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size=32000,
|
||||
hidden_size=4096,
|
||||
intermediate_size=11008,
|
||||
num_hidden_layers=32,
|
||||
num_attention_heads=32,
|
||||
num_key_value_heads=None,
|
||||
hidden_act="silu",
|
||||
max_position_embeddings=2048,
|
||||
initializer_range=0.02,
|
||||
rms_norm_eps=1e-6,
|
||||
use_cache=True,
|
||||
pad_token_id=None,
|
||||
bos_token_id=1,
|
||||
eos_token_id=2,
|
||||
pretraining_tp=1,
|
||||
tie_word_embeddings=True,
|
||||
rope_theta=10000.0,
|
||||
rope_scaling=None,
|
||||
attention_bias=False,
|
||||
attention_dropout=0.0,
|
||||
scale_emb=1,
|
||||
dim_model_base=1,
|
||||
scale_depth=1,
|
||||
start_layer=8,
|
||||
head_multi=True,
|
||||
head_type="simple",
|
||||
**kwargs,
|
||||
):
|
||||
self.vocab_size = vocab_size
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.hidden_size = hidden_size
|
||||
self.intermediate_size = intermediate_size
|
||||
self.num_hidden_layers = num_hidden_layers
|
||||
self.num_attention_heads = num_attention_heads
|
||||
|
||||
# for backward compatibility
|
||||
if num_key_value_heads is None:
|
||||
num_key_value_heads = num_attention_heads
|
||||
|
||||
self.num_key_value_heads = num_key_value_heads
|
||||
self.hidden_act = hidden_act
|
||||
self.initializer_range = initializer_range
|
||||
self.rms_norm_eps = rms_norm_eps
|
||||
self.pretraining_tp = pretraining_tp
|
||||
self.use_cache = use_cache
|
||||
self.rope_theta = rope_theta
|
||||
self.rope_scaling = rope_scaling
|
||||
self._rope_scaling_validation()
|
||||
self.attention_bias = attention_bias
|
||||
self.attention_dropout = attention_dropout
|
||||
self.scale_emb = scale_emb
|
||||
self.dim_model_base = dim_model_base
|
||||
self.scale_depth = scale_depth
|
||||
|
||||
self.start_layer = start_layer
|
||||
self.head_multi = head_multi
|
||||
self.head_type = head_type
|
||||
|
||||
super().__init__(
|
||||
pad_token_id=pad_token_id,
|
||||
bos_token_id=bos_token_id,
|
||||
eos_token_id=eos_token_id,
|
||||
tie_word_embeddings=tie_word_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
try:
|
||||
import flash_attn
|
||||
self._attn_implementation = "flash_attention_2"
|
||||
except:
|
||||
pass
|
||||
|
||||
def _rope_scaling_validation(self):
|
||||
"""
|
||||
Validate the `rope_scaling` configuration.
|
||||
"""
|
||||
if self.rope_scaling is None:
|
||||
return
|
||||
|
||||
if not isinstance(self.rope_scaling, dict) or len(self.rope_scaling) != 2:
|
||||
raise ValueError(
|
||||
"`rope_scaling` must be a dictionary with with two fields, `type` and `factor`, "
|
||||
f"got {self.rope_scaling}"
|
||||
)
|
||||
rope_scaling_type = self.rope_scaling.get("type", None)
|
||||
rope_scaling_factor = self.rope_scaling.get("factor", None)
|
||||
if rope_scaling_type is None or rope_scaling_type not in ["linear", "dynamic"]:
|
||||
raise ValueError(
|
||||
f"`rope_scaling`'s type field must be one of ['linear', 'dynamic'], got {rope_scaling_type}"
|
||||
)
|
||||
if rope_scaling_factor is None or not isinstance(rope_scaling_factor, float) or rope_scaling_factor <= 1.0:
|
||||
raise ValueError(f"`rope_scaling`'s factor field must be a float > 1, got {rope_scaling_factor}")
|
||||
@@ -0,0 +1,232 @@
|
||||
import os
|
||||
import re
|
||||
import logging
|
||||
from torch import nn
|
||||
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
|
||||
from peft import LoraConfig, TaskType, get_peft_model, PeftModel
|
||||
|
||||
from FlagEmbedding.finetune.reranker.decoder_only.layerwise.arguments import RerankerModelArguments
|
||||
|
||||
from .modeling_minicpm_reranker import LayerWiseMiniCPMForCausalLM, LayerWiseHead
|
||||
from .configuration_minicpm_reranker import LayerWiseMiniCPMConfig
|
||||
|
||||
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, only_for_one_logit):
|
||||
"""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_type == 'from_raw_model':
|
||||
config.use_cache = False
|
||||
config.start_layer = config.num_hidden_layers
|
||||
config.head_multi = False
|
||||
config.head_type = 'raw'
|
||||
|
||||
model = LayerWiseMiniCPMForCausalLM.from_pretrained(
|
||||
model_args.model_name_or_path,
|
||||
trust_remote_code=model_args.trust_remote_code,
|
||||
# torch_dtype=torch.float16 if training_args.fp16 else 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,
|
||||
)
|
||||
|
||||
config.start_layer = model_args.start_layer
|
||||
config.head_multi = model_args.head_multi
|
||||
config.head_type = model_args.head_type
|
||||
model.config = config
|
||||
|
||||
if model.config.head_type == 'complex':
|
||||
if model.config.head_multi == True:
|
||||
lm_head = nn.ModuleList([LayerWiseHead(
|
||||
model.config.hidden_size, model.config.vocab_size) for _ in range(
|
||||
model.config.start_layer,
|
||||
model.config.num_hidden_layers + 1)])
|
||||
for i in range(len(lm_head)):
|
||||
lm_head[i].linear_head.load_state_dict(model.lm_head.state_dict())
|
||||
model.set_output_embeddings(lm_head)
|
||||
else:
|
||||
lm_head = LayerWiseHead(model.config.hidden_size, 1)
|
||||
state_dict_back = model.lm_head.state_dict()
|
||||
state_dict_back['weight'] = state_dict_back['weight'][only_for_one_logit: only_for_one_logit + 1, :]
|
||||
lm_head.linear_head.load_state_dict(state_dict_back)
|
||||
model.set_output_embeddings(lm_head)
|
||||
else:
|
||||
if only_for_one_logit is None:
|
||||
raise ValueError('`only for one logit` cannot be None.')
|
||||
if model.config.head_multi == True:
|
||||
lm_head = nn.ModuleList([LayerWiseHead(
|
||||
model.config.hidden_size, 1) for _ in range(
|
||||
model.config.start_layer,
|
||||
model.config.num_hidden_layers + 1)])
|
||||
state_dict_back = model.lm_head.state_dict()
|
||||
state_dict_back['weight'] = state_dict_back['weight'][only_for_one_logit: only_for_one_logit + 1, :]
|
||||
for i in range(len(lm_head)):
|
||||
lm_head[i].linear_head.load_state_dict(state_dict_back)
|
||||
model.set_output_embeddings(lm_head)
|
||||
else:
|
||||
lm_head = LayerWiseHead(model.config.hidden_size, 1)
|
||||
state_dict_back = model.lm_head.state_dict()
|
||||
state_dict_back['weight'] = state_dict_back['weight'][only_for_one_logit: only_for_one_logit + 1, :]
|
||||
lm_head.linear_head.load_state_dict(state_dict_back)
|
||||
model.set_output_embeddings(lm_head)
|
||||
# modules_to_save = model_args.modules_to_save
|
||||
# target_modules = model_args.target_modules
|
||||
else:
|
||||
config.use_cache = False
|
||||
|
||||
model = LayerWiseMiniCPMForCausalLM.from_pretrained(
|
||||
model_args.model_name_or_path,
|
||||
# torch_dtype=torch.float16 if training_args.fp16 else 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,
|
||||
)
|
||||
# target_modules = model_args.target_modules
|
||||
# target_modules.extend(model_args.modules_to_save)
|
||||
# modules_to_save = None
|
||||
|
||||
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,
|
||||
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_type == 'from_raw_model':
|
||||
config = LayerWiseMiniCPMConfig.from_pretrained('BAAI/bge-reranker-v2-minicpm-layerwise',
|
||||
cache_dir=model_args.cache_dir,
|
||||
token=model_args.token,
|
||||
trust_remote_code=model_args.trust_remote_code)
|
||||
config.start_layer = model_args.start_layer
|
||||
config.head_multi = model_args.head_multi
|
||||
config.head_type = model_args.head_type
|
||||
|
||||
model = LayerWiseMiniCPMForCausalLM.from_pretrained(model_args.model_name_or_path,
|
||||
config=config,
|
||||
cache_dir=model_args.cache_dir,
|
||||
token=model_args.token,
|
||||
trust_remote_code=model_args.trust_remote_code)
|
||||
|
||||
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,89 @@
|
||||
import torch
|
||||
from transformers import PreTrainedModel, AutoTokenizer
|
||||
import logging
|
||||
from typing import List, Union, Dict, Optional
|
||||
from torch import Tensor
|
||||
|
||||
from FlagEmbedding.abc.finetune.reranker import AbsRerankerModel, RerankerOutput
|
||||
|
||||
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``.
|
||||
start_layer (int, optional): Starting layer for layerwise. Defaults to ``8``.
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
base_model: PreTrainedModel,
|
||||
tokenizer: AutoTokenizer = None,
|
||||
train_batch_size: int = 4,
|
||||
start_layer: int = 8
|
||||
):
|
||||
super().__init__(
|
||||
base_model,
|
||||
tokenizer=tokenizer,
|
||||
train_batch_size=train_batch_size,
|
||||
)
|
||||
|
||||
self.start_layer = start_layer
|
||||
|
||||
def encode(self, features):
|
||||
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)
|
||||
all_logits = outputs.logits
|
||||
all_scores = []
|
||||
for logits in all_logits:
|
||||
all_scores.append(logits[:, -1].contiguous())
|
||||
return all_scores
|
||||
|
||||
def forward(self, pair: Union[Dict[str, Tensor], List[Dict[str, Tensor]]] = None, teacher_scores: Optional[Tensor] = None):
|
||||
ranker_logits = self.encode(pair) # (batch_size * num, dim)
|
||||
|
||||
if self.training:
|
||||
loss = 0
|
||||
for logits in ranker_logits:
|
||||
grouped_logits = logits.view(self.train_batch_size, -1)
|
||||
target = torch.zeros(self.train_batch_size, device=grouped_logits.device, dtype=torch.long)
|
||||
loss += self.compute_loss(grouped_logits, target)
|
||||
|
||||
if teacher_scores is None:
|
||||
teacher_scores = ranker_logits[-1].view(
|
||||
self.train_batch_size,
|
||||
-1
|
||||
)
|
||||
teacher_targets = torch.softmax(teacher_scores.detach(), dim=-1)
|
||||
for logits in ranker_logits[:-1]:
|
||||
student_scores = logits.view(
|
||||
self.train_batch_size,
|
||||
-1
|
||||
)
|
||||
loss += - torch.mean(torch.sum(torch.log_softmax(student_scores, dim=-1) * teacher_targets, dim=-1))
|
||||
else:
|
||||
teacher_scores = torch.Tensor(teacher_scores)
|
||||
teacher_scores = teacher_scores.view(self.train_batch_size, -1)
|
||||
teacher_targets = torch.softmax(teacher_scores.detach(), dim=-1).to(ranker_logits[-1].device)
|
||||
for logits in ranker_logits:
|
||||
student_scores = logits.view(
|
||||
self.train_batch_size,
|
||||
-1
|
||||
)
|
||||
loss += - torch.mean(torch.sum(torch.log_softmax(student_scores, dim=-1) * teacher_targets, dim=-1))
|
||||
else:
|
||||
loss = None
|
||||
|
||||
# print(loss)
|
||||
return RerankerOutput(
|
||||
loss=loss,
|
||||
scores=ranker_logits,
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,109 @@
|
||||
import os
|
||||
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 FlagEmbedding.finetune.reranker.decoder_only.layerwise.modeling import CrossDecoderModel
|
||||
from FlagEmbedding.finetune.reranker.decoder_only.layerwise.arguments import RerankerModelArguments
|
||||
from FlagEmbedding.finetune.reranker.decoder_only.layerwise.trainer import DecoderOnlyRerankerTrainer
|
||||
from FlagEmbedding.finetune.reranker.decoder_only.layerwise.load_model import get_model, save_merged_model
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class DecoderOnlyRerankerRunner(AbsRerankerRunner):
|
||||
"""
|
||||
Decoder only layerwise 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.
|
||||
"""
|
||||
# print(self.model_args.model_name_or_path)
|
||||
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, tokenizer('Yes', add_special_tokens=False)['input_ids'][-1])
|
||||
|
||||
model = CrossDecoderModel(
|
||||
base_model,
|
||||
tokenizer=tokenizer,
|
||||
train_batch_size=self.training_args.per_device_train_batch_size,
|
||||
start_layer=self.model_args.start_layer
|
||||
)
|
||||
|
||||
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 peft import get_peft_model_state_dict
|
||||
|
||||
from FlagEmbedding.abc.finetune.reranker import AbsRerankerTrainer
|
||||
|
||||
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