Add SWIFT coding agent probe experiment scripts
This commit is contained in:
88
scripts/swift_train_common.sh
Executable file
88
scripts/swift_train_common.sh
Executable file
@@ -0,0 +1,88 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
|
||||
cd "${ROOT_DIR}"
|
||||
|
||||
export http_proxy="${http_proxy:-http://100.72.0.101:8888}"
|
||||
export https_proxy="${https_proxy:-http://100.72.0.101:8888}"
|
||||
export HTTP_PROXY="${HTTP_PROXY:-${http_proxy}}"
|
||||
export HTTPS_PROXY="${HTTPS_PROXY:-${https_proxy}}"
|
||||
export HF_ENDPOINT="${HF_ENDPOINT:-https://hf-mirror.com}"
|
||||
export TOKENIZERS_PARALLELISM="${TOKENIZERS_PARALLELISM:-false}"
|
||||
|
||||
if [[ -f .venv/bin/activate ]]; then
|
||||
source .venv/bin/activate
|
||||
elif [[ "${DRY_RUN:-0}" != "1" ]]; then
|
||||
echo "Missing .venv. Run ./scripts/setup_env.sh first." >&2
|
||||
exit 2
|
||||
fi
|
||||
mkdir -p outputs runs logs
|
||||
|
||||
TRAIN_JSONL="${TRAIN_JSONL:-data/processed/training_probe/train.jsonl}"
|
||||
VAL_JSONL="${VAL_JSONL:-data/processed/training_probe/validation.jsonl}"
|
||||
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}"
|
||||
LORA_RANK="${LORA_RANK:-32}"
|
||||
|
||||
require_file() {
|
||||
if [[ ! -f "$1" ]]; then
|
||||
echo "Missing required file: $1" >&2
|
||||
exit 2
|
||||
fi
|
||||
}
|
||||
|
||||
run_swift_train() {
|
||||
local model_path="$1"
|
||||
local train_type="$2"
|
||||
local run_name="$3"
|
||||
local output_dir="outputs/${run_name}"
|
||||
local tb_dir="runs/${run_name}"
|
||||
local log_file="logs/${run_name}.log"
|
||||
|
||||
require_file "${TRAIN_JSONL}"
|
||||
require_file "${VAL_JSONL}"
|
||||
mkdir -p "${output_dir}" "${tb_dir}" logs
|
||||
|
||||
local cmd=(
|
||||
swift sft
|
||||
--model "${model_path}"
|
||||
--dataset "${TRAIN_JSONL}"
|
||||
--val_dataset "${VAL_JSONL}"
|
||||
--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}"
|
||||
--warmup_ratio "${WARMUP_RATIO}"
|
||||
--max_length "${MAX_LENGTH}"
|
||||
--save_steps "${SAVE_STEPS}"
|
||||
--eval_steps "${EVAL_STEPS}"
|
||||
--logging_steps "${LOGGING_STEPS}"
|
||||
--report_to tensorboard
|
||||
--logging_dir "${tb_dir}"
|
||||
--output_dir "${output_dir}"
|
||||
--save_total_limit "${SAVE_TOTAL_LIMIT:-3}"
|
||||
--dataloader_num_workers "${DATALOADER_NUM_WORKERS:-4}"
|
||||
)
|
||||
|
||||
if [[ "${train_type}" == "lora" ]]; then
|
||||
cmd+=(--lora_rank "${LORA_RANK}")
|
||||
fi
|
||||
|
||||
printf '%q ' "${cmd[@]}" | tee "${log_file}.cmd"
|
||||
echo
|
||||
if [[ "${DRY_RUN:-0}" == "1" ]]; then
|
||||
return 0
|
||||
fi
|
||||
"${cmd[@]}" 2>&1 | tee "${log_file}"
|
||||
}
|
||||
Reference in New Issue
Block a user