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

如何使用ONNX算子实现正则表达式字符串替换?

问题:Scikit-Learn管道转ONNX时自定义字符串正则替换转换器的实现困境

背景

我正尝试将文本输入的Scikit-Learn机器学习管道导出为ONNX格式。管道里的TfIdfVectorizer、TruncatedSVD等转换器已有标准支持,但第一步需要实现一个通过正则表达式修改输入文本的自定义转换器。

根据scikitlearn-onnx文档,添加自定义转换器需要编写自定义形状函数和转换函数,且转换函数必须基于ONNX标准预定义算子实现,但现有ONNX算子无法完成基础字符串操作。

核心需求

要实现的正则替换示例为单位转换:12m -> 12 meters,用Python的re包可以轻松实现,但没法通过现有ONNX算子完成。

已尝试的方法及问题

  • 用类Python正则算子:ONNX没有正则相关算子;
  • 逐字符遍历替换数字后的"m":OnnxEqual不支持字符串比较;
  • 转ASCII值后处理:ONNX没有类似GNU tr的转换算子,OnnxCast不支持非严格数值字符串转换;
  • 用OnnxUnique的inverse_indicies转换:OnnxSplit处理字符串张量报错,OnnxSequenceInsert无法将字符串合并为单元素张量。

测试代码

import re
import numpy as np
from sklearn.base import BaseEstimator, TransformerMixin
from skl2onnx import to_onnx, update_registered_converter
from skl2onnx.common.data_types import StringTensorType
from skl2onnx.algebra.onnx_ops import OnnxSplit, OnnxConstant
from onnxruntime import InferenceSession

class MyTransformer(BaseEstimator, TransformerMixin):
    def fit_transform(self, X, y=None):
        return re.sub("(?<=[0-9])m ", " meters ", X)

def shape_function(operator):
    input = StringTensorType([1])
    output = StringTensorType([None, 1])
    operator.inputs[0].type = input
    operator.outputs[0].type = output

def converter_function(scope, operator, container):
    op = operator.raw_operator
    opv = container.target_opset
    out = operator.outputs

    X = operator.inputs[0]

    one_tensor = OnnxConstant(value_int=1, op_version=opv)
    string_tensor = OnnxConstant(value_strings=["ab"], op_version=opv)
    string_split_tensor = OnnxSplit(string_tensor, one_tensor, op_version=opv, output_names=out[:1])

    string_split_tensor.add_to(scope, container)

update_registered_converter(MyTransformer, "MyTransformer", shape_function, converter_function)
my_transformer = MyTransformer()
onnx_model = to_onnx(my_transformer, initial_types=[["X", StringTensorType([None, 1])]])

test_string = "The Empire State Building is 443m tall."
sess = InferenceSession(onnx_model.SerializeToString())
output = sess.run(None, {"X": np.array([test_string])})

运行报错信息

2022-08-16 12:35:46.235861185 [W:onnxruntime:, graph.cc:106 MergeShapeInfo] Error merging shape info for output. 'variable' source:{1} target:{,1}. Falling back to lenient
merge.
2022-08-16 12:35:46.237767860 [E:onnxruntime:, inference_session.cc:1530 operator()] Exception during initialization: /onnxruntime_src/onnxruntime/core/optimizer/optimizer_execution_frame.cc:75 onnxruntime::OptimizerExecutionFrame::Info::Info(const std::vector<const onnxruntime::Node*>&, const InitializedTensorSet&, const onnxruntime::Path&,
const onnxruntime::IExecutionProvider&, const std::function<bool(const std::__cxx11::basic_string<char>&)>&) [ONNXRuntimeError] : 2 : INVALID_ARGUMENT : string tensor can not use pre-allocated buffer

疑问

如何利用现有ONNX算子正确完成这类字符串操作?


内容的提问来源于stack exchange,提问作者NolantheNerd

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 20:57:19