You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何以官方方式使用wandb Sweep结合Hugging Face Transformers并兼容HF特性?

Hugging Face Transformers 与 wandb Sweep 官方结合方案

核心实现思路

合并命令行参数与wandb Sweep配置时,遵循Sweep配置优先级高于命令行参数的原则(超参搜索需覆盖固定参数)。利用Hugging Face Trainer的report_to='wandb'参数自动处理wandb初始化,无需手动调用wandb.init(),避免代码冗余混乱。

完整示例代码

import argparse
import wandb
from transformers import Trainer, TrainingArguments, AutoModelForSequenceClassification, AutoTokenizer
from datasets import load_dataset

def parse_args():
    parser = argparse.ArgumentParser()
    # 基础固定参数
    parser.add_argument("--model_name", type=str, default="bert-base-uncased")
    parser.add_argument("--dataset_name", type=str, default="imdb")
    parser.add_argument("--num_train_epochs", type=int, default=3)
    parser.add_argument("--per_device_train_batch_size", type=int, default=16)
    # 超参搜索可变参数(由Sweep覆盖)
    parser.add_argument("--learning_rate", type=float, default=5e-5)
    parser.add_argument("--weight_decay", type=float, default=0.01)
    return parser.parse_args()

def train():
    # 1. 解析命令行参数
    args = parse_args()
    
    # 2. 合并wandb Sweep配置(仅在Sweep运行时生效)
    run = wandb.init()
    sweep_config = run.config
    # 用Sweep配置覆盖命令行参数
    for k, v in sweep_config.items():
        setattr(args, k, v)
    
    # 3. 加载数据集与模型
    tokenizer = AutoTokenizer.from_pretrained(args.model_name)
    model = AutoModelForSequenceClassification.from_pretrained(args.model_name, num_labels=2)
    dataset = load_dataset(args.dataset_name)
    
    def tokenize_function(examples):
        return tokenizer(examples["text"], padding="max_length", truncation=True)
    tokenized_datasets = dataset.map(tokenize_function, batched=True)
    small_train_dataset = tokenized_datasets["train"].shuffle(seed=42).select(range(1000))
    small_eval_dataset = tokenized_datasets["test"].shuffle(seed=42).select(range(1000))
    
    # 4. 设置TrainingArguments,指定report_to='wandb'
    training_args = TrainingArguments(
        output_dir="./results",
        learning_rate=args.learning_rate,
        per_device_train_batch_size=args.per_device_train_batch_size,
        num_train_epochs=args.num_train_epochs,
        weight_decay=args.weight_decay,
        report_to="wandb",  # 让Trainer自动关联wandb训练日志
        logging_dir="./logs",
        logging_steps=10,
    )
    
    # 5. 初始化Trainer并启动训练
    trainer = Trainer(
        model=model,
        args=training_args,
        train_dataset=small_train_dataset,
        eval_dataset=small_eval_dataset,
    )
    trainer.train()

if __name__ == "__main__":
    # 兼容命令行直接运行与Sweep运行
    train()

关键细节说明

  • 参数合并逻辑:先解析命令行参数,再通过wandb.run.config获取Sweep的超参配置,遍历覆盖命令行参数,既支持单独运行,也支持Sweep自动替换超参。
  • report_to='wandb'的作用:设置该参数后,Trainer会自动完成wandb.init()并关联训练过程,无需手动初始化。如果需要自定义wandb项目名、实体等参数,可在train()开头手动调用wandb.init(project="your-project", entity="your-entity"),Trainer会复用已有的wandb运行实例。
  • Sweep配置文件示例:
    program: train.py
    method: grid
    parameters:
      learning_rate:
        values: [1e-5, 2e-5, 5e-5]
      weight_decay:
        values: [0.0, 0.01]
    
    运行Sweep时,执行wandb sweep sweep_config.yaml生成Sweep ID,再启动代理wandb agent <sweep-id>即可开始超参搜索。

内容的提问来源于stack exchange,提问作者Charlie Parker

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.17 13:19:55