# 开发自定义算子
当预置算子无法满足您的数据处理需求时，您可以基于ModelArts数据精炼SDK开发自定义算子。本文介绍如何搭建工程目录、编写算子代码、配置依赖与模型文件并打包算子。
#### 前提条件
已完成[准备本地开发环境](https://support.huaweicloud.com/dataprepare-modelarts/dataprepare-modelarts-0040.html)。
 #### 算子工程目录说明
在本地创建一个算子工程目录，结构如下：
```
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的顺序依次执行，前一步的输出作为后一步的输入。
![](https://support.huaweicloud.com/dataprepare-modelarts/figure/zh-cn_image_0000002683291228.png "点击放大")
算子执行流程如下表所示。
| 步骤 | **处理阶段** | **类名**      | **是否必填** | **输入**                                                       | **输出**   | **说明**                           |
|:---|:---|:---|:---|:---|:---|:---|
| ①  | 加载数据集      | 系统自动处理    | 系统执行     | 系统自动加载                                                      | 原始数据      | 系统自动加载数据集，读取原始数据。                |
| ②  | 数据预处理   | PreProcess   | 可选        | 原始数据                                                         | 预处理后的数据  | 过滤、转换、去重、缺失值填充等，结果作为Process的输入。 |
| ③  | 算子核心逻辑  | Process     | 必填        | 预处理后的数据 （或原始数据） | 处理后的数据 | 执行核心处理逻辑，如打标、过滤、增强等。            |
| ④ | 数据后处理    | PostProcess | 可选       | 处理后的数据                                                     | 最终处理结果  | 结果过滤、格式标准化、冗余列清理等。               |
| ⑤  | 输出结果     | 系统自动处理       | 系统执行     | 最终处理结果                                                      | 系统自动输出  | 系统自动输出最终处理结果。                  |
   
1. 加载数据集，读取原始数据。
2. 如果定义了PreProcess，对每条数据执行预处理。返回False的样本被过滤，返回dict或pa.Array的样本新增列。
3. 对PreProcess输出的数据执行Process.process()，执行算子核心逻辑。
4. 如果定义了PostProcess，对Process输出的数据执行后处理。返回False的样本被过滤，返回dict或pa.Array的样本新增列。
5. 输出最终处理结果。
 
#### 步骤一：编写算子代码
根据您的数据处理需求，选择合适的算子基类并实现process方法。关于各算子基类的详细说明和差异对比，请参见[算子基类与差异详解](https://support.huaweicloud.com/dataprepare-modelarts/dataprepare-modelarts-0042.html)。
三个处理阶段可选的算子基类及关键信息如下表所示，帮助您快速选择合适的基类组合。
| **算子基类**       | **处理粒度** | **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中的算子基类（关于基类差异详见[算子基类与差异详解](https://support.huaweicloud.com/dataprepare-modelarts/dataprepare-modelarts-0042.html)），可选基类及适用场景如下：
- 继承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/目录。
 
