news 2026/9/7 6:21:13

Attention OCR 深度解析:TensorFlow 街景文字识别模型从 FSNS 训练到 SavedModel 导出的完整实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Attention OCR 深度解析:TensorFlow 街景文字识别模型从 FSNS 训练到 SavedModel 导出的完整实战

Attention OCR 深度解析:TensorFlow 街景文字识别模型从 FSNS 训练到 SavedModel 导出的完整实战

【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models

Attention OCR 是 TensorFlow Models 仓库中一个面向真实世界街景图像文字提取的经典研究项目,其论文提出的 "ConvNets + RNN + 注意力机制" 架构在 FSNS(French Street Name Signs)数据集上取得了 84.2% 的完整序列准确率,超越了此前 72.46% 的基准。本文以 research/attention_ocr/README.md 为核心,完整覆盖环境准备、FSNS 数据下载、训练/微调命令、自有数据接入以及预训练模型推理与导出的全部实操流程,并结合 model.py、sequence_layers.py、data_provider.py 等源码深入解析多视角卷积塔、注意力解码器与字符序列损失的具体实现。读完后你可以独立复现该模型的训练与评估,并掌握将其接入自有街景/文字图像数据的方法。

一、项目概览与模型架构

README 将该项目定位为 "A TensorFlow model for real-world image text extraction problems",即面向真实图像(而非印刷体、干净文本)的文字提取。论文的核心贡献有三点:

  1. 基于 ConvNets、RNN 与一种新型注意力机制的模型,在 FSNS 上达到 84.2% 完整序列准确率(此前基准 72.46%);
  2. 研究了使用不同深度 CNN 特征提取器时的速度/精度权衡
  3. 提供的开源版本与论文版本的主要差异(见 README 的 Disclaimer):论文使用 50 个 K80 GPU 分布式异步训练,而随代码发布的 checkpoint 是单 GPU(Titan X)训练约 6 天(约 400k 步)所得,400k 步时达到 83.79% 完整序列准确率(24 小时约 81%),并且坐标编码(coordinate encoding)默认关闭

从源码结构看,模型的计算图由 model.py 中的Model.create_base方法串联而成(model.py第 480~554 行),整体数据流为:

输入图像 [batch, 600, 150, 3] → 像素归一化 (x - 0.5) * 2.5 → tf.split 沿宽度切成 4 个视角 (views) → 4 个共享权重的 InceptionV3 卷积塔 (final_endpoint=Mixed_5d) → (可选) 坐标 one-hot 编码拼接 → 拼接展平为 [batch, seq_length, features] → 注意力 + 自回归 LSTM 解码器 (256 单元) → 每步线性层输出字符 logits → predicted_chars / predicted_scores / predicted_text / predicted_length / normalized_seq_conf 等输出端点

其中几个关键设计点在源码中清晰可见:

  • 像素归一化create_base开头执行images = tf.subtract(images, 0.5); images = tf.multiply(images, 2.5),把 [0,1] 的像素映射到约 [-0.5, 1.5] 的对称区间,与 Inception 预训练分布一致(见 model.py)。
  • 四视角机制:FSNS 图像宽 600、高 150,实际上是把同一街牌在 4 个不同视角(拍摄时刻/位置)下的画面横向拼接成一张宽图。num_of_views=4由数据集配置决定(datasets/fsns.py 的DEFAULT_CONFIG),4 个视角共享同一套 Inception 权重(reuse=(i != 0)),随后pool_views_fn把 4 个特征图在高度维度堆叠拼接(model.py第 405~423 行)。
  • 特征数量与序列长度的匹配_create_lstm_inputsmodel.py第 348~371 行)会断言卷积塔输出的特征维数不少于seq_length(FSNS 为 37),若超过则截取前seq_length个,即"每个时间步对应图像上的一个水平位置段",这是注意力机制得以逐字符对齐空间位置的基础。
  • 输出端点OutputEndpoints命名元组包含chars_logitpredicted_charspredicted_scorespredicted_textpredicted_lengthpredicted_confnormalized_seq_conf等(model.py第 37~41 行),字符 ID 到 UTF-8 文本的转换由CharsetMapper基于tf.contrib.lookup.index_to_string_table_from_tensor完成(model.py第 70~96 行)。

二、环境与数据集准备(继承 README 完整步骤)

2.1 安装要求

  1. TensorFlow 1.15(README 明确要求,代码基于 TF 1.x 的contribslim组件):
python3 -m venv ~/.tensorflow source ~/.tensorflow/bin/activate pip install --upgrade pip pip install --upgrade tensorflow-gpu=1.15
  1. 磁盘空间:下载 FSNS 数据集至少需要 158GB 空闲空间:
cd research/attention_ocr/python/datasets aria2c -c -j 20 -i ../../../street/python/fsns_urls.txt cd ..
  1. 内存:16GB 以上,推荐 32GB。
  2. 计算设备train.py同时支持 CPU 和 GPU,但推荐使用 GPU;作者已在 Titan X 与 GTX980 上测试。

2.2 FSNS 数据集构成

FSNS 数据集按子集切分为多个 TFRecord 文件,各子集规模(摘自 README):

子集文件数单文件约样本数(源码配置)
Train512 个300MB1,044,868
Validation64 个40MB16,150
Test64 个50MB20,040
testdata若干小数据集较小用于验证模型能"学到东西"
合计约 158GB

样本数依据 datasets/fsns.py 中DEFAULT_CONFIGsplits配置:train/test/validation 的size分别为 1044868 / 20404 / 16150,文件模式分别为train/train*test/test*validation/validation*。下载链接列表(download.tensorflow.org/data/fsns-20160927/...下的 charset、train、validation、test、testdata 各文件)在 README 中完整列出,集中存储在仓库research/street目录的python/fsns_urls.txt中。

数据集每个样本的关键特征字段(由get_split中的keys_to_features定义,fsns.py):

  • image/encoded:PNG 编码图像,image/format默认 png;
  • image/widthimage/orig_width:用于反推图像内含几个视角(_NumOfViewsHandlernum_of_views * orig_width / width);
  • image/class:定长[max_sequence_length](37)的字符编码序列;
  • image/unpadded_class:变长真实字符编码;
  • image/text:Unicode 文本。

字符集由charset_size=134.txt定义(134 个字符类,null_code=133),read_charset代码\t字符的制表符格式解析,其中<nul>会被替换为浅灰色块字符(fsns.py)。

三、训练、测试与微调(README 命令 + 源码级参数)

3.1 运行单元测试

cd research/attention_ocr/python find . -name "*_test.py" -printf '%P\n' | xargs python3 -m unittest

仓库中对应的测试文件包括 model_test.py、sequence_layers_test.py、data_provider_test.py、datasets/fsns_test.py、model_export_test.py 等。

3.2 从零训练

python train.py

train.py的主流程(train.py):创建数据集与模型 → 在replica_device_setter设备域内构建数据管道 →model.create_base建图 →model.create_loss得到总损失 →slim.learning.train进入训练循环(支持 checkpoint 恢复、定时保存)。

3.3 使用 Inception 预训练权重初始化

wget http://download.tensorflow.org/models/inception_v3_2016_08_28.tar.gz tar xf inception_v3_2016_08_28.tar.gz python train.py --checkpoint_inception=./inception_v3.ckpt

从源码看,--checkpoint_inception会在create_init_fn_to_restore中只恢复AttentionOcr_v1/conv_tower_fn/INCE作用域下的变量(model.py),即卷积塔单独从 ImageNet 预训练权重初始化,RNN/注意力部分随机初始化。

3.4 用官方 checkpoint 微调

wget http://download.tensorflow.org/models/attention_ocr_2017_08_09.tar.gz tar xf attention_ocr_2017_08_09.tar.gz python train.py --checkpoint=model.ckpt-399731

3.5 关键命令行参数(源码实测默认值)

train.pyeval.py共享 common_flags.py 中定义的参数,结合 train.py 补充如下:

参数默认值说明
batch_size32批大小
crop_width/crop_heightNone中心裁剪尺寸,用于缩减每个视角的宽度
train_log_dir/tmp/attention_ocr/train事件日志与 checkpoint 目录
dataset_namefsns数据集模块名,须为datasets包内可导入模块
split_nametrain数据集切分:train / test / validation
dataset_dirNone数据集根目录,默认用fsns.py中的DEFAULT_DATASET_DIR
checkpoint''恢复完整模型权重的 checkpoint 路径
checkpoint_inception''仅恢复 Inception 卷积塔权重的 checkpoint 路径
learning_rate0.004学习率
optimizermomentum可选 momentum / adam / adadelta / adagrad / rmsprop(见create_optimizer,train.py)
momentum0.9momentum 与 rmsprop 优化器的动量
use_augment_inputTrue是否启用图像增强
final_endpointMixed_5dInceptionV3 截断端点,论文借此研究 CNN 深度与速度/精度的权衡
use_attentionTrue是否使用注意力机制
use_autoregressionTrue是否使用自回归(反馈上一字符)
num_lstm_units256序列 LSTM 单元数
weight_decay0.00004字符预测全连接层的权重衰减
lstm_state_clip_value10.0LSTM cell state 裁剪值
label_smoothing0.1label smoothing 权重
ignore_nullsTrue计算损失时忽略 null 字符
average_across_timestepsFalse是否按 label 权重总和平均损失
clip_gradient_norm2.0梯度裁剪范数(train.py 独有)
save_summaries_secs60写 summary 的频率(秒)
save_interval_secs600保存 checkpoint 的频率(秒)
max_number_of_steps1e10最大梯度步数
task/ps_tasks/total_num_replicas0 / 0 / 1多 worker 与参数服务器配置;sync_replicas=True时启用同步复制
reset_train_dirFalseTrue 时清空train_log_dir重新开始
show_graph_statsFalse输出模型参数量统计

eval.py另有:num_batches(100,评估批次数)、eval_log_dir/tmp/attention_ocr/eval)、eval_interval_secs(60)、number_of_steps(评估次数上限),评估循环通过slim.evaluation.evaluation_loop运行,且强制device_count={"GPU": 0}在 CPU 上评估(eval.py)。

3.6 训练循环与损失函数的实现要点

  • 优化器与梯度slim.learning.create_train_op负责梯度计算,summarize_gradients=True并默认按clip_gradient_norm=2.0裁剪(train.py)。
  • 序列损失sequence_loss_fn将 logits/labels 沿时间维 unstack 后调用tf.contrib.legacy_seq2seq.sequence_loss;当label_smoothing>0时用稠密 softmax 交叉熵(对 one-hot 标签做1-weightweight/num_classes平滑),否则用稀疏版本;ignore_nulls=True时所有位置权重为 1,否则非 null 字符位置权重为 1(model.py)。损失通过slim.losses集合汇总为总损失(含正则项)。
  • 字符集与置信度null_based_length_prediction依据预测序列中 null 字符的计数推断文本长度,predicted_conf为各字符最大对数概率的累加,normalized_seq_conf将其换算为几何平均置信度(model.py)。
  • 评估指标create_summaries在 eval 模式下注册CharacterAccuracySequenceAccuracy(流式指标,序列准确率在第一个 null 处截断序列后比较),实现见 metrics.py。

四、注意力与自回归解码器:核心创新在源码中的位置

序列解码层由 sequence_layers.py 实现,提供四种可切换的组合(get_layer_class,第 400~422 行):

use_attentionuse_autoregression说明
AttentionWithAutoregression默认配置,论文主模型
Attention纯注意力解码
NetSliceWithAutoregression按固定特征切片 + 自回归
NetSlice固定特征切片基线

实现细节值得注意的几处:

  1. 注意力解码Attention.unroll_cell直接调用tf.contrib.legacy_seq2seq.attention_decoder,把卷积塔特征self._net[batch, num_features, feature_size])作为attention_states,即注意力在空间特征序列上"池化"出每个字符对应的视觉证据(sequence_layers.py)。
  2. 自回归输入:训练时用 ground-truth one-hot 标签(teacher forcing,get_train_input返回self._labels_one_hot[:, i-1, :]),推理时用上一时刻 argmax 字符的 one-hot(get_eval_input),首步输入为零向量。
  3. 正交初始化orthogonal_initializer用 SVD 生成正交矩阵初始化 LSTM 权重与 softmax 权重,作者引用了关于 RNN 正交初始化的文献注释(sequence_layers.py)。
  4. LSTM 配置LSTMCell(256, use_peepholes=False, cell_clip=10.0, state_is_tuple=True, initializer=orthogonal_initializer)(第 250~255 行),字符 logits 由共享的softmax_w [256, num_char_classes]/softmax_b线性层按时间步复用产生,并施加0.5 * weight_decay的 L2 正则。

五、数据管道:增强、裁剪与批处理

data_provider.py 定义了模型训练所需的输入端点InputEndpoints(images、images_orig、labels、labels_one_hot,README "Using your own image data" 一节引用的 data_provider.py 第 33 行 即此定义),其形状约定为:

  • images:[batch_size, H, W, 3]
  • labels:[batch_size, seq_length]
  • labels_one_hot:[batch_size, seq_length, num_char_classes]

管道各阶段的实现:

  • 预处理preprocess_imageconvert_image_dtype转 float32;若开启增强/裁剪,先把图像按num_towers=4沿宽度拆成 4 个视角,分别做中心裁剪(central_crop,会断言原图不小于目标尺寸)与增强后拼接回去(data_provider.py)。
  • 数据增强augment_image:随机裁剪(min_object_covered=0.8、宽高比 0.8~1.2、面积 0.8~1.0)→ 随机选择 4 种插值方法之一 resize 回原尺寸 → 随机 4 选 1 的颜色扰动(distort_color)→ 裁剪到 [-1.5, 1.5](第 49~89 行)。
  • 批量与洗牌shuffle_batch默认配置num_batching_threads=8queue_capacity=3000min_after_dequeue=1000DEFAULT_SHUFFLE_CONFIG,第 45~46 行),配合slim.dataset_data_provider的队列读取 TFRecord。
  • one-hot 编码labels_one_hot = slim.one_hot_encoding(label, dataset.num_char_classes),134 个字符类。

六、接入你自己的图像数据(README 两条路线 + 源码佐证)

README 给出两种方案:

路线 1:数据存成 FSNS 同款格式,复用datasets/fsns.py

新建datasets/newtextdataset.py

import fsns DEFAULT_DATASET_DIR = 'path/to/the/dataset' DEFAULT_CONFIG = { 'name': 'MYDATASET', 'splits': { 'train': { 'size': 123, 'pattern': 'tfexample_train*' }, 'test': { 'size': 123, 'pattern': 'tfexample_test*' } }, 'charset_filename': 'charset_size.txt', 'image_shape': (150, 600, 3), 'num_of_views': 4, 'max_sequence_length': 37, 'null_code': 42, 'items_to_descriptions': { 'image': 'A [150 x 600 x 3] color image.', 'label': 'Characters codes.', 'text': 'A unicode string.', 'length': 'A length of the encoded text.', 'num_of_views': 'A number of different views stored within the image.' } } def get_split(split_name, dataset_dir=None, config=None): if not dataset_dir: dataset_dir = DEFAULT_DATASET_DIR if not config: config = DEFAULT_CONFIG return fsns.get_split(split_name, dataset_dir, config)

然后做两件事:

  1. 把新模块加入 datasets/init.py(当前内容即from datasets import fsns__all__注册表);
  2. 命令行指定数据集名:
python train.py --dataset_name=newtextdataset

注意eval.py也需要相同的--dataset_name标志。配置字典中各字段的含义可从 fsns.py 的DEFAULT_CONFIG对照理解:splits[*].size/pattern用于TFRecordReader的文件 glob 与样本数,charset_filename定位字符集文件,null_code必须等于字符集中<nul>的编码(FSNS 中为 133,示例中的 42 是自定义值)。数据以 FSNS 格式存储的方法,README 指向了作者维护的 Stack Overflow 说明(此处不再外部跳转,字段格式直接参考get_splitkeys_to_features的五个字段定义即可)。

路线 2:定义全新的数据集格式

模型训练只需要三路输入(images、labels、labels_one_hot,形状见上文第五节),README 指引以data_provider.py第 33 行的InputEndpoints为参照、以 datasets/fsns.py 为范例自行实现get_split。这条路线的落点是让新数据集对象同样携带num_char_classesmax_sequence_lengthnum_of_viewsnull_codecharsetimage_shape等属性——因为train.pycommon_flags.create_model(dataset.num_char_classes, dataset.max_sequence_length, dataset.num_of_views, dataset.null_code)data_provider.get_data(dataset, ...)都依赖这些属性(train.py)。

七、使用预训练模型:推理与 SavedModel 导出

README 说明官方未单独发布推理脚本,但用 Serving 基础设施导出 SavedModel 是推荐做法,也提供了手工建图的 5 步替代方案。

7.1 推荐方式:导出 SavedModel

python model_export.py \ --checkpoint=model.ckpt-399731 \ --export_dir=/tmp/attention_ocr_export

model_export.py 的实际行为:

  • dataset_nameexport_dir为必填参数;checkpoint 也可直接用--train_log_dir指向训练目录取最新 checkpoint;
  • --export_for_serving(默认 True)时,输入为序列化的 tf.Example proto(placeholder 名tf_example),并在图内附加图像解码归一化;设为 False 时输入为uint8图像张量(名images),且必须指定--batch_size
  • 输出张量非常丰富:predictions(字符 ID)、scoreschars_logitpredicted_lengthpredicted_textpredicted_confnormalized_seq_conf,以及attention_mask_0..36(37 个时间步的注意力掩码,可用于可视化模型"看哪里",见 model_export_lib.py);
  • 导出前会检查字符集文件存在,否则直接报错 "export will fail"。

7.2 手工建图推理(5 步法)

  1. 为图像定义 placeholder(或直接使用 numpy 数组);
  2. 建图(参照 eval.py 第 60 行):
endpoints = model.create_base(images_placeholder, labels_one_hot=None)

注意labels_one_hot=None即切换到推理分支(自回归输入取自上一时刻预测而非真值标签); 3. 加载预训练权重(参照 model.py 的create_init_fn_to_restore,或如model_export.py那样用tf.train.Saver+saver.restore); 4. 运行图计算:

predictions = sess.run(endpoints.predicted_chars, feed_dict={images_placeholder: images_actual_data})
  1. 用字符集文件把字符 ID 转成 UTF-8(离线脚本可用read_charset,图内转换则对应导出的predicted_text输出)。

README 特别提醒:张量名可能随代码演进而变化,旧 checkpoint 可能因此无法加载;一次性实验可用 checkpoint 变量名替换或assign_from_checkpoint_fn配合自定义 var_list 修复,任何长期服务都建议走 TensorFlow Serving 的 SavedModel 方案

八、版本差异与注意事项(Disclaimer 全解)

  • 当前开源版本达到83.79% 完整序列准确率(400k 步训练后),与论文 84.2% 的差距主要来自训练规模:论文用 50 个 K80 GPU 异步分布式训练,开源 checkpoint 是单张 Titan X 训练约 6 天所得(24 小时约 81%);
  • 坐标编码默认关闭encode_coordinates_fnenabled=False,见 model.py 的default_mparams):开启后会给特征图每个位置拼接 x/y 坐标的 one-hot 编码,属于可选消融项;
  • 代码基于 TensorFlow 1.15(tf.contribslimtf.compat.v1混用),迁移到 TF2 需替换slim.datasetlegacy_seq2seqcontrib.legacy_seq2seq.attention_decoder等组件,仓库内未提供迁移实现;
  • 训练与评估分离:eval.py强制 CPU 运行评估(device_count={"GPU": 0}),train.py默认开启数据增强而eval.py固定augment=False,两者共用common_flags保证模型结构参数一致。

九、文件地图与延伸阅读

文件职责
research/attention_ocr/README.md项目说明、环境/数据/训练/推理操作手册(本文主体)
research/attention_ocr/python/model.py模型建图:卷积塔、视角池化、损失、置信度、checkpoint 恢复
research/attention_ocr/python/sequence_layers.py注意力/自回归序列解码层与正交初始化
research/attention_ocr/python/data_provider.py数据读取、增强、中心裁剪、批量洗牌
research/attention_ocr/python/datasets/fsns.pyFSNS 数据集配置与 TFRecord 解码
research/attention_ocr/python/train.py训练入口与超参数
research/attention_ocr/python/eval.py评估入口(CPU 评估循环)
research/attention_ocr/python/model_export.pycheckpoint → SavedModel 导出
research/attention_ocr/python/demo_inference.py推理演示
research/attention_ocr/python/datasets/testdata/fsns/download_data.py小规模 testdata 数据下载脚本

按以上步骤,你可以在满足 158GB 磁盘与 GPU 环境的前提下完整复现 Attention OCR 的训练-评估-导出链路,也可以基于路线 1/2 将其迁移到自有街景文字识别任务上;对"模型为什么能逐字符对齐图像位置"这一核心问题,sequence_layers.py 的attention_decoder调用与 model.py 的_create_lstm_inputs空间切片逻辑是最直接的源码答案。

【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/7 6:19:46

AI-Edge边缘计算实战:从模型压缩到TensorRT部署的完整指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/7 6:17:45

用Excel VBA打造进销存系统:从表结构到自动记账全解析

简介&#xff1a;一套面向中小企业与Excel进阶用户的进销存管理系统VBA实现资源&#xff0c;围绕采购、销售、库存和报表四大核心流程&#xff0c;提供低成本且可灵活定制的业务管理方案。压缩包内共3个文件&#xff0c;包括可直接运行的主工作簿xlsm文件、用于关联关系的rels文…

作者头像 李华
网站建设 2026/9/7 6:14:53

Unity游戏开发:宝可梦机甲变身盲盒系统完整实现指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/7 6:14:50

CUDA统一内存深度解析:从原理到性能优化实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/7 6:12:10

STK中文教程详解:从卫星轨道设计到覆盖分析的航天仿真实践

简介&#xff1a;STK中文教程.zip是一份面向航天工程、遥感及军事仿真初学者的系统教程合集&#xff0c;围绕STK软件从基础概念到任务规划的典型学习路径做了完整整理。压缩包共48个文件&#xff0c;以14个PDF教程、2个PPT演示、8个GIF操作演示及多个场景工程文件&#xff08;M…

作者头像 李华