如何将TrainingArguments对象转为JSON并正确加载用于训练?
问题:TrainingArguments序列化/反序列化后无法用于Trainer初始化
问题背景
为了更直观地管理训练参数,希望将TrainingArguments对象保存为JSON文件,训练时直接从JSON加载参数。保存代码如下:
import json import dataclasses from transformers import TrainingArguments class EnhancedJSONEncoder(json.JSONEncoder): def default(self, obj): if dataclasses.is_dataclass(obj): return dataclasses.asdict(obj) def save_json(json_path, file_args): file_str = json.dumps(file_args, cls=EnhancedJSONEncoder) file_json = json.loads(file_str) with open(json_path, "w") as f: json.dump(file_json, f, indent=4) if __name__ == "__main__": training_args = TrainingArguments( output_dir = "train_model", do_eval=True, do_predict=False, do_train=True, eval_steps=20, save_steps=-1, per_device_eval_batch_size=8, per_device_train_batch_size=4, max_steps=1000, learning_rate=7e-05, evaluation_strategy="steps", dataloader_num_workers=1, ) training_args._n_gpu = 1 save_json("training_args.json",training_args)
加载时使用SimpleNamespace:
from types import SimpleNamespace import json with open("training_args.json", 'r') as f: training_json = json.load(f) training_args = SimpleNamespace(**training_json)
此时能正常访问参数属性,但初始化Trainer时出现错误:
'types.SimpleNamespace' object has no attribute 'get_process_log_level'
原因分析
SimpleNamespace只是一个简单的属性容器,仅保存了参数的键值对,但没有TrainingArguments类内置的方法(比如get_process_log_level)和内部逻辑。而Trainer要求传入的必须是真正的TrainingArguments实例,而非普通的命名空间对象。
解决方案
方法1:使用官方提供的from_dict方法加载(推荐)
Hugging Face Transformers库为TrainingArguments提供了from_dict()方法,专门用于从字典重建完整的TrainingArguments实例,包含所有类方法和内部属性。
修改加载代码:
import json from transformers import TrainingArguments with open("training_args.json", 'r') as f: training_json = json.load(f) # 用from_dict直接构建TrainingArguments实例 training_args = TrainingArguments.from_dict(training_json) # 若有需要手动设置的属性(如_n_gpu),可在此补充 training_args._n_gpu = 1
方法2:简化保存逻辑(配合方法1使用)
其实TrainingArguments自带to_dict()方法,无需自定义JSON编码器,保存代码可以简化为:
import json from transformers import TrainingArguments def save_json(json_path, training_args): # 直接使用官方to_dict方法序列化 args_dict = training_args.to_dict() with open(json_path, "w") as f: json.dump(args_dict, f, indent=4) if __name__ == "__main__": training_args = TrainingArguments( output_dir = "train_model", do_eval=True, do_predict=False, do_train=True, eval_steps=20, save_steps=-1, per_device_eval_batch_size=8, per_device_train_batch_size=4, max_steps=1000, learning_rate=7e-05, evaluation_strategy="steps", dataloader_num_workers=1, ) training_args._n_gpu = 1 save_json("training_args.json", training_args)
总结
不要用SimpleNamespace替代TrainingArguments实例,官方提供的to_dict()和from_dict()方法是专门为序列化/反序列化设计的,能完美适配Trainer的需求,同时保证参数的完整性和正确性。
内容的提问来源于stack exchange,提问作者4daJKong
相关产品推荐
相关产品推荐

