onnx.register_custom_op_symbolic
Функция onnx.register_custom_op_symbolic используется для регистрации
пользовательских символических функций при экспорте модели в формат ONNX.
Она позволяет определить, как пользовательские операции PyTorch должны быть
преобразованы в стандартные операции ONNX. Первым параметром функция принимает
имя символической функции, вторым - функцию, которая будет вызываться во время
экспорта для преобразования операции.
Синтаксис
torch.onnx.register_custom_op_symbolic(symbolic_name, symbolic_fn, opset_version)
Пример 1
Давайте зарегистрируем пользовательскую операцию для операции custom_abs:
import torch
import torch.onnx
def custom_abs_symbolic(g, input):
return g.op("Abs", input)
torch.onnx.register_custom_op_symbolic(
"custom_ops::custom_abs",
custom_abs_symbolic,
opset_version=13
)
Пример 2
Давайте зарегистрируем пользовательскую операцию с несколькими аргументами:
import torch
import torch.onnx
def custom_add_symbolic(g, input1, input2, alpha):
alpha_tensor = g.op("Constant", value_t=torch.tensor(alpha))
return g.op("Add", input1, g.op("Mul", input2, alpha_tensor))
torch.onnx.register_custom_op_symbolic(
"custom_ops::custom_add",
custom_add_symbolic,
opset_version=13
)
Пример 3
Давайте зарегистрируем пользовательскую операцию для пользовательского модуля:
import torch
import torch.onnx
import torch.nn as nn
class CustomModule(nn.Module):
def forward(self, x):
return torch.sigmoid(x) * 2
def custom_module_symbolic(g, input):
sigmoid = g.op("Sigmoid", input)
scale = g.op("Constant", value_t=torch.tensor(2.0))
return g.op("Mul", sigmoid, scale)
torch.onnx.register_custom_op_symbolic(
"custom_ops::CustomModule",
custom_module_symbolic,
opset_version=13
)
Пример 4
Давайте зарегистрируем операцию для функции с несколькими типами аргументов:
import torch
import torch.onnx
def custom_clamp_symbolic(g, input, min_val, max_val):
min_tensor = g.op("Constant", value_t=torch.tensor(min_val))
max_tensor = g.op("Constant", value_t=torch.tensor(max_val))
clamped = g.op("Max", min_tensor, input)
return g.op("Min", clamped, max_tensor)
torch.onnx.register_custom_op_symbolic(
"custom_ops::custom_clamp",
custom_clamp_symbolic,
opset_version=13
)
Пример 5
Давайте зарегистрируем операцию для функции с именованными аргументами:
import torch
import torch.onnx
def custom_linear_symbolic(g, input, weight, bias=None):
output = g.op("MatMul", input, weight)
if bias is not None:
output = g.op("Add", output, bias)
return output
torch.onnx.register_custom_op_symbolic(
"custom_ops::custom_linear",
custom_linear_symbolic,
opset_version=13
)
Смотрите также
-
функцию
onnx.export,
которая экспортирует модель PyTorch в формат ONNX -
функцию
onnx.dynamo_export,
которая экспортирует модель с использованием Dynamo -
функцию
export.export,
которая экспортирует модель в стандартный формат -
функцию
jit.script,
которая компилирует модель в TorchScript