# ctc_jax **Repository Path**: fmscole/ctc_jax ## Basic Information - **Project Name**: ctc_jax - **Description**: No description available - **Primary Language**: Unknown - **License**: MulanPSL-2.0 - **Default Branch**: master - **Homepage**: None - **GVP Project**: No ## Statistics - **Stars**: 0 - **Forks**: 0 - **Created**: 2026-06-04 - **Last Updated**: 2026-06-06 ## Categories & Tags **Categories**: Uncategorized **Tags**: None ## README # CTC — CRNN 中文文字识别 (纯JAX实现,stax规范) 纯 JAX 实现的 CTC 文字识别训练框架,使用 stax_plus 构建 CRNN 风格模型。 ## 依赖 - Python ≥ 3.10 - JAX (GPU 版推荐) - optax - Pillow - OpenCV (`cv2`) - torch (仅用作 DataLoader) - fonttools ```bash pip install jax[cuda12] optax pillow opencv-python torch fonttools ``` ## 项目结构 ``` ├── train.py # 训练入口 ├── stax_plus.py # 神经网络层库(Conv/BN/Dense/…) ├── generator.py # 图像生成 + PyTorch DataLoader ├── fontutils.py # 系统字体扫描与字符检测 ├── config.py # 配置 ├── ctcloss.py # CTC Loss(自定义 JAX 实现) └── data/ ├── words.py # 词表定义 └── char_std_5990.txt # 5990 个汉字/符号 ``` ## 快速开始 ```bash python train.py ``` 首次运行会触发 JIT 编译,耗时约 30–60 秒。之后每个 epoch 约 … 秒(取决于 GPU)。 ## 训练细节 - **模型**: 7 层 Conv + BN + ReLU + MaxPool → Dense - **输入**: 灰度图像 32×512(水平文字) - **输出**: 每帧在 5990 类上的 logits,CTC 解码后得到文字序列 - **优化器**: Adagrad, lr=0.01 - **Batch**: 100 - **Epochs**: 100 ## CTC Loss 实现 `ctcloss.py` 包含两个版本: - **`alpha()`** — 纯前向 α 递推(使用 JAX autograd 做反向传播) - **`ctcloss()`** — 手写 custom VJP,前向 α 递推 + 反向 β 递推 ### 算法:增广序列统一递推 将 blank 插入 label 序列形成增广序列 `S = [∅, l₀, ∅, l₁, ..., ∅, l_{L-1}, ∅]`(长度 2L+1),递推统一为单一 α/β 数组,不区分 blank/char: ``` 前向 α(t, s) = lp(t, S[s]) + logadd( α(t-1, s), α(t-1, s-1), α(t-1, s-2)·mask[s] ) 反向 β(t-1, s) = logadd( β(t,s)+lp(t,s), β(t,s+1)+lp(t,s+1), β(t,s+2)+lp(t,s+2)+mask[s+2] ) 梯度 ∂loss/∂logits = softmax - Σ_{s:S[s]=k} exp(α(t,s) + β(t,s) - log_P) ``` ### Custom VJP vs JAX Autograd | | JAX autograd | 手写 custom VJP | |---|---|---| | **反向实现** | JAX 自动为 `scan` 生成反向 scan | 显式 β 递推 + posterior 计算 | | **内存** | 基准 | 略优(~20-30%),省去 step 内部中间张量(pad 结果、logaddexp 分支) | | **速度** | 200 batch 耗时 ~6.05s | 200 batch 耗时 ~5.03s(**快 ~17%**) | | **数值稳定** | `logaddexp` 在 `(-inf,-inf)` 处梯度可能 NaN | 使用 `_safe_logaddexp` 保护,梯度始终有效 | | **正确性** | 理论正确 | 与 autograd 梯度一致(rtol < 1e-6) | **实测结论**:B=100, T=128, K=5990 的 CRNN 训练中,手写 custom VJP 每 200 batch 比 JAX autograd 快约 1 秒(~17%)。这个差距来自省去了 step 函数内部中间张量的保存和反向 trace。当 `T` 更大时优势更明显。 ## 数据 - 实时合成图像:随机背景 + 随机字体 + 随机文字 - 字体来自系统 `/usr/share/fonts/`(自动扫描 .ttf) - 字符集 5990 个,包含常用汉字、字母、数字、标点