如何以官方方式使用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配置文件示例:
运行Sweep时,执行program: train.py method: grid parameters: learning_rate: values: [1e-5, 2e-5, 5e-5] weight_decay: values: [0.0, 0.01]wandb sweep sweep_config.yaml生成Sweep ID,再启动代理wandb agent <sweep-id>即可开始超参搜索。
内容的提问来源于stack exchange,提问作者Charlie Parker
相关产品推荐
相关产品推荐

