开发自定义算子
当预置算子无法满足您的数据处理需求时,您可以基于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 | 目录/文件 | 是否必填 | 说明 |
|---|---|---|
| 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 | 可选 | 处理后的数据 | 最终处理结果 | 结果过滤、格式标准化、冗余列清理等。 |
| ⑤ | 输出结果 | 系统自动处理 | 系统执行 | 最终处理结果 | 系统自动输出 | 系统自动输出最终处理结果。 |
- 加载数据集,读取原始数据。
- 如果定义了PreProcess,对每条数据执行预处理。返回False的样本被过滤,返回dict或pa.Array的样本新增列。
- 对PreProcess输出的数据执行Process.process(),执行算子核心逻辑。
- 如果定义了PostProcess,对Process输出的数据执行后处理。返回False的样本被过滤,返回dict或pa.Array的样本新增列。
- 输出最终处理结果。
步骤一:编写算子代码
根据您的数据处理需求,选择合适的算子基类并实现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/目录。