mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
deprecate modules_to_save + saveing eval results more robustly
This commit is contained in:
parent
1d9acec4f3
commit
9671fde108
7 changed files with 120 additions and 57 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue