doc-to-lora/intx_sft.py

569 lines
20 KiB
Python
Executable file

import logging
import os
import random
import string
import time
from collections import defaultdict
from copy import copy, deepcopy
from functools import partial
from importlib.resources import read_binary
from typing import Callable, Optional
import numpy as np
import torch
import wandb
import yaml
from ctx_to_lora.data_utils import (
convert_ctx_prompt_response_to_messages,
get_preprocessing_fn,
get_sft_prompt_formatting_fn,
get_tokenized_dataset,
tokenize_chat_messages,
tokenize_ctx_text,
)
from datasets import (
concatenate_datasets,
interleave_datasets,
disable_caching,
load_dataset,
IterableDataset,
)
from ctx_to_lora.model_loading import (
get_lora_config,
get_model_and_tokenizer,
get_tokenizer,
)
from ctx_to_lora.modeling_utils import (
EarlyExit,
HyperLoRA,
ModulatedPretrainedModel,
get_hypernet_config,
)
from rouge_score import rouge_scorer
from ctx_to_lora.training_utils import TRAINING_TASK, train_model
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
AutoConfig,
DataCollatorForSeq2Seq,
EvalPrediction,
HfArgumentParser,
set_seed,
)
from peft import PeftModel
from transformers.utils import is_liger_kernel_available
from torch.utils.data import DataLoader
from ctx_to_lora.utils import (
extract_cli_args,
get_base_model,
get_run_name,
log_num_train_params,
save_yaml,
setup_logging,
validate_args,
)
from ctx_to_lora.configs import (
ArgumentParser,
CtxTrainingArguments,
TrainingArguments,
DataArguments,
ExperimentSetup,
LoRAArguments,
ModelArguments,
HypernetArguments,
AggregatorArguments,
CtxEncoderArguments,
)
logger = logging.getLogger()
LOCAL_RANK = int(os.getenv("LOCAL_RANK", "0"))
def compute_per_token_acc(shift_logits, shift_labels, valid_masks):
indices = torch.where(valid_masks)
# acc = (shift_logits.argmax(-1) == shift_labels)[indices].float().mean().item()
# return {"per_token_acc": acc}
acc = (shift_logits.argmax(-1) == shift_labels)[indices].float()
return {
"per_token_accs": acc.flatten().tolist(),
"n_per_token_accs": valid_masks.sum().item(),
}
def compute_prefix_matching(shift_logits, shift_labels, valid_masks):
lengths = valid_masks.sum(dim=1)
is_wrong = (shift_logits.argmax(-1) != shift_labels) * valid_masks
is_correct = (shift_logits.argmax(-1) == shift_labels) * valid_masks
# NOTE: not reliable for multi-turn conversations
# ie, all tokens in the following user's turn will be correct
# still monotonically correlate with perf though
wrong_pos = torch.argmax(is_wrong, dim=1) - torch.argmax(valid_masks, dim=1)
perf = wrong_pos / lengths
# if all tokens are correct, set to 1
perf = torch.where(is_correct.sum(dim=1) == lengths, 1, perf)
# return {"prefix_matching": perf.mean().item()}
return {
"prefix_matchings": perf.tolist(),
"n_prefix_matchings": valid_masks.shape[0],
}
@torch.no_grad()
def compute_perplexity(shift_logits, shift_labels, valid_masks):
loss_fct = torch.nn.CrossEntropyLoss(reduction="none")
loss = loss_fct(shift_logits.transpose(1, 2), shift_labels)
loss = (loss * valid_masks).sum(dim=1) / valid_masks.sum(dim=1)
# perplexity = torch.exp(loss).mean().item()
# return {"perplexity": perplexity}
preplexities = torch.exp(loss)
return {
"perplexities": preplexities.tolist(),
"n_perplexities": valid_masks.shape[0],
}
class Evaluator:
def __init__(self, metric_fns: list[Callable]):
self.metric_fns = metric_fns
self.reset()
def reset(self):
self.accum_metrics = defaultdict(list)
self.count = defaultdict(list)
def update(self, shift_logits, shift_labels, valid_masks):
for metric_fn in self.metric_fns:
metric = metric_fn(shift_logits, shift_labels, valid_masks)
for k, v in metric.items():
if k.startswith("n_"):
self.count[k[2:]].append(v)
else:
self.accum_metrics[k] += v
def compute(self):
# Get result across entire eval set
result = {
k: np.sum(v) / np.sum(self.count[k]) for k, v in self.accum_metrics.items()
}
# Reset batch statistics
self.reset()
return result
@torch.no_grad()
def compute_metrics(
eval_pred: EvalPrediction,
compute_result: bool,
evaluator: Evaluator,
) -> Optional[dict]:
logits, labels = eval_pred.predictions, eval_pred.label_ids
shift_logits = logits[..., :-1, :]
shift_labels = labels[..., 1:]
valid_masks = torch.where(shift_labels != -100, 1, 0)
evaluator.update(shift_logits, shift_labels, valid_masks)
if compute_result:
return evaluator.compute()
# def compute_metrics(eval_pred: EvalPrediction) -> dict:
# """
# Custom metrics function for the trainer
# Args:
# eval_pred: tuple of predictions and labels
# Returns:
# dictionary containing metric names (str) and values (Any)
# """
# # compute per token accuracy
# # logits, labels = eval_pred.predictions, eval_pred.label_ids
# # shift_logits = logits[..., :-1, :]
# pred_ids, labels = eval_pred.predictions, eval_pred.label_ids
# shift_pred_ids = pred_ids[..., :-1]
# shift_labels = labels[..., 1:]
# valid_masks = np.where(shift_labels != -100, 1, 0)
# per_token_acc = compute_per_token_acc(shift_pred_ids, shift_labels, valid_masks)
# prefix_matching = compute_prefix_matching(shift_pred_ids, shift_labels, valid_masks)
# # entropy = compute_entropy(shift_logits, shift_labels, valid_masks)
# return dict(
# **per_token_acc,
# **prefix_matching,
# # **entropy,
# num_valid_tokens=valid_masks.sum(),
# num_samples=valid_masks.shape[0],
# )
# # https://discuss.huggingface.co/t/cuda-out-of-memory-when-using-trainer-with-compute-metrics/2941/13
# def preprocess_logits_for_metrics(logits, labels):
# """
# Original Trainer may have a memory leak.
# This is a workaround to avoid storing too many tensors that are not needed.
# """
# pred_ids = torch.argmax(logits, dim=-1)
# return pred_ids
def main():
############ Argument parsing
parser = ArgumentParser(
(
DataArguments,
CtxTrainingArguments,
ModelArguments,
LoRAArguments,
TrainingArguments,
HypernetArguments,
AggregatorArguments,
CtxEncoderArguments,
)
)
(
data_args,
ctx_args,
model_args,
lora_args,
training_args,
hypernet_args,
aggregator_args,
ctx_encoder_args,
) = parser.parse()
# there shouldn't be overlap between args
validate_args(
[
data_args,
ctx_args,
model_args,
lora_args,
training_args,
hypernet_args,
aggregator_args,
ctx_encoder_args,
]
)
checkpoint_dir = training_args.resume_from_checkpoint
if checkpoint_dir and not os.path.isdir(checkpoint_dir):
raise NotADirectoryError(f"Checkpoint{checkpoint_dir} is not a directory")
# should be the same across processes
# still possible to have a name crash though
# logging_dir is just "runs/DATE_TIME_HOSTNAME"
run_name = (
get_run_name(seed_str=training_args.logging_dir.strip("runs/"))
if not checkpoint_dir
else checkpoint_dir.strip("/").split("/")[-2]
)
output_dir = f"train_outputs/runs/{run_name}"
setup_logging(output_dir, debug=os.getenv("DEBUG", False))
logger.debug(f"CMD: {' '.join(os.sys.argv)}")
cli_args = extract_cli_args(os.sys.argv)
save_yaml(cli_args, f"{output_dir}/cli_args.yaml")
if "config" in cli_args:
config_name = os.path.basename(cli_args["config"]).split(".yaml")[0]
os.environ["WANDB_TAGS"] = config_name
run_name = os.path.basename(output_dir)
training_args.run_name = run_name
training_args.output_dir = output_dir
training_args.logging_dir = output_dir
if (
training_args.lr_scheduler_type == "cosine_with_min_lr"
and training_args.lr_scheduler_kwargs is None
):
training_args.lr_scheduler_kwargs = {"min_lr": 1e-7}
args = {
**vars(deepcopy(data_args)),
**vars(deepcopy(ctx_args)),
**vars(deepcopy(model_args)),
**vars(deepcopy(lora_args)),
**vars(deepcopy(training_args)),
**vars(deepcopy(hypernet_args)),
**vars(deepcopy(aggregator_args)),
**vars(deepcopy(ctx_encoder_args)),
}
args["deepspeed_plugin"] = None
logger.debug(f"args: {args}")
save_yaml(args, f"{output_dir}/args.yaml")
set_seed(training_args.seed)
############ Model setup
if not ctx_args.from_pretrained_checkpoint:
model_name = model_args.model_name_or_path
model, tokenizer = get_model_and_tokenizer(
**vars(model_args),
train=True,
requires_grad=ctx_args.exp_setup == ExperimentSetup.FULL_FINETUNE,
peft_config=get_lora_config(model_name, **vars(lora_args)),
)
ctx_name = ctx_encoder_args.ctx_encoder_model_name_or_path
if ctx_name is not None:
ctx_encoder_model_config = AutoConfig.from_pretrained(
ctx_name, trust_remote_code=True
)
if "Llama" in ctx_name and "Vision" in ctx_name:
ctx_encoder_model_config = ctx_encoder_model_config.text_config
ctx_tokenizer = get_tokenizer(ctx_name)
else:
ctx_encoder_model_config = model.config
ctx_tokenizer = tokenizer
if ctx_args.exp_setup == ExperimentSetup.HYPER_LORA:
logger.info("Using HyperLoRA")
if not ctx_args.from_pretrained_checkpoint:
hypernet_config = get_hypernet_config(
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
)
logger.info(
f"Using the first {ctx_encoder_args.layer_idx} layers"
" as the context encoder"
)
model = ModulatedPretrainedModel(
model,
hypernet_config,
ctx_encoder_args,
ctx_args.use_kl_loss,
)
else:
logger.info(
f"Loading from checkpoint: {ctx_args.from_pretrained_checkpoint}"
)
model = ModulatedPretrainedModel.from_state_dict(
torch.load(open(ctx_args.from_pretrained_checkpoint, "rb")),
train=True,
)
tokenizer = get_tokenizer(model.base_model.config.name_or_path)
ctx_name = model.ctx_encoder_args.ctx_encoder_model_name_or_path
if ctx_name is None:
ctx_name = model.base_model.config.name_or_path
ctx_tokenizer = get_tokenizer(ctx_name)
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")
model.set_adapter("default")
model.train()
logger.debug(model)
log_num_train_params(model)
############ Dataset setup
logger.info("Loading dataset...")
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,
use_kl_loss=ctx_args.use_kl_loss,
# streaming=data_args.streaming,
)
tokenized_ds = {}
for split, ds_names in zip(
["train", "validation", "test"],
[data_args.train_ds_names, data_args.val_ds_names, data_args.test_ds_names],
):
if not ds_names:
continue
tokenized_ds[split] = {
os.path.basename(ds_name): _get_tokenized_dataset(ds_name, split)
for ds_name in ds_names
}
train_ds = tokenized_ds["train"]
val_ds = dict()
if "validation" in tokenized_ds:
n_val_samples = data_args.max_val_samples_per_ds
for ds_name, ds in tokenized_ds["validation"].items():
if ds is None:
# take some samples from train_ds
ds = train_ds[ds_name].take(n_val_samples)
train_ds[ds_name] = train_ds[ds_name].skip(n_val_samples)
val_ds[ds_name] = ds
val_indices = np.random.permutation(len(ds))[:n_val_samples]
val_ds[ds_name] = val_ds[ds_name].select(val_indices)
# train_ds = concatenate_datasets(list(train_ds.values()))
# total_len = sum(len(ds) for ds in train_ds.values())
train_ds_len = [len(ds) for ds in train_ds.values()]
total_len = sum(train_ds_len)
train_ds = interleave_datasets(
list(train_ds.values()),
probabilities=[l / total_len for l in train_ds_len],
seed=training_args.seed,
)
# val_train_indices = np.random.permutation(len(train_ds))[:500]
# val_ds["train"] = train_ds.select(val_train_indices)
test_ds = dict()
if "test" in tokenized_ds:
n_test_samples = data_args.max_test_samples_per_ds
for ds_name, ds in tokenized_ds["test"].items():
test_ds[ds_name] = ds
test_indices = np.random.permutation(len(ds))[:n_test_samples]
test_ds[ds_name] = test_ds[ds_name].select(test_indices)
logger.info(f"train_ds: {train_ds}")
logger.info(f"val_ds: {val_ds}")
logger.info(f"test_ds: {test_ds}")
# TODO: change to a faster collator? e.g.,
# https://huggingface.co/blog/packing-with-FA2
# data_collator = DataCollatorForSeq2Seq(tokenizer, model, pad_to_multiple_of=8)
def train_collator(inp_list, tokenizer):
# input is a list of tokenized sequences
padding_kwargs = dict(
padding=True,
padding_side="right",
pad_to_multiple_of=8,
return_tensors="pt",
)
ctx_ids = None
if "ctx_ids" in inp_list[0]:
# have to be manual since it has [ctx_len, features] shape
# pad to the longest ctx_len in the batch
# which can have a different length from the input_ids, attn_mask, labels
ctx_ids = [example.pop("ctx_ids") for example in inp_list]
ctx_ids = torch.nn.utils.rnn.pad_sequence(
ctx_ids,
batch_first=True,
padding_value=0,
)
# exotic keys won't be padded, so we need to pad them as well
ctx_attn_mask = [example.pop("ctx_attn_mask") for example in inp_list]
ctx_attn_mask = torch.nn.utils.rnn.pad_sequence(
ctx_attn_mask,
batch_first=True,
padding_value=0,
)
chat_ids = None
if "chat_ids" in inp_list[0]:
chat_ids = [x.pop("chat_ids") for x in inp_list]
chat_ids = torch.nn.utils.rnn.pad_sequence(
chat_ids,
batch_first=True,
padding_value=0,
)
chat_attn_mask = [x.pop("chat_attn_mask") for x in inp_list]
chat_attn_mask = torch.nn.utils.rnn.pad_sequence(
chat_attn_mask,
batch_first=True,
padding_value=0,
)
chat_labels = [x.pop("chat_labels") for x in inp_list]
chat_labels = torch.nn.utils.rnn.pad_sequence(
chat_labels,
batch_first=True,
padding_value=-100,
)
chat_labels = torch.where(chat_attn_mask == 0, -100, chat_labels)
labels = [x.pop("labels") for x in inp_list]
padded_seq = tokenizer.pad(inp_list, **padding_kwargs)
# hacky explicit padding since the labels are not padded by default
labels = tokenizer.pad({"input_ids": labels}, **padding_kwargs)["input_ids"]
labels = torch.where(padded_seq["attention_mask"] == 0, -100, labels)
out = {**padded_seq, "labels": labels}
if ctx_ids is not None:
out["ctx_ids"] = ctx_ids
out["ctx_attn_mask"] = ctx_attn_mask
if chat_ids is not None:
out["chat_ids"] = chat_ids
out["chat_attn_mask"] = chat_attn_mask
out["chat_labels"] = chat_labels
return out
# TODO: use SFTTrainer instead? https://huggingface.co/docs/trl/en/sft_trainer
# TODO: use packing with SFTTrainer
# HACK [local patch]: see transformers/trainer_seq2seq.py for supressing
# "Trainer.tokenizer is now deprecated. You should use Trainer.processing_class instead."
# HACK [local patch]: deepspeed model loading problem (for resume training)
# see https://github.com/microsoft/DeepSpeed/pull/6626/files
if training_args.use_liger_kernel and is_liger_kernel_available():
from liger_kernel.transformers import _apply_liger_kernel_to_instance
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_name is not None:
logger.info("Applying liger-kernel to 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)
if LOCAL_RANK == 0:
wandb.init(
project=os.getenv("WANDB_PROJECT"),
name=run_name,
group=run_name,
config=args,
tags=os.getenv("WANDB_TAGS").split(","),
notes=ctx_args.notes,
resume="allow",
)
else:
wandb.init(mode="disabled")
train_model(
model,
# tokenizer,
training_args,
train_ds,
val_ds,
test_ds,
partial(train_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,
)
logger.info(f"Training run finished and saved to {output_dir}")
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"
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
if os.getenv("DEBUG", False):
disable_caching()
main()