更新时间:2026-09-28 GMT+08:00
分享

开发自定义算子

当预置算子无法满足您的数据处理需求时,您可以基于ModelArts数据精炼SDK开发自定义算子。本文介绍如何搭建工程目录、编写算子代码、配置依赖与模型文件并打包算子。

前提条件

已完成准备本地开发环境。

算子工程目录说明

在本地创建一个算子工程目录,结构如下:

my_custom_op/
├── process.py            # 【必填】算子入口文件,不能修改文件名,定义PreProcess/Process/PostProcess类
├── dependency/           # 【按需】算子依赖的第三方Python包及requirements.txt
│   ├── requirements.txt
│   ├── some_package-1.0.0-py3-none-any.whl
│   └── another_package-0.5.0.tar.gz
└── weights/              # 【按需】算子执行依赖的模型文件
    └── model.bin
表1 目录结构说明

目录/文件

是否必填

说明

process.py

必填

算子核心逻辑文件,至少定义Process类。

dependency/

按需

存放算子额外依赖的.whl / .tar.gz包,以及requirements.txt。

weights/

按需

存放算子运行时依赖的模型权重文件。

算子执行流程说明

自定义算子通过在process.py中定义PreProcess、Process、PostProcess这3个Python类来实现数据处理逻辑。这3个类分别对应数据预处理、算子核心逻辑和数据后处理,均需继承SDK中的算子基类(Mapper、BatchMapper、Filter、BatchFilter、Generator)。Process类为必填,PreProcess和PostProcess按需定义。

系统运行自定义算子时,按照PreProcess→Process→PostProcess的顺序依次执行,前一步的输出作为后一步的输入。

算子执行流程如下表所示。

步骤

处理阶段

类名

是否必填

输入

输出

说明

①

加载数据集

系统自动处理

系统执行

系统自动加载

原始数据

系统自动加载数据集,读取原始数据。

②

数据预处理

PreProcess

可选

原始数据

预处理后的数据

过滤、转换、去重、缺失值填充等,结果作为Process的输入。

③

算子核心逻辑

Process

必填

预处理后的数据

(或原始数据)

处理后的数据

执行核心处理逻辑,如打标、过滤、增强等。

④

数据后处理

PostProcess

可选

处理后的数据

最终处理结果

结果过滤、格式标准化、冗余列清理等。

⑤

输出结果

系统自动处理

系统执行

最终处理结果

系统自动输出

系统自动输出最终处理结果。

  1. 加载数据集,读取原始数据。
  2. 如果定义了PreProcess,对每条数据执行预处理。返回False的样本被过滤,返回dict或pa.Array的样本新增列。
  3. 对PreProcess输出的数据执行Process.process(),执行算子核心逻辑。
  4. 如果定义了PostProcess,对Process输出的数据执行后处理。返回False的样本被过滤,返回dict或pa.Array的样本新增列。
  5. 输出最终处理结果。

步骤一:编写算子代码

根据您的数据处理需求,选择合适的算子基类并实现process方法。关于各算子基类的详细说明和差异对比,请参见算子基类与差异详解。

三个处理阶段可选的算子基类及关键信息如下表所示,帮助您快速选择合适的基类组合。

算子基类

处理粒度

PreProcess(可选)

Process(必填)

PostProcess(可选)

process方法返回值

典型场景

Filter

单条样本

支持

支持

支持

bool

行级过滤

Mapper

单条样本

支持

支持

支持

dict[str, Any]

行级打标/转换

BatchMapper

批量样本

支持

支持

不支持

pa.Array

批量打标/转换

BatchFilter

批量样本

不支持

支持

支持

pa.Array(bool)

批量过滤

Generator

单条样本

不支持

支持

不支持

list[dict[str, Any]]

一生多

PreProcess(数据预处理,可选)

PreProcess是数据处理流程的第1步,用于对原始数据进行预处理(如格式转换、去重、缺失值填充、异常值过滤等),执行结果作为Process的输入。PreProcess为按需定义,必须继承Filter、Mapper或BatchMapper基类。可选基类及适用场景如下:

  • 继承Filter:用于在核心处理之前过滤掉不符合条件的样本(如空值、格式异常等),process方法返回bool决定样本是否保留。
  • 继承Mapper:用于在核心处理之前对样本进行预处理转换(如字段重命名、编码转换等),process方法返回dict为新增或修改的列。
  • 继承BatchMapper:用于批量预处理场景(如批量特征提取),process方法返回pa.Array为批量新增或修改的列。

PreProcess代码示例如下:

  • 示例1:继承Filter,预处理阶段过滤掉空文本和过短文本。
    from typing import Any
    
    from modelarts.data.refiner.dataset.dpt import PreTrainText
    from modelarts.data.refiner.op.base_op import Filter
    
    class PreProcess(Filter):
        """数据预处理:过滤掉空文本和长度不足的样本"""
    
        def __init__(self, op_args: dict[str, Any], **kwargs):
            super().__init__(op_args, **kwargs)
            self.min_length = op_args.get("min_length", 50)
    
        def process(self, sample: dict[str, Any]) -> bool:
            pre_train_text = PreTrainText.from_dict(sample)
            if not pre_train_text.text or not pre_train_text.text.strip():
                return False
            return len(pre_train_text.text) >= self.min_length
  • 示例2:继承Mapper,预处理阶段为样本新增文本长度列。
    from typing import Any
    
    import daft
    from modelarts.data.refiner.dataset.sft import Messages
    from modelarts.data.refiner.op.base_op import Mapper
    
    class PreProcess(Mapper):
        """数据预处理:为样本新增文本总长度列"""
    
        def __init__(self, op_args: dict[str, Any], **kwargs):
            super().__init__(op_args, **kwargs)
    
        def process(self, sample: dict[str, Any]) -> dict[str, Any]:
            messages = Messages.from_dict(sample)
            total_len = sum(len(msg.content) for msg in messages.messages)
            return {'text_length': total_len}
    
        @staticmethod
        def __return_extend_columns_type__() -> daft.schema.Schema:        return daft.schema.Schema.from_pydict({'text_length': daft.DataType.int64()})

Process(算子核心逻辑,必填)

Process是数据处理流程的第2步,也是必填步骤,承载算子的核心处理逻辑。Process接收PreProcess的输出作为输入(如果未定义PreProcess,则接收原始数据),执行核心处理后将结果传递给PostProcess。Process必须继承SDK中的算子基类(关于基类差异详见算子基类与差异详解),可选基类及适用场景如下:

  • 继承Mapper:行级打标/转换,逐条处理样本并新增列。适用于文本分类打标、质量评分、字段变换等场景。process方法返回dict[str, Any]为新增列的键值对;必须实现__return_extend_columns_type__()声明新增列类型。
  • 继承BatchMapper:批量打标/转换,以列为单位批量处理。适用于需要模型批量推理、GPU加速等场景。process方法接收dict[str, pa.Array],返回pa.Array;必须实现__return_extend_columns_type__()声明新增列类型。
  • 继承Filter:行级过滤,逐条判断样本是否保留。适用于按条件过滤低质量样本的场景。process方法返回bool,True保留/False过滤;无需实现__return_extend_columns_type__()。
  • 继承BatchFilter:批量过滤,以列为单位批量判断样本是否保留。适用于批量质量检测过滤的场景。process方法接收dict[str, pa.Array],返回pa.Array(bool);无需实现__return_extend_columns_type__()。
  • 继承Generator:一条样本生成多条新样本。适用于视频切分、数据增强等场景。process方法返回list[dict[str, Any]],每个元素是完整的新样本字典;无需实现__return_extend_columns_type__()。

Process代码示例如下:

  • 示例1:继承Mapper,为SFT对话数据打质量评分标签。
    from typing import Any
    
    import daft
    from modelarts.data.refiner.dataset.sft import Messages
    from modelarts.data.refiner.op.base_op import Mapper
    
    class Process(Mapper):
        """算子核心逻辑:为SFT对话数据打质量评分标签"""
    
        def __init__(self, op_args: dict[str, Any], **kwargs):
            super().__init__(op_args, **kwargs)
            self.high_threshold = op_args.get("high_threshold", 5000)
            self.medium_threshold = op_args.get("medium_threshold", 1000)
    
        def process(self, sample: dict[str, Any]) -> dict[str, Any]:
            messages = Messages.from_dict(sample)
            total_len = sum(len(msg.content) for msg in messages.messages)
            if total_len > self.high_threshold:
                return {'quality': 'high'}
            elif total_len > self.medium_threshold:
                return {'quality': 'medium'}
            else:
                return {'quality': 'low'}
    
        @staticmethod
        def __return_extend_columns_type__() -> daft.schema.Schema:
            return daft.schema.Schema.from_pydict({'quality': daft.DataType.string()})
  • 示例2:继承Filter,过滤文本长度不足的预训练数据。
    from typing import Any
    
    from modelarts.data.refiner.dataset.dpt import PreTrainText
    from modelarts.data.refiner.op.base_op import Filter
    
    class Process(Filter):
        """算子核心逻辑:过滤文本长度不超过阈值的预训练样本"""
    
        def __init__(self, op_args: dict[str, Any], **kwargs):
            super().__init__(op_args, **kwargs)
            self.text_len = op_args.get("text_len", 1000)
    
        def process(self, sample: dict[str, Any]) -> bool:
            pre_train_text = PreTrainText.from_dict(sample)
            return len(pre_train_text.text) > self.text_len
  • 示例3:继承Generator,视频切分生成多个片段。
    from typing import Any
    
    from modelarts.data.refiner.dataset.video import Video
    from modelarts.data.refiner.op.base_op import Generator
    
    class Process(Generator):
        """算子核心逻辑:将视频切分为多个小片段"""
    
        def __init__(self, op_args: dict[str, Any], **kwargs):
            super().__init__(op_args, **kwargs)
            self.clip_duration = op_args.get("clip_duration", 10)
    
        def process(self, sample: dict[str, Any]) -> list[dict[str, Any]]:
            video = Video.from_dict(sample)
            video_clips = self.cut_video(video.file_path)
            return [Video.to_dict(clip) for clip in video_clips]
    
        # 实现视频切分逻辑
        def cut_video(self, video_path: str) -> list[Video]:
            ...

PostProcess(算子后处理,可选)

PostProcess是数据处理流程的第3步,用于对Process处理后的结果进行后处理(如结果过滤、格式标准化、冗余列清理等)。PostProcess为按需定义,必须继承Filter、BatchFilter或Mapper基类。可选基类及适用场景如下:

  • 继承Filter:用于对Process输出的结果进行二次过滤(如过滤低质量打标结果),process方法返回bool决定样本是否保留。
  • 继承BatchFilter:用于对Process输出的结果进行批量二次过滤,process方法返回pa.Array(bool)批量决定样本是否保留。
  • 继承Mapper:用于对Process输出的结果进行后处理转换(如字段重命名、格式标准化等),process方法返回dict为新增或修改的列。

PostProcess代码示例如下:

  • 示例1:继承Filter,后处理阶段过滤掉质量评分为low的样本。
    from typing import Any
    
    from modelarts.data.refiner.op.base_op import Filter
    
    class PostProcess(Filter):
        """算子后处理:过滤掉质量评分为low的样本"""
    
        def __init__(self, op_args: dict[str, Any], **kwargs):
            super().__init__(op_args, **kwargs)
    
        def process(self, sample: dict[str, Any]) -> bool:
            quality = sample.get('quality', 'low')
            return quality != 'low'
  • 示例2:继承BatchFilter,批量后处理过滤短文本样本。
    from typing import Any
    
    import pyarrow as pa
    from modelarts.data.refiner.dataset.dpt import TEXT
    from modelarts.data.refiner.op.base_op import BatchFilter
    
    class PostProcess(BatchFilter):
        """算子后处理:批量过滤文本长度不足的样本"""
    
        def __init__(self, op_args: dict[str, Any], **kwargs):
            super().__init__(op_args, **kwargs)
            self.text_len = op_args.get("text_len", 100)
    
        def process(self, sample: dict[str, pa.Array]) -> pa.Array:
            texts = sample.get(TEXT)
            results = []
            for text in texts:
                text_value = text.as_py()
                results.append(len(text_value) > self.text_len)
                return pa.array(results, type=pa.bool_())

完整process.py示例

以下是一个包含PreProcess、Process、PostProcess三个阶段的完整process.py示例,演示对预训练文本数据的“过滤→打标→过滤”流水线:PreProcess过滤空文本和过短文本,Process为通过过滤的样本打上长度等级标签,PostProcess批量过滤掉short级别的样本。

from typing import Any

import daft
import pyarrow as pa
from modelarts.data.refiner.dataset.dpt import PreTrainText
from modelarts.data.refiner.op.base_op import Filter, Mapper, BatchFilter

class PreProcess(Filter):
    """第1步数据预处理:过滤掉空文本和过短文本"""

    def __init__(self, op_args: dict[str, Any], **kwargs):
        super().__init__(op_args, **kwargs)
        self.min_length = op_args.get("min_length", 50)

    def process(self, sample: dict[str, Any]) -> bool:
        pre_train_text = PreTrainText.from_dict(sample)
        if not pre_train_text.text or not pre_train_text.text.strip():
            return False
        return len(pre_train_text.text) >= self.min_length

class Process(Mapper):
    """第2步算子核心逻辑:为文本打上长度等级标签"""

    def __init__(self, op_args: dict[str, Any], **kwargs):
        super().__init__(op_args, **kwargs)

    def process(self, sample: dict[str, Any]) -> dict[str, Any]:
        pre_train_text = PreTrainText.from_dict(sample)
        text_len = len(pre_train_text.text)
        if text_len > 5000:
            level = "long"
        elif text_len > 500:
            level = "medium"
        else:
            level = "short"
        return {"length_level": level}

    @staticmethod
    def __return_extend_columns_type__() -> daft.schema.Schema:
        return daft.schema.Schema.from_pydict({"length_level": daft.DataType.string()})

class PostProcess(BatchFilter):
    """第3步算子后处理:批量过滤掉short级别的样本"""

    def __init__(self, op_args: dict[str, Any], **kwargs):
        super().__init__(op_args, **kwargs)

    def process(self, sample: dict[str, pa.Array]) -> pa.Array:
        levels = sample.get("length_level")
        results = []
        for level in levels:
            results.append(level.as_py() != "short")
            return pa.array(results, type=pa.bool_())

步骤二:配置依赖与模型文件

完成算子代码编写后,需要配置运行时依赖和模型文件等,确保算子在不同环境正常运行。

  • 算子入参(op_args)

    算子构造器接收op_args参数,这是一个字典,用于传递用户自定义的运行时参数。在process.py的类中,您可以通过op_args获取运行时参数。使用方式如下:

    class Process(Mapper):
        def __init__(self, op_args: dict[str, Any], **kwargs):
            super().__init__(op_args, **kwargs)
            self.threshold = op_args.get("threshold", 1000)
            self.mode = op_args.get("mode", "strict")
            self.model_path = op_args.get("model_path", "./weights/model.bin")

    建议为所有参数提供合理的默认值,确保算子在未指定参数时也能正常运行。系统参数(如content_type、input_columns)通过kwargs注入,不要在op_args中定义同名参数。

  • 第三方依赖

    如果算子依赖了SDK之外的Python包,需要将依赖放置在dependency/目录下。requirements.txt中列出的包,其对应的.whl或.tar.gz文件必须同时存在于dependency/目录下。SDK自身已包含的依赖(daft、pyarrow、pandas、opencv-python、numpy等)无需重复添加。安装时会优先从dependency/目录本地安装,无需公网下载。

  • 模型文件

    如果算子运行时需要加载模型(如推理模型、NLP模型等),将模型文件放置在weights/目录下。在算子代码中可通过op_args传入模型路径,或使用相对路径加载。

    import os
    class Process(BatchMapper):
        def __init__(self, op_args: dict[str, Any], **kwargs):
            super().__init__(op_args, **kwargs)
            self.model_path = op_args.get("model_path", os.path.join(os.path.dirname(__file__), "weights", "model.bin"))
  • 加速器配置

    如果算子需要NPU加速,可以在类中覆盖_accelerator属性:

    class Process(BatchMapper):
        _accelerator = "npu"  # 使用NPU加速,可选值: "cpu", "npu"
        _max_concurrency = 4  # 降低并发数,避免显存溢出
    
        def __init__(self, op_args: dict[str, Any], **kwargs):
            super().__init__(op_args, **kwargs)

步骤三:打包算子及依赖

开发完成后,需要将算子工程打包为tar格式。算子工程目录请参见算子工程目录说明。

注意事项如下:

  • 必须包含process.py文件,且其中至少定义了Process类。
  • dependency/和weights/目录如无内容可不放,但如果放置则目录名必须准确。
  • 如果算子仅依赖SDK已包含的库(daft、pyarrow、pandas、numpy、opencv-python等),则无需创建dependency/目录。

相关文档