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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 22:47:49