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",即面向真实图像(而非印刷体、干净文本)的文字提取。论文的核心贡献有三点:
- 基于 ConvNets、RNN 与一种新型注意力机制的模型,在 FSNS 上达到 84.2% 完整序列准确率(此前基准 72.46%);
- 研究了使用不同深度 CNN 特征提取器时的速度/精度权衡;
- 提供的开源版本与论文版本的主要差异(见 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_inputs(model.py第 348~371 行)会断言卷积塔输出的特征维数不少于seq_length(FSNS 为 37),若超过则截取前seq_length个,即"每个时间步对应图像上的一个水平位置段",这是注意力机制得以逐字符对齐空间位置的基础。 - 输出端点:
OutputEndpoints命名元组包含chars_logit、predicted_chars、predicted_scores、predicted_text、predicted_length、predicted_conf、normalized_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 安装要求
- TensorFlow 1.15(README 明确要求,代码基于 TF 1.x 的
contrib与slim组件):
python3 -m venv ~/.tensorflow source ~/.tensorflow/bin/activate pip install --upgrade pip pip install --upgrade tensorflow-gpu=1.15- 磁盘空间:下载 FSNS 数据集至少需要 158GB 空闲空间:
cd research/attention_ocr/python/datasets aria2c -c -j 20 -i ../../../street/python/fsns_urls.txt cd ..- 内存:16GB 以上,推荐 32GB。
- 计算设备:
train.py同时支持 CPU 和 GPU,但推荐使用 GPU;作者已在 Titan X 与 GTX980 上测试。
2.2 FSNS 数据集构成
FSNS 数据集按子集切分为多个 TFRecord 文件,各子集规模(摘自 README):
| 子集 | 文件数 | 单文件约 | 样本数(源码配置) |
|---|---|---|---|
| Train | 512 个 | 300MB | 1,044,868 |
| Validation | 64 个 | 40MB | 16,150 |
| Test | 64 个 | 50MB | 20,040 |
| testdata | 若干小数据集 | 较小 | 用于验证模型能"学到东西" |
| 合计 | — | — | 约 158GB |
样本数依据 datasets/fsns.py 中DEFAULT_CONFIG的splits配置: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/width、image/orig_width:用于反推图像内含几个视角(_NumOfViewsHandler:num_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.pytrain.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-3997313.5 关键命令行参数(源码实测默认值)
train.py与eval.py共享 common_flags.py 中定义的参数,结合 train.py 补充如下:
| 参数 | 默认值 | 说明 |
|---|---|---|
batch_size | 32 | 批大小 |
crop_width/crop_height | None | 中心裁剪尺寸,用于缩减每个视角的宽度 |
train_log_dir | /tmp/attention_ocr/train | 事件日志与 checkpoint 目录 |
dataset_name | fsns | 数据集模块名,须为datasets包内可导入模块 |
split_name | train | 数据集切分:train / test / validation |
dataset_dir | None | 数据集根目录,默认用fsns.py中的DEFAULT_DATASET_DIR |
checkpoint | '' | 恢复完整模型权重的 checkpoint 路径 |
checkpoint_inception | '' | 仅恢复 Inception 卷积塔权重的 checkpoint 路径 |
learning_rate | 0.004 | 学习率 |
optimizer | momentum | 可选 momentum / adam / adadelta / adagrad / rmsprop(见create_optimizer,train.py) |
momentum | 0.9 | momentum 与 rmsprop 优化器的动量 |
use_augment_input | True | 是否启用图像增强 |
final_endpoint | Mixed_5d | InceptionV3 截断端点,论文借此研究 CNN 深度与速度/精度的权衡 |
use_attention | True | 是否使用注意力机制 |
use_autoregression | True | 是否使用自回归(反馈上一字符) |
num_lstm_units | 256 | 序列 LSTM 单元数 |
weight_decay | 0.00004 | 字符预测全连接层的权重衰减 |
lstm_state_clip_value | 10.0 | LSTM cell state 裁剪值 |
label_smoothing | 0.1 | label smoothing 权重 |
ignore_nulls | True | 计算损失时忽略 null 字符 |
average_across_timesteps | False | 是否按 label 权重总和平均损失 |
clip_gradient_norm | 2.0 | 梯度裁剪范数(train.py 独有) |
save_summaries_secs | 60 | 写 summary 的频率(秒) |
save_interval_secs | 600 | 保存 checkpoint 的频率(秒) |
max_number_of_steps | 1e10 | 最大梯度步数 |
task/ps_tasks/total_num_replicas | 0 / 0 / 1 | 多 worker 与参数服务器配置;sync_replicas=True时启用同步复制 |
reset_train_dir | False | True 时清空train_log_dir重新开始 |
show_graph_stats | False | 输出模型参数量统计 |
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-weight与weight/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 模式下注册CharacterAccuracy与SequenceAccuracy(流式指标,序列准确率在第一个 null 处截断序列后比较),实现见 metrics.py。
四、注意力与自回归解码器:核心创新在源码中的位置
序列解码层由 sequence_layers.py 实现,提供四种可切换的组合(get_layer_class,第 400~422 行):
| 类 | use_attention | use_autoregression | 说明 |
|---|---|---|---|
AttentionWithAutoregression | 是 | 是 | 默认配置,论文主模型 |
Attention | 是 | 否 | 纯注意力解码 |
NetSliceWithAutoregression | 否 | 是 | 按固定特征切片 + 自回归 |
NetSlice | 否 | 否 | 固定特征切片基线 |
实现细节值得注意的几处:
- 注意力解码:
Attention.unroll_cell直接调用tf.contrib.legacy_seq2seq.attention_decoder,把卷积塔特征self._net([batch, num_features, feature_size])作为attention_states,即注意力在空间特征序列上"池化"出每个字符对应的视觉证据(sequence_layers.py)。 - 自回归输入:训练时用 ground-truth one-hot 标签(teacher forcing,
get_train_input返回self._labels_one_hot[:, i-1, :]),推理时用上一时刻 argmax 字符的 one-hot(get_eval_input),首步输入为零向量。 - 正交初始化:
orthogonal_initializer用 SVD 生成正交矩阵初始化 LSTM 权重与 softmax 权重,作者引用了关于 RNN 正交初始化的文献注释(sequence_layers.py)。 - 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_image:convert_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=8、queue_capacity=3000、min_after_dequeue=1000(DEFAULT_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)然后做两件事:
- 把新模块加入 datasets/init.py(当前内容即
from datasets import fsns与__all__注册表); - 命令行指定数据集名:
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_split中keys_to_features的五个字段定义即可)。
路线 2:定义全新的数据集格式
模型训练只需要三路输入(images、labels、labels_one_hot,形状见上文第五节),README 指引以data_provider.py第 33 行的InputEndpoints为参照、以 datasets/fsns.py 为范例自行实现get_split。这条路线的落点是让新数据集对象同样携带num_char_classes、max_sequence_length、num_of_views、null_code、charset、image_shape等属性——因为train.py中common_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_exportmodel_export.py 的实际行为:
dataset_name与export_dir为必填参数;checkpoint 也可直接用--train_log_dir指向训练目录取最新 checkpoint;--export_for_serving(默认 True)时,输入为序列化的 tf.Example proto(placeholder 名tf_example),并在图内附加图像解码归一化;设为 False 时输入为uint8图像张量(名images),且必须指定--batch_size;- 输出张量非常丰富:
predictions(字符 ID)、scores、chars_logit、predicted_length、predicted_text、predicted_conf、normalized_seq_conf,以及attention_mask_0..36(37 个时间步的注意力掩码,可用于可视化模型"看哪里",见 model_export_lib.py); - 导出前会检查字符集文件存在,否则直接报错 "export will fail"。
7.2 手工建图推理(5 步法)
- 为图像定义 placeholder(或直接使用 numpy 数组);
- 建图(参照 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})- 用字符集文件把字符 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_fn的enabled=False,见 model.py 的default_mparams):开启后会给特征图每个位置拼接 x/y 坐标的 one-hot 编码,属于可选消融项; - 代码基于 TensorFlow 1.15(
tf.contrib、slim、tf.compat.v1混用),迁移到 TF2 需替换slim.dataset、legacy_seq2seq、contrib.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.py | FSNS 数据集配置与 TFRecord 解码 |
| research/attention_ocr/python/train.py | 训练入口与超参数 |
| research/attention_ocr/python/eval.py | 评估入口(CPU 评估循环) |
| research/attention_ocr/python/model_export.py | checkpoint → 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),仅供参考