Tune default SFT and LoRA hyperparameters
This commit is contained in:
@@ -25,12 +25,27 @@ MAX_LENGTH="${MAX_LENGTH:-262144}"
|
||||
SAVE_STEPS="${SAVE_STEPS:-1000}"
|
||||
EVAL_STEPS="${EVAL_STEPS:-1000}"
|
||||
LOGGING_STEPS="${LOGGING_STEPS:-1}"
|
||||
GRAD_ACCUM_STEPS="${GRAD_ACCUM_STEPS:-1}"
|
||||
PER_DEVICE_BATCH_SIZE="${PER_DEVICE_BATCH_SIZE:-1}"
|
||||
NUM_EPOCHS="${NUM_EPOCHS:-1}"
|
||||
LEARNING_RATE="${LEARNING_RATE:-1e-5}"
|
||||
WARMUP_RATIO="${WARMUP_RATIO:-0.1}"
|
||||
LR_SCHEDULER_TYPE="${LR_SCHEDULER_TYPE:-cosine}"
|
||||
LORA_RANK="${LORA_RANK:-32}"
|
||||
DEFAULT_PER_DEVICE_BATCH_SIZE="${DEFAULT_PER_DEVICE_BATCH_SIZE:-1}"
|
||||
DEFAULT_GRAD_ACCUM_STEPS="${DEFAULT_GRAD_ACCUM_STEPS:-1}"
|
||||
DEFAULT_EVAL_BATCH_SIZE="${DEFAULT_EVAL_BATCH_SIZE:-1}"
|
||||
|
||||
env_key() {
|
||||
printf '%s' "$1" | tr '[:lower:]-' '[:upper:]_' | sed 's/[^A-Z0-9_]/_/g'
|
||||
}
|
||||
|
||||
env_or_default() {
|
||||
local name="$1"
|
||||
local fallback="$2"
|
||||
if [[ -n "${!name:-}" ]]; then
|
||||
printf '%s' "${!name}"
|
||||
else
|
||||
printf '%s' "${fallback}"
|
||||
fi
|
||||
}
|
||||
|
||||
require_file() {
|
||||
if [[ ! -f "$1" ]]; then
|
||||
@@ -46,6 +61,28 @@ run_swift_train() {
|
||||
local output_dir="outputs/${run_name}"
|
||||
local tb_dir="runs/${run_name}"
|
||||
local log_file="logs/${run_name}.log"
|
||||
local run_key
|
||||
run_key="$(env_key "${run_name}")"
|
||||
|
||||
local default_lr
|
||||
if [[ "${train_type}" == "lora" ]]; then
|
||||
default_lr="${LORA_LEARNING_RATE:-5e-5}"
|
||||
else
|
||||
default_lr="${FULL_LEARNING_RATE:-1e-5}"
|
||||
fi
|
||||
local learning_rate
|
||||
learning_rate="$(env_or_default "${run_key}_LEARNING_RATE" "${LEARNING_RATE:-${default_lr}}")"
|
||||
|
||||
local type_key
|
||||
type_key="$(env_key "${train_type}")"
|
||||
local type_bsz_var="${type_key}_PER_DEVICE_BATCH_SIZE"
|
||||
local type_accum_var="${type_key}_GRAD_ACCUM_STEPS"
|
||||
local per_device_batch_size
|
||||
local grad_accum_steps
|
||||
local eval_batch_size
|
||||
per_device_batch_size="$(env_or_default "${run_key}_PER_DEVICE_BATCH_SIZE" "${PER_DEVICE_BATCH_SIZE:-${!type_bsz_var:-${DEFAULT_PER_DEVICE_BATCH_SIZE}}}")"
|
||||
grad_accum_steps="$(env_or_default "${run_key}_GRAD_ACCUM_STEPS" "${GRAD_ACCUM_STEPS:-${!type_accum_var:-${DEFAULT_GRAD_ACCUM_STEPS}}}")"
|
||||
eval_batch_size="$(env_or_default "${run_key}_EVAL_BATCH_SIZE" "${EVAL_PER_DEVICE_BATCH_SIZE:-${DEFAULT_EVAL_BATCH_SIZE}}")"
|
||||
|
||||
require_file "${TRAIN_JSONL}"
|
||||
require_file "${VAL_JSONL}"
|
||||
@@ -59,11 +96,12 @@ run_swift_train() {
|
||||
--train_type "${train_type}"
|
||||
--torch_dtype bfloat16
|
||||
--num_train_epochs "${NUM_EPOCHS}"
|
||||
--per_device_train_batch_size "${PER_DEVICE_BATCH_SIZE}"
|
||||
--per_device_eval_batch_size 1
|
||||
--gradient_accumulation_steps "${GRAD_ACCUM_STEPS}"
|
||||
--learning_rate "${LEARNING_RATE}"
|
||||
--per_device_train_batch_size "${per_device_batch_size}"
|
||||
--per_device_eval_batch_size "${eval_batch_size}"
|
||||
--gradient_accumulation_steps "${grad_accum_steps}"
|
||||
--learning_rate "${learning_rate}"
|
||||
--warmup_ratio "${WARMUP_RATIO}"
|
||||
--lr_scheduler_type "${LR_SCHEDULER_TYPE}"
|
||||
--max_length "${MAX_LENGTH}"
|
||||
--save_steps "${SAVE_STEPS}"
|
||||
--eval_steps "${EVAL_STEPS}"
|
||||
|
||||
Reference in New Issue
Block a user