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

如何在ruamel.yaml中嵌套使用representer_mappings实现自定义类的YAML序列化?

如何在ruamel.yaml中嵌套使用representer_mappings实现自定义类的YAML序列化?

我完全明白你遇到的困扰——你用工厂模式重构了LearningTask,就是为了能优雅地序列化到YAML,但自己写的representer因为嵌套处理逻辑不对报错了。咱们一步步把这个问题解决掉。

首先先分析你之前的报错原因:你犯了两个小但关键的错误:

  1. 笔误:把ruamel.yaml的represent_mapping方法写成了representer_mapping(多了一个r),这是直接触发报错的原因;
  2. 手动构建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

关键优化点说明

  1. 修正方法名笔误:把错误的representer_mapping改成ruamel.yaml原生的represent_mapping;
  2. 统一标签生成逻辑:为自定义类和第三方类生成不同格式的标签,保证YAML的可读性和可反序列化性;
  3. 委托representer处理子组件:让ruamel.yaml的内置逻辑处理每个子组件的序列化,避免手动构建Node带来的兼容性问题。

如果之后需要反序列化这个YAML,只需要给每个标签注册对应的构造函数即可,逻辑和序列化对称,这里就不展开了。

备注:内容来源于stack exchange,提问作者Felix Benning

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 14:44:37