Skip to content
zibranxoPublic

About

scene recognition system (crnn+ bahdanau attention)

Resources

Stars

0 stars

Watchers

0 watching

Forks

Repository files navigation

OCR with CTC Loss

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.


Architecture

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

Why this architecture

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

Expected Results

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

Installation

git clone https://github.com/yourname/ocr-ctc
cd ocr-ctc
pip install -r requirements.txt

GPU 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

Datasets

MJSynth / Synth90k (training — 9 M images)

# Download and extract to /data/mjsynth
bash scripts/download_mjsynth.sh /data/mjsynth

Expected layout:

/data/mjsynth/
    annotation_train.txt
    annotation_val.txt
    mnt/ramdisk/max/90kDICT32px/
        1/1_Hello_12345.jpg
        ...

IIIT5K (evaluation)

bash scripts/download_iiit5k.sh /data/iiit5k

Expected layout:

/data/iiit5k/
    traindata.mat
    testdata.mat
    data/
        1_1_1.png
        ...

SVT — Street View Text (evaluation)

bash scripts/download_svt.sh /data/svt

Expected layout:

/data/svt/
    train.xml
    test.xml
    img/
        00_00.jpg
        ...

Training

Quick smoke-test (no dataset needed)

python train.py --use_attention false --max_epochs 2

Uses an in-memory synthetic dataset with random word images.

CTC-only

python train.py \
    --config configs/ctc_only.json \
    --mjsynth_path /data/mjsynth \
    --iiit5k_path  /data/iiit5k  \
    --wandb_api_key YOUR_KEY

CTC + Attention hybrid (recommended)

python train.py \
    --config configs/hybrid.json \
    --mjsynth_path /data/mjsynth \
    --iiit5k_path  /data/iiit5k  \
    --wandb_api_key YOUR_KEY

Resume from checkpoint

python train.py \
    --config configs/hybrid.json \
    --mjsynth_path /data/mjsynth \
    --iiit5k_path  /data/iiit5k  \
    --resume outputs/hybrid/checkpoints/last.ckpt

Memory-constrained (6 GB GPU)

The 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

Key Training Arguments

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


Evaluation

# 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.csv

Sample 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
=======================================================

Inference

# 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 10

Output:

  Image          : word.jpg
  Predicted text : "london"
  Confidence     : 0.934
  Decoder used   : greedy

Vocabulary

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.


Repository Structure

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

Monitoring

TensorBoard

tensorboard --logdir outputs/default/logs

W&B

Set --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

Output Artefacts

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

Troubleshooting

CUDA out of memory

  1. Reduce --batch_size to 64, add --accumulate_grads 2 (keeps effective batch = 128).
  2. Ensure --gradient_checkpointing true (default).
  3. Ensure --precision 16-mixed (default).
  4. Set num_workers=0 on Windows if workers cause OOM.

NaN loss

  • Ensure zero_infinity=True in 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 not installing

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.

Windows DataLoader workers crash

Set --num_workers 0 to disable multiprocessing:

python train.py --num_workers 0 ...

Very slow first epoch with --compile

torch.compile compiles on first use — expect a 2-5 min delay before training starts. After the first epoch it is significantly faster.


References

  • 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).

About

scene recognition system (crnn+ bahdanau attention)

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages