SentenceTransformerTrainer 通过 transformers 支持超参数优化。transformers 支持四种超参数搜索后端:Optuna、SigOpt、Ray Tune 和 W&B。使用之前,需要先安装所选后端。原文用下面这一行并列表示可选后端:
pip install optuna/sigopt/wandb/ray[tune]
这行中的斜线表示备选项,不是应当照抄执行的安装命令。本页使用 Optuna,应安装 optuna,例如使用 pip install optuna;同时按你的环境准备 Sentence Transformers 等依赖。
接下来将展示如何使用 optuna 后端进行超参数优化。其他后端的用法相近,更详细的信息请查看相应后端文档或 Transformers HPO 文档。
HPO 的组成部分
超参数优化包含以下四个部分:
| 组成部分 | 作用 |
|---|---|
| 超参数搜索空间 | 指定超参数取值范围。 |
| 模型初始化 | 为一次试验初始化 SentenceTransformer 模型。 |
| 损失函数初始化 | 根据模型初始化损失函数。 |
| 计算目标值 | 确定要最小化或最大化的数值。 |
超参数搜索空间
用一个函数定义超参数搜索空间:该函数返回字典,字典中包含超参数及各自的搜索范围。下面是使用 optuna 为 SentenceTransformer 模型定义搜索空间的例子:
def hpo_search_space(trial):
return {
"num_train_epochs": trial.suggest_int("num_train_epochs", 1, 2),
"per_device_train_batch_size": trial.suggest_int("per_device_train_batch_size", 32, 128),
"warmup_steps": trial.suggest_float("warmup_steps", 0, 0.3),
"learning_rate": trial.suggest_float("learning_rate", 1e-6, 1e-4, log=True),
}
核验补充:当前 Transformers TrainingArguments 文档说明,小于 1 的浮点型 warmup_steps 表示总训练步数的比例;整数则表示确切步数。因此这个示例的 0 到 0.3 搜索范围应结合所用版本的 API 解释,不能直接当作整数步数。
模型初始化
模型初始化函数接收当前 trial(一次试验)的超参数,并返回 SentenceTransformer 模型。这个函数通常很简单:
def hpo_model_init(trial):
return SentenceTransformer("distilbert/distilbert-base-uncased")
损失函数初始化
损失函数初始化函数接收当前试验初始化的模型,并返回损失函数。例如:
def hpo_loss_init(model):
return CosineSimilarityLoss(model)
这段展示函数接口的形状;后文的完整 AllNLI 示例使用 MultipleNegativesRankingLoss,并给出了相应导入。
计算目标值
目标函数接收评估得到的 metrics,返回需要最小化或最大化的浮点值。例如:
def hpo_compute_objective(metrics):
return metrics["eval_sts-dev_spearman_cosine"]
metrics 的字典键带有 eval_ 前缀。如果要最大化某个评估器的指标,评估器的 name 也会进入键名,与指标名通过连字符连接。另一个常见目标是 eval_loss。
原文在这里把评估器名称写成 name="stsb_dev",却同时给出 eval_sts-dev_spearman_cosine 这个目标键,二者不一致。下面的完整示例使用 name="sts-dev",与其目标键相对应。这里保留原文代码;迁移时应检查实际输出的 metrics 键名。
把各部分组合起来
任何常规训练循环都可以进行 HPO。区别是用 SentenceTransformerTrainer.hyperparameter_search 替代 SentenceTransformerTrainer.train。下面的完整示例串起了所有组件。
相关文档包括 AllNLI 数据集、EmbeddingSimilarityEvaluator、前面的搜索空间与初始化函数、SentenceTransformerTrainingArguments、SentenceTransformerTrainer 及其 hyperparameter_search 方法。
from sentence_transformers import SentenceTransformer, SentenceTransformerTrainer, SentenceTransformerTrainingArguments
from sentence_transformers.sentence_transformer.evaluation import EmbeddingSimilarityEvaluator, SimilarityFunction
from sentence_transformers.sentence_transformer.losses import MultipleNegativesRankingLoss
from sentence_transformers.sentence_transformer.training_args import BatchSamplers
from datasets import load_dataset
# 1. Load the AllNLI dataset: https://huggingface.co/datasets/sentence-transformers/all-nli, only 10k train and 1k dev
train_dataset = load_dataset("sentence-transformers/all-nli", "triplet", split="train[:10000]")
eval_dataset = load_dataset("sentence-transformers/all-nli", "triplet", split="dev[:1000]")
# 2. Create an evaluator to perform useful HPO
stsb_eval_dataset = load_dataset("sentence-transformers/stsb", split="validation")
dev_evaluator = EmbeddingSimilarityEvaluator(
sentences1=stsb_eval_dataset["sentence1"],
sentences2=stsb_eval_dataset["sentence2"],
scores=stsb_eval_dataset["score"],
main_similarity=SimilarityFunction.COSINE,
name="sts-dev",
)
# 3. Define the Hyperparameter Search Space
def hpo_search_space(trial):
return {
"num_train_epochs": trial.suggest_int("num_train_epochs", 1, 2),
"per_device_train_batch_size": trial.suggest_int("per_device_train_batch_size", 32, 128),
"warmup_steps": trial.suggest_float("warmup_steps", 0, 0.3),
"learning_rate": trial.suggest_float("learning_rate", 1e-6, 1e-4, log=True),
}
# 4. Define the Model Initialization
def hpo_model_init(trial):
return SentenceTransformer("distilbert/distilbert-base-uncased")
# 5. Define the Loss Initialization
def hpo_loss_init(model):
return MultipleNegativesRankingLoss(model)
# 6. Define the Objective Function
def hpo_compute_objective(metrics):
"""
Valid keys are: 'eval_loss', 'eval_sts-dev_pearson_cosine', 'eval_sts-dev_spearman_cosine',
'eval_sts-dev_pearson_manhattan', 'eval_sts-dev_spearman_manhattan', 'eval_sts-dev_pearson_euclidean',
'eval_sts-dev_spearman_euclidean', 'eval_sts-dev_pearson_dot', 'eval_sts-dev_spearman_dot',
'eval_sts-dev_pearson_max', 'eval_sts-dev_spearman_max', 'eval_runtime', 'eval_samples_per_second',
'eval_steps_per_second', 'epoch'
due to the evaluator that we're using.
"""
return metrics["eval_sts-dev_spearman_cosine"]
# 7. Define the training arguments
args = SentenceTransformerTrainingArguments(
# Required parameter:
output_dir="checkpoints",
# Optional training parameters:
# max_steps=10000, # We might want to limit the number of steps for HPO
fp16=True, # Set to False if you get an error that your GPU can't run on FP16
bf16=False, # Set to True if you have a GPU that supports BF16
batch_sampler=BatchSamplers.NO_DUPLICATES, # MultipleNegativesRankingLoss benefits from no duplicate samples in a batch
# Optional tracking/debugging parameters:
eval_strategy="no", # We don't need to evaluate/save during HPO
save_strategy="no",
logging_steps=10,
run_name="hpo", # Will be used in W&B if `wandb` is installed
)
# 8. Create the trainer with model_init rather than model
trainer = SentenceTransformerTrainer(
model=None,
args=args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
evaluator=dev_evaluator,
model_init=hpo_model_init,
loss=hpo_loss_init,
)
# 9. Perform the HPO
best_trial = trainer.hyperparameter_search(
hp_space=hpo_search_space,
compute_objective=hpo_compute_objective,
n_trials=20,
direction="maximize",
backend="optuna",
)
print(best_trial)
示例会加载 10,000 条训练数据、1,000 条开发集数据和 STSb 验证集,为每次试验重新初始化模型与损失,搜索 20 次,并最大化验证集余弦相似度对应的 Spearman 相关系数。以下是原文在 2024 年 5 月 17 日记录的输出,包含 20 次试验和最终 BestRun,本稿没有执行这些训练或复测数值:
[I 2024-05-17 15:10:47,844] Trial 0 finished with value: 0.7889856589698055 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 123, 'warmup_steps': 0.07380948785410107, 'learning_rate': 2.686331417509812e-06}. Best is trial 0 with value: 0.7889856589698055.
[I 2024-05-17 15:12:13,283] Trial 1 finished with value: 0.7927780672090986 and parameters: {'num_train_epochs': 2, 'per_device_train_batch_size': 69, 'warmup_steps': 0.2927897848007451, 'learning_rate': 5.885372118095137e-06}. Best is trial 1 with value: 0.7927780672090986.
[I 2024-05-17 15:12:43,896] Trial 2 finished with value: 0.7684829743509601 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 114, 'warmup_steps': 0.0739429232666916, 'learning_rate': 7.344415188959276e-05}. Best is trial 1 with value: 0.7927780672090986.
[I 2024-05-17 15:14:49,730] Trial 3 finished with value: 0.7873032743147989 and parameters: {'num_train_epochs': 2, 'per_device_train_batch_size': 43, 'warmup_steps': 0.15184370143796674, 'learning_rate': 9.703232080395476e-06}. Best is trial 1 with value: 0.7927780672090986.
[I 2024-05-17 15:15:39,597] Trial 4 finished with value: 0.7759251781929949 and parameters: {'num_train_epochs': 2, 'per_device_train_batch_size': 127, 'warmup_steps': 0.263946220093495, 'learning_rate': 1.231454337152625e-06}. Best is trial 1 with value: 0.7927780672090986.
[I 2024-05-17 15:17:02,191] Trial 5 finished with value: 0.7964580509886684 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 34, 'warmup_steps': 0.2276865359631089, 'learning_rate': 7.889007438884571e-06}. Best is trial 5 with value: 0.7964580509886684.
[I 2024-05-17 15:18:55,559] Trial 6 finished with value: 0.7901878917859169 and parameters: {'num_train_epochs': 2, 'per_device_train_batch_size': 48, 'warmup_steps': 0.23228838664572948, 'learning_rate': 2.883013292682523e-06}. Best is trial 5 with value: 0.7964580509886684.
[I 2024-05-17 15:20:27,027] Trial 7 finished with value: 0.7935671067660925 and parameters: {'num_train_epochs': 2, 'per_device_train_batch_size': 62, 'warmup_steps': 0.22061123927198237, 'learning_rate': 2.95413457610349e-06}. Best is trial 5 with value: 0.7964580509886684.
[I 2024-05-17 15:22:23,147] Trial 8 finished with value: 0.7848123114933252 and parameters: {'num_train_epochs': 2, 'per_device_train_batch_size': 45, 'warmup_steps': 0.23071701022961139, 'learning_rate': 9.793681667449783e-06}. Best is trial 5 with value: 0.7964580509886684.
[I 2024-05-17 15:22:52,826] Trial 9 finished with value: 0.7909708416168918 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 121, 'warmup_steps': 0.22440506724181647, 'learning_rate': 4.0744671365843346e-05}. Best is trial 5 with value: 0.7964580509886684.
[I 2024-05-17 15:23:30,395] Trial 10 finished with value: 0.7928991732385567 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 89, 'warmup_steps': 0.14607293301068847, 'learning_rate': 2.5557492055039498e-05}. Best is trial 5 with value: 0.7964580509886684.
[I 2024-05-17 15:24:18,024] Trial 11 finished with value: 0.7991870087507459 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 66, 'warmup_steps': 0.16886154348739527, 'learning_rate': 3.705926066938032e-06}. Best is trial 11 with value: 0.7991870087507459.
[I 2024-05-17 15:25:44,198] Trial 12 finished with value: 0.7923304174306207 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 33, 'warmup_steps': 0.15953772535423974, 'learning_rate': 1.8076298025704224e-05}. Best is trial 11 with value: 0.7991870087507459.
[I 2024-05-17 15:26:20,739] Trial 13 finished with value: 0.8020260244040395 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 90, 'warmup_steps': 0.18105202625281253, 'learning_rate': 5.513908793512551e-06}. Best is trial 13 with value: 0.8020260244040395.
[I 2024-05-17 15:26:57,783] Trial 14 finished with value: 0.7571110256860063 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 95, 'warmup_steps': 0.00122391151793258, 'learning_rate': 1.0432486633629492e-06}. Best is trial 13 with value: 0.8020260244040395.
[I 2024-05-17 15:27:32,581] Trial 15 finished with value: 0.8009013936824717 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 101, 'warmup_steps': 0.1761274711346081, 'learning_rate': 4.5918293464430035e-06}. Best is trial 13 with value: 0.8020260244040395.
[I 2024-05-17 15:28:05,850] Trial 16 finished with value: 0.8017668050806169 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 103, 'warmup_steps': 0.10766501647726355, 'learning_rate': 5.0309795522333e-06}. Best is trial 13 with value: 0.8020260244040395.
[I 2024-05-17 15:28:37,393] Trial 17 finished with value: 0.7769412380909586 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 108, 'warmup_steps': 0.1036610178950246, 'learning_rate': 1.7747598626081271e-06}. Best is trial 13 with value: 0.8020260244040395.
[I 2024-05-17 15:29:19,340] Trial 18 finished with value: 0.8011921300048339 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 80, 'warmup_steps': 0.117014165550441, 'learning_rate': 1.238558867958792e-05}. Best is trial 13 with value: 0.8020260244040395.
[I 2024-05-17 15:29:59,508] Trial 19 finished with value: 0.8027501854704168 and parameters: {'num_train_epochs': 1, 'per_device_train_batch_size': 84, 'warmup_steps': 0.014601112207929548, 'learning_rate': 5.627813947769514e-06}. Best is trial 19 with value: 0.8027501854704168.
BestRun(run_id='19', objective=0.8027501854704168, hyperparameters={'num_train_epochs': 1, 'per_device_train_batch_size': 84, 'warmup_steps': 0.014601112207929548, 'learning_rate': 5.627813947769514e-06}, run_summary=None)
在原文的这一实验中,最优超参数在 STS 开发集上取得约 0.802 的 Spearman 相关系数。作为比较,默认训练参数 per_device_train_batch_size=8、learning_rate=5e-5 的结果是 0.736;依据经验选择 per_device_train_batch_size=64、learning_rate=2e-5 的结果是 0.783。这说明 HPO 在该实验中有效改善了模型表现;这些是原文报告的结果,不能当作其他数据集或设备的保证。
应用边界补充:20 次试验会重复进行训练并消耗计算资源。选定超参数后,应重新训练,并在未用于调参的独立测试集上确认效果。
示例脚本
hpo_nli.py 是在 AllNLI 数据集上进行超参数优化的示例脚本。











暂无评论内容