HuggingFace Trainer无法向wandb上报问题求助
问题:设置Hugging Face Trainer的report_to为wandb时触发ValueError错误
我按照文档配置Trainer的report_to参数为wandb,代码如下:
training_args = TrainingArguments( output_dir="test_trainer", evaluation_strategy="steps", learning_rate=config.learning_rate, num_train_epochs=config.epochs, weight_decay=config.weight_decay, logging_dir=config.logging_dir, report_to="wandb", save_total_limit=1, per_device_train_batch_size=config.batch_size, per_device_eval_batch_size=config.batch_size, fp16=True, load_best_model_at_end=True, seed=42 )
初始化Trainer时:
trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=eval_dataset, compute_metrics=compute_metrics )
出现如下报错:
--------------------------------------------------------------------------- ValueError Traceback (most recent call last) <ipython-input-68-b009351ab52d> in <module> 4 train_dataset=train_dataset, 5 eval_dataset=eval_dataset, ----> 6 compute_metrics=compute_metrics 7 ) ~/.virtualenvs/transformers_lab/lib/python3.7/site-packages/transformers/trainer.py in __init__(self, model, args, data_collator, train_dataset, eval_dataset, tokenizer, model_init, compute_metrics, callbacks, optimizers) 286 "You should subclass `Trainer` and override the `create_optimizer_and_scheduler` method." 287 ) --> 288 default_callbacks = DEFAULT_CALLBACKS + get_reporting_integration_callbacks(self.args.report_to) 289 callbacks = default_callbacks if callbacks is None else default_callbacks + callbacks 290 self.callback_handler = CallbackHandler( ~/.virtualenvs/transformers_lab/lib/python3.7/site-packages/transformers/integrations.py in get_reporting_integration_callbacks(report_to) 794 if integration not in INTEGRATION_TO_CALLBACK: 795 raise ValueError( --> 796 f"{integration} is not supported, only {', '.join(INTEGRATION_TO_CALLBACK.keys())} are supported." 797 ) 798 return [INTEGRATION_TO_CALLBACK[integration] for integration in report_to] ValueError: w is not supported, only azure_ml, comet_ml, mlflow, tensorboard, wandb are supported.
请问有没有人遇到过相同的错误?
解决方法
- 错误原因:旧版本
transformers库中,report_to参数被设计为接受列表类型,传入字符串时会被当成可迭代对象逐个遍历字符,导致把"wandb"拆成了"w"、"a"等单个字符,触发不支持的报错。 - 修复步骤:将
report_to的值改为列表格式:training_args = TrainingArguments( # 其他参数保持不变 report_to=["wandb"], # 其他参数保持不变 ) - 验证:修改后重新初始化Trainer,即可正常加载wandb集成,不会再出现该错误。
内容的提问来源于stack exchange,提问作者Weber Huang
相关产品推荐
相关产品推荐

