custom ctx_encoder!

This commit is contained in:
51616 2025-01-13 16:45:26 +00:00
parent 284d35fb12
commit 273b35378f
7 changed files with 100 additions and 23 deletions

View file

@ -27,6 +27,8 @@ neftune_noise_alpha: 1
weight_decay: 0.01
warmup_ratio: 0.1
dataloader_prefetch_factor: 16
dataloader_num_workers: 16
# LoRA
lora_r: 8
lora_dropout: 0.05

View file

@ -252,6 +252,10 @@ class CtxTrainingArguments:
default=2**13,
metadata={"help": "Maximum base length for training."},
)
max_ctx_len: Optional[int] = field(
default=2**13,
metadata={"help": "Maximum context length for training."},
)
max_new_tokens: Optional[int] = field(
default=2**10,
metadata={"help": "Maximum new tokens for generation-based evaluation."},
@ -320,9 +324,16 @@ class HypernetArguments:
@dataclass
class CtxEncoderArguments:
layer_idx: int = field(
default=4,
metadata={"help": "Layer index for context encoder."},
ctx_encoder_model_name_or_path: str = field(
default=None,
metadata={"help": "Context encoder model name or path."},
)
layer_idx: Optional[int] = field(
default=None,
metadata={
"help": "Layer index for context encoder. "
"Default to L//4 where L is the number of layers of the ctx model"
},
)

View file

@ -171,6 +171,8 @@ def get_tokenized_dataset(
split: str,
tokenizer: PreTrainedTokenizerBase,
tokenizer_kwargs: dict[str, Any],
ctx_tokenizer: PreTrainedTokenizerBase,
ctx_tokenizer_kwargs: dict[str, Any],
add_ctx_to_chat: bool,
add_repeat_prompt: bool,
add_negative_prompt: bool,
@ -199,13 +201,27 @@ def get_tokenized_dataset(
if add_repeat_prompt and "context_numbers" not in ds_name:
ds = ds.map(add_repeat_prompt_fn, batched=True, batch_size=None)
tokenized_ds = construct_and_tokenize_ctx_qa(
tokenizer, tokenizer_kwargs, add_ctx_to_chat, use_kl_loss, need_ctx_ids, ds
tokenizer,
tokenizer_kwargs,
ctx_tokenizer,
ctx_tokenizer_kwargs,
add_ctx_to_chat,
use_kl_loss,
need_ctx_ids,
ds,
)
return tokenized_ds
def construct_and_tokenize_ctx_qa(
tokenizer, tokenizer_kwargs, add_ctx_to_chat, use_kl_loss, need_ctx_ids, ds
tokenizer,
tokenizer_kwargs,
ctx_tokenizer,
ctx_tokenizer_kwargs,
add_ctx_to_chat,
use_kl_loss,
need_ctx_ids,
ds,
):
# for sft + chat_model, we need to convert the dataset to chat format
# add "messages" field
@ -229,6 +245,7 @@ def construct_and_tokenize_ctx_qa(
# for use_kl_loss, we need "chat_ids" and "chat_attn_mask"
if use_kl_loss:
raise NotImplementedError("KL loss deprecated")
tokenized_ds = tokenized_ds.map(
convert_ctx_prompt_response_to_messages,
fn_kwargs={"add_ctx_to_chat": True},
@ -255,7 +272,7 @@ def construct_and_tokenize_ctx_qa(
# tokenize the ctx_text to get ctx_ids and ctx_attn_mask
tokenized_ds = tokenized_ds.map(
tokenize_ctx_text,
fn_kwargs={"tokenizer": tokenizer},
fn_kwargs={"tokenizer": ctx_tokenizer},
batched=True,
num_proc=16,
)

View file

@ -262,16 +262,33 @@ def main():
requires_grad=ctx_args.exp_setup == ExperimentSetup.FULL_FINETUNE,
peft_config=get_lora_config(model_name, **vars(lora_args)),
)
if ctx_encoder_args.ctx_encoder_model_name_or_path is not None:
ctx_encoder_model, ctx_tokenizer = get_model_and_tokenizer(
ctx_encoder_args.ctx_encoder_model_name_or_path,
train=False,
requires_grad=False,
)
else:
ctx_encoder_model = model
ctx_tokenizer = tokenizer
if ctx_args.exp_setup == ExperimentSetup.HYPER_LORA:
logger.info("Using HyperLoRA")
hypernet_config = get_hypernet_config(model, hypernet_args, aggregator_args)
hypernet_config = get_hypernet_config(
model, ctx_encoder_model, hypernet_args, aggregator_args
)
# hypernet = HyperLoRA(
# get_hypernet_config(model, hypernet_args, aggregator_args),
# model,
# ).to(model.device)
# ctx_encoder = EarlyExit(get_base_model(model), ctx_encoder_args.layer_idx)
if ctx_encoder_args.layer_idx is None:
ctx_encoder_args.layer_idx = ctx_encoder_model.config.num_hidden_layers // 4
logger.info(
f"Using the first {ctx_encoder_args.layer_idx} layers"
" as the context encoder"
)
model = ModulatedPretrainedModel(
model,
hypernet_config,
@ -293,10 +310,13 @@ def main():
add_ctx_to_chat = not isinstance(model, ModulatedPretrainedModel)
tokenizer_kwargs = {"max_length": ctx_args.max_base_len}
ctx_tokenizer_kwargs = {"max_length": ctx_args.max_ctx_len} # not used for now
_get_tokenized_dataset = partial(
get_tokenized_dataset,
tokenizer=tokenizer,
tokenizer_kwargs=tokenizer_kwargs,
ctx_tokenizer=ctx_tokenizer,
ctx_tokenizer_kwargs=ctx_tokenizer_kwargs,
add_ctx_to_chat=add_ctx_to_chat,
add_repeat_prompt=ctx_args.add_repeat_prompt,
add_negative_prompt=ctx_args.add_negative_prompt,
@ -471,6 +491,9 @@ def main():
if isinstance(model, ModulatedPretrainedModel):
logger.info("Applying liger-kernel to ModulatedPretrainedModel")
_apply_liger_kernel_to_instance(model=model.base_model.base_model.model)
if ctx_encoder_args.ctx_encoder_model_name_or_path is not None:
logger.info("Applying liger-kernel to ctx_encoder_model")
_apply_liger_kernel_to_instance(model=ctx_encoder_model)
elif isinstance(model, PeftModel):
logger.info("Applying liger-kernel to PeftModel")
_apply_liger_kernel_to_instance(model=model.base_model.model)
@ -485,21 +508,21 @@ def main():
train_model(
model,
tokenizer,
# tokenizer,
training_args,
train_ds,
val_ds,
test_ds,
partial(train_collator, tokenizer=tokenizer),
partial(generation_collator, tokenizer=tokenizer),
# partial(generation_collator, tokenizer=tokenizer),
compute_metrics=partial(
compute_metrics,
evaluator=Evaluator(
[compute_per_token_acc, compute_prefix_matching, compute_perplexity]
),
),
max_new_tokens=ctx_args.max_new_tokens,
gen_per_device_eval_batch_size=ctx_args.gen_per_device_eval_batch_size,
# max_new_tokens=ctx_args.max_new_tokens,
# gen_per_device_eval_batch_size=ctx_args.gen_per_device_eval_batch_size,
)
logger.info(f"Training run finished and saved to {output_dir}")

View file

@ -5,7 +5,12 @@ import torch
from peft import LoraConfig, PeftConfig, PeftModel, VeraConfig
from peft import get_peft_config as _get_peft_config
from peft.utils import PeftType
from transformers import AutoModel, AutoModelForCausalLM, AutoTokenizer
from transformers import (
AutoModel,
AutoModelForCausalLM,
AutoTokenizer,
MllamaForConditionalGeneration,
)
logger = logging.getLogger()
@ -118,15 +123,21 @@ def get_model(
# load_in_4bit=True,
# load_in_8bit=True,
)
is_vision_model = "Llama" in model_name_or_path and "Vision" in model_name_or_path
if model_kwargs is not None:
model_init_kwargs.update(model_kwargs)
if use_flash_attn:
model_init_kwargs["attn_implementation"] = "flash_attention_2"
if is_vision_model:
model_init_kwargs["attn_implementation"] = "sdpa"
# for training disable cache
if train:
if train and not is_vision_model:
model_init_kwargs["use_cache"] = False
logger.debug(f"Model init kwargs: {model_init_kwargs}")
model = AutoModelForCausalLM.from_pretrained(**model_init_kwargs)
if not is_vision_model:
model = AutoModelForCausalLM.from_pretrained(**model_init_kwargs)
else:
model = MllamaForConditionalGeneration.from_pretrained(**model_init_kwargs)
if peft_config is not None:
model = PeftModel(model, peft_config)
model.train(train)

View file

@ -70,11 +70,14 @@ class AggregatorConfig:
def get_aggregator_config(
model: PreTrainedModel, output_size: int, aggregator_args: AggregatorArguments
model: PreTrainedModel,
ctx_encoder_model: PreTrainedModel,
output_size: int,
aggregator_args: AggregatorArguments,
):
lora_config = model.peft_config["default"]
return AggregatorConfig(
feature_size=model.config.hidden_size,
feature_size=ctx_encoder_model.config.hidden_size,
output_size=output_size,
num_layers=get_num_layers(model),
num_modules=len(lora_config.target_modules),
@ -97,6 +100,7 @@ class HypernetConfig:
def get_hypernet_config(
model: PreTrainedModel,
ctx_encoder_model: PreTrainedModel,
hypernet_args: HypernetArguments,
aggregator_args: AggregatorArguments,
):
@ -110,6 +114,7 @@ def get_hypernet_config(
feature_sizes=get_peft_in_out_features(model, peft_config=lora_config),
aggregator_config=get_aggregator_config(
model,
ctx_encoder_model,
hypernet_args.latent_size,
aggregator_args,
),
@ -577,8 +582,16 @@ class ModulatedPretrainedModel(nn.Module):
def _init_model(self):
self.hypernet = HyperLoRA(self.hypernet_config).to(self.device)
# TODO: allow ctx_encoder to be other models
if self.ctx_encoder_args.ctx_encoder_model_name_or_path is not None:
encoder_model = get_model(
self.ctx_encoder_args.ctx_encoder_model_name_or_path,
train=True,
requires_grad=False,
)
else:
encoder_model = self.base_model
self.ctx_encoder = EarlyExit(
get_base_model(self.base_model), self.ctx_encoder_args.layer_idx
get_base_model(encoder_model), self.ctx_encoder_args.layer_idx
)
def to(self, *args, **kwargs):

View file

@ -108,17 +108,17 @@ def eval_generation(eval_trainer, tokenizer, dataset, split, gen_kwargs):
def train_model(
model,
tokenizer,
# tokenizer,
training_args,
train_dataset=None,
val_dataset=None,
test_dataset=None,
train_collator=None,
generation_collator=None,
# generation_collator=None,
compute_metrics=None,
preprocess_logits_for_metrics=None,
max_new_tokens=2**13,
gen_per_device_eval_batch_size=1,
# preprocess_logits_for_metrics=None,
# max_new_tokens=2**13,
# gen_per_device_eval_batch_size=1,
):
checkpoint = None
if training_args.resume_from_checkpoint is not None:
@ -132,7 +132,7 @@ def train_model(
eval_dataset=val_dataset,
data_collator=train_collator,
compute_metrics=compute_metrics,
preprocess_logits_for_metrics=preprocess_logits_for_metrics,
# preprocess_logits_for_metrics=preprocess_logits_for_metrics,
)
# Trainer loads the best model after training