import logging import os import random import string import time from math import ceil 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 ( DS_LEN, 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, PretrainedConfig, DataCollatorForSeq2Seq, EvalPrediction, HfArgumentParser, set_seed, ) from peft import PeftModel from transformers.utils import is_liger_kernel_available from transformers.data import DataCollatorWithFlattening 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, ] ) set_seed(training_args.seed) 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" slurm_job_id = f"_{os.getenv('SLURM_JOB_ID')}" if os.getenv("SLURM_JOB_ID") else "" run_name = ( get_run_name(seed_str=training_args.logging_dir.strip("runs/") + slurm_job_id) 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-8} 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") ############ 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_name = model.base_model.config.name_or_path ctx_encoder_model_config = model.config ctx_tokenizer = tokenizer if ctx_args.exp_setup == ExperimentSetup.HYPER_LORA: # TODO: handle only extra_modules case (no target_modules) 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, ) training_args.gen_lora_l1_reg_coef = ctx_args.gen_lora_l1_reg_coef else: logger.info( f"Loading from checkpoint: {ctx_args.from_pretrained_checkpoint}" ) model = ModulatedPretrainedModel.from_state_dict( torch.load(ctx_args.from_pretrained_checkpoint), train=True, use_flash_attn=model_args.use_flash_attn, ) 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 base_model_config = AutoConfig.from_pretrained( model_args.model_name_or_path, trust_remote_code=True ) base_model_config.save_pretrained(output_dir) 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, set_format=None if ctx_args.use_sequence_packing else "pt", # 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 streaming = (split == "train") and data_args.streaming tokenized_ds[split] = { os.path.basename(ds_name): _get_tokenized_dataset( ds_name, split, streaming=streaming ) for ds_name in ds_names } train_ds = tokenized_ds["train"] logging.info(f"train_ds: {train_ds}") 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) max_steps = ceil( sum([DS_LEN[ds] for ds in train_ds]) * training_args.num_train_epochs / training_args.per_device_train_batch_size / training_args.gradient_accumulation_steps / training_args.world_size ) training_args.max_steps = max_steps # interleaving streaming datasets # simplify the probs for smaller datasets # slightly upsample those datasets probs = [ 0.01 if "fw_qa" not in ds_name else 1 + 0.01 - 0.01 * len(train_ds) for ds_name in train_ds ] train_ds = interleave_datasets( list(train_ds.values()), probabilities=probs, stopping_strategy="all_exhausted", seed=training_args.seed, ) logger.info(f"train_ds: {train_ds}") logger.info(f"val_ds: {val_ds}") flattener = DataCollatorWithFlattening() def train_packed_collator(inp_list): # no padding packed_inputs = flattener(inp_list, return_tensors="pt") if "ctx_ids" in inp_list[0]: ctx_ids = [{"input_ids": example["ctx_ids"]} for example in inp_list] packed_ctx = flattener(ctx_ids, return_tensors="pt") packed_inputs["ctx_ids"] = packed_ctx["input_ids"] packed_inputs["ctx_position_ids"] = packed_ctx["position_ids"] return packed_inputs 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, ) 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 return out train_collator = ( train_packed_collator if ctx_args.use_sequence_packing else partial(train_collator, tokenizer=tokenizer) ) # TODO: use SFTTrainer instead? https://huggingface.co/docs/trl/en/sft_trainer # TODO: use packing with SFTTrainer # HACK [local patch]: deepspeed model loading problem (for resume training) # see https://github.com/microsoft/DeepSpeed/pull/6626/files # /home/rujikorn_sakana_ai/.conda/envs/ctx-to-lora/lib/python3.10/site-packages/deepspeed/runtime/engine.py 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") if isinstance(model.base_model, PeftModel): _apply_liger_kernel_to_instance(model=model.base_model.base_model.model) else: _apply_liger_kernel_to_instance(model=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.base_model) 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, train_collator, # 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" os.environ["KMP_AFFINITY"] = "disabled" # fixing iterable dataset stuck os.environ["OMP_NUM_THREADS"] = "16" if os.getenv("DEBUG", False): disable_caching() # randomly sleep to avoid run_name collision # time.sleep(random.random() * 13) main()