import argparse
import os
import random
import string
from pathlib import Path
import torch
import torch.nn as nn
import torch.nn.functional as F
from PIL import Image, ImageDraw, ImageFont, ImageFilter
from torch.utils.data import Dataset, DataLoader
CHARSET = string.digits + string.ascii_uppercase BLANK_CHAR = "-"
NUM_CLASSES = len(CHARSET) + 1
IMG_WIDTH = 160
IMG_HEIGHT = 60
MIN_LEN = 4
MAX_LEN = 6
BATCH_SIZE = 64
EPOCHS = 50
LR = 0.001
def generate_captcha(text: str, width: int = IMG_WIDTH, height: int = IMG_HEIGHT) -> Image.Image:
bg_color = (
random.randint(200, 255),
random.randint(200, 255),
random.randint(200, 255),
)
img = Image.new("RGB", (width, height), bg_color)
draw = ImageDraw.Draw(img)
font_paths = [
"/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf",
"/usr/share/fonts/truetype/liberation/LiberationSans-Bold.ttf",
"/usr/share/fonts/truetype/freefont/FreeSansBold.ttf",
"/System/Library/Fonts/Helvetica.ttc",
]
font = None
for fp in font_paths:
if os.path.exists(fp):
try:
font = ImageFont.truetype(fp, random.randint(28, 36))
break
except Exception:
pass
if font is None:
font = ImageFont.load_default()
for _ in range(random.randint(3, 8)):
x1, y1 = random.randint(0, width), random.randint(0, height)
x2, y2 = random.randint(0, width), random.randint(0, height)
color = (random.randint(100, 200), random.randint(100, 200), random.randint(100, 200))
draw.line([(x1, y1), (x2, y2)], fill=color, width=random.randint(1, 2))
for _ in range(random.randint(50, 150)):
x, y = random.randint(0, width), random.randint(0, height)
color = (random.randint(100, 200), random.randint(100, 200), random.randint(100, 200))
draw.point((x, y), fill=color)
x_offset = random.randint(10, 25)
for ch in text:
color = (
random.randint(20, 80),
random.randint(20, 80),
random.randint(20, 80),
)
y_offset = random.randint(8, 18)
draw.text((x_offset, y_offset), ch, font=font, fill=color)
x_offset += random.randint(20, 28)
if random.random() > 0.3:
img = img.filter(ImageFilter.GaussianBlur(radius=random.uniform(0.3, 1.0)))
if random.random() > 0.5:
img = img.transform(
img.size,
Image.Transform.AFFINE,
(1, random.uniform(-0.1, 0.1), 0, random.uniform(-0.05, 0.05), 1, 0),
)
return img
def random_text(min_len: int = MIN_LEN, max_len: int = MAX_LEN) -> str:
length = random.randint(min_len, max_len)
return "".join(random.choices(CHARSET, k=length))
class CaptchaDataset(Dataset):
def __init__(self, size: int = 10_000):
self.size = size
self.samples = [random_text() for _ in range(size)]
def __len__(self):
return self.size
def __getitem__(self, idx):
text = self.samples[idx]
img = generate_captcha(text)
img = img.convert("RGB")
img = img.resize((IMG_WIDTH, IMG_HEIGHT))
arr = torch.tensor(list(img.getdata()), dtype=torch.float32).view(IMG_HEIGHT, IMG_WIDTH, 3)
arr = arr.permute(2, 0, 1) / 255.0
label = torch.tensor([CHARSET.index(c) for c in text], dtype=torch.long)
label_len = torch.tensor(len(text), dtype=torch.long)
return arr, label, label_len
def collate_fn(batch):
images, labels, label_lengths = zip(*batch)
images = torch.stack(images)
labels = torch.cat(labels)
label_lengths = torch.stack(label_lengths)
return images, labels, label_lengths
class CRNN(nn.Module):
def __init__(self, num_classes: int = NUM_CLASSES):
super().__init__()
self.cnn = nn.Sequential(
nn.Conv2d(3, 64, 3, 1, 1), nn.ReLU(), nn.MaxPool2d(2, 2), nn.Conv2d(64, 128, 3, 1, 1), nn.ReLU(), nn.MaxPool2d(2, 2), nn.Conv2d(128, 256, 3, 1, 1), nn.BatchNorm2d(256), nn.ReLU(),
nn.Conv2d(256, 256, 3, 1, 1), nn.ReLU(), nn.MaxPool2d((2, 2), (2, 1), (0, 1)), nn.Conv2d(256, 512, 3, 1, 1), nn.BatchNorm2d(512), nn.ReLU(),
nn.Conv2d(512, 512, 3, 1, 1), nn.ReLU(), nn.MaxPool2d((2, 2), (2, 1), (0, 1)), nn.Conv2d(512, 512, 2, 1, 0), nn.BatchNorm2d(512), nn.ReLU(), )
self.rnn_input_size = 512 * 2 self.rnn_hidden = 256
self.num_classes = num_classes
self.lstm = nn.LSTM(
self.rnn_input_size,
self.rnn_hidden,
num_layers=2,
bidirectional=True,
dropout=0.2,
batch_first=False,
)
self.fc = nn.Linear(self.rnn_hidden * 2, num_classes)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.cnn(x) b, c, h, w = x.size()
x = x.permute(3, 0, 1, 2) x = x.reshape(w, b, c * h)
x, _ = self.lstm(x) x = self.fc(x) x = F.log_softmax(x, dim=2)
return x
class CTCLossWrapper(nn.Module):
def __init__(self):
super().__init__()
self.ctc = nn.CTCLoss(blank=len(CHARSET), reduction="mean", zero_infinity=True)
def forward(self, preds, labels, label_lengths, input_lengths):
return self.ctc(preds, labels, input_lengths, label_lengths)
def decode(pred: torch.Tensor) -> str:
pred = pred.argmax(dim=1).cpu().numpy()
chars = []
prev = -1
for p in pred:
if p != prev and p != len(CHARSET):
chars.append(CHARSET[p] if p < len(CHARSET) else "")
prev = p
return "".join(chars)
def export_onnx_from_checkpoint(output_dir: Path, device=None) -> Path:
if device is None:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
ckpt = output_dir / "crnn_text.pt"
if not ckpt.exists():
raise FileNotFoundError(f"no checkpoint at {ckpt}, train first (omit --export-only)")
model = CRNN().to(device)
model.load_state_dict(torch.load(ckpt, map_location=device))
model.eval()
dummy_input = torch.randn(1, 3, IMG_HEIGHT, IMG_WIDTH).to(device)
onnx_path = output_dir / "crnn_text.onnx"
torch.onnx.export(
model,
dummy_input,
onnx_path,
input_names=["input"],
output_names=["output"],
dynamic_axes={"input": {0: "batch_size"}, "output": {1: "batch_size"}},
opset_version=11,
dynamo=False,
)
print(f"ONNX model saved to {onnx_path}")
return onnx_path
def train(output_dir: Path, epochs: int = EPOCHS, batch_size: int = BATCH_SIZE):
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Training on {device}")
model = CRNN().to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=LR)
criterion = CTCLossWrapper().to(device)
dataset = CaptchaDataset(size=20_000)
loader = DataLoader(dataset, batch_size=batch_size, shuffle=True, collate_fn=collate_fn, num_workers=4)
val_dataset = CaptchaDataset(size=1_000)
val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn, num_workers=4)
best_acc = 0.0
for epoch in range(1, epochs + 1):
model.train()
total_loss = 0.0
for images, labels, label_lengths in loader:
images = images.to(device)
labels = labels.to(device)
label_lengths = label_lengths.to(device)
preds = model(images) T, B, _ = preds.size()
input_lengths = torch.full((B,), T, dtype=torch.long, device=device)
loss = criterion(preds, labels, label_lengths, input_lengths)
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
avg_loss = total_loss / len(loader)
print(f"Epoch {epoch:02d} | loss={avg_loss:.4f}", flush=True)
model.eval()
correct = 0
total = 0
with torch.no_grad():
for images, labels, label_lengths in val_loader:
images = images.to(device)
preds = model(images) preds = preds.permute(1, 0, 2)
label_ptr = 0
for b in range(preds.size(0)):
text = decode(preds[b])
true_len = label_lengths[b].item()
true_text = "".join(CHARSET[labels[label_ptr + i].item()] for i in range(true_len))
label_ptr += true_len
if text == true_text:
correct += 1
total += 1
acc = correct / total if total > 0 else 0.0
print(f"Epoch {epoch:02d} | loss={avg_loss:.4f} | val_acc={acc:.4f}", flush=True)
if acc > best_acc:
best_acc = acc
output_dir.mkdir(parents=True, exist_ok=True)
ckpt = output_dir / "crnn_text.pt"
torch.save(model.state_dict(), ckpt)
print(f" -> saved checkpoint (acc={acc:.4f})")
print("\nExporting to ONNX...")
export_onnx_from_checkpoint(output_dir, device)
print(f"Final validation accuracy: {best_acc:.4f}")
def default_output_dir() -> Path:
cache = os.environ.get("XDG_CACHE_HOME")
base = Path(cache) if cache else Path.home() / ".cache"
return base / "captchaforge" / "models" / "crnn_text"
def emit_eval_set(out_dir: Path, count: int, seed: int = 1234) -> Path:
import json
random.seed(seed)
eval_dir = out_dir / "eval"
eval_dir.mkdir(parents=True, exist_ok=True)
labels: dict[str, str] = {}
for i in range(count):
text = random_text()
generate_captcha(text).save(eval_dir / f"img_{i:04d}.png")
labels[f"img_{i:04d}.png"] = text
(eval_dir / "labels.json").write_text(json.dumps(labels, indent=2, sort_keys=True))
print(f"wrote {count} eval samples + labels.json -> {eval_dir}")
return eval_dir
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Train CRNN for text CAPTCHA recognition")
parser.add_argument("--output-dir", type=Path, default=default_output_dir())
parser.add_argument("--epochs", type=int, default=EPOCHS)
parser.add_argument("--batch-size", type=int, default=BATCH_SIZE)
parser.add_argument(
"--emit-eval-count",
type=int,
default=0,
help="Emit N labeled eval CAPTCHAs to <output-dir>/eval and exit (no training).",
)
parser.add_argument("--emit-eval-seed", type=int, default=1234)
parser.add_argument(
"--export-only",
action="store_true",
help="Skip training; re-export crnn_text.onnx from an existing crnn_text.pt checkpoint.",
)
args = parser.parse_args()
if args.emit_eval_count > 0:
emit_eval_set(args.output_dir, args.emit_eval_count, args.emit_eval_seed)
elif args.export_only:
export_onnx_from_checkpoint(args.output_dir)
else:
train(args.output_dir, args.epochs, args.batch_size)