注:为确保您的 C++ 自定义运算与 TensorFlow 的官方 pip 软件包 ABI 兼容,请遵循自定义运算仓库中的指南。指南包含端到端代码示例以及用于构建和分发自定义运算的 Docker 镜像。
如果您想创建的运算不在现有 TensorFlow 库的涵盖范围内,我们建议您首先尝试以现有 Python 运算或函数组合的形式使用 Python 编写该运算。如果无法做到这一点,则可以创建自定义 C++ 运算。由于以下几点原因,您可能希望创建自定义 C++ 运算:
- 无法轻易或根本无法将您的运算表示为现有运算的组合。
- 将您的运算表示为现有基元的组合并不高效。
- 您想以手动方式融合未来的编译器难以融合的基元组合。
例如,假设您想要实现诸如“中值池化”之类的功能,与“MaxPool”算子类似,但需要计算滑动窗口期间的中值而不是最大值。可以使用运算组合来实现这一目的(例如,使用 ExtractImagePatches 和 TopK),但在性能或内存效率方面可能不如原生运算那样出色,对于原生运算,您可以利用单个融合运算实现更巧妙的过程。和往常一样,通常有必要首先尝试使用算子组合来表示您想要的运算,只有在这被证实难以实现或效率低下时,才选择添加新运算。
要整合自定义运算,您需要执行以下操作:
- 在 C++ 文件中注册新运算。运算注册会定义运算功能的接口(规范),此接口与运算的实现无关。例如,运算注册会定义运算的名称及运算的输入和输出,还会定义用于张量形状推断的形状函数。
- 使用 C++ 实现运算。运算的实现称为内核,它是您在第 1 步中注册的规范的具体实现。可以有多个内核用于不同的输入/输出类型或架构(例如,CPU、GPU)。
- 创建一个 Python 封装容器(可选)。此封装容器是用于以 Python 创建运算的公共 API。默认封装容器是根据运算注册生成的,用户可以直接使用它或向其中添加内容。
- 编写一个函数来计算运算的梯度(可选)。
- 测试运算。为方便起见,我们通常在 Python 中进行测试,但您也可以在 C++ 中测试运算。如果您要定义梯度,可以使用 Python
tf.test.compute_gradient_error验证梯度。要了解如何测试 ReLu 之类的算子及其梯度的前向函数,请参阅relu_op_test.py。
前提条件
- 对 C++ 有一定的了解。
- 必须已安装 TensorFlow 二进制文件,或者必须已下载 TensorFlow 源代码,并且能够构建。
定义运算接口
您可以通过将接口注册到 TensorFlow 系统来定义运算的接口。在注册中,您需要指定运算的名称、输入(类型和名称)和输出(类型和名称),以及文档字符串和该运算可能需要的任意特性。
要了解这一过程的工作原理,假设您想要创建一个接受 int32 张量并输出该张量副本(将第一个元素之外的所有其他元素都设置为零)的运算。为此,请先创建一个名为 zero_out.cc 的文件,然后添加对 REGISTER_OP 宏的调用,该宏可以定义运算的接口:
#include "tensorflow/core/framework/op.h"
#include "tensorflow/core/framework/shape_inference.h"
using namespace tensorflow;
REGISTER_OP("ZeroOut")
.Input("to_zero: int32")
.Output("zeroed: int32")
.SetShapeFn([](::tensorflow::shape_inference::InferenceContext* c) {
c->set_output(0, c->input(0));
return Status::OK();
});
ZeroOut 运算会将一个包含 32 位整数的张量 to_zero 作为输入,并输出一个包含 32 位整数的张量 zeroed。该运算还使用形状函数来确保输出张量与输入张量的形状相同。例如,如果输入是形状为 [10, 20] 的张量,则此形状函数会指定输出形状也是 [10, 20]。
注:运算名称必须采用驼峰命名法,并且对于在二进制文件中注册的所有其他运算,该名称必须唯一。
实现运算的内核
在定义接口后,您需要为运算提供一个或多个实现。要创建其中一个内核,请先创建一个扩展 OpKernel 并重写 OpKernel 方法的类。Compute 方法提供了一个类型为 OpKernelContext* 的 context 参数,您可以从中访问输入张量和输出张量等有用信息。
将内核添加到您在上面创建的文件中。内核可能如下所示:
#include "tensorflow/core/framework/op_kernel.h"
using namespace tensorflow;
class ZeroOutOp : public OpKernel {
public