ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

TFLM_day3

2026/9/29 7:31:28 拓冰建站 浏览量
TFLM_day3 轻松学习 TFLM Day 3:以fully_connected为例写自定义优化 kernel摘要:本文以FULLY_CONNECTED为例,讲解如何在 TFLM(TensorFlow Lite Micro)中编写自定义优化 kernel,把普通算子替换为你自己芯片的加速实现。文章从 reference 版本的Init/Prepare/Eval生命周期讲起,逐步拆解自定义 kernel 的骨架、OpData设计、Prepare阶段(类型检查、Shape 计算、量化参数、scratch buffer)和Eval阶段(含 fallback 写法),并给出完整的your_nnlib_fully_connected_s8参考实现、main.cc注册示例、Makefile 构建接入示例,以及三层测试方案(Kernel 单测、Example 测试、benchmark)。最后总结常见坑和推荐实现顺序,帮助你从 reference 平滑迁移到your_chip后端。Day 2 我们梳理了 TFLM 推理调用链。今天进入硬件厂商最关心的地方:怎么把一个普通 TFLM 算子换成你自己芯片的优化实现。本文以FULLY_CONNECTED为例。这个算子很适合学习,因为全连接本质上是矩阵乘加:output = activation(input * weights + bias)如果你是 AI 芯片设计商,fully_connected通常可以映射到:int8 矩阵乘单元。DSP MAC 阵列。NPU dense layer。你们自己的 NN library,比如your_nnlib_fully_connected_s8()。建议配合这些源码阅读:tensorflow/lite/micro/kernels/fully_connected.htensorflow/lite/micro/kernels/fully_connected.cctensorflow/lite/micro/kernels/fully_connected_common.cctensorflow/lite/micro/kernels/cmsis_nn/fully_connected.cctensorflow/lite/micro/kernels/ceva/fully_connected.cctensorflow/lite/micro/kernels/fully_connected_test.cc1. 先看 reference 版本长什么样TFLM 的 reference fully connected 在:tensorflow/lite/micro/kernels/fully_connected.cc它的结构非常典型:Register_FULLY_CONNECTED() | v 返回 TFLMRegistration | ├── init - FullyConnectedInit ├── prepare - FullyConnectedPrepare └── invoke - FullyConnectedEval也就是说,一个 TFLM kernel 不是一个单独函数,而是一组生命周期函数。1.1 Init 做什么reference 版本的FullyConnectedInit很短:void*FullyConnectedInit(TfLiteContext*context,constchar*buffer,size_t length){TFLITE_DCHECK(context-AllocatePersistentBuffer!=nullptr);returncontext-AllocatePersistentBuffer(context,sizeof(OpDataFullyConnected));}它只做一件事:给OpDataFullyConnected分配持久内存。OpDataFullyConnected定义在:tensorflow/lite/micro/kernels/fully_connected.h里面保存:字段作用output_multiplier量化输出乘子。output_shift量化输出位移。output_activation_min/maxfused activation 的量化裁剪范围。input_zero_point输入 zero point。filter_zero_point权重 zero point。output_zero_point输出 zero point。filter_buffer_indexint4 权重展开时使用的 scratch index。per_channel_output_multiplierper-channel 量化乘子数组。per_channel_output_shiftper-channel 量化 shift 数组。is_per_channel是否使用 per-channel 量化。新手要记住:OpData是 Prepare 阶段算好、Invoke 阶段反复用的缓存。1.2 Prepare 做什么FullyConnectedPrepare做这些事:FullyConnectedPrepare | ├── 取 input/filter/bias/output 临时 TfLiteTensor ├── 检查 input 和 output 类型一致 ├── 检查 filter 类型是否支持 ├── 如果 filter 是 int4,申请 scratch buffer 用于展开 ├── CalculateOpDataFullyConnected(...) │ | │ ├── 计算量化 multiplier/shift │ ├── 判断 per-tensor 或 per-channel quantization │ ├── 计算 activation min/max │ └── 保存 zero point │ └── 释放临时 TfLiteTensor重点是这句:CalculateOpDataFullyConnected(context,params-activation,input-type,input,filter,bias,output,data);它定义在:tensorflow/lite/micro/kernels/fully_connected_common.cc自定义优化 kernel 应该尽量复用这个 common 函数,因为量化参数算错,性能再快也没用。1.3 Eval 做什么FullyConnectedEval是真正计算的地方:FullyConnectedEval | ├── GetEvalInput(context, node, input_index) ├── GetEvalInput(context, node, filter_index) ├── GetEvalInput(context, node, bias_index) ├── GetEvalOutput(context, node, output_index) ├── 读取 OpDataFullyConnected ├── 根据 input-type 分支 │ | │ ├── float32 - reference_ops::FullyConnected │ ├── int8 - reference_integer_ops::FullyConnected │ └── int16 - reference_integer_ops::FullyConnected │ └── 写 output tensor你的优化 kernel 通常就是替换 Eval 里的计算部分。2. 自定义优化 kernel 的目标假设我们有一个 AI 芯片,叫your_chip,它提供了一个 NN library:intyour_nnlib_fully_connected_s8(constint8_t*input,constint8_t*weights,constint32_t*bias,int8_t*output,intbatches,intinput_depth,intoutput_depth,int32_tinput_offset,int32_toutput_offset,int32_toutput_multiplier,intoutput_shift,int32_tactivation_min,int32_tactivation_max,void*scratch,intscratch_size);那么我们的 TFLM 优化 kernel 要做的事就是:TFLM tensor / OpData | | 参数转换 v your_nnlib_fully_connected_s8(...) | v AI 芯片执行 | v 写回 TFLM output tensor注意:TFLM 不应该知道你的硬件寄存器、DMA、command queue 细节。那些应封装在your_nnlib或 driver 里。3. 建议目录结构可以新增一个目录:tensorflow/lite/micro/kernels/your_chip/ ├── fully_connected.cc ├── your_chip_common.h ├── your_chip_common.cc └── README.md如果以后继续优化其他算子,可以扩展成:tensorflow/lite/micro/kernels/your_chip/ ├── add.cc ├── conv.cc ├── depthwise_conv.cc ├── fully_connected.cc ├── pooling.cc ├── softmax.cc ├── your_chip_common.cc └── your_chip_common.h参考现有目录:tensorflow/lite/micro/kernels/cmsis_nn/tensorflow/lite/micro/kernels/ceva/tensorflow/lite/micro/kernels/xtensa/tensorflow/lite/micro/kernels/arc_mli/4. 自定义 fully_connected 的整体骨架下面是一个适合学习的骨架。它展示 TFLM kernel 应该怎样写,但your_nnlib_*是你们芯片 SDK 里的函数,需要你们自己实现。#include"tensorflow/lite/micro/kernels/fully_connected.h"#include"tensorflow/lite/c/builtin_op_data.h"#include"tensorflow/lite/c/common.h"#include"tensorflow/lite/kernels/internal/reference/integer_ops/fully_connected.h"#include"tensorflow/lite/micro/kernels/kernel_util.h"#include"tensorflow/lite/micro/micro_log.h"#include"your_chip/your_nnlib.h"namespacetflite{namespace{structOpData{OpDataFullyConnected reference_op_data;intscratch_index;intscratch_size;intbatches;intinput_depth;intoutput_depth;};void*Init(TfLiteContext*context,constchar*buffer,size_t length){TFLITE_DCHECK(context-AllocatePersistentBuffer!=nullptr);returncontext-AllocatePersistentBuffer(context,sizeof(OpData));}TfLiteStatusPrepare(TfLiteContext*context,TfLiteNode*node){TFLITE_DCHECK(node-user_data!=nullptr);TFLITE_DCHECK(node-builtin_data!=nullptr);auto*data=static_castOpData*(node-user_data);data-scratch_index=-1;data-scratch_size=0;constauto*params=static_castconstTfLiteFullyConnectedParams*(node-builtin_data);MicroContext*micro_context=GetMicroContext(context);TfLiteTensor*input=micro_context-AllocateTempInputTensor(node,kFullyConnectedInputTensor);TfLiteTensor*filter=micro_context-AllocateTempInputTensor(node,kFullyConnectedWeightsTensor);TfLiteTensor*bias=micro_context-AllocateTempInputTensor(node,kFullyConnectedBiasTensor);TfLiteTensor*output=micro_context-AllocateTempOutputTensor(node,kFullyConnectedOutputTensor);TF_LITE_ENSURE(context,input!=nullptr);TF_LITE_ENSURE(context,filter!=nullptr);TF_LITE_ENSURE(context,output!=nullptr);TF_LITE_ENSURE_EQ(context,input-type,output-type);constboolsupported=input-type==kTfLiteInt8filter-type==kTfLiteInt8;if(!supported){MicroPrintf("your_chip fully_connected only supports int8 input/filter.");returnkTfLiteError;}TF_LITE_ENSURE_STATUS(CalculateOpDataFullyConnected(context,params-activation,input-type,input,filter,bias,output,data-reference_op_data));constRuntimeShape input_shape=tflite::micro::GetTensorShape(input);constRuntimeShape filter_shape=tflite::micro::GetTensorShape(filter);constRuntimeShape output_shape=tflite::micro::GetTensorShape(output);data-input_depth=filter_shape.Dims(filter_shape.DimensionsCount()-1);data-output_depth=output_shape