TensorFlow LSTM 视频目标检测实战指南:research/lstm_object_detection 的架构原理、训练评估与 TFLite 部署
【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models
本文以 TensorFlow 仓库research/lstm_object_detection模块为对象,系统讲解面向移动端视频流的 LSTM 目标检测方案:它对应 CVPR 2018 的"Temporally-Aware Feature Maps"与后续"Memory-Guided(Looking Fast and Slow)"两条技术路线。读完本文,你将掌握该模块的代码组织结构、LSTM-SSD 元架构与关键配置字段,能够基于样例 pipeline 配置完成视频序列数据的训练与逐帧评估,并按官方导出手册将 checkpoint 逐步导出为可在端侧运行的 TFLite FlatBuffer 模型。
LSTM 视频目标检测的核心思路是"让检测网络对时间维度建模":与逐帧独立检测相比,它复用帧间的时序上下文来提升检测质量、同时通过稀疏计算(后续的 Interleaved 方案)显著降低平均每帧开销,因而天然适配移动端视频场景。
一、模块背景:两篇驱动论文与技术路线
本模块在官方 README 中明确了其实现对应的两篇论文,它们构成了整个代码库的设计骨架:
- Mobile Video Object Detection with Temporally-Aware Feature Maps(Liu, Mason 与 Zhu, Menglong,CVPR 2018)—— 即"LSTM-SSD",把 SSD 检测器与卷积 LSTM 状态结合,使特征图携带历史帧信息(temporally-aware);
- Looking Fast and Slow: Memory-Guided Mobile Video Object Detection(Liu, Mason、Zhu, Menglong、White, Marie、Li, Yinxiao 与 Kalenichenko, Dmitry)—— 即"Interleaved LSTM-SSD",通过一个轻量网络快速处理大部分帧、一个稍重网络配合记忆门控处理关键帧,在精度与每帧计算量之间做权衡。
这两条路线的实现均可以在仓库代码中直接验证。元架构文件开头的 docstring 即声明"这是带 LSTM 状态的卷积 Multibox/SSD 检测模型在视频数据上的通用 TensorFlow 实现,同时支持常规 LSTM-SSD 与 interleaved LSTM-SSD 两种框架"(参见 lstm_ssd_meta_arch.py)。
二、代码地图:十分钟看懂模块结构
research/lstm_object_detection/目录将数据、模型、训练评估与端侧导出分成清晰的层次:
| 子目录/文件 | 职责 | 关键内容 |
|---|---|---|
| inputs/ | 视频序列数据输入 | seq_dataset_builder.py构建按video_length展开的 batch;tf_sequence_example_decoder.py解码 TF SequenceExample |
| lstm/ | 循环单元实现 | lstm_cells.py(卷积 LSTM 单元)、rnn_decoder.py(按时间步展开解码) |
| meta_architectures/ | 整体元架构 | lstm_ssd_meta_arch.py中的LSTMSSDMetaArch |
| models/ | 特征提取器 | Mobilenet V1 版与 Interleaved Mobilenet V2 版 LSTM-SSD 特征提取器 |
| configs/ | 完整训练配置 | 两份 ImageNet-Vid 样例 pipeline(V1 与 Interleaved V2) |
| protos/ | 配置扩展协议 | pipeline.proto定义LstmModel扩展消息 |
| metrics/ | 视频评估 | coco_evaluation_all_frames.py对所有帧做 COCO 式评测 |
| train.py / eval.py | 训练/评估入口 | 命令行入口,读取 pipeline 配置 |
| export_tflite_lstd_graph.py / export_tflite_lstd_model.py | 端侧导出 | checkpoint → frozen graph → TFLite FlatBuffer |
| tflite/ | C++ 端侧推理客户端 | mobile_lstd_tflite_client、mobile_ssd_tflite_client等 |
三、核心原理:LSTMSSDMetaArch 如何"用 LSTM 做检测"
LSTMSSDMetaArch直接继承自 Object Detection API 的SSDMetaArch(见 lstm_ssd_meta_arch.py),因此 SSD 的 anchor 生成、box predictor、NMS 后处理、分类/定位损失等组件全部复用,差异集中在两点:
- 引入
unroll_length(时间展开长度):模型每次消费的是一个长度为unroll_length的视频片段而非单帧,该值由配置中的train_unroll_length/eval_unroll_length决定,并被写入元架构(见 model_builder.py 中"若 feature extractor 类型含lstm则从 lstm 配置取 unroll length"的逻辑)。 predict接收并维护 LSTM 状态:在predict()中,特征提取器以states与state_name为输入展开时序特征提取,预测字典里额外携带states_and_outputs(以及非空状态时的step),供循环过程传递状态(见 lstm_ssd_meta_arch.py)。注释同时表明:在导出模型等场景中状态恒为零,因此忽略 step。
时序建模的"状态"具体由 lstm/lstm_cells.py 提供卷积 LSTM 单元、由 lstm/rnn_decoder.py 完成逐时间步的循环解码;它们配有 lstm_cells_test.py 与 rnn_decoder_test.py 可验证行为。
四、两种特征提取器与模型注册
模型工厂通过扩充SSD_FEATURE_EXTRACTOR_CLASS_MAP注册两种新特征提取器类型(见 model_builder.py):
feature_extractor.type字符串 | 对应类 | 特征 |
|---|---|---|
lstm_ssd_mobilenet_v1 | LSTMSSDMobileNetV1FeatureExtractor(源码) | 常规 LSTM-SSD,逐帧共享同一主干 |
lstm_ssd_interleaved_mobilenet_v2 | LSTMSSDInterleavedMobilenetV2FeatureExtractor(源码) | Interleaved 版本,多个主干按策略交错 |
构建时,除标准 SSD 参数(depth_multiplier、min_depth、use_depthwise、卷积超参等)外,还会把 LSTM 特有配置注入特征提取器,包括lstm_state_depth、flatten_state、clip_state、scale_state、is_quantized、low_res;对 interleaved 类型还会设置pre_bottleneck、多档depth_multipliers,并根据is_training分别选择train_interleave_method或eval_interleave_method(见 model_builder.py)。
五、配置体系:理解 lstm_model 扩展字段
LSTM-SSD 的配置在标准TrainEvalPipelineConfig之上通过 proto 扩展追加lstm_model消息(扩展号为 205743444,见 protos/pipeline.proto)。仓库两份样例配置即采用该结构。所有字段的含义与默认值在 proto 中均有注释,汇总如下:
| 字段 | 含义 | 默认值 | 样例配置中的取值 |
|---|---|---|---|
train_unroll_length | 训练时的时间展开长度 | 无 | 4 |
eval_unroll_length | 评估/导出时的时间展开长度 | 无 | 4 |
lstm_state_depth | LSTM 状态特征图深度 | 256 | 320(Interleaved V2 版) |
depth_multipliers | 多个特征提取器的深度倍率(interleaved/ensemble) | 无 | 1.4、0.35 |
train_interleave_method | 训练时模型交错策略,取值RANDOM/RANDOM_SKIP_SMALL | RANDOM | RANDOM_SKIP_SMALL |
eval_interleave_method | 评估时交错策略,取值RANDOM/RANDOM_SKIP/SKIPK | SKIP9 | SKIP3 |
lstm_state_stride | LSTM 状态的步长 | 32 | — |
flatten_state | 是否摊平 LSTM 状态与输出(仅供 tfmini/tflite 导出内部使用,pipeline 中一般不要设置) | false | — |
pre_bottleneck | 是否在进入 LSTM 门控前加瓶颈层,使多个主干可各自拥有瓶颈、不必强制输出同维度 | false | true(Interleaved V2 版) |
scale_state | 是否归一化 LSTM 状态 | false | — |
clip_state | 是否将 LSTM 状态裁剪到 [0, 6] | true | — |
is_quantized | 是否量化训练(由graph_rewriter覆盖,无需手动设置) | false | — |
low_res | interleaved 模型用较小网络时是否对输入降采样 | false | true(Interleaved V2 版) |
pre_bottleneck字段的注释形象解释了 interleaved 设计的动机:模型 1 输出深度为 d1 的特征图、模型 2 输出深度为 d2 的特征图,pre-bottleneck 允许 LSTM 输入表现为conv(concat([f_1, h]))或conv(concat([f_2, h])),避免两路特征被强制对齐到同一维度。
两份官方样例配置对照
- lstm_ssd_mobilenet_v1_imagenet.config:常规 LSTM-SSD + Mobilenet V1,用于 ImageNet-Vid。典型设定包括
num_classes: 30、Faster R-CNN box coder(y_scale/x_scale: 10.0、height_scale/width_scale: 5.0)、5 层 SSD anchor(min_scale 0.2、max_scale 0.95、宽高比 1/2/0.5/3/0.3333)、fixed_shape_resizer缩放到 256×256;训练采用 RMSProp(初始学习率 0.002、200000 步衰减 0.95、momentum 0.9),from_detection_checkpoint: true做检测模型微调;输入为TF_SEQUENCE_EXAMPLE类型的 tf_record 视频,video_length: 4。 - lstm_ssd_interleaved_mobilenet_v2_imagenet.config:在 V1 配置基础上把
feature_extractor.type换为lstm_ssd_interleaved_mobilenet_v2、输入分辨率提到 320×320,并加上第五节表格中的 interleaved 专属参数。
两份配置都要求:label map 路径与输入数据路径需替换为真实路径;评估部分使用metrics_set: "coco_evaluation_all_frames",即对所有帧做统一评测(而非只评测稀疏帧),对应实现见 metrics/coco_evaluation_all_frames.py。
六、训练与评估:命令行实操
模块同时支持"单一 pipeline 文件"与"模型/训练/输入三个配置分开"两种方式。以下命令均在research/目录下执行,并需保证lstm_object_detection与object_detection两个 Python 包可被导入(代码基于tensorflow.compat.v1编写,请使用匹配的 TensorFlow 1.x 兼容环境)。
6.1 训练
# 方式一:单一 pipeline 配置(推荐,train.py 会忽略其他 config 参数) python lstm_object_detection/train.py \ --logtostderr \ --train_dir=path/to/train_dir \ --pipeline_config_path=lstm_object_detection/configs/lstm_ssd_mobilenet_v1_imagenet.config # 方式二:分别给出 model / train / input 三份配置 python lstm_object_detection/train.py \ --logtostderr \ --train_dir=path/to/train_dir \ --model_config_path=model.config \ --train_config_path=train.config \ --input_config_path=train_input.config从 train.py 源码可以看到完整流程:解析配置得到model、lstm_model、train_config、train_input_config四组对象;用model_builder.build(model_config, lstm_config, is_training=True)构造检测模型;把配置里的数据增强选项逐一交给preprocessor_builder构建;调用seq_dataset_builder.build(...)并以lstm_config.train_unroll_length作为展开长度构造输入;最后交给trainer.train(...)。注意 train.py 中对分布式训练的限制:若存在多个 worker 而没有 ps 任务会直接报错("At least 1 ps task is needed for distributed training."),说明该训练器面向单机或 master-worker+ps 的传统分布式模式。
6.2 评估
python lstm_object_detection/eval.py \ --logtostderr \ --checkpoint_dir=path/to/checkpoint_dir \ --eval_dir=path/to/eval_dir \ --pipeline_config_path=lstm_object_detection/configs/lstm_ssd_mobilenet_v1_imagenet.configeval.py 的关键行为:校验checkpoint_dir与eval_dir必填;把解析出的完整 pipeline 写入eval_dir/pipeline.config留档;评估用模型以is_training=False构建、按eval_unroll_length展开;从 label map 加载类别信息;若设置--run_once则把max_evals强制置为 1(单轮评估即退出),否则按配置持续监听新 checkpoint 评估。另有--eval_training_data开关:置真时用训练输入替换评估输入,并把评估展开长度对齐训练展开长度(见 eval.py),便于在训练集上自查拟合程度。
七、TFLite 导出:checkpoint 到端侧模型的两步走
本模块最实用的部分是模型部署流水线,官方导出步骤记录在 g3doc/exporting_models.md(README 中唯一展开的实操链接,也是模块"开箱可用"的入口),核心是先导 frozen graph、再转 FlatBuffer两步。
第一步:从 checkpoint 导出 TFLite frozen graph
# 在 research/ 目录下执行 PIPELINE_CONFIG_PATH={pipeline 配置路径} TRAINED_CKPT_PREFIX=/{path/to/model.ckpt} EXPORT_DIR={导出目录} python lstm_object_detection/export_tflite_lstd_graph.py \ --pipeline_config_path ${PIPELINE_CONFIG_PATH} \ --trained_checkpoint_prefix ${TRAINED_CKPT_PREFIX} \ --output_directory ${EXPORT_DIR} \ --add_preprocessing_op导出成功后${EXPORT_DIR}内将出现两个文件:tflite_graph.pb(二进制冻结图)与tflite_graph.pbtxt(文本图)。
第二步:从 frozen graph 转 TFLite FlatBuffer
FROZEN_GRAPH_PATH={上一步导出的 tflite_graph.pb} EXPORT_PATH={输出 FlatBuffer 文件名} PIPELINE_CONFIG_PATH={pipeline 配置路径} python lstm_object_detection/export_tflite_lstd_model.py \ --export_path ${EXPORT_PATH} \ --frozen_graph_path ${FROZEN_GRAPH_PATH} \ --pipeline_config_path ${PIPELINE_CONFIG_PATH}输出${EXPORT_PATH}即为可供应用加载的 TFLite FlatBuffer 模型。
导出背后的实现约束(源码级解读)
结合 export_tflite_lstd_graph_lib.py 可以把两个命令的真实语义讲透,这些限制直接决定你的模型能否成功导出:
- 只支持 SSD 模型:非 SSD 配置会抛出
ValueError(见 export_tflite_lstd_graph_lib.py)。 - 只支持
fixed_shape_resizer:输入占位符input_video_tensor的形状取[eval_unroll_length, height, width, channels],其中高宽来自fixed_shape_resizer,通道数在convert_to_grayscale时为 1、否则为 3;其他 resizer 一律报错(见 export_tflite_lstd_graph_lib.py)。这就是为什么评估/导出时的输入是一段视频切片而非单帧。 - 后处理在端侧完成:导出器把
raw_outputs/box_encodings、raw_outputs/class_predictions(score conversion 之后)和常量张量anchors冻结为输出,然后追加名为TFLite_Detection_PostProcess的 custom op,并写入 NMS 阈值、每类/总数最大框数、box coder 的 y/x/h/w scale 等属性,最后用strip_unused_nodes剪掉无关节点(见 export_tflite_lstd_graph_lib.py)。这意味着解码框与 NMS 都由 TFLite runtime 的自定义算子执行。 - 移动平均参数的处理:若
eval_config.use_moving_averages为真,会先把 checkpoint 中的变量替换成滑动平均值再冻结,这也是两个官方配置在 eval 段开启use_moving_averages的原因(见 export_tflite_lstd_graph_lib.py)。 - 量化支持:导出阶段会检查 pipeline 中是否存在
graph_rewriter,存在则按量化图重写器执行(is_quantized);从结构可以推断,配合 protos/quant_overrides.proto 可实现量化感知训练的导出链路。
而 export_tflite_lstd_model.py 中,转换器指定输入数组input_video_tensor、输出为TFLite_Detection_PostProcess的 4 个输出张量(框/类别/分数/框数),并显式设置converter.allow_custom_ops = True以放行上面的端侧后处理算子;注意该脚本用配置里的eval_unroll_length拼输入形状,其input_shapes字典中空间尺寸固定写为 320×320,若你的配置使用了其他分辨率,需要留意此处的一致性。
八、端侧运行:C++ 推理客户端
导出并不是终点。tflite/目录提供了在移动端加载该模型进行推理的参考实现:核心是 mobile_lstd_tflite_client.h 与 mobile_ssd_tflite_client.h(LSTD 专用版与通用移动 SSD 版),配套 tflite/utils/ 下的转换与 SSD 解码工具(conversion_utils.cc、ssd_utils.cc),以及一组描述 anchors、box encodings、检测结果与 label map 的 protobuf(tflite/protos/)。如需在仓库内快速验证导出的 TFLite 模型可用性,可以参考 test_tflite_model.py 的调用方式。
九、实战建议与延伸阅读
- 从样例配置起步:两份官方配置注释明确标注"For training on Imagenet Video",
num_classes: 30对应 ImageNet-Vid;迁移到自有视频数据集时,需同步替换 label map、序列化 TFRecord 数据与类别数,并注意tf_record_video_input_reader中video_length(数据侧)与unroll_length(模型侧)的匹配。 - 数据格式是前提:输入必须是带时间维度的
TF_SEQUENCE_EXAMPLE序列,由 inputs/tf_sequence_example_decoder.py 解码、inputs/seq_dataset_builder.py 组 batch;普通单帧检测数据集需要先完成序列化改造。 - 导出前核对三项配置:模型必须是 SSD 系、resizer 必须是
fixed_shape_resizer、eval_unroll_length需与预期视频切片长度一致,否则导出脚本会直接报错或产出形状不符的模型。 - 进一步阅读:模块 README、模型导出手册、LSTM-SSD 元架构源码、LstmModel 配置协议、元架构测试(可当作架构行为的可执行说明)。仓库内 legacy/detection/ 等目录属于其他检测实现,与本文主题无关,阅读时注意区分。
需要提醒的是,模块代码整体基于 TensorFlow 1.x 语义(大量tensorflow.compat.v1与tf.gfile调用),在现代 TF2 环境中运行需以 TF1 兼容模式为前提;本文所述命令与配置均以当前仓库research/lstm_object_detection内的实际文件为准。
【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考