LESS 核心思想
此前我们详细介绍了LESS方法的原理——通过仅选择5%有影响力的数据,即可超越全量指令微调的效果。今天直接落地实践,看看这套流水线具体怎么跑通。

LESS 核心思想
LESS的核心思路是:给定少量体现特定能力的示例,从海量指令数据中精准筛选出5%最具影响力的样本,用于目标微调。结果不仅优于全量数据集,而且选出的子集在不同模型规模和系列中都能保持有效。
数据选择流水线
- 使用LoRA进行热身训练。
- 构建投影低维梯度特征的梯度数据存储,可供不同任务重复使用。
- 利用数据选择算法从存储中构建训练数据集。
- 用选出的数据训练模型。
实验关键结果
- LESS在不同模型上均有效。
- 仅选择5%的数据,效果通常优于完整数据集。
- 用小模型选择的数据,能提升大模型及不同模型族的性能。
LESS的局限性
- 需要先用候选数据的随机5%进行热身训练,这对获取有效梯度特征至关重要,但增加了计算开销。
- 默认使用补全token的平均梯度,会偏向短序列,导致性能下降。缓解方法是归一化梯度特征,用余弦相似度替代点积来估计影响。
- 最小化验证损失(交叉熵)并不总能单调提升准确率。
- 一阶近似忽略了多个数据点叠加的影响——比如两个重复点得分都高,但未必能带来双重收益。
LESS 应用
环境安装
pip3 install torch==2.1.2 torchvision torchaudio
cd LESS
pip install -r requirement.txt
# 以可编辑模式安装 `less` 包
pip install -e .
数据准备
按照open-instruct库准备指令调优数据集。这里组合使用四个训练集:Flan v2、COT、Dolly和Open Assistant。评估阶段另加三个数据集:MMLU、Tydiqa和BBH。下面提供这些文件的处理版本。
数据选择流水线
1. 热身训练
这是提升数据选择性能的关键步骤。取整个数据集的一小部分,用LoRA进行训练。
DATA_DIR=../data
MODEL_PATH=meta-llama/Llama-2-7b-hf
PERCENTAGE=0.05 # 训练数据占比,可在脚本内指定具体文件
DATA_SEED=3
JOB_NAME=llama2-7b-p${PERCENTAGE}-lora-seed${DATA_SEED}
./less/scripts/train/warmup_lora_train.sh "$DATA_DIR" "$MODEL_PATH" "$PERCENTAGE" "$DATA_SEED" "$JOB_NAME"
2. 构建梯度数据存储
热身训练完成后,收集整个训练数据集的梯度。每个检查点都需获取目标训练数据的梯度。
CKPT=105
TRAINING_DATA_NAME=dolly
TRAINING_DATA_FILE=../data/train/processed/dolly/dolly_data.jsonl # 更换数据时同步修改路径
GRADIENT_TYPE="adam"
MODEL_PATH=../out/llama2-7b-p0.05-lora-seed3/checkpoint-${CKPT}
OUTPUT_PATH=../grads/llama2-7b-p0.05-lora-seed3/${TRAINING_DATA_NAME}-ckpt${CKPT}-${GRADIENT_TYPE}
DIMS="8192"
./less/scripts/get_info/get_train_lora_grads.sh
"$TRAINING_DATA_FILE"
"$MODEL_PATH"
"$OUTPUT_PATH"
"$DIMS"
"$GRADIENT_TYPE"
这样就创建了一个数据存储,包含了后续选择所需的所有检查点和训练数据的梯度。
3. 为任务选择数据
针对特定下游任务选择数据前,先用与训练时相同的指令调优提示格式,准备该任务的数据。这里已为BBH、TydiQA和MMLU设置了数据加载模块。如需其他任务,可扩展less/data_selection/get_validation_dataset.py脚本。
获取验证数据梯度的过程与训练数据类似,区别在于此处生成的是用于影响力估计的SGD梯度。
CKPT=105
TASK=tydiqa
MODEL_PATH=../out/llama2-7b-p0.05-lora-seed3/checkpoint-${CKPT}
OUTPUT_PATH=../grads/llama2-7b-p0.05-lora-seed3/${TASK}-ckpt${CKPT}-sgd # 验证数据统一使用sgd
DATA_DIR=../data
DIMS="4096 8192" # 默认投影维度为8192
./less/scripts/get_info/get_eval_lora_grads.sh "$TASK" "$DATA_DIR" "$MODEL_PATH" $OUTPUT_PATH "$DIMS"
正常来说,需要获得上一步中所有检查点的验证数据梯度。拿到之后,即可运行以下脚本计算每个训练数据点的影响力得分,并选出得分最高的前k个。
DIM=8192
CKPTS="105 211 317 420"
CHECKPOINT_WEIGHTS="1.6877e-05 1.2859e-05 7.7030e-06 2.5616e-06"
GRADIENT_PATH=../grads/llama2-7b-p0.05-lora-seed3/{}-ckpt{}-adam/dim${DIM}
TRAIN_FILE_NAMES="flan_v2 cot dolly oasst1"
VALIDATION_GRADIENT_PATH=../grads/llama2-7b-p0.05-lora-seed3/{}-ckpt{}-sgd/dim${DIM}
TARGET_TASK_NAMES="tydiqa"
SELECTED_DATA_OUTPUT_PATH="../selected_data"
./less/scripts/data_selection/matching.sh
"$GRADIENT_PATH"
"$TRAIN_FILE_NAMES"
"$CKPTS"
"$CHECKPOINT_WEIGHTS"
"$VALIDATION_GRADIENT_PATH"
"$TARGET_TASK_NAMES"
"$SELECTED_DATA_OUTPUT_PATH"
每个训练数据点的影响力得分会保存在OUTPUT_PATH目录下,再用下面这个脚本选出top k。
python3 -m less.data_selection.write_selected_data
--target_task_names ${TARGET_TASK_NAMES}
--train_file_names ${TRAIN_FILE_NAMES}
--train_files ../data/train/processed/dolly/dolly_data.jsonl ../data/train/processed/oasst1/oasst1_data.jsonl
--output_path $SELECTED_DATA_OUTPUT_PATH
--percentage 0.05
4. 使用选择的数据进行训练
选好数据后,运行以下脚本进行模型训练:
TARGET_TASK_NAME="tydiqa"
PERCENTAGE=0.05
TRAIN_FILES=../selected_data/${TARGET_TASK_NAME}/top_p${PERCENTAGE}.jsonl
MODEL_PATH=meta-llama/Llama-2-7b-hf
JOB_NAME=llama2-7b-less-p${PERCENTAGE}-lora
./less/scripts/train/lora_train.sh "$TRAIN_FILES" "$MODEL_PATH" "$JOB_NAME"
注意:若想全参数微调,只需去掉LoRA训练参数。
评估
使用MMLU、Tydiqa和BBH三个评估数据集来检验数据选择流水线的效果。评估依赖open-instruct库,具体步骤如下:
1:安装 Open-Instruct
git clone https://github.com/allenai/open-instruct.git
cd open-instruct
pip install -e .
2:评估
evaluation目录下提供了三个评估脚本:eval_mmlu.sh、eval_tydiqa.sh和eval_bbh.sh。下面是eval_bbh.sh的示例:
source eval.sh
eval_bbh() {
cd $n/space10/open-instruct
mdir=$1
type=$2
set_sa ve_dir $mdir bbh
mkdir -p $sa ve_dir
cmd="python -m eval.bbh.run_eval
--data_dir $DATA_DIR/bbh
--sa ve_dir $sa ve_dir
--model $mdir
--tokenizer $mdir
--eval_batch_size 10
--convert_to_bf16
--max_num_examples_per_task 40"
eval "$cmd"
}
valid_bbh() {
cd $n/space10/open-instruct
mdir=$1
type=$2
set_valid_dir $mdir bbh
echo $sa ve_dir
mkdir -p $sa ve_dir
cmd="python -m eval.bbh.run_eval
--data_dir $DATA_DIR/bbh-valid
--sa ve_dir $sa ve_dir
--model $mdir
--tokenizer $mdir
--eval_batch_size 10
--convert_to_bf16
--eval_valid
--max_num_examples_per_task 3"
}
extract_bbh() {
mdir=$1
set_sa ve_dir $mdir bbh-nonchat
result=$(jq .a verage_exact_match $sa ve_dir/metrics.json)
result=$(echo "$result * 100" | bc)
echo $result
}
extract_valid_bbh() {
mdir=$1
set_valid_dir $mdir bbh-nonchat
result=$(jq .a verage_exact_match $sa ve_dir/metrics.json)
result=$(echo "$result * 100" | bc)
echo $result
}
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名