10. MUSA AI 框架算子适配及 CV 图像处理开发示例
10.1. muPyTorch - muDNN 对接扩展
muPyTorch - muDNN 对接扩展,可以让 PyTorch 框架使用 muDNN 算子库进行扩展,从而可以让用户在不对模型进行较大改动的情况下,使用 musa - plugin 在 MTGPU 上运行并计算 AI 模型。本节主要关注:PyTorch 的 dispatch 机制、算子调用以及开发者如何进行 muPyTorch 与 muDNN 算子库的对接。对接完成后,PyTorch 中的算子就可以通过 muDNN 算子库,运行在 MTGPU 设备上,进行算子计算和模型训练及推理。
10.1.1. 算子派发机制
PyTorch 中有很多算子,同一种算子又分为很多类型。以 add 算子为例,按 device 类型可以分为 CPU、GPU、NPU、TPU、FPGA 等设备;按 layout 类型可以分为普通张量和稀疏张量,不同的张量有不同的布局;按 dtype 类型可以分为 fp64、fp32、fp16、int8,甚至 bf16、hf32 等不同类型。那么当 python 端调用一个 torch.add 算子时,最后实际执行的算子该怎么选择呢?此时,就要提到 PyTorch 的 dispatch 机制了,接下来便对 dispatch 机制进行具体介绍。
10.1.1.1. dispatch 机制
dispatcher 可以理解为分发器,当执行一个 operator 时,dispatcher 会根据 tensor 输入的一些信息和一些其他信息(参数个数、返回值类型等)计算得到一个 dispatch key,然后根据 dispatch key 从 dispatch table 中找到对应的 kernel 函数指针,最后回调执行。dispatch table 信息如下图所示。

图 1 Dispatch Table 表示意图
可以看到 dispatch key 不仅有硬件后端,还有一些更抽象的概念如 autograd、tracing 等。dispatch key 的计算是通过一个 dispatch key set 的结构来实现的,dispatch key set 可以理解为一个 64 bit 的数组,每一个 bit 都代表了一个 dispatch key,从左到右有优先级关系。同样一算子,有针对不同 dispatch key 的实现,针对这些散落 在 PyTorch 各处的注册实现,这个数组将所有可用实现都结合到一起,调用其中优先级最高的 dispatch key 对应的 kernel 实现。那么这个 dispatch table 是如何形成的呢?此时,就要提到算子的注册机制了,接下来介绍算子注册的交互方式。
10.1.1.2. dispatch table 注册
首选,需要定义一个关于算子的运算符模式 schema。此时并没有提供操作符的具体实现,只是提供了一个模式字符串,指定操作符的类型签名。后续其他内核具体实现时,都将遵守这种模式。以 add 算子为例,具体定义方式如下:
TORCH_LIBRARY(myops, m) {
m.def("add(Tensor self, Tensor other) -> Tensor");
}
如上所示,已经定义了 add 的 schema,接下来需要提供这个操作符的具体后端实现。最简单的注册方法是 def("add", add_cpu),这会将内核注册为在所有情况下运行,即使这个张量不是 CPU 张量。为了确保我们注册的特定算子在特定的设备上运行,可以使用 TORCH_LIBRARY_IMPL 宏。在这个宏中使用 m.impl 就可以将带有 dispatch key 信息的算子实现注册到 dispatch table 中。具体定义方式如下:
TORCH_LIBRARY_IMPL(myops, CPU, m) {
m.impl("add", add_cpu_fun_ptr);
}
如上所示,本例中的 CPU 便是 dispatch key,add_cpu_fun_ptr 便是 add 算子 CPU 对应的函数指针。同理,也可以注册一个 CUDA 后端或 MTGPU 后端的 item,这些注册可以跨文件分割或者跨库边界分割。通常来说,注册结构为一个单独的 TORCH_LIBRARY 文件,里面集中的列出了命名空间中每个自定义的操作符。然后每个 dispatch key 都有一个 TORCH_LIBRARY_IMPL 块,当然也可以将 TORCH_LIBRARY_IMPL 块按 operator 类型切分为不同的块。特别地,当每个算子都有一个分离的文件时,同时不想在头文件中暴露这些算子时,可以直接将注册文件放在函数定义的 cpp 文件中。
10.1.2. 算子调用流程
通过上面介绍,可以了解到算子的注册方式和算子的 dispatch 机制。那么,在 python 端调用一个 torch.add 算子是如何衔接到到我们的 C++ 代码执行呢?这就要说到 PyTorch 中的 C++/CUDA 扩展了,C++/CUDA 扩展一般有预编译和实时编译 (JIT) 模式,这里主要介绍预编译模式,即假设我们的 C++ 拓展代码已经编译打包成库文件。Python 中有一个 pybind11 的宏,主要用来在 C++ 代码中创建 Python 的链接库。这里以创建名为 ext 的扩展库 (extension) 为例,通过 pybind11 绑定 C++ 代码到 Python 示例如下:
Tensor my_add(Tensor a, Tensor b);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m ){
m.def("my_add", &my_add, "my_add (CPU/CUDA) ", py::arg("a"), py::arg("b"));
}
如上所示,PYBIND11_MODULE 的作用是为 C++ 代码接入 Python 解释器提供入口,TORCH_EXTENSION_NAME 是编译器在编译扩展库过程中出现的宏,对应为 extension 中的 name 变量,在这里会被解释为 ext。m 代表 TORCH_EXTENSION_NAME 所对应的实例模块,{}中的每个 m.def 都定义了一个 ext 的成员函数,其一般形式为 m.def("函数名",函数指针, "文档", 参数列表)。通过这种形式,my_add 也就顺利地成为了 ext 的成员函数,其具体实现为已经定义好的 my_add 函数。在 Python 脚本中通过"from ext import my_add"导入 ext 模块中的 my_add,就可以调用自定义算子 my_add 了。
10.1.3. muDNN 算子库对接
本节开始主要介绍了 muDNN 算子库与 muPyTorch 框架的具体对接流程,涉及设备注册、算子对接、算子正确性测试等内容。
10.1.3.1. 设备注册流程
通过上述介绍,我们更关心如何在 PyTorch 中注册新的设备后端,即如何让 PyTorch 支持我们的 MTGPU 设备。注册的详细步骤如下:
- 步骤 1:添加
backend
torch.backends控制 PyTorch 支持的各种后端的行为,如:torch.backends.cuda,torch.backends.mkl。- 在
PyTorch/c10/core/Backend.h头文件的 "enum class Backend" 中添加后端MUSA,并且在backendToDispatchKey()中将backend与DispatchKey(PrivateUse1)绑定,在backendToDeviceType()将backend与DeviceType(MTGPU)绑定。
- 步骤 2:添加
device
- 在
PyTorch/c10/core/DeviceType.h头文件的 "enum class DeviceType" 中添加 MTGPU。 - 在
PyTorch/c10/core/Device.h头文件的 "struct Device" 中添加is_mtgpu(),获得对应的DeviceType。 - 在
PyTorch/c10/core/DispatchKey.cpp源文件的函数toBackendComponent()中添加 MTGPU 对应的BackendComponent。 - 在
PyTorch/c10/core/TensorOptions.h头文件的函数computeDispatchKey()中根据 MTGPU 获得对应的DispatchKey。 - 在
PyTorch/torch/library.h头文件的dispach中将Device_type MTGPU与DispatchKey::PrivateUse1绑定。
10.1.3.2. 算子对接流程
算子对接可以将 PyTorch 框架与 muDNN 算子库进行对接,使得在不对用户模型进行较大改动的情况下,将 AI 模型通过 muDNN 算子库在 mtgpu 上进行运行和计算。使用方式也十分简单,仅需将数据和模型搬运到 mtgpu 上,计算完成后再搬回 cpu 即可。这里以 relu 算子为例,介绍算子库对接的基本流程。
首先,需要查看 muDNN 算子库是否支持 relu 算子,而 relu 算子属于 Unary 类算子,查看 muDNN 头文件 mudnn_math.h 可知,Unary.Mode 支持 relu 算子。接下来,需要查看 PyTorch 框架的 C++ 拓展接口中 relu 算子对应的 dispatch 函数,在 aten/src/ATen/native/native_functions.yaml 文件中,搜索 relu 可以得到如下代码:
- func: relu(Tensor self) -> Tensor
device_check: NoCheck # TensorIterator
variants: function, method
dispatch:
CPU, CUDA: relu
MPS: relu_mps
MkldnnCPU: mkldnn_relu
QuantizedCPU: relu_quantized_cpu
NestedTensorCPU, NestedTensorCUDA: NestedTensor_relu
由上述 dispatch 可知,在 CPU 和 CUDA 平台调用该算子时,会 dispatch 到名称为 relu 的函数。因此,下一步需要查看 relu 函数的函数原型,该算子的函数实现在 torch/include/ATen/Functions.h 文件中可以找到,可知具体实现在 torch/include/ATen/ops/relu.h 文件中,由此得到了函数原型 Tensor relu(const at::Tensor & self)。
在得到函数原型后,我们就知道了该算子具体调用的函数的运算符模式,在注册自定义算子的时候,需要遵守这种模式。接下来,定义 MTGPU 上的 relu 算子时,可以将函数 Tensor mtgpu_relu(const at::Tensor & self) 作为 MTGPU 上注册的 relu 函数名,函数内部实现由我们自定义,最后使用上文所述的 TORCH_LIBRARY_IMPL 注册 MTGPU 上的 relu 算子,如 m.impl("relu", &mtgpu),至此算子对接流程基本完成。
10.1.3.3. 自定义函数实现
上文较详细地介绍了算子对接的基本流程,但并未具体介绍自定义函数内部的实现细节,本小节将具体介绍函数内部实现的一些注意事项。
由于函数声明中使用的数据类型是 PyTorch 框架 C++ 扩展模块 ATen 中的相关数据类型,而算字库有自己的一套数据类型。因此,自定义函数中首先需要根据框架中的 Tensor 数据信息来创建 muDNN 中的 Tensor 数据类型,这个操作可以借助 CreatMTensor 函数来 完成,该函数主要提取了框架中 Tensor 的 dim、size、type 和 addr 等信息,返回一个 muDNN 的 mTensor 数据类型。mTensor 数据类型都创建完成后,接着需要创建 mHandle 句柄和具体 op 的实例对象,并根据相关 op 的具体情况,设置相应的如 mode、axis、alpha 等超参信息,最后将符合 muDNN 类型要求的参数传入 op 的 Run 方法中运行得到输出结果。
另外,由于 PyTorch 的数据类型的 storage 具有 offset 属性,当多个 tensor 共用一块 storage 时,tensor 的 offset 可能不为零,此时 PyTorch 中打印 tensor 的地址信息可以发现该地址是经过 PyTorch 计算后得到的地址信息,即 storage 首地址 + offset * sizeof(type)。另外,算子库目前也不支持不连续的 tensor,因此在自定义函数中,往往需要进行 tensor 的连续性检测和连续性处理,这个操作可以借助 MusaContiguous 完成。
最后,在编写自定义算子时,有时需要进行一些边界检查,这时我们可以参照 native 目录下的 CPU 代码或者 CUDA 代码进行编写。
注册伪代码如下:
Tensor musa_relu(const Tensor& input) {
// step1: create output tensor first
Tensor result = at::native::empty_mtgpu(input.sizes(), DeviceType::MTGPU, ...);
// step2: create mTensor(in&out), set Attribute
using mTensor = ::musa::dnn::Tensor;
mTensor in_tensor, out_tensor;
in_tensor.SetAddr(input.data_ptr());
out_tensor.SetAddr(result.data_ptr());
in_tensor.SetNdInfo(mTensor::Type::FLOAT, input.size(), ...);
out_tensor.SetNdInfo(mTensor::Type::FLOAT, result.size(), ...);
// step3: create operation descriptor
using mHandle = ::musa::dnn::Handle;
using Unary = ::musa::dnn::Unary;
Unary op;
op.SetMode(Unary::Mode::RELU);
// step4: run op kernel function
op.Run(mHandle, out_tensor, in_tensor);
return result;
}