自定义算子开发手册
百度胜算提供了一套完善的自定义算子开发流程,支持将Python脚本封装为算子像使用系统算子一样编排为工作流对外提供数据处理能力。本文以“开发一个标点替换 Transformer 算子”为例,介绍如何在百度胜算中完成自定义算子的开发、测试、打包、上线和使用。本文仅展示一条可以快速验证的实践路径。
前提条件
- 当前用户拥有创建数据目录、数据模式和自定义算子的权限。
- 已准备自定义算子的 Python 代码和测试数据。
产品开通
如果提供的镜像无法满足算子使用要求,需要在提供的镜像上安装新的python库或者软件包,那么需要开通CCR(镜像仓库)产品。
可以先跳过该章节,如果后续在开发算子时候,发现需要在已有镜像的基础上安装新的python库或软件包,可以再回来看该章节进行操作。开通流程如下:
开通镜像仓库
- 在百度胜算页面,搜索CCR。

- 创建实例,根据自身需求可以选择对应的购买种类。

配置镜像仓库
- 进入实例,新建命名空间;

- 新建镜像仓库,选择命名空间,并给仓库取名称;

- 获取访问凭证,可以选择设置固定密码或者使用临时密码,用来登录镜像仓库,推送镜像;

- 配置白名单。

推送镜像
单击镜像仓库的快捷指令。

- 在自己机器的终端,使用docker login命令来登录镜像仓库,如果遇到链接不上的问题,可能是配置白名单的问题
- 然后通过docker tag的命令,将已经打好的镜像进行重命名
- 最后使用docker push 的命令,推送到镜像仓库,推送成功可以在镜像仓库看到该镜像,后续在自定义算子的时候,可以选择该镜像作为算子运行的镜像

场景概览
自定义算子的完整流程如下:
1准备开发环境
2 ↓
3编写并测试 Python 算子
4 ↓
5打包为 whl 文件
6 ↓
7在百度胜算中创建算子和算子版本
8 ↓
9将自定义算子加入工作流
10 ↓
11运行工作流并查看结果
准备自定义算子开发环境
1. 获取开发镜像和样例包
| 物料项 | 链接 | 描述 |
|---|---|---|
| 镜像 | https://cloud.baidu.com/doc/DataBuilder/s/0mtbcfa9d | 镜像中的重点软件版本如下: nvcc==12.4 python==3.10.12 torch==2.5.1 transformers==4.47.1 paddlepaddle==3.1.0 onnxruntime==1.21.1 opencv-python==4.11.0.86 |
| 自定义算子样例包 | https://operator.bj.bcebos.com/release/operator/databuilder_dev-v0.3.2.tar?authorization=bce-auth-v1%2FALTAKRDyrIQOaBuxqyREAShe7e%2F2025-12-26T03%3A29%3A01Z%2F-1%2Fhost%2Fd62a50cc27701605487f9f59ea7b5192f31fc8a7151b52f0d5889c5bc52f457d |
|
2. 加载基础镜像
在服务器上下载镜像压缩包,并执行以下命令:
1# 1. 下载模型包
2wget -O db_operator_dev.tar "https://offical-images.bj.bcebos.com/ray/ray-2.1.7.tar?authorization=bce-auth-v1/33bda7b5e8ea4d669c2c7944449665c6/2026-08-27T11:13:55Z/-1/host/9e2b6bc9d3ab41c2d4b140ea9f2509dac30d44d23a8029967190593a52fb3589"
3# 2. 进入算子镜像tar包所在目录
4cd /xxx
5# 3. 展开镜像
6docker load < ./db_operator_dev.tar
镜像加载可能需要几分钟。可以使用以下命令查看镜像是否加载成功:
1docker image ls | grep iregistry.baidu-int.com/doris-rdw/algorithm/operator
镜像成功展开后会有类似如下输出:

3. 创建开发容器
使用以下命名创建算子开发容器:
1docker run -itd --gpus=all \
2 --privileged=true \
3 --shm-size=64g \
4 --net=host \
5 --name db_op_dev \
6 iregistry.baidu-int.com/doris-rdw/algorithm/operator:0.7 \
7 /bin/bash
4. 进入容器并验证环境
算子开发容器创建完毕后,通过如下命令进入容器:
1docker exec -it db_op_dev /bin/bash
进入容器,可以通过镜像预置的自定义算子测试代码验证开发环境:
1# 1. 进入阈值自定义算子测试目录
2cd /root/databuilder_dev/tests
3# 2. 运行标点替换算子测试脚本
4sh run_punctuation_replace.sh
算子运行成功后输出如下内容:

了解核心概念
算子
算子是python编写的具有一定功能的代码,代码中可以是调用opencv进行图片、视频等的处理,也可以包含模型推理(本地模型或远端服务)对图片、文本等进行处理。为了自动会算子代码进行分布式加速和降低自定义算子编写难度,百度胜算的算子代码编写有一定的规范,需要继承算子框架SDK定义的5类算子基类,只需要实现核心逻辑代码。
工作流
工作流是由多个算子通过串联和并联的方式组成的DAG图,由计算引擎调度运行。
算子种类
| 算子类型 | 说明 |
|---|---|
| Extractor | 数据提取类算子,如pdf文档内容提取、html文档内容提取等,将一行输入数据变更为一行或多行数据。 |
| Filter | 数据过滤算子,如文本句子长度过滤、图片宽高过滤等,将不符合条件的数据行过滤丢弃。 |
| Transformer | 数据转换算子,如中文繁简体转换、图片人脸模糊等,对输入的数据进行转换处理。 |
| Deduplicator | 数据去重算子,如文本hash去重、图片hash去重等,以行为单位去除数据中的重复行。 |
| Embedding | 数据特征嵌入,如文本特征嵌入、图片特征嵌入等,对文本和图片数据提取特征。 |
了解接口和参数
1.Operator

百度胜算算子基类,定义了算子的一些公共变量值和函数。
2.Extractor

子类算子需要实现 extract()函数。
3.Filter

子类算子需要实现 _compute_stats()和compare_stats()函数。
4.Transformer

子类算子需要实现 transform_batched()函数。
5.Deduplicator

子类需要实现 run()函数。
6.Embedding

子类需要实现 _compute_vector()函数。
自定义算子开发
步骤一:环境准备
参考准备自定义算子开发环境,完成自定义算子容器创建,后面的算子开发、测试、打包都在自定义算子容器内完成。
步骤二:预置自定义算子介绍
- 镜像中内置了3个自定义算子(image_resizer/punctuation_filter/punctuation_replacer)

- 3个算子的测试脚本

步骤三:Transformer算子
1.PunctuationReplacer 算子类结构详解
1.1.类属性 - 算子元信息定义

这些属性定义了算子的基本特性和执行环境配置,是算子框架识别和管理该算子的关键元数据。
1.2.初始化方法 - 参数设置与预处理

初始化时完成:
-
参数传递
- 通过
*args, **kwargs隐式支持父类参数(如text_key指定文本字段名) - 无显式参数,简化调用接口
- 通过
-
预处理优化
- 正则预编译:在初始化时编译
r'([^\w\s])+'模式,提升运行时性能 - 模式语义:匹配所有非字母数字(
\w)且非空白(\s)的连续字符
- 正则预编译:在初始化时编译
-
资源准备
- 无外部资源依赖(如模型加载),轻量级初始化
1.3.核心处理方法 - 单样本处理逻辑

处理流程:
-
输入输出结构
- 输入:要求
samples为字典,且包含self.text_key指定的文本列表 - 输出:保持原数据结构,仅更新文本内容
- 输入:要求
-
核心逻辑
- 正则替换:使用预编译模式匹配标点符号
- 等量空格替换:通过
lambda动态生成与匹配项等长的空格(如...→) - 列表推导式:高效实现批量处理
2.PunctuationReplacer 算子接口详解
2.1.公共接口总览
| 接口类型 | 名称/属性 | 说明 |
|---|---|---|
| 类属性 | _op_type="transform" |
标识为文本转换类算子 |
_batched_op=True |
支持批量文本处理 | |
_ray_execute_mode |
声明Ray流水线任务模式(PIPELINE_TASK) |
|
| 初始化方法 | __init__ |
预编译正则表达式模式 |
| 核心方法 | transform_batched |
执行标点替换的主逻辑 |
| 正则工具 | punctuation_pattern |
匹配非字母数字和非空白字符的正则表达式(r'([^\w\s])+') |
2.2.核心接口细节
初始化方法 __init__****

设计要点:
通过 re.compile 提前优化正则匹配性能
支持通过 *args, **kwargs 传递父类参数(如 text_key)
批处理方法 transform_batched****

关键行为:
使用正则替换所有非\w(字母数字)和\s(空白)字符
连续标点会被替换为等量空格(如"!!"→" ")
2.3.接口调用示例
1# 1. input data
2 samples = [
3 {
4 'text': '特殊的康熙部首或者扩展部首会被去除,⼏几⺇'
5 },
6 {
7 'text': '请问你是谁dasoidhao@1264fg.45om'
8 },
9 {
10 'text': '匹配汉字包括繁體字'
11 }
12 ]
13
14 input_dataset = RayDataset.from_list(samples)
15
16 # 2. create operator
17 op = PunctuationReplacer()
18
19 # 3. run
20 output_dataset = input_dataset.run(op)
21
22 # 4. get & check result
23 text_data = output_dataset.get_column(column=op.text_key)
24
25 print(f'---result: {text_data}')
3.ImageResizer 算子类结构详解
3.1.类属性 - 算子元信息定义

3.2.初始化方法 - 参数设置与预处理

初始化流程:
- 参数接收与默认值设置
- 父类初始化,
super().init(*args, **kwargs) # 调用Transformer基类的初始化方法 - 实例属性赋值
- 工具对象初始化,
self.smart_file = SmartFile() # 创建文件路径处理工具实例 - 工作目录设置
3.3.核心处理方法 - transform_batched

处理流程:
-
路径转换:
volume_2_local():处理分布式存储路径local_2_volume():结果回传存储系统
- 图像处理:
Image.open(path).resize((w, h)) # 核心缩放操作 -
命名控制:
- 使用
make_unique_name_by_dict生成唯一文件名 - 支持原始文件名保留(
need_hash_name=False)
- 使用
4.ImageResizer 算子接口详解
4.1.公共接口总览
| 接口类型 | 名称/属性 | 说明 |
|---|---|---|
| 类属性 | _op_type |
算子类型标识(固定为"transform") |
_batched_op |
批量处理标志(默认为True) |
|
_ray_execute_mode |
Ray执行模式("PIPELINE_TASK") |
|
| 初始化方法 | __init__ |
配置目标尺寸、输出路径等参数 |
| 核心方法 | transform_batched |
执行批量图像缩放处理 |
| 工具对象 | smart_file |
处理本地与分布式存储路径转换 |
4.2.核心接口细节
初始化方法

参数说明:
width/height:目标尺寸(像素),默认224x224dst_path:输出目录路径(空表示不保存)need_hash_name:是否对输出文件哈希命名(避免冲突)
关键行为:
- 自动创建临时工作目录(
download_path) - 初始化文件处理器(
SmartFile)
核心方法 transform_batched

输入输出:
- 输入:
samples字典需包含self.image_key指定的图像路径列表 - 输出:更新后的
samples字典(含处理后的路径列表)
处理流程:
- 路径转换(分布式→本地)
- 使用PIL进行图像缩放
- 结果回传(本地→分布式存储)
4.3.接口调用示例
1# 1. input data
2 samples = [
3 {
4 'images': './images/cat.jpg'
5 },
6 {
7 'images': './images/cat2.jpg'
8 },
9 {
10 'images': './images/lena.jpg'
11 }
12 ]
13
14 input_dataset = RayDataset.from_list(samples)
15
16 # 2. create operator
17 dst_path = './ret'
18 os.makedirs(dst_path, exist_ok=True)
19 op = ImageResizer(width=1024, height=1024, dst_path=dst_path, need_hash_name=False)
20
21 # 3. run
22 output_dataset = input_dataset.run(op)
23
24 # 4. get & check result
25 image_path = output_dataset.get_column(column=op.image_key)
26
27 for name in image_path:
28 image = Image.open(name)
29 print(f'---result image {name} shape: {image.size}')
5.自定义加载模型并推理算子开发示例
- 算子定义
1import pandas as pd
2import ray
3from typing import Dict
4import numpy as np
5import torch
6
7from palette.ops.base import Transformer
8from palette.core.ray_dataset import RayDataset
9
10# 1.写自定义算子
11class TorchPredictor(Transformer):
12 ##算子设置为transform类型
13 _op_type = "transform"
14 ##算子设置为批处理,即__call__中接收一批数据,返回一批数据的处理结果
15 _batched_op = True
16 ## 算子的名称
17 _name = "operator_transform"
18 ## 设置算子执行模式为PIPELINE_ACTOR,这样init只会执行一次
19 _ray_execute_mode = "PIPELINE_ACTOR"
20 ## 设置算子输入数据格式为numpy
21 _ray_batch_format = "numpy"
22 ## 若算子需要使用GPU,则设置算子处理器为cuda
23 _processor = "cuda"
24
25
26 def __init__(self, *args, **kwargs):
27 ##调用父类的初始化函数
28 super().__init__(*args, **kwargs)
29
30 ##初始化模型,也可以使用模型权重文件加载
31 self.model = torch.nn.Identity().cuda()
32 self.model.eval()
33
34 ## 实现算子的__call__方法,接收一批数据,返回一批数据的处理结果,在__call__中调用初始化函数中的模型进行推理
35 def __call__(self, batch: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
36 inputs = torch.as_tensor(batch["data"], dtype=torch.float32).cuda()
37 with torch.inference_mode():
38 ## 调用模型进行推理
39 batch["output"] = self.model(inputs).detach().cpu().numpy()
40 return batch
-
算子测试
- 测试代码:编写test.py,其中包含自定义算子的定义和初始化、输入数据的构造、输出数据的展示。
1import pandas as pd
2import ray
3from typing import Dict
4import numpy as np
5import torch
6
7from palette.ops.base import Transformer
8from palette.core.ray_dataset import RayDataset
9
10# 1.写自定义算子
11class TorchPredictor(Transformer):
12 ##算子设置为transform类型
13 _op_type = "transform"
14 ##算子设置为批处理,即__call__中接收一批数据,返回一批数据的处理结果
15 _batched_op = True
16 ## 算子的名称
17 _name = "operator_transform"
18 ## 设置算子执行模式为PIPELINE_ACTOR,这样init只会执行一次
19 _ray_execute_mode = "PIPELINE_ACTOR"
20 ## 设置算子输入数据格式为numpy
21 _ray_batch_format = "numpy"
22 ## 若算子需要使用GPU,则设置算子处理器为GPU
23 _processor = "cuda"
24
25
26 def __init__(self, *args, **kwargs):
27 ##调用父类的初始化函数
28 super().__init__(*args, **kwargs)
29
30 ##初始化模型,也可以使用模型权重文件加载
31 self.model = torch.nn.Identity().cuda()
32 self.model.eval()
33
34 ## 实现算子的__call__方法,接收一批数据,返回一批数据的处理结果,在__call__中调用初始化函数中的模型进行推理
35 def __call__(self, batch: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
36 inputs = torch.as_tensor(batch["data"], dtype=torch.float32).cuda()
37 with torch.inference_mode():
38 ## 调用模型进行推理
39 batch["output"] = self.model(inputs).detach().cpu().numpy()
40 return batch
41
42# 2.构造数据测试自定义算子
43# 2.1. create dataset
44samples = ray.data.from_numpy(np.ones((32, 100)))
45input_dataset = RayDataset(ray_dataset=samples)
46
47# 2.2. create operator
48op = TorchPredictor()
49
50# 2.3. run
51output_dataset = input_dataset.run(op)
52
53# 2.4. show result
54output_dataset.data.show()
1* 测试结果:在本地执行test.py

步骤四:Filter算子
1.PunctuationFilter 算子类结构详解
1.1.类属性 - 算子元信息定义

1.2.初始化方法 - 参数设置与预处理

初始化流程:
- 调用父类Filter的初始化方法
- 匹配标点符号
1.3.核心处理方法 - 统计方法 _compute_stats

关键设计解析:
-
逐字符检测机制
- 采用
for char in text逐字符扫描,确保不遗漏任何位置 - 使用集合
punctuation_chars实现O(1)复杂度的查找 - 发现第一个标点立即
break,优化处理效率
- 采用
1.4.过滤方法 compare_stats

关键设计解析:
-
双模式处理
- 批量模式:返回惰性求值的
map对象(兼容大数据场景) - 单条模式:直接返回布尔值(适配Ray分布式任务)
- 批量模式:返回惰性求值的
2.PunctuationFilter 算子接口详解
2.1.公共接口总览
| 接口类型 | 名称 | 说明 |
|---|---|---|
| 类属性 | 同上表 | 定义算子元信息 |
| 初始化方法 | __init__ |
配置标点符号集合 |
| 统计方法 | _compute_stats |
生成标点检测结果 |
| 过滤方法 | compare_stats |
执行过滤逻辑 |
2.2.核心接口细节
统计方法 _compute_stats

功能说明:
为每个文本生成统计信息
检测文本是否包含标点符号
结果存储在Fields.stats字段中
过滤方法 compare_stats

过滤逻辑:
批量处理时:返回每个文本的过滤决策列表
单条处理时:返回单个布尔值
过滤条件:不包含标点符号的文本
2.3.接口调用示例
基本使用
1# 1. input data
2 samples = [
3 {
4 'text': '特殊的康熙部首或者扩展部首会被去除,⼏几⺇'
5 },
6 {
7 'text': '请问你是谁dasoidhao@1264fg.45om'
8 },
9 {
10 'text': '匹配汉字包括繁體字'
11 }
12 ]
13
14 input_dataset = RayDataset.from_list(samples)
15
16 # 2. create operator
17 op = PunctuationFilter()
18
19 # 3. run
20 output_dataset = input_dataset.run(op)
21
22 # 4. get & check result
23 text_data = output_dataset.get_column(column=op.text_key)
24
25 print(f'---result: {text_data}')
自定义算子测试
标点替换算子测试 (test_punctuation_replace.py)
输入数据:

输入转换:

输出数据:

图像缩放算子测试 (test_image_resize.py)
输入数据:

-
格式规范:
- 每个样本为字典类型
- 必须包含
images键(对应op.image_key) - 值为图像文件路径(相对/绝对路径)
输入转换:

-
转换后特征:
- 分布式可处理格式(兼容Ray引擎)
- 保留原始字段结构
- 自动分区处理(对用户透明)
输出结果:
输出结构:
1* 保持与输入相同的记录数(3条)
2* 每个输出包含处理后的图像路径
3* 路径格式:`./ret/原始文件名.jpg`
基础功能验证:
| 测试重点 | 验证方式 | 预期结果 |
|---|---|---|
| 尺寸调整准确性 | 检查输出图像的image.size |
应全部变为(1024, 1024) |
| 路径处理正确性 | 检查image_path是否在dst_path下 |
输出路径必须包含./ret |
| 文件名保留 | need_hash_name=False时 |
文件名应与输入保持一致 |
异常场景验证:
| 测试类型 | 测试用例设计 | 预期处理 |
|---|---|---|
| 非法输入 | 传入非图像文件(如.txt) | 应抛出明确的异常提示 |
| 空路径 | 'images': '' |
应跳过或报错 |
| 网络存储 | 使用volume://前缀路径 |
需验证SmartFile转换正确性 |
标点过滤算子测试 (test_db_punctuation_filter.py)
输入数据:

格式转换:

输出结果:

基础功能验证:
| 测试维度 | 验证方法 | 通过标准 |
|---|---|---|
| 过滤准确性 | 输出记录数 == 1 |
仅保留纯净文本记录 |
| 字段完整性 | 'text' in output_dataset.schema() |
保留原始字段结构 |
| 内容保留 | '繁體字' in text_data[0] |
确保未修改保留文本 |
异常处理规范:
| 异常类型 | 测试用例 | 预期行为 |
|---|---|---|
| 空文本 | {'text': ''} |
应保留(视作无标点) |
| 非法编码 | {'text': b'\xff\xfe'} |
抛出UnicodeDecodeError |
| 缺失字段 | {'id': 123} |
抛出KeyError |
自定义算子打包
1. 执行打包命令
在开发容器中进入样例包目录并执行:
1cd /root/databuilder_dev
2sh ./build.sh

打包完成后,whl 文件会生成在以下目录:
1/root/databuilder_dev/output/dist

2. 在容器内安装并复测
将生成的 whl 包安装到开发容器中:
1cd /root/databuilder_dev/output/dist
2pip install databuilder_vendor_operators-<版本>-py3-none-any.whl
安装完成后,再次运行对应测试脚本。只有本地测试通过后,才建议上传到百度胜算创建算子版本。
自定义算子上线
1. 创建算子
- 登录百度胜算控制台,进入目标工作空间。
- 进入元数据,选择目标数据目录和数据模式。
- 在数据模式页面单击立即创建 > 创建算子。

-
填写以下信息:
- 算子名称:填写自定义算子名称,只能使用字母、数字和下划线,并确保在当前数据模式中不重复;
- 算子别名:填写便于用户理解的名称;
- 描述:说明算子的处理能力;
- 使用说明:说明输入字段、输出字段和使用方式。

- 选择提交并创建算子版本继续创建版本。算子说明、控制台配置项的完整说明,请参考算子。
2. 创建算子版本
-
创建算子版本时,依次填写以下配置:
- 基本信息:配置版本描述和算子类型,本示例现在
TRANSFORM;Filter示例选择FILTER。

- 代码配置:代码语言选择Python,然后添加上述步骤打包生成的whl文件。

- 运行约束:支持引擎选择Ray,资源类型选择CPU,类名输入catalog_op.default.image_resizer.v1.ImageResizer,然后选择官方镜像,基础镜像。

注意: 如果选择自定义镜像,需要先准备镜像仓库并将镜像推送到仓库。镜像仓库、镜像凭证和白名单配置请以项目环境为准。
- 参数配置:
- 基本信息:配置版本描述和算子类型,本示例现在

- 检查上传的 whl 文件、类名、算子类型、支持引擎和资源类型。
- 单击保存,完成算子版本创建。
保存后,可以在对应的数据目录和数据模式下找到该自定义算子。
使用自定义算子
1.创建工作流
- 在左侧导航栏选择数据处理 > 工作流。
- 单击创建,填写工作流名称、所属项目和工作流类型。
- 然后单击确定。
2.添加并配置算子
- 在工作流编辑页面,选择创建空白工作流,从左侧任务列表中选择算子任务。
- 可以通过左侧的算子构建工作流。
- 将自定义算子添加到画布中,在右侧算子信息配置单击浏览按钮,选择已创建的算子。
- 再将系统算子和自定义算子按数据处理顺序串联。

3.运行并查看运行结果
- 核对算子的执行资源和依赖配置。
- 单击保存,然后单击立即运行,启动工作流。

- 在工作流页面单击运行记录。
- 找到刚刚运行的工作流,在操作列单击查看。

- 选择自定义算子任务,单击右侧的任务结果,查看输出数据或输出路径。

- 复制该路径可以在元数据-catalog_op-数据卷-output中找到输出结果:

- 可以下载该文件本地打开,可以看到:

- 选择其中一个路径寻找验证,在元数据-catalog_op-数据卷-output_image/2中:

- 打开后同样查找路径,下载图片。

- 下载后打开,可以看到图片尺寸。

自定义算子的最佳实践
该章节提供了一部分自定义算子的最佳实践,可以借助这些最佳实践,开发高效运行的自定义算子。
使用vllm加载qwen-vl-2b模型进行道路图片标注的Transformer算子
1import base64
2import json
3import time
4import os
5import re
6
7from palette.ops.base import Transformer
8from palette.util.file_utils import SmartFile
9from databuilder.model.util import get_model_path
10
11system_prompt = """
12You are an expert driving-scene annotator. You classify what is visible in a single image and produce a JSON object with fixed keys.xxx
13"""
14user_prompt = """You will receive a single forward-driving scene image (e.g. base64 or URL). xxxxx"""
15
16def run_inference(image_path, system_prompt_text, user_prompt_text, model, sampling_params):
17 inference_start_time = time.perf_counter()
18 print(f"Start to annotation input image {image_path}. Start timestamp: {inference_start_time}#############")
19 user_content = [
20 {"type": "text", "text": user_prompt_text},
21 ]
22
23 with open(image_path, "rb") as image_file:
24 image_data = base64.b64encode(image_file.read()).decode("utf-8")
25 new_image = {"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{image_data}"}}
26 user_content.append(new_image)
27
28 messages = [{"role": "system", "content": [{"type": "text", "text": system_prompt_text}]},
29 {"role": "user", "content": user_content}]
30
31 outputs = model.chat([messages], sampling_params)
32 inference_end_time = time.perf_counter()
33 inference_elapsed_time = inference_end_time - inference_start_time
34 print(f"Finish annotation input image {image_path}. End timestamp: {inference_end_time}, use {inference_elapsed_time:.4f}s#############")
35 return outputs[0].outputs[0].text
36
37def parse_llm_json(llm_output):
38 try:
39 # 1. 尝试使用正则表达式提取 ```json ... ``` 或 ``` ... ``` 中的内容
40 # re.DOTALL 让 . 也能匹配换行符
41 pattern = r"```(?:json)?\s*(.*?)```"
42 match = re.search(pattern, llm_output, re.DOTALL)
43
44 if match:
45 # 如果找到了 Markdown 代码块,提取其中的内容
46 json_str = match.group(1)
47 else:
48 # 2. 如果没找到 Markdown 标记,尝试直接寻找最外层的 {} 或 []
49 # 这一步是为了处理 "好的,这是你的 JSON:{...}" 这种情况
50 # 寻找第一个 { 或 [ 开始,到最后一个 } 或 ] 结束
51
52 # 简单的启发式搜索:找到第一个 { 或 [
53 start_index = -1
54 end_index = -1
55
56 # 寻找最早出现的 { 或 [
57 first_curly = llm_output.find('{')
58 first_square = llm_output.find('[')
59
60 if first_curly == -1 and first_square == -1:
61 # 既没有 { 也没有 [,可能不是 JSON
62 raise ValueError("未找到 JSON 开始符号")
63
64 # 确定开始位置
65 if first_curly != -1 and (first_square == -1 or first_curly < first_square):
66 start_index = first_curly
67 end_index = llm_output.rfind('}') + 1
68 else:
69 start_index = first_square
70 end_index = llm_output.rfind(']') + 1
71
72 if start_index != -1 and end_index != -1:
73 json_str = llm_output[start_index:end_index]
74 else
75 json_str = llm_output
76
77 # 3. 解析 JSON
78 return json.loads(json_str)
79
80 except json.JSONDecodeError as e:
81 print(f"JSON 解析失败: {e}")
82 return None
83 except Exception as e:
84 print(f"发生错误: {e}")
85 return None
86
87#定义算子类ImageAnnotationVLLM,继承于Transformer
88class ImageAnnotationVLLM(Transformer):
89 _op_type = "transform" # 算子类型:转换类
90 _batched_op = True # 支持批量处理
91 _processor = 'cuda' # 默认使用CPU加速,设置为cuda表示算子需要使用GPU资源运行
92 _name = "image_annotation_vllm" # 算子名称
93 _ray_execute_mode = "PIPELINE_ACTOR" # Ray的Actor执行模式
94 _ray_batch_format = "numpy" # 批处理格式
95
96 def __init__(self,
97 max_tokens: int = 4096,
98 max_model_len: int = 8192,
99 temperature: float = 0.7,
100 top_p: float = 0.9,
101 model_name: str = None,
102 preload: bool = False, #若算子有模型加载等一系列数据处理过程中只运行一次的操作,可以在构造器中加入preload参数,当preload为True完成模型加载等一系列操作。
103 *args, **kwargs):
104 super().__init__(*args, **kwargs)
105 self.model = None
106 self.sampling_params = None
107 self.smart_file = SmartFile()
108 self.download_path = os.path.join(os.getcwd(), os.path.splitext(os.path.basename(__file__))[0])
109 self.model_name = model_name
110 self.max_tokens = max_tokens
111 self.max_model_len = max_model_len
112 self.temperature = temperature
113 self.top_p = top_p
114
115 #自定义算子创建时,可设置依赖的模型版本,Databuilder会自动下载和复用模型文件。此处可以通过get_model_path来获取依赖模型的下载地址
116 self.model_path = get_model_path(self.model_name)
117
118 # 当preload为True时,完成模型加载和下载目录的初始化,此处可以保证整个数据处理过程中,模型只会加载一次。
119 if preload:
120 os.makedirs(self.download_path, exist_ok=True)
121
122 print("preload model")
123 self.init_model()
124
125 def init_model(self):
126 from vllm import LLM, SamplingParams
127 process_start_time = time.perf_counter()
128
129 # 使用LLM加载本地的模型文件
130 self.model = LLM(
131 model=self.model_path,
132 trust_remote_code=True,
133 max_model_len=self.max_model_len,
134 )
135
136 self.sampling_params = SamplingParams(
137 max_tokens=self.max_tokens,
138 temperature=self.temperature,
139 top_p=self.top_p,
140 )
141
142 process_elapsed_time = time.perf_counter() - process_start_time
143 print(f"Finish loading {self.model_name} model by vllm. Use {process_elapsed_time:.4f}s#############")
144
145 def download_image(self, image_path):
146 local_path = self.smart_file.volume_2_local(image_path, self.download_path)
147 return local_path
148
149 #实现transform算子的transform_batched函数,因为_ray_batch_format设置为numpy,所以sample的类型为Dict[str, numpy.ndarray]
150 def transform_batched(self, samples):
151 #取出samples中的images列
152 src_images = samples[self.image_key]
153 #最后返回图片名及标注信息,由此创建两个list。
154 file_name = []
155 annotations = []
156 #对于每张图片,先从volume下载到本地,随后传入图片下载路径、模型等参数完成模型推理。
157 for image_path in src_images:
158 local_path = self.download_image(image_path)
159 annotation = run_inference(local_path, system_prompt, user_prompt, self.model, self.sampling_params)
160 ret = parse_llm_json(annotation)
161 #将图片名加入file_name
162 file_name.append(image_path)
163 #将标注信息加入annotations
164 annotations.append(ret)
165
166 #最后返回一个Dict[str, list]也可返回Dict[str, numpy.ndarray],包含两个元素分别是file_name和annotations。
167 return {"file_name": file_name, "annotations": annotations}
使用transformer加载qwen-vl-2b模型进行道路图片标注的Transformer算子
1import base64
2import json
3import re
4import time
5import os
6
7from databuilder.model.util import get_model_path
8from palette.ops.base import Transformer
9from palette.util.file_utils import SmartFile, make_unique_name_by_dict
10
11system_prompt = """
12You are an expert driving-scene annotator. You classify what is visible in a single image and produce a JSON object with fixed keys.xxx
13"""
14user_prompt = """You will receive a single forward-driving scene image (e.g. base64 or URL). Classify all attributes below.xxx
15"""
16
17def run_inference(image_path, system_prompt_text, user_prompt_text, model, processor):
18 inference_start_time = time.perf_counter()
19 print(f"Start to annotation input image {image_path}. Start timestamp: {inference_start_time}#############")
20 with open(image_path, "rb") as image_file:
21 image_data = base64.b64encode(image_file.read()).decode("utf-8")
22 # prepare Qwen-VL message template
23 user_content = [
24 {"type": "text", "text": user_prompt_text},
25 {"type": "image", "image": image_data}
26 ]
27 messages = [{"role": "system", "content": [{"type": "text", "text": system_prompt_text}]}, {"role": "user", "content": user_content}]
28
29 inputs = processor.apply_chat_template(
30 messages, tokenize=True, add_generation_prompt=True, return_dict=True, return_tensors="pt"
31 )
32
33 inputs = inputs.to(model.device)
34
35 generated_ids = model.generate(**inputs, max_new_tokens=4096)
36 generated_ids_trimmed = [
37 out_ids[len(in_ids):] for in_ids, out_ids in zip(inputs.input_ids, generated_ids)
38 ]
39 output_text = processor.batch_decode(
40 generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False
41 )
42
43 inference_end_time = time.perf_counter()
44 inference_elapsed_time = inference_end_time - inference_start_time
45 print(f"Finish annotation input image {image_path}. End timestamp: {inference_end_time}, use {inference_elapsed_time:.4f}s#############")
46 return output_text[0]
47
48def parse_llm_json(llm_output):
49 try:
50 # 1. 尝试使用正则表达式提取 ```json ... ``` 或 ``` ... ``` 中的内容
51 # re.DOTALL 让 . 也能匹配换行符
52 pattern = r"```(?:json)?\s*(.*?)```"
53 match = re.search(pattern, llm_output, re.DOTALL)
54
55 if match:
56 # 如果找到了 Markdown 代码块,提取其中的内容
57 json_str = match.group(1)
58 else:
59 # 2. 如果没找到 Markdown 标记,尝试直接寻找最外层的 {} 或 []
60 # 这一步是为了处理 "好的,这是你的 JSON:{...}" 这种情况
61 # 寻找第一个 { 或 [ 开始,到最后一个 } 或 ] 结束
62
63 # 简单的启发式搜索:找到第一个 { 或 [
64 start_index = -1
65 end_index = -1
66
67 # 寻找最早出现的 { 或 [
68 first_curly = llm_output.find('{')
69 first_square = llm_output.find('[')
70
71 if first_curly == -1 and first_square == -1:
72 # 既没有 { 也没有 [,可能不是 JSON
73 raise ValueError("未找到 JSON 开始符号")
74
75 # 确定开始位置
76 if first_curly != -1 and (first_square == -1 or first_curly < first_square):
77 start_index = first_curly
78 end_index = llm_output.rfind('}') + 1
79 else:
80 start_index = first_square
81 end_index = llm_output.rfind(']') + 1
82
83 if start_index != -1 and end_index != -1:
84 json_str = llm_output[start_index:end_index]
85 else:
86 json_str = llm_output
87
88 # 3. 解析 JSON
89 return json.loads(json_str)
90
91 except json.JSONDecodeError as e:
92 print(f"JSON 解析失败: {e}")
93 return None
94 except Exception as e:
95 print(f"发生错误: {e}")
96 return None
97
98#定义算子类ImageAnnotationVLLM,继承于Transformer
99class ImageAnnotationTransform(Transformer):
100 _op_type = "transform" # 算子类型:转换类
101 _batched_op = True # 支持批量处理
102 _processor = 'cuda' # 默认使用CPU加速,设置为cuda表示算子需要使用GPU资源运行
103 _name = "image_annotation_transform" # 算子名称
104 _ray_execute_mode = "PIPELINE_ACTOR" # Ray的Actor执行模式
105 _ray_batch_format = "numpy" # 批处理格式
106
107 def __init__(self,
108 model_name: str = None,
109 preload: bool = False, #若算子有模型加载等一系列数据处理过程中只运行一次的操作,可以在构造器中加入preload参数,当preload为True完成模型加载等一系列操作。
110 *args, **kwargs):
111 super().__init__(*args, **kwargs)
112 self.processor = None
113 self.model = None
114 self.smart_file = SmartFile()
115 self.download_path = os.path.join(os.getcwd(), os.path.splitext(os.path.basename(__file__))[0])
116 self.model_name = model_name
117
118 #自定义算子创建时,可设置依赖的模型版本,Databuilder会自动下载和复用模型文件。此处可以通过get_model_path来获取依赖模型的下载地址
119 self.model_path = get_model_path(self.model_name)
120
121 # 当preload为True时,完成模型加载和下载目录的初始化,此处可以保证整个数据处理过程中,模型只会加载一次。
122 if preload:
123 print("preload model")
124 self.init_model()
125
126 def init_model(self):
127 # 使用transformers加载本地的模型文件
128 from transformers import AutoModelForImageTextToText, AutoProcessor
129 process_start_time = time.perf_counter()
130 self.model = AutoModelForImageTextToText.from_pretrained(
131 self.model_path,
132 dtype="auto",
133 device_map="auto"
134 )
135 self.processor = AutoProcessor.from_pretrained(self.model_path)
136 process_elapsed_time = time.perf_counter() - process_start_time
137 print(f"Finish loading {self.model_name} model. Use {process_elapsed_time:.4f}s#############")
138
139 def download_image(self, image_path):
140 local_path = self.smart_file.volume_2_local(image_path, self.download_path)
141 return local_path
142
143 #实现transform算子的transform_batched函数,因为_ray_batch_format设置为numpy,所以sample的类型为Dict[str, numpy.ndarray]
144 def transform_batched(self, samples):
145 #取出samples中的images列
146 src_images = samples[self.image_key]
147 #最后返回图片名及标注信息,由此创建两个list。
148 file_name = []
149 annotations = []
150 #对于每张图片,先从volume下载到本地,随后传入图片下载路径、模型等参数完成模型推理。
151 for image_path in src_images:
152 os.makedirs(self.download_path, exist_ok=True)
153 local_path = self.download_image(image_path)
154 annotation = run_inference(local_path, system_prompt, user_prompt, self.model, self.processor)
155 ret = parse_llm_json(annotation)
156 #将图片名加入file_name
157 file_name.append(image_path)
158 #将标注信息加入annotations
159 annotations.append(ret)
160
161 #最后返回一个Dict[str, list]也可返回Dict[str, numpy.ndarray],包含两个元素分别是file_name和annotations。
162 return {"file_name": file_name, "annotations": annotations}
FAQ
算子whl包的版本号如何更新?
算子代码编写完毕,线下环境测试没有问题需要更新到线上DataBuilder环境时,需要对whl的版本号进行更新,修改setup.py中的version字段值。
1* 推荐版本以4段式表示,每个版本段数字以递增的方式修改,正式版本更新时更新前3个版本段数字,非正式版本仅更新第4个版本段数字
2 * 例如 0.0.1.12
3
4* 不管正式版本还是非正式版本,上线到DataBuilder平台时,都需要更新版本号,更新前后相同的版本号whl包DataBuilder平台不会拉取更新
评价此篇文章
