mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
better logging
This commit is contained in:
parent
d4f01c999e
commit
92135a94af
5 changed files with 178 additions and 80 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue