deprecate modules_to_save + saveing eval results more robustly

This commit is contained in:
51616 2025-01-15 18:10:29 +00:00
parent 1d9acec4f3
commit 9671fde108
7 changed files with 120 additions and 57 deletions

View file

@ -187,7 +187,7 @@ class TrainingArguments(TrainingArguments):
# metadata={"help": "Whether to load the best model at the end of training."},
# )
save_total_limit: int = field(
default=1,
default=5,
metadata={"help": "Total number of checkpoints to save."},
)
save_strategy: str = field(
@ -248,10 +248,10 @@ class LoRAArguments:
default=None,
metadata={"help": ("LoRA target modules.")},
)
modules_to_save: Optional[list[str]] = field(
default=None,
metadata={"help": ("Modules to save.")},
)
# modules_to_save: Optional[list[str]] = field(
# default=None,
# metadata={"help": ("Modules to save.")},
# )
@dataclass

View file

@ -113,6 +113,7 @@ def compute_metrics(
def save_generated_text(samples, output_dir, split):
os.makedirs(output_dir, exist_ok=True)
with open(f"{output_dir}/{split}_generated_text.jsonl", "w") as f:
for sample in samples:
f.write(json.dumps(sample) + "\n")
@ -408,16 +409,18 @@ def evaluate(checkpoint_path, args, split, generative):
if __name__ == "__main__":
os.environ["TRANSFORMERS_NO_ADVISORY_WARNINGS"] = "true"
os.environ["WANDB_MODE"] = "disabled"
checkpoint_path = sys.argv[1]
checkpoint_dir = "/".join(checkpoint_path.split("/")[:-1])
run_dir = "/".join(checkpoint_path.split("/")[:-2])
cur_it = int(checkpoint_path.split("checkpoint-")[1].split("/")[0])
args = Namespace(**yaml.unsafe_load(open(f"{run_dir}/args.yaml", "r")))
print(f"checkpoint_path: {checkpoint_path}")
print(f"run_dir: {run_dir}")
# print(f"args: {args}")
args.output_dir = checkpoint_dir
args.logging_dir = checkpoint_dir
args.output_dir = f"{run_dir}/eval-results-{cur_it}"
args.logging_dir = f"{run_dir}/eval-results-{cur_it}"
args.run_name = run_dir.split("/")[-1]
evaluate(checkpoint_path, args, split="validation", generative=False)

View file

@ -29,7 +29,7 @@ from datasets import (
load_dataset,
IterableDataset,
)
from model_loading import get_lora_config, get_model_and_tokenizer
from model_loading import get_lora_config, get_model_and_tokenizer, get_tokenizer
from modeling_utils import (
EarlyExit,
HyperLoRA,
@ -41,6 +41,7 @@ from training_utils import TRAINING_TASK, train_model
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
AutoConfig,
DataCollatorForSeq2Seq,
EvalPrediction,
HfArgumentParser,
@ -279,22 +280,21 @@ def main():
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,
ctx_encoder_model_config = AutoConfig.from_pretrained(
ctx_encoder_args.ctx_encoder_model_name_or_path
)
ctx_tokenizer = get_tokenizer(ctx_encoder_args.ctx_encoder_model_name_or_path)
else:
ctx_encoder_model = model
ctx_encoder_model_config = model.config
ctx_tokenizer = tokenizer
if ctx_args.exp_setup == ExperimentSetup.HYPER_LORA:
logger.info("Using HyperLoRA")
hypernet_config = get_hypernet_config(
model, ctx_encoder_model, hypernet_args, aggregator_args
model, ctx_encoder_model_config, hypernet_args, aggregator_args
)
if ctx_encoder_args.layer_idx is None:
ctx_encoder_args.layer_idx = ctx_encoder_model.config.num_hidden_layers // 4
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"
@ -305,6 +305,10 @@ def main():
ctx_encoder_args,
ctx_args.use_kl_loss,
)
if len([p for p in model.ctx_encoder.parameters() if p.requires_grad]):
raise ValueError("ctx_encoder contains trainable parameters")
if len([p for p in model.base_model.parameters() if p.requires_grad]):
raise ValueError("base model contains trainable parameters")
else:
# activate LoRA
logger.info("Using LoRA")
@ -467,7 +471,7 @@ def main():
_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)
_apply_liger_kernel_to_instance(model=model.ctx_encoder)
elif isinstance(model, PeftModel):
logger.info("Applying liger-kernel to PeftModel")
_apply_liger_kernel_to_instance(model=model.base_model.model)

View file

@ -146,6 +146,7 @@ def get_model(
if "modules_to_save" not in name:
param.requires_grad = requires_grad
else:
logging.debug(f"modules_to_save: {name} to be trained")
# always train "modules_to_save"
if not "layernorm" in name:
raise NotImplementedError(

View file

@ -23,13 +23,20 @@ from peft import (
PeftConfig,
PeftModel,
LoraRuntimeConfig,
get_peft_model_state_dict,
set_peft_model_state_dict,
)
from peft.utils import PeftType, TaskType
from peft.tuners._buffer_dict import BufferDict
from peft.tuners.tuners_utils import BaseTunerLayer, check_target_module_exists
from pooling import POOL_FN, get_pooling_fn
from torch import Tensor, nn
from transformers import PerceiverConfig, PerceiverModel, PreTrainedModel
from transformers import (
PerceiverConfig,
PerceiverModel,
PreTrainedModel,
PretrainedConfig,
)
from transformers.models.perceiver.modeling_perceiver import (
PerceiverBasicDecoder,
)
@ -71,13 +78,13 @@ class AggregatorConfig:
def get_aggregator_config(
model: PreTrainedModel,
ctx_encoder_model: PreTrainedModel,
ctx_encoder_model_config: PretrainedConfig,
output_size: int,
aggregator_args: AggregatorArguments,
):
lora_config = model.peft_config["default"]
return AggregatorConfig(
feature_size=ctx_encoder_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),
@ -100,7 +107,7 @@ class HypernetConfig:
def get_hypernet_config(
model: PreTrainedModel,
ctx_encoder_model: PreTrainedModel,
ctx_encoder_model_config: PretrainedConfig,
hypernet_args: HypernetArguments,
aggregator_args: AggregatorArguments,
):
@ -114,7 +121,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,
ctx_encoder_model_config,
hypernet_args.latent_size,
aggregator_args,
),
@ -332,7 +339,8 @@ class EarlyExit(nn.Module):
def __init__(self, base_model: PreTrainedModel, exit_layer: int):
super().__init__()
self.base_model = base_model
self.exit_layer = exit_layer
self.base_model.layers = base_model.layers[:exit_layer]
# self.exit_layer = exit_layer
@torch.no_grad()
def forward(self, **kwargs):
@ -342,7 +350,7 @@ class EarlyExit(nn.Module):
# kwargs["attention_mask"] = kwargs["attention_mask"].unsqueeze(0)
with (
early_exit(self.base_model, self.exit_layer),
# early_exit(self.base_model, self.exit_layer),
maybe_add_batch_dim(kwargs) as (batched_input, batched_attn_mask),
):
model_outputs = self.base_model(**kwargs)
@ -583,15 +591,25 @@ 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
# ctx_encoder_name =
# 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
ctx_model_name = self.ctx_encoder_args.ctx_encoder_model_name_or_path
if ctx_model_name is None:
ctx_model_name = self.base_model.config.name_or_path
# use an explicit copy of the base model
# for using with "modules_to_save"
encoder_model = get_model(
ctx_model_name,
train=True,
requires_grad=False,
)
self.ctx_encoder = EarlyExit(
get_base_model(encoder_model), self.ctx_encoder_args.layer_idx
)
@ -649,7 +667,20 @@ class ModulatedPretrainedModel(nn.Module):
# self.hypernet.head.bias.requires_grad = False
def state_dict(self, *args, **kwargs):
# we assume ctx_encoder and base model is frozen here
if len([p for p in self.ctx_encoder.parameters() if p.requires_grad]):
raise ValueError("ctx_encoder contains trainable parameters")
if len([p for p in self.base_model.parameters() if p.requires_grad]):
raise ValueError("base model contains trainable parameters")
state_dict = self.hypernet.state_dict(*args, **kwargs)
# base_state_dict = dict()
# if self.base_model.modules_to_save:
# for k, v in get_peft_model_state_dict(self.base_model).items():
# if any(module in k for module in self.base_model.modules_to_save):
# # save only "modules_to_save"
# base_state_dict[k] = v
# state_dict["base_model"] = base_state_dict
state_dict["hypernet_config"] = self.hypernet_config
state_dict["ctx_encoder_args"] = self.ctx_encoder_args
return state_dict
@ -668,6 +699,16 @@ class ModulatedPretrainedModel(nn.Module):
f"but the hypernet config is for: {self.hypernet_config.lora_config.base_model_name_or_path}"
)
self._init_model()
# base_state_dict = state_dict.pop("base_model", None)
# if base_state_dict:
# logging.info("Loading base model state dict")
# peft_load_result = set_peft_model_state_dict(
# self.base_model,
# base_state_dict,
# *args,
# **kwargs,
# )
# logging.info(f"peft_load_result: {peft_load_result}")
return self.hypernet.load_state_dict(state_dict, *args, **kwargs)
# @torch.no_grad()
@ -869,11 +910,12 @@ if __name__ == "__main__":
model_name,
train=True,
requires_grad=False,
peft_config=get_lora_config(model_name),
peft_config=get_lora_config(
model_name, modules_to_save=["input_layernorm", "post_attention_layernorm"]
),
)
# TODO: base model shoul be init with target lora config
print(base_model)
ctx_model = base_model
print(base_model.modules_to_save)
ctx_model_config = base_model.config
device = base_model.device
# lora_config = base_model.peft_config["default"]
# d_in, d_out = get_peft_in_out_features(base_model, peft_config=lora_config)
@ -881,7 +923,7 @@ if __name__ == "__main__":
hypernet_args = HypernetArguments(latent_size=512)
aggregator_args = AggregatorArguments(aggregator_type=AGGREGATOR_TYPE.PERCEIVER)
hypernet_config = get_hypernet_config(
base_model, ctx_model, hypernet_args, aggregator_args
base_model, ctx_model_config, hypernet_args, aggregator_args
)
ctx_encoder_args = CtxEncoderArguments(layer_idx=4)
# ctx_encoder = EarlyExit(get_base_model(base_model), 4)
@ -922,15 +964,15 @@ if __name__ == "__main__":
modelout = model(ctx_ids, ctx_attn_mask, **prompt_inputs)
print(modelout)
# model.load_state_dict(
# torch.load(
# open(
# "train_outputs/runs/Jan14_11-10-56_slurm0-a3nodeset-1_30c9e93d/checkpoint-1720/pytorch_model.bin",
# "rb",
# )
# )
# )
# modelout = model(ctx_ids, ctx_attn_mask, **prompt_inputs)
# print(modelout)
state_dict = torch.load(
open(
"train_outputs/runs/Jan15_17-40-28_slurm0-a3nodeset-0_fd03d041/checkpoint-1000/pytorch_model.bin",
"rb",
)
)
breakpoint()
model.load_state_dict(state_dict)
modelout = model(ctx_ids, ctx_attn_mask, **prompt_inputs)
print(modelout)
breakpoint()

View file

@ -60,15 +60,16 @@ def train_model(
# Trainer loads the best model after training
# is done when load_best_model_at_end=True
train_result = trainer.train(resume_from_checkpoint=checkpoint)
trainer.save_model() # just in case OOM when run trainer.evaluate()
trainer.log_metrics("train", train_result.metrics)
clear_gpu()
metrics = trainer.evaluate(dict(**val_dataset, test=test_dataset))
trainer.log_metrics("eval", metrics)
trainer.save_metrics("eval", metrics)
trainer.save_model()
clear_gpu()
# metrics = trainer.evaluate(dict(**val_dataset, test=test_dataset))
# trainer.log_metrics("eval", metrics)
# trainer.save_metrics("eval", metrics)
# trainer.save_model()
# clear_gpu()
# ############## Evaluation
# # TODO: eval does not work when using with deepspeed
# # make a separate eval script

View file

@ -7,12 +7,14 @@ import shutil
import time
import os
import argparse
import gc
from glob import glob
import numpy as np
import pandas as pd
import wandb
import yaml
import torch
from eval import evaluate
@ -24,6 +26,13 @@ def flatten(l):
return itertools.chain.from_iterable(l)
def clear_gpu():
gc.collect()
torch.cuda.empty_cache()
torch.cuda.reset_max_memory_allocated()
torch.cuda.reset_max_memory_cached()
class Watcher:
def __init__(self, patterns):
self.patterns = patterns
@ -53,7 +62,6 @@ class Watcher:
if __name__ == "__main__":
os.environ["TRANSFORMERS_NO_ADVISORY_WARNINGS"] = "true"
os.environ["TOKENIZERS_PARALLELISM"] = "true"
os.environ["WANDB_PROJECT"] = "ctx_to_lora"
os.environ["WANDB_WATCH"] = "" # "all"
os.environ["WANDB_CONSOLE"] = "off"
@ -82,12 +90,13 @@ if __name__ == "__main__":
# cp is delete before we can read it
continue
run_dir = file.split("/checkpoint")[0]
cur_it = int(file.split("checkpoint-")[1].split("/")[0])
checkpoint_dir = os.path.dirname(file)
args = argparse.Namespace(
**yaml.unsafe_load(open(f"{run_dir}/args.yaml", "r"))
)
args.output_dir = checkpoint_dir
args.logging_dir = checkpoint_dir
args.output_dir = f"{run_dir}/eval-results-{cur_it}"
args.logging_dir = f"{run_dir}/eval-results-{cur_it}"
args.run_name = run_dir.split("/")[-1]
print("Evaluating...")
print(f"checkpoint_dir: {checkpoint_dir}")
@ -101,9 +110,12 @@ if __name__ == "__main__":
}
wandb.init(**wandb_kwargs)
metrics = evaluate(file, args, split="validation", generative=False)
gen_metrics = evaluate(file, args, split="validation", generative=True)
metrics.update(gen_metrics)
# gen_metrics = evaluate(file, args, split="validation", generative=True)
# metrics.update(gen_metrics)
wandb.log(metrics, step=curstep)
wandb.finish()
print(f"Logged metrics: {metrics}")
print("=" * 80)
clear_gpu()
watcher.save_state()