如何在ruamel.yaml中嵌套使用representer_mappings实现自定义类的YAML序列化?
如何在ruamel.yaml中嵌套使用representer_mappings实现自定义类的YAML序列化?
我完全明白你遇到的困扰——你用工厂模式重构了LearningTask,就是为了能优雅地序列化到YAML,但自己写的representer因为嵌套处理逻辑不对报错了。咱们一步步把这个问题解决掉。
首先先分析你之前的报错原因:你犯了两个小但关键的错误:
- 笔误:把ruamel.yaml的
represent_mapping方法写成了representer_mapping(多了一个r),这是直接触发报错的原因; - 手动构建
MappingNode后直接传给上级mapping,不符合ruamel.yaml的内部处理逻辑——正确的做法是让representer帮你处理每个子组件的序列化,而不是手动组装节点。
修正后的完整实现方案
下面是能生成你想要的YAML格式的完整代码,我会标注关键细节:
import ruamel.yaml from attr import dataclass import torch import pytorch_lightning as lightning from typing import Any, Type # 初始化ruamel.yaml实例,配置美观的缩进格式 yaml = ruamel.yaml.YAML() yaml.indent(mapping=2, sequence=4, offset=2) # 定义参数存储类 @dataclass class Params: model: dict[str, Any] loss: dict[str, Any] data: dict[str, Any] # 工厂模式的任务类 @dataclass class BetterLearningTask: model_class: Type[torch.nn.Module] loss_class: Type[torch.nn.Module] data_class: Type[lightning.LightningDataModule] params: Params @property def model(self): return self.model_class(**self.params.model) @property def loss(self): return self.loss_class(**self.params.loss) @property def data(self): return self.data_class(**self.params.data) # 为BetterLearningTask注册序列化逻辑 @yaml.register_class class BetterLearningTask(BetterLearningTask): @classmethod def to_yaml(cls, representer, node): # 辅助函数:统一生成「带类标签的参数映射」节点 def build_component_node(cls_inst, params): # 为第三方类(如torch的内置损失函数)生成带完整模块路径的标签 if cls_inst.__module__ != '__main__': tag = f'!{cls_inst.__module__}.{cls_inst.__name__}' # 自定义类直接用类名当标签 else: tag = f'!{cls_inst.__name__}' # 让representer处理参数映射的生成 return representer.represent_mapping(tag, params) # 组装BetterLearningTask的核心映射 task_mapping = { 'model': build_component_node(node.model_class, node.params.model), 'loss': build_component_node(node.loss_class, node.params.loss), 'data': build_component_node(node.data_class, node.params.data) } # 生成最终的BetterLearningTask节点 return representer.represent_mapping(f'!{cls.__name__}', task_mapping)
测试与效果
我们用自定义模型和数据模块测试一下:
# 自定义测试用模型 class MyModelInheritingFromPytorchModule(torch.nn.Module): def __init__(self, conv_kernel_size: int = 3): super().__init__() self.conv = torch.nn.Conv2d(3, 16, kernel_size=conv_kernel_size) # 自定义测试用数据模块 class MyLightningDataModule(lightning.LightningDataModule): def __init__(self, batch_size: int = 32): super().__init__() self.batch_size = batch_size # 创建任务工厂实例 task_factory = BetterLearningTask( model_class=MyModelInheritingFromPytorchModule, loss_class=torch.nn.CrossEntropyLoss, data_class=MyLightningDataModule, params=Params( model={"conv_kernel_size": 3}, loss={"reduction": "mean"}, data={"batch_size": 64} ) ) # 序列化到YAML import io stream = io.StringIO() yaml.dump(task_factory, stream) stream.seek(0) print(stream.read())
运行后会输出你想要的格式:
!BetterLearningTask model: !MyModelInheritingFromPytorchModule conv_kernel_size: 3 loss: !torch.nn.CrossEntropyLoss reduction: mean data: !MyLightningDataModule batch_size: 64
关键优化点说明
- 修正方法名笔误:把错误的
representer_mapping改成ruamel.yaml原生的represent_mapping; - 统一标签生成逻辑:为自定义类和第三方类生成不同格式的标签,保证YAML的可读性和可反序列化性;
- 委托representer处理子组件:让ruamel.yaml的内置逻辑处理每个子组件的序列化,避免手动构建Node带来的兼容性问题。
如果之后需要反序列化这个YAML,只需要给每个标签注册对应的构造函数即可,逻辑和序列化对称,这里就不展开了。
备注:内容来源于stack exchange,提问作者Felix Benning
相关产品推荐
相关产品推荐

