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, 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, ] ) 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_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, ) 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 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) 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) ) batch_sampler = None if ctx_args.use_multipack_sampler: from multipack_sampler.multipack_sampler import MultipackDistributedBatchSampler """ sampler = MultipackDistributedBatchSampler( batch_max_length=batch_max_len, lengths=lengths, seed=0 ) dataloader = DataLoader(data, batch_sampler=sampler) """ lengths = np.array([len(x["ctx_ids"]) for x in train_ds]) batch_sampler = MultipackDistributedBatchSampler( batch_max_length=ctx_args.per_device_train_max_batch_len, lengths=lengths, seed=training_args.seed, ) # 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, test_ds, train_collator, train_batch_sampler=batch_sampler, # 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() # randomly sleep to avoid run_name collision time.sleep(random.random() * 13) main()