ืnew default + save best model

This commit is contained in:
51616 2025-01-05 13:25:59 +00:00
parent c0cf32c8af
commit 2d23c7dc0d
6 changed files with 121 additions and 48 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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