diff --git a/hyperlora/hooks.py b/hyperlora/hooks.py index eb12772..8945811 100644 --- a/hyperlora/hooks.py +++ b/hyperlora/hooks.py @@ -10,7 +10,7 @@ from torch import Tensor from torch.utils.hooks import RemovableHandle from utils import get_layers -logger = logging.getLogger(__name__) +logger = logging.getLogger() def remove_hook_handles(handles: list[RemovableHandle]) -> None: diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 0540372..b52e148 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -2,12 +2,14 @@ from copy import copy from functools import partial from importlib.resources import read_binary import logging +import os import random import string import time import numpy as np import torch +import yaml from data_utils import ( convert_ctx_prompt_response_to_messages, get_preprocessing_fn, @@ -27,7 +29,15 @@ from transformers import ( HfArgumentParser, TrainingArguments, ) -from utils import log_num_train_params +from utils import ( + extract_cli_args, + get_run_name, + log_num_train_params, + save_yaml, + setup_logging, + validate_args, + validate_columns, +) from configs import ( ArgumentParser, @@ -37,7 +47,7 @@ from configs import ( ModelArguments, ) -logger = logging.getLogger(__name__) +logger = logging.getLogger() def compute_per_token_acc(shifted_logits, shifted_labels): @@ -90,32 +100,34 @@ def compute_metrics(eval_pred: EvalPrediction) -> dict: return dict(**per_token_acc, **prefix_matching, **entropy) -def get_run_name(): - uuid = "".join( - [random.choice(string.ascii_letters + string.digits) for _ in range(8)] - ) - run_name = time.strftime("%Y%m%d-%H%M%S") + f"_{uuid}" - return run_name - - -def main(): - # Set logging verbosity to INFO - logging.basicConfig(level=logging.INFO) - +def main(output_dir: str): + ############ Argument parsing parser = ArgumentParser( (CtxTrainingArguments, ModelArguments, LoRAArguments, TrainingArguments) ) ctx_args, model_args, lora_args, training_args = parser.parse() - run_name = get_run_name() - training_args.run_name = run_name - training_args.output_dir = f"train_outputs/{run_name}" - training_args.logging_dir = f"train_outputs/{run_name}" + # there shouldn't be overlap between args + validate_args([ctx_args, model_args, lora_args, training_args]) - logger.info(f"Run name: {run_name}") + args = { + **vars(ctx_args), + **vars(model_args), + **vars(lora_args), + **vars(training_args), + } + + 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 + logger.info(f"run_name: {run_name}") logger.info(f"ctx_args: {ctx_args}") logger.info(f"model_args: {model_args}") logger.info(f"lora_args: {lora_args}") + logger.debug(f"args: {args}") + + ############ Model setup model_name = model_args.model_name_or_path model, tokenizer = get_model_and_tokenizer( @@ -144,12 +156,12 @@ def main(): logger.info("Using LoRA") model.set_adapter("default") - print(model) + logger.debug(model) log_num_train_params(model) - # max_seq_len = 1024 + ############ Dataset setup - print("Loading dataset...") + logger.info("Loading dataset...") train_file = "data/raw_datasets/context_numbers/train.jsonl" eval_file = "data/raw_datasets/context_numbers/val.jsonl" ds = load_dataset("json", data_files={"train": train_file, "eval": eval_file}) @@ -199,6 +211,9 @@ def main(): "val": tokenized_ds["eval"], } + logger.debug(f"train_ds: {train_ds}") + logger.debug(f"eval_ds: {eval_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) @@ -241,6 +256,8 @@ def main(): # HACK: see transformers/trainer.py for liger-kernel patch # slows down training speed w/ short inputs # might improve/decrease training speed w/ longer inputs + # TODO: add wandb notes somewhere + # wandb.init(project="ctx_to_lora", name=run_name, notes=args.notes) train_model( model, train_ds, @@ -251,15 +268,10 @@ def main(): ) -def validate_columns(tokenized_ds): - cols = ["input_ids", "attention_mask", "labels"] - if "ctx_features" in tokenized_ds["train"].column_names: - cols += ["ctx_features", "ctx_attn_mask"] - ref_cols = set(cols) - assert ( - set(tokenized_ds["train"].column_names) == ref_cols - ), f"Columns mismatch: {set(tokenized_ds['train'].column_names)} != {ref_cols}" - - if __name__ == "__main__": - main() + run_name = get_run_name() + output_dir = f"train_outputs/{run_name}" + setup_logging(output_dir, debug=os.environ.get("DEBUG", False)) + logger.debug(f"CMD: {' '.join(os.sys.argv)}") + save_yaml(extract_cli_args(os.sys.argv), f"{output_dir}/config.yaml") + main(output_dir) diff --git a/hyperlora/modeling_utils.py b/hyperlora/modeling_utils.py index eaa5e97..631c844 100644 --- a/hyperlora/modeling_utils.py +++ b/hyperlora/modeling_utils.py @@ -18,7 +18,7 @@ from transformers import PreTrainedModel from transformers.modeling_outputs import ModelOutput from utils import get_lora_module_names, get_num_layers, get_peft_in_out_features -logger = logging.getLogger(__name__) +logger = logging.getLogger() AGGREGATOR_TYPE = Enum("AGGREGATOR_TYPE", ["POOLER", "PERCEIVER"]) diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index d3c39ad..3b69258 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -1,14 +1,8 @@ -import math -import os -import random -from enum import Enum, auto +from enum import Enum -import torch from transformers import Seq2SeqTrainer, Trainer from transformers.trainer_utils import get_last_checkpoint -device = torch.device("cuda" if torch.cuda.is_available() else "cpu") - TRAINING_TASK = Enum("TRAINING_TASK", ["CAUSAL_LM", "COMPLETION"]) @@ -22,40 +16,40 @@ def train_model( compute_metrics=None, ): - last_checkpoint = None - if ( - os.path.isdir(training_args.output_dir) - and not training_args.overwrite_output_dir - ): - last_checkpoint = get_last_checkpoint(training_args.output_dir) - if last_checkpoint is None and len(os.listdir(training_args.output_dir)) > 0: - raise ValueError( - f"Output directory ({training_args.output_dir})" - " already exists and is not empty. " - "Use --overwrite_output_dir to overcome." - ) - elif ( - last_checkpoint is not None and training_args.resume_from_checkpoint is None - ): - print( - f"Checkpoint detected, resuming training at {last_checkpoint}. " - "To avoid this behavior, change " - "the `--output_dir` or add `--overwrite_output_dir` to train from scratch." - ) + # last_checkpoint = None + # if ( + # os.path.isdir(training_args.output_dir) + # and not training_args.overwrite_output_dir + # ): + # last_checkpoint = get_last_checkpoint(training_args.output_dir) + # if last_checkpoint is None and len(os.listdir(training_args.output_dir)) > 0: + # raise ValueError( + # f"Output directory ({training_args.output_dir})" + # " already exists and is not empty. " + # "Use --overwrite_output_dir to overcome." + # ) + # elif ( + # last_checkpoint is not None and training_args.resume_from_checkpoint is None + # ): + # print( + # f"Checkpoint detected, resuming training at {last_checkpoint}. " + # "To avoid this behavior, change " + # "the `--output_dir` or add `--overwrite_output_dir` to train from scratch." + # ) - if ( - max( - training_args.per_device_train_batch_size, - training_args.per_device_eval_batch_size, - ) - == 1 - ): - data_collator = None + # if ( + # max( + # training_args.per_device_train_batch_size, + # training_args.per_device_eval_batch_size, + # ) + # == 1 + # ): + # data_collator = None - # print training_args at local_rank 0 - local_rank = int(os.getenv("LOCAL_RANK", "0")) - if local_rank == 0: - print(training_args) + # # print training_args at local_rank 0 + # local_rank = int(os.getenv("LOCAL_RANK", "0")) + # if local_rank == 0: + # print(training_args) # Seq2SeqTrainer is actually just the same as Trainer # (although it uses a different data collator, i.e., explicit prompt/answer separation) @@ -73,12 +67,12 @@ def train_model( checkpoint = None - if training_args.resume_from_checkpoint is not None: - checkpoint = training_args.resume_from_checkpoint - elif last_checkpoint is not None: - checkpoint = last_checkpoint + # if training_args.resume_from_checkpoint is not None: + # checkpoint = training_args.resume_from_checkpoint + # elif last_checkpoint is not None: + # checkpoint = last_checkpoint - print(f"Loaded from the checkpoint: {checkpoint}") + # print(f"Loaded from the checkpoint: {checkpoint}") # TODO: save the best model based on eval loss? train_result = trainer.train(resume_from_checkpoint=checkpoint) diff --git a/hyperlora/utils.py b/hyperlora/utils.py index 41f7ba4..3ad07de 100644 --- a/hyperlora/utils.py +++ b/hyperlora/utils.py @@ -1,13 +1,20 @@ +import ast +import os +import random +import string +import time +import yaml import logging from contextlib import contextmanager from typing import Iterable, Optional + import torch from peft import PeftConfig, PeftModel from peft.tuners.tuners_utils import BaseTunerLayer, check_target_module_exists from peft.utils import get_peft_model_state_dict -logger = logging.getLogger(__name__) +logger = logging.getLogger() # taken from https://discuss.pytorch.org/t/opinion-eval-should-be-a-context-manager/18998/3 @@ -61,6 +68,91 @@ def log_num_train_params(model): ) +def get_run_name(): + uuid = "".join( + [random.choice(string.ascii_letters + string.digits) for _ in range(8)] + ) + run_name = time.strftime("%Y%m%d-%H%M%S") + f"_{uuid}" + return run_name + + +def try_convert(s): + try: + return ast.literal_eval(s) + except: + return s + + +def extract_cli_args(argv: list[str]): + out = dict() + for elem in argv: + if elem.endswith(".yaml"): + out["config"] = elem + + elif elem.startswith("--"): + k, v = elem.split("=") + k = k.split("--")[1] + v = try_convert(v) + # if k.startswith('env_'): + # k = k.split('_')[1] + out[k] = v + return out + + +def setup_logging(output_dir, debug=False): + global logger + + os.makedirs(output_dir, exist_ok=True) + + log_formatter = logging.Formatter( + fmt="%(asctime)s %(levelname)s: %(message)s", datefmt="%Y-%m-%d %H:%M:%S" + ) + stream_level = logging.DEBUG if debug else logging.INFO + stream_handler = logging.StreamHandler() + stream_handler.setFormatter(log_formatter) + stream_handler.setLevel(stream_level) + logger.addHandler(stream_handler) + + log_path = f"{output_dir}/debug.log" + debug_handler = logging.FileHandler(log_path, delay=True) + debug_handler.setLevel(logging.DEBUG) + debug_handler.setFormatter(log_formatter) + logger.addHandler(debug_handler) + logger.setLevel(logging.DEBUG) + logger.info(f"Logging to: {log_path}") + + +def validate_columns(tokenized_ds): + cols = ["input_ids", "attention_mask", "labels"] + if "ctx_features" in tokenized_ds["train"].column_names: + cols += ["ctx_features", "ctx_attn_mask"] + ref_cols = set(cols) + assert ( + set(tokenized_ds["train"].column_names) == ref_cols + ), f"Columns mismatch: {set(tokenized_ds['train'].column_names)} != {ref_cols}" + + +def validate_args(args_list): + # there shouldn't be overlap between args + keys = set() + for args in args_list: + args_keys = set(vars(args).keys()) + assert len(keys & args_keys) == 0, "Overlap between args" + keys |= args_keys + + +def save_yaml(data, path): + # Filter out non-primitive fields + data = { + k: v + for k, v in data.items() + if isinstance(v, (int, float, str, bool, list, dict, type(None))) + } + + with open(path, "w") as file: + yaml.dump(data, file) + + def get_peft_in_out_features( model: PeftModel, peft_config: Optional[PeftConfig] = None,