From ef91544fbcef04467e26233df37fd8346ad314c9 Mon Sep 17 00:00:00 2001 From: 51616 Date: Tue, 14 Jan 2025 11:58:10 +0000 Subject: [PATCH] deepspeed trainable w/ shared output dir + cosine schedule + data kwargs + explicit training when use modules_to_save --- hyperlora/configs.py | 8 +++++ hyperlora/data_utils.py | 61 ++++++++++++++++++++++++++---------- hyperlora/intx_sft.py | 63 +++++++++++++++++++++++--------------- hyperlora/model_loading.py | 7 +++++ hyperlora/utils.py | 17 +++++++--- 5 files changed, 110 insertions(+), 46 deletions(-) diff --git a/hyperlora/configs.py b/hyperlora/configs.py index da42059..212efed 100644 --- a/hyperlora/configs.py +++ b/hyperlora/configs.py @@ -154,6 +154,14 @@ class TrainingArguments(TrainingArguments): default=0.95, metadata={"help": "Adam beta 2."}, ) + lr_scheduler_type: str = field( + default="cosine_with_min_lr", + metadata={"help": "Learning rate scheduler type."}, + ) + lr_scheduler_kwargs: dict = field( + default=None, + metadata={"help": "Learning rate scheduler kwargs."}, + ) eval_on_start: bool = field( default=True, metadata={"help": "Whether to evaluate on the start of training."}, diff --git a/hyperlora/data_utils.py b/hyperlora/data_utils.py index f70add5..b81f221 100644 --- a/hyperlora/data_utils.py +++ b/hyperlora/data_utils.py @@ -1,5 +1,6 @@ import logging import numpy as np +from glob import glob from typing import Any, Callable, Iterator, Optional from datasets import load_dataset, IterableDataset @@ -9,25 +10,21 @@ from transformers import PreTrainedTokenizerBase IGNORE_INDEX = -100 logger = logging.getLogger() - +FW_QA_PATHS = [ + f"data/raw_datasets/fw_qa/{i:05d}.parquet" for i in [0, 1, 6, 7, 8, 10, 22, 30, 35] +] DS_KWARGS = { "hotpot_qa": dict( - train=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="train[1000:]"), - validation=dict( - path="hotpotqa/hotpot_qa", name="fullwiki", split="train[:1000]" - ), - test=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="validation"), + train=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="train"), + validation=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="validation"), ), "hotpot_qa_tiny": dict( - train=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="train[1000:2000]"), - validation=dict( - path="hotpotqa/hotpot_qa", name="fullwiki", split="train[:1000]" - ), - test=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="validation"), + train=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="train[:1000]"), + validation=dict(path="hotpotqa/hotpot_qa", name="fullwiki", split="validation"), ), "pwc": dict( - train=dict(path="sggetao/PwC", split="train[1000:]"), - validation=dict(path="sggetao/PwC", split="train[:1000]"), + train=dict(path="sggetao/PwC", split="train[900:]"), + validation=dict(path="sggetao/PwC", split="train[:900]"), test=dict(path="sggetao/PwC", split="test"), ), "pwc_tiny": dict( @@ -47,6 +44,30 @@ DS_KWARGS = { split="train", ), ), + "fw_qa": dict( + train=dict( + path="parquet", + data_files=FW_QA_PATHS, + split="train", + ), + validation=dict( + path="parquet", + data_files="data/raw_datasets/fw_qa/*_val.parquet", + split="train", + ), + ), + "ctx_qa": dict( + train=dict( + path="parquet", + data_files=glob("data/raw_datasets/ctx_qa/*res.parquet"), + split="train", + ), + validation=dict( + path="parquet", + data_files=glob("data/raw_datasets/ctx_qa/*val.parquet"), + split="train", + ), + ), } @@ -78,6 +99,8 @@ def filter_long_samples(samples): def add_repeat_prompt_fn(samples): + # TODO: maybe repeat only short context (<5K chars) + # TODO: also check batch size and mem util unique_contexts = set() ctxs, prompts, responses = [], [], [] for ctx, prompt, response in zip( @@ -191,10 +214,10 @@ def get_tokenized_dataset( cols_to_remove = [ col for col in ds.column_names if col not in ["context", "prompt", "response"] ] - ds = ds.map(get_preprocessing_fn(ds_name)) + ds = ds.map(get_preprocessing_fn(ds_name), num_proc=16) ds = ds.remove_columns(cols_to_remove) - ds = ds.filter(filter_none, batched=True) - ds = ds.filter(filter_long_samples, batched=True) + ds = ds.filter(filter_none, batched=True, num_proc=16) + ds = ds.filter(filter_long_samples, batched=True, num_proc=16) if split == "train": if add_negative_prompt: ds = ds.map(add_negative_prompt_fn, batched=True, batch_size=None) @@ -228,9 +251,12 @@ def construct_and_tokenize_ctx_qa( ds = ds.map( convert_ctx_prompt_response_to_messages, fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat}, + num_proc=16, ) # add "chat" field - ds = ds.map(get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer)) + ds = ds.map( + get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer), num_proc=16 + ) # tokenize the chat + mask the assistant inputs # add "input_ids", "attention_mask", "labels" @@ -241,6 +267,7 @@ def construct_and_tokenize_ctx_qa( "mask_assistant_inputs": True, "tokenizer_kwargs": tokenizer_kwargs, }, + num_proc=16, ) # for use_kl_loss, we need "chat_ids" and "chat_attn_mask" diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index e5896cd..ea9dee5 100755 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -4,7 +4,7 @@ import random import string import time from collections import defaultdict -from copy import copy +from copy import copy, deepcopy from functools import partial from importlib.resources import read_binary from typing import Callable, Optional @@ -74,6 +74,8 @@ from configs import ( 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) @@ -221,21 +223,30 @@ def main(): ) args = { - **vars(data_args), - **vars(ctx_args), - **vars(model_args), - **vars(lora_args), - **vars(training_args), - **vars(hypernet_args), - **vars(aggregator_args), - **vars(ctx_encoder_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 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") - run_name = get_run_name() if not checkpoint_dir else checkpoint_dir.split("/")[-2] + # 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) + if not checkpoint_dir + else checkpoint_dir.split("/")[-2] + ) + output_dir = f"train_outputs/{run_name}" setup_logging(output_dir, debug=os.getenv("DEBUG", False)) logger.debug(f"CMD: {' '.join(os.sys.argv)}") @@ -249,6 +260,12 @@ def main(): 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} logger.debug(f"args: {args}") save_yaml(args, f"{output_dir}/args.yaml") set_seed(training_args.seed) @@ -277,12 +294,6 @@ def main(): hypernet_config = get_hypernet_config( model, ctx_encoder_model, hypernet_args, aggregator_args ) - # hypernet = HyperLoRA( - # get_hypernet_config(model, hypernet_args, aggregator_args), - # model, - # ).to(model.device) - - # ctx_encoder = EarlyExit(get_base_model(model), ctx_encoder_args.layer_idx) if ctx_encoder_args.layer_idx is None: ctx_encoder_args.layer_idx = ctx_encoder_model.config.num_hidden_layers // 4 logger.info( @@ -497,14 +508,18 @@ def main(): elif isinstance(model, PeftModel): logger.info("Applying liger-kernel to PeftModel") _apply_liger_kernel_to_instance(model=model.base_model.model) - 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, - ) + + 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, + ) + else: + wandb.init(mode="disabled") train_model( model, diff --git a/hyperlora/model_loading.py b/hyperlora/model_loading.py index 48740a2..9f164c2 100644 --- a/hyperlora/model_loading.py +++ b/hyperlora/model_loading.py @@ -145,6 +145,13 @@ def get_model( for name, param in model.named_parameters(): if "modules_to_save" not in name: param.requires_grad = requires_grad + else: + # always train "modules_to_save" + if not "layernorm" in name: + raise NotImplementedError( + f"modules_to_save should only be layernorm, got {name}" + ) + param.requires_grad = True return model diff --git a/hyperlora/utils.py b/hyperlora/utils.py index 223405b..e30b2a9 100644 --- a/hyperlora/utils.py +++ b/hyperlora/utils.py @@ -4,6 +4,7 @@ import os import random import string import time +import hashlib from contextlib import contextmanager from typing import Iterable, Optional @@ -73,11 +74,17 @@ 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}" +def get_run_name(seed_str: Optional[str] = None): + if not seed_str: + 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}" + else: + # Generate a UUID from the seed string + hash_object = hashlib.sha256(seed_str.encode()) + uuid = hash_object.hexdigest()[:8] # Take the first 8 characters of the hash + run_name = seed_str + f"_{uuid}" return run_name