From 2d23c7dc0d64a0540163ad4f560e8594b9ad223b Mon Sep 17 00:00:00 2001 From: 51616 Date: Sun, 5 Jan 2025 13:25:59 +0000 Subject: [PATCH] =?UTF-8?q?=E0=B8=B7new=20default=20+=20save=20best=20mode?= =?UTF-8?q?l?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- configs/context_numbers_10.yaml | 28 +++++++------ configs/context_numbers_128.yaml | 28 ++++++------- configs/context_numbers_256.yaml | 28 ++++++------- hyperlora/configs.py | 68 +++++++++++++++++++++++++++++++- hyperlora/intx_sft.py | 5 ++- hyperlora/training_utils.py | 12 ++++-- 6 files changed, 121 insertions(+), 48 deletions(-) diff --git a/configs/context_numbers_10.yaml b/configs/context_numbers_10.yaml index dcd9710..87cdeb4 100644 --- a/configs/context_numbers_10.yaml +++ b/configs/context_numbers_10.yaml @@ -2,28 +2,30 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] -eval_on_start: True -eval_strategy: "steps" -eval_steps: 500 -save_strategy: "no" -# save_steps: 500 -logging_strategy: "steps" -logging_steps: 100 -use_liger_kernel: true -remove_unused_columns: false +# eval_on_start: True +# eval_strategy: "steps" +# eval_steps: 500 +# save_strategy: "no" +# # save_steps: 500 +# logging_strategy: "steps" +# logging_steps: 100 +# use_liger_kernel: true +# remove_unused_columns: false # needed to avoid OOM by compute the metrics batch by batch # w/o this the trainer stores logits of all sample in memory... -batch_eval_metrics: true +# batch_eval_metrics: true -per_device_train_batch_size: 128 +per_device_train_batch_size: 64 per_device_eval_batch_size: 128 +max_new_tokens: 64 +gen_per_device_eval_batch_size: 128 # optim: schedule_free_adamw -learning_rate: 0.00001 +learning_rate: 0.0001 # lr_scheduler_type: "constant_with_warmup" neftune_noise_alpha: 1 weight_decay: 0.1 -warmup_ratio: 0.1 +warmup_ratio: 0.05 # LoRA lora_r: 16 diff --git a/configs/context_numbers_128.yaml b/configs/context_numbers_128.yaml index 3d5e073..8c3c06c 100644 --- a/configs/context_numbers_128.yaml +++ b/configs/context_numbers_128.yaml @@ -2,28 +2,28 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] -eval_on_start: True -eval_strategy: "steps" -eval_steps: 500 -save_strategy: "no" -# save_steps: 500 -logging_strategy: "steps" -logging_steps: 100 -use_liger_kernel: true -remove_unused_columns: false +# eval_on_start: True +# eval_strategy: "steps" +# eval_steps: 500 +# save_strategy: "no" +# # save_steps: 500 +# logging_strategy: "steps" +# logging_steps: 100 +# use_liger_kernel: true +# remove_unused_columns: false # needed to avoid OOM by compute the metrics batch by batch # w/o this the trainer stores logits of all sample in memory... -batch_eval_metrics: true +# batch_eval_metrics: true -per_device_train_batch_size: 128 -per_device_eval_batch_size: 128 +per_device_train_batch_size: 64 +per_device_eval_batch_size: 64 # optim: schedule_free_adamw -learning_rate: 0.00001 +learning_rate: 0.0001 # lr_scheduler_type: "constant_with_warmup" neftune_noise_alpha: 1 weight_decay: 0.1 -warmup_ratio: 0.1 +warmup_ratio: 0.05 # LoRA lora_r: 16 diff --git a/configs/context_numbers_256.yaml b/configs/context_numbers_256.yaml index 6cd940c..2f4b079 100644 --- a/configs/context_numbers_256.yaml +++ b/configs/context_numbers_256.yaml @@ -2,28 +2,28 @@ output_dir: "" # just a placeholder bf16: true model_name_or_path: meta-llama/Llama-3.2-1B-Instruct label_names: ["labels"] -eval_on_start: True -eval_strategy: "steps" -eval_steps: 500 -save_strategy: "no" -# save_steps: 500 -logging_strategy: "steps" -logging_steps: 100 -use_liger_kernel: true -remove_unused_columns: false +# eval_on_start: True +# eval_strategy: "steps" +# eval_steps: 500 +# save_strategy: "no" +# # save_steps: 500 +# logging_strategy: "steps" +# logging_steps: 100 +# use_liger_kernel: true +# remove_unused_columns: false # needed to avoid OOM by compute the metrics batch by batch # w/o this the trainer stores logits of all sample in memory... -batch_eval_metrics: true +# batch_eval_metrics: true -per_device_train_batch_size: 128 -per_device_eval_batch_size: 128 +per_device_train_batch_size: 32 +per_device_eval_batch_size: 32 # optim: schedule_free_adamw -learning_rate: 0.00001 +learning_rate: 0.0001 # lr_scheduler_type: "constant_with_warmup" neftune_noise_alpha: 1 weight_decay: 0.1 -warmup_ratio: 0.1 +warmup_ratio: 0.05 # LoRA lora_r: 16 diff --git a/hyperlora/configs.py b/hyperlora/configs.py index 28a4d74..752a4a0 100644 --- a/hyperlora/configs.py +++ b/hyperlora/configs.py @@ -6,7 +6,7 @@ from enum import Enum, auto from typing import Any, Dict, List, Literal, NewType, Optional, Tuple import yaml -from transformers import MODEL_FOR_CAUSAL_LM_MAPPING, HfArgumentParser +from transformers import MODEL_FOR_CAUSAL_LM_MAPPING, HfArgumentParser, TrainingArguments MODEL_CONFIG_CLASSES = list(MODEL_FOR_CAUSAL_LM_MAPPING.keys()) MODEL_TYPES = tuple(conf.model_type for conf in MODEL_CONFIG_CLASSES) @@ -119,6 +119,64 @@ class ExperimentSetup(str, Enum): FULL_FINETUNE = "full_finetune" +@dataclass +class TrainingArguments(TrainingArguments): + eval_on_start: bool = field( + default=True, + metadata={"help": "Whether to evaluate on the start of training."}, + ) + eval_strategy: str = field( + default="steps", + metadata={"help": "Evaluation strategy."}, + ) + eval_steps: int = field( + default=10_000, + metadata={"help": "Evaluation steps."}, + ) + metric_for_best_model: str = field( + default="val_loss", + metadata={"help": "Metric for best model."}, + ) + greater_is_better: bool = field( + default=False, + metadata={"help": "Whether the metric is better when it is greater."}, + ) + load_best_model_at_end: bool = field( + default=True, + metadata={"help": "Whether to load the best model at the end of training."}, + ) + save_total_limit: int = field( + default=1, + metadata={"help": "Total number of checkpoints to save."}, + ) + save_strategy: str = field( + default="steps", + ) + save_steps: int = field( + default=10_000, + ) + save_safetensors: bool = field( + default=False, + ) + logging_strategy: str = field( + default="steps", + ) + logging_steps: int = field( + default=100, + ) + use_liger_kernel: bool = field( + default=True, + ) + remove_unused_columns: bool = field( + default=False, + ) + # needed to avoid OOM by compute the metrics batch by batch + # w/o this the trainer stores logits of all sample in memory... + batch_eval_metrics: bool = field( + default=True, + ) + + @dataclass class ModelArguments: """ @@ -161,6 +219,14 @@ class CtxTrainingArguments: default=2**13, metadata={"help": "Maximum base length for training."}, ) + max_new_tokens: Optional[int] = field( + default=2**13, + metadata={"help": "Maximum new tokens for generation-based evaluation."}, + ) + gen_per_device_eval_batch_size: Optional[int] = field( + default=1, + metadata={"help": "Per device evaluation batch size for generation."}, + ) @dataclass diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index c2928d9..2b72c42 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -36,7 +36,6 @@ from transformers import ( DataCollatorForSeq2Seq, EvalPrediction, HfArgumentParser, - TrainingArguments, ) from utils import ( extract_cli_args, @@ -51,6 +50,7 @@ from utils import ( from configs import ( ArgumentParser, CtxTrainingArguments, + TrainingArguments, DataArguments, ExperimentSetup, LoRAArguments, @@ -221,7 +221,6 @@ def main(output_dir): training_args.run_name = run_name training_args.output_dir = output_dir training_args.logging_dir = output_dir - training_args.save_safetensors = False logger.debug(f"args: {args}") save_yaml(args, f"{output_dir}/args.yaml") @@ -402,6 +401,8 @@ def main(output_dir): [compute_per_token_acc, compute_prefix_matching, compute_entropy] ), ), + max_new_tokens=ctx_args.max_new_tokens, + gen_per_device_eval_batch_size=ctx_args.gen_per_device_eval_batch_size, # compute_metrics, # preprocess_logits_for_metrics, ) diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index 142533d..6990be5 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -113,6 +113,8 @@ def train_model( compute_metrics=None, preprocess_logits_for_metrics=None, per_sample_loss_avg=False, + max_new_tokens=2**13, + gen_per_device_eval_batch_size=1, ): # last_checkpoint = None @@ -171,7 +173,8 @@ def train_model( # print(f"Loaded from the checkpoint: {checkpoint}") - # TODO: save the best model based on eval loss? + # Trainer loads the best model after training + # is done when load_best_model_at_end=True (our default) train_result = trainer.train(resume_from_checkpoint=checkpoint) trainer.log_metrics("train", train_result.metrics) metrics = trainer.evaluate(dict(**val_dataset, test=test_dataset)) @@ -182,8 +185,6 @@ def train_model( clear_gpu() ############## Evaluation - # TODO: generalize gen_kwargs for validation - max_new_tokens = 100 # max_input_len=2**13 # for input truncation gen_kwargs = dict(do_sample=False, max_new_tokens=max_new_tokens) @@ -197,8 +198,11 @@ def train_model( for attr, value in training_args.to_dict().items(): if attr in seq2seq_training_args_fields: eval_trainer_args[attr] = value + eval_trainer_args["eval_strategy"] = "no" + eval_trainer_args["save_strategy"] = "no" eval_trainer_args["overwrite_output_dir"] = True + eval_trainer_args["per_device_eval_batch_size"] = gen_per_device_eval_batch_size # NOTE: could also set kv_cache implementation here eval_trainer_args = Seq2SeqTrainingArguments( @@ -232,6 +236,6 @@ def train_model( val_dataset = val_dataset["val"] eval_generation(eval_trainer, tokenizer, val_dataset, "val", gen_kwargs) clear_gpu() - + if test_dataset is not None: eval_generation(eval_trainer, tokenizer, test_dataset, "test", gen_kwargs)