confia-captcha-ocr-pt / modeling_captcha.py
confiacompany's picture
Upload folder using huggingface_hub
f60c41e verified
Raw
History Blame Contribute Delete
3.19 kB
#!/usr/bin/env python3
"""
Positional Slot-Head CNN for 150x50 5-Character Captcha Recognition.
Developed by Confia Company.
Architecture Details:
- 3 Convolutional blocks with BatchNorm & ReLU.
- 2 MaxPool2d operations preserving high-resolution stroke features.
- AdaptiveAvgPool2d to a (6, 20) spatial feature grid.
- 5 positional slot heads (each focusing on a 4-column / 30px spatial window).
"""
import json
import os
import torch
import torch.nn as nn
from PIL import Image
from torchvision import transforms as T
DEFAULT_ALPHABET = "23456789abcdefghijklmnopqrstuvwxyz"
N_CHARS = 5
BASE_TRANSFORM = T.Compose([
T.Grayscale(),
T.ToTensor(),
T.Normalize(0.5, 0.5)
])
def conv_block(in_channels: int, out_channels: int) -> nn.Sequential:
return nn.Sequential(
nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU()
)
class PositionalSlotHeadCNN(nn.Module):
"""
Positional Slot-Head CNN model architecture.
Each of the 5 heads reads its corresponding 4-column window in the 6x20 spatial feature grid.
"""
def __init__(self, n_cls: int = 32, n_chars: int = N_CHARS):
super().__init__()
self.n_chars = n_chars
self.n_cls = n_cls
self.b1 = nn.Sequential(conv_block(1, 32), conv_block(32, 32), nn.MaxPool2d(2))
self.b2 = nn.Sequential(conv_block(32, 64), conv_block(64, 64), nn.MaxPool2d(2))
self.b3 = nn.Sequential(conv_block(64, 128), conv_block(128, 128))
self.pool = nn.AdaptiveAvgPool2d((6, 20))
self.drop = nn.Dropout(0.35)
self.slots = nn.ModuleList([
nn.Sequential(
nn.Linear(128 * 6 * 4, 256),
nn.ReLU(),
nn.Linear(256, n_cls)
)
for _ in range(self.n_chars)
])
def forward(self, x: torch.Tensor):
x = self.pool(self.b3(self.b2(self.b1(x))))
x = self.drop(x)
return [self.slots[s](x[:, :, :, 4 * s:4 * s + 4].flatten(1)) for s in range(self.n_chars)]
def load_model(weights_path: str, alphabet: str = DEFAULT_ALPHABET, device: str = None) -> torch.nn.Module:
if device is None:
device = "cuda" if torch.cuda.is_available() else "cpu"
model = PositionalSlotHeadCNN(n_cls=len(alphabet))
state_dict = torch.load(weights_path, map_location=device)
model.load_state_dict(state_dict)
model.to(device)
model.eval()
return model
def predict_captcha(model: torch.nn.Module, image_input, alphabet: str = DEFAULT_ALPHABET, device: str = None) -> str:
if device is None:
device = "cuda" if torch.cuda.is_available() else "cpu"
if isinstance(image_input, str):
img = Image.open(image_input).convert("RGB")
elif isinstance(image_input, Image.Image):
img = image_input.convert("RGB")
else:
raise ValueError("image_input must be a file path (str) or PIL Image")
x = BASE_TRANSFORM(img).unsqueeze(0).to(device)
with torch.no_grad():
logits = model(x)
prediction = "".join(alphabet[head_logits.argmax().item()] for head_logits in logits)
return prediction