首页 > 教程攻略 > ai资讯 >LESS 核心思想

LESS 核心思想

来源:互联网 时间:2026-08-02 13:24:29

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

LESS 核心思想

LESS 核心思想

LESS的核心思路是:给定少量体现特定能力的示例,从海量指令数据中精准筛选出5%最具影响力的样本,用于目标微调。结果不仅优于全量数据集,而且选出的子集在不同模型规模和系列中都能保持有效。

数据选择流水线

  1. 使用LoRA进行热身训练。
  2. 构建投影低维梯度特征的梯度数据存储,可供不同任务重复使用。
  3. 利用数据选择算法从存储中构建训练数据集。
  4. 用选出的数据训练模型。

实验关键结果

  1. LESS在不同模型上均有效。
  2. 仅选择5%的数据,效果通常优于完整数据集。
  3. 用小模型选择的数据,能提升大模型及不同模型族的性能。

LESS的局限性

  1. 需要先用候选数据的随机5%进行热身训练,这对获取有效梯度特征至关重要,但增加了计算开销。
  2. 默认使用补全token的平均梯度,会偏向短序列,导致性能下降。缓解方法是归一化梯度特征,用余弦相似度替代点积来估计影响。
  3. 最小化验证损失(交叉熵)并不总能单调提升准确率。
  4. 一阶近似忽略了多个数据点叠加的影响——比如两个重复点得分都高,但未必能带来双重收益。

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.sheval_tydiqa.sheval_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
}