Scene Text Recognition (STR) system — reads text from natural images (street signs, book covers, license plates) using a CRNN architecture trained with CTC Loss, with an optional Bahdanau attention decoder for higher accuracy.
Image [B, 1, 32, W]
│
CNN Backbone (VGG-style, 8 conv layers, BN everywhere)
│ → [B, T, 512] T = W // 4
│
BiLSTM Encoder (2 layers, hidden=256/dir → 512 total)
│ → [B, T, 512]
│
├─────────────────────────────────┐
│ │
CTC Head Attention Decoder (toggleable)
[T, B, 37] log-probs [B, 25, 39] logits
│ │
CTCLoss CrossEntropyLoss
└──────────┬──────────────────────┘
│
Joint Loss = 0.3·CTC + 0.7·Attention
| Component | Role |
|---|---|
| CNN backbone | Extracts rich local visual features per image column |
| BiLSTM | Captures long-range context (knowing Q usually follows U) |
| CTC head | Handles variable-length output without explicit alignment |
| Attention decoder | Attends to specific feature positions for each output character |
| Joint loss | CTC stabilises early training; attention refines accuracy |
| Configuration | IIIT5K Word Accuracy | Train Time (RTX 4050) |
|---|---|---|
| CTC only, greedy decode | ~80-83 % | ~8 hrs / 10 epochs |
| CTC only, beam search | ~82-85 % | ~8 hrs / 10 epochs |
| CTC + Attention hybrid | ~85-89 % | ~12 hrs / 10 epochs |
git clone https://github.com/yourname/ocr-ctc
cd ocr-ctc
pip install -r requirements.txtGPU note: Install the CUDA-enabled PyTorch build that matches your driver before running
pip install -r requirements.txt:pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121
# Download and extract to /data/mjsynth
bash scripts/download_mjsynth.sh /data/mjsynthExpected layout:
/data/mjsynth/
annotation_train.txt
annotation_val.txt
mnt/ramdisk/max/90kDICT32px/
1/1_Hello_12345.jpg
...
bash scripts/download_iiit5k.sh /data/iiit5kExpected layout:
/data/iiit5k/
traindata.mat
testdata.mat
data/
1_1_1.png
...
bash scripts/download_svt.sh /data/svtExpected layout:
/data/svt/
train.xml
test.xml
img/
00_00.jpg
...
python train.py --use_attention false --max_epochs 2Uses an in-memory synthetic dataset with random word images.
python train.py \
--config configs/ctc_only.json \
--mjsynth_path /data/mjsynth \
--iiit5k_path /data/iiit5k \
--wandb_api_key YOUR_KEYpython train.py \
--config configs/hybrid.json \
--mjsynth_path /data/mjsynth \
--iiit5k_path /data/iiit5k \
--wandb_api_key YOUR_KEYpython train.py \
--config configs/hybrid.json \
--mjsynth_path /data/mjsynth \
--iiit5k_path /data/iiit5k \
--resume outputs/hybrid/checkpoints/last.ckptThe default config already uses --precision 16-mixed and
--gradient_checkpointing true. If you still get OOM:
python train.py \
--config configs/hybrid.json \
--batch_size 64 \
--accumulate_grads 2 \ # effective batch = 128
--mjsynth_path /data/mjsynth \
--iiit5k_path /data/iiit5k| Argument | Default | Description |
|---|---|---|
--use_attention |
true |
Attach attention decoder (hybrid mode) |
--attn_lambda |
0.3 |
CTC weight in joint loss |
--hidden_size |
256 |
BiLSTM hidden size per direction |
--batch_size |
128 |
Samples per GPU step |
--max_epochs |
20 |
Training epochs |
--lr |
1e-3 |
Initial AdamW learning rate |
--precision |
16-mixed |
Mixed-precision mode |
--gradient_checkpointing |
true |
Save VRAM at cost of ~30% extra compute |
--compile |
false |
torch.compile — ~20% speedup, slow startup |
--grad_clip |
5.0 |
Gradient clipping norm (essential for BiLSTM) |
--resume |
None |
Checkpoint path to resume from |
--wandb_api_key |
env | W&B key (or set WANDB_API_KEY) |
Full list: python train.py --help
# IIIT5K with greedy decoder
python evaluate.py \
--checkpoint outputs/hybrid/checkpoints/best.ckpt \
--dataset iiit5k \
--iiit5k_path /data/iiit5k
# SVT with beam search
python evaluate.py \
--checkpoint outputs/hybrid/checkpoints/best.ckpt \
--dataset svt \
--svt_path /data/svt \
--decoder beam --beam_width 10
# Save per-sample results to CSV
python evaluate.py \
--checkpoint outputs/hybrid/checkpoints/best.ckpt \
--dataset iiit5k \
--iiit5k_path /data/iiit5k \
--output results.csvSample output:
=======================================================
Evaluation — IIIT5K | decoder: greedy
=======================================================
Samples : 3000
Word Accuracy : 85.43 %
Char Accuracy : 94.12 %
CER : 5.88 %
WER : 14.57 %
Accuracy by word length:
len 3 : 91.20 %
len 4 : 88.50 %
len 5 : 85.10 %
len 6 : 82.30 %
len >7 : 74.60 %
Top char confusions (gt→pred, count):
'o' → '0' : 42
'l' → '1' : 38
=======================================================
# Single image
python inference.py \
--checkpoint outputs/hybrid/checkpoints/best.ckpt \
--image path/to/word.jpg
# With display
python inference.py \
--checkpoint outputs/hybrid/checkpoints/best.ckpt \
--image path/to/word.jpg --show
# Directory → CSV
python inference.py \
--checkpoint outputs/hybrid/checkpoints/best.ckpt \
--image path/to/images/ \
--output predictions.csv
# Beam search
python inference.py \
--checkpoint outputs/hybrid/checkpoints/best.ckpt \
--image word.jpg --decoder beam --beam_width 10Output:
Image : word.jpg
Predicted text : "london"
Confidence : 0.934
Decoder used : greedy
36 characters + 1 blank = 37 CTC classes.
| Token | Index | Notes |
|---|---|---|
<blank> |
0 | CTC blank — PyTorch hard requirement |
a … z |
1 … 26 | Lowercase letters |
0 … 9 |
27 … 36 | Digits |
<SOS> |
37 | Attention decoder only |
<EOS> |
38 | Attention decoder only |
All labels are lowercased. Non-alphanumeric characters are dropped at encode time.
ocr-ctc/
├── configs/ # JSON experiment configs
│ ├── base_config.json
│ ├── ctc_only.json
│ └── hybrid.json
├── data/
│ ├── dataset.py # MJSynthDataset, IIIT5KDataset, SVTDataset
│ ├── transforms.py # ResizeToHeight, GaussianNoise, pipelines
│ └── collate.py # Variable-width batch collation, CTCBatch
├── models/
│ ├── backbone.py # VGG-style CNN (8 conv layers, BN, gradient checkpointing)
│ ├── encoder.py # 2-layer BiLSTM
│ ├── ctc_head.py # Linear → log-softmax → time-first permute
│ ├── attention_decoder.py # Bahdanau attention + GRU decoder
│ └── crnn.py # Master model + from_args()
├── losses/
│ ├── ctc_loss.py # CTCLossWrapper (zero_infinity=True)
│ └── joint_loss.py # λ·CTC + (1-λ)·Attention, attention target builder
├── metrics/
│ ├── word_accuracy.py # Exact match + per-length breakdown
│ ├── char_accuracy.py # 1 - CER
│ ├── cer_wer.py # Levenshtein CER, WER, general WER
│ └── confusion.py # CharConfusionMatrix (36×36, difflib alignment)
├── training/
│ ├── trainer.py # OCRLightningModule (PL LightningModule)
│ ├── callbacks.py # SamplePredictionCallback, build_callbacks()
│ └── scheduler.py # ReduceLROnPlateau config
├── decoder/
│ ├── greedy.py # GreedyDecoder (argmax + collapse + strip-blank)
│ └── beam.py # BeamDecoder (pure-Python CTC beam search + ctcdecode)
├── utils/
│ ├── vocab.py # CHARS, VOCAB, encode(), decode(), decode_ctc()
│ ├── logger.py # TensorBoard + W&B logger setup
│ ├── checkpoint.py # build_checkpoint_callbacks(), load_module()
│ └── viz.py # Confusion matrix, sample grids, training curves
├── scripts/ # Dataset download shell scripts
├── outputs/ # Auto-created: checkpoints/, logs/, plots/
├── args.py # All argparse definitions (45 args) + JSON loader
├── train.py # Training entry point
├── evaluate.py # Evaluation entry point
└── inference.py # Single/batch inference entry point
tensorboard --logdir outputs/default/logsSet --wandb_api_key YOUR_KEY or export WANDB_API_KEY=YOUR_KEY.
Each run automatically logs:
- Loss curves (train/val, CTC/attention split)
- Word accuracy, char accuracy, CER, WER
- Confusion matrix heatmap (per epoch)
- Sample prediction grid (every 500 steps)
- Learning rate schedule
- GPU memory utilisation
outputs/<experiment_name>/
checkpoints/
best-epoch=005-val_acc=0.8543.ckpt ← top-3 by val_word_accuracy
last.ckpt ← always overwritten
logs/
tensorboard/ ← TensorBoard event files
plots/
confusion_epoch_005.png ← 36×36 confusion matrix
samples_step_0002000.png ← sample prediction grid
- Reduce
--batch_sizeto 64, add--accumulate_grads 2(keeps effective batch = 128). - Ensure
--gradient_checkpointing true(default). - Ensure
--precision 16-mixed(default). - Set
num_workers=0on Windows if workers cause OOM.
- Ensure
zero_infinity=Truein CTCLoss (this is the default in this codebase). - Ensure
--grad_clip 5.0(default) — NaN is almost always caused by gradient explosion in the BiLSTM.
ctcdecode requires a C++ build toolchain and a matching torch version.
The pure-Python beam search fallback in decoder/beam.py is used automatically and produces identical results.
Set --num_workers 0 to disable multiprocessing:
python train.py --num_workers 0 ...torch.compile compiles on first use — expect a 2-5 min delay before training starts.
After the first epoch it is significantly faster.
- Shi, B., Bai, X., & Yao, C. (2016). An end-to-end trainable neural network for image-based sequence recognition and its application to scene text recognition. IEEE TPAMI.
- Graves, A., et al. (2006). Connectionist temporal classification: Labelling unsegmented sequence data with recurrent neural networks. ICML.
- Sheng, F., et al. (2019). ASTER: An attentional scene text recognizer with flexible rectification. IEEE TPAMI.
- Jaderberg, M., et al. (2014). Synthetic data and artificial neural networks for natural scene text recognition. (MJSynth dataset).