better logging

This commit is contained in:
51616 2024-12-23 04:03:19 +00:00
parent d4f01c999e
commit 92135a94af
5 changed files with 178 additions and 80 deletions

View file

@ -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:

View file

@ -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)

View file

@ -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"])

View file

@ -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)

View file

@ -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,