Upload files to "/"

fonts
Romans Juškevičs 2026-09-18 08:53:40 +00:00
parent d2a97824f3
commit 2397988447
2 changed files with 226 additions and 0 deletions

126
synthetic_data.py 100644
View File

@ -0,0 +1,126 @@
"""
Synthetic data generator.
Maps to the flowchart's "Sintētiskie dati" branch: since we don't have scanned
handwriting yet, we render Latvian text as images and distort them (rotation,
shear, noise, blur, varying stroke width) to roughly approximate handwriting
variability. This is a placeholder, not a replacement for real data - swap in
scanned samples via `RealDataset` in dataset.py as soon as you have them.
For much better realism, drop a handwriting-style .ttf (e.g. any free cursive/
print handwriting font that supports Latvian diacritics: ā č ē ģ ī ķ ļ ņ š ū ž)
into FONT_DIR and it will be picked up automatically.
"""
import glob
import os
import random
import numpy as np
from PIL import Image, ImageDraw, ImageFilter, ImageFont
FONT_DIR = os.path.join(os.path.dirname(__file__), "fonts")
FALLBACK_FONT = "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf"
TARGET_HEIGHT = 32 # fixed line height fed into the CNN
def _available_fonts() -> list[str]:
fonts = glob.glob(os.path.join(FONT_DIR, "*.ttf")) + glob.glob(os.path.join(FONT_DIR, "*.otf"))
return fonts if fonts else [FALLBACK_FONT]
def _default_corpus() -> list[str]:
"""Small built-in Latvian word/phrase list, used if data/corpus.txt is absent."""
return [
"Labdien", "paldies", "lūdzu", "sveiki", "atā", "jā", "nē",
"Rīga", "Latvija", "valoda", "grāmata", "skola", "pilsēta",
"saule", "lietus", "sniegs", "vējš", "koks", "upe", "jūra",
"draugs", "ģimene", "māja", "ceļš", "darbs", "laiks",
"Es mācos latviešu valodu.", "Šodien ir skaista diena.",
"Viņš dzīvo Rīgā.", "Mums ir daudz darba.", "Kur ir tuvākā aptieka?",
"Cik tas maksā?", "Es gribu kafiju.", "Rīt būs saulains laiks.",
"Bērni spēlējas parkā.", "Šī grāmata ir ļoti interesanta.",
]
def load_corpus() -> list[str]:
corpus_path = os.path.join(os.path.dirname(__file__), "data", "corpus.txt")
if os.path.exists(corpus_path):
with open(corpus_path, encoding="utf-8") as f:
lines = [ln.strip() for ln in f if ln.strip()]
if lines:
return lines
return _default_corpus()
def render_text_line(text: str, font_path: str | None = None, augment: bool = True) -> Image.Image:
"""Render `text` as a single grayscale line image of fixed height."""
font_path = font_path or random.choice(_available_fonts())
font_size = random.randint(26, 40)
font = ImageFont.truetype(font_path, font_size)
# Measure text to size the canvas, with padding
dummy = Image.new("L", (10, 10), color=255)
bbox = ImageDraw.Draw(dummy).textbbox((0, 0), text, font=font)
w = max(1, bbox[2] - bbox[0]) + 20
h = max(1, bbox[3] - bbox[1]) + 20
img = Image.new("L", (w, h), color=255)
draw = ImageDraw.Draw(img)
draw.text((10 - bbox[0], 10 - bbox[1]), text, font=font, fill=0)
if augment:
img = _augment(img)
img = _resize_keep_ratio(img, TARGET_HEIGHT)
return img
def _augment(img: Image.Image) -> Image.Image:
# Slight random rotation
angle = random.uniform(-3, 3)
img = img.rotate(angle, expand=True, fillcolor=255)
# Slight shear via affine transform, mimics slanted handwriting
shear = random.uniform(-0.25, 0.25)
w, h = img.size
img = img.transform(
(w + int(abs(shear) * h), h),
Image.AFFINE,
(1, shear, -shear * h if shear < 0 else 0, 0, 1, 0),
fillcolor=255,
)
if random.random() < 0.5:
img = img.filter(ImageFilter.GaussianBlur(radius=random.uniform(0.2, 0.8)))
if random.random() < 0.5:
arr = np.array(img).astype(np.float32)
noise = np.random.normal(0, random.uniform(3, 10), arr.shape)
arr = np.clip(arr + noise, 0, 255).astype(np.uint8)
img = Image.fromarray(arr)
return img
def _resize_keep_ratio(img: Image.Image, target_height: int) -> Image.Image:
w, h = img.size
new_w = max(1, int(w * (target_height / h)))
return img.resize((new_w, target_height), Image.BILINEAR)
def generate_batch(n: int, corpus: list[str] | None = None) -> list[tuple[Image.Image, str]]:
corpus = corpus or load_corpus()
samples = []
for _ in range(n):
text = random.choice(corpus)
samples.append((render_text_line(text), text))
return samples
if __name__ == "__main__":
os.makedirs("samples", exist_ok=True)
for i, (img, text) in enumerate(generate_batch(500)):
img.save(f"samples/sample_{i}.png")
print(f"sample_{i}.png -> {text!r} size={img.size}")

100
train.py 100644
View File

@ -0,0 +1,100 @@
"""
Training script.
Usage:
python train.py --steps 2000 # synthetic data only
python train.py --real-data path/to/real_dataset # mix in real scans once you have them
Maps to the flowchart's "Projektu struktūra -> AI" step, trained against data
prepared in "Rokrakstu bāzes sagatavošana".
"""
import argparse
import os
import random
import torch
from torch.utils.data import ConcatDataset, DataLoader
from alphabet import NUM_CLASSES, encode
from dataset import RealHTRDataset, SyntheticHTRDataset, collate_batch
from model import CRNN
def build_dataset(real_data_path: str | None, synth_length: int):
datasets = [SyntheticHTRDataset(length=synth_length)]
if real_data_path:
datasets.append(RealHTRDataset(real_data_path))
return datasets[0] if len(datasets) == 1 else ConcatDataset(datasets)
def train(args):
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
dataset = build_dataset(args.real_data, args.synth_per_epoch)
loader = DataLoader(
dataset, batch_size=args.batch_size, shuffle=True,
collate_fn=collate_batch, num_workers=args.num_workers,
)
model = CRNN(num_classes=NUM_CLASSES).to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=args.lr)
ctc_loss = torch.nn.CTCLoss(blank=0, zero_infinity=True)
os.makedirs(args.checkpoint_dir, exist_ok=True)
step = 0
model.train()
while step < args.steps:
for images, texts, widths in loader:
images = images.to(device)
targets, target_lengths = [], []
for t in texts:
enc = encode(t)
targets.extend(enc)
target_lengths.append(len(enc))
targets = torch.tensor(targets, dtype=torch.long)
target_lengths = torch.tensor(target_lengths, dtype=torch.long)
log_probs = model(images) # [T, B, C]
T = log_probs.size(0)
input_lengths = torch.full((images.size(0),), T, dtype=torch.long)
loss = ctc_loss(log_probs, targets, input_lengths, target_lengths)
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)
optimizer.step()
step += 1
if step % args.log_every == 0:
print(f"step {step}/{args.steps} loss={loss.item():.4f}")
if step % args.checkpoint_every == 0 or step == args.steps:
ckpt_path = os.path.join(args.checkpoint_dir, f"crnn_step{step}.pt")
torch.save(model.state_dict(), ckpt_path)
print(f"saved checkpoint -> {ckpt_path}")
if step >= args.steps:
break
print("Training finished.")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--steps", type=int, default=2000)
parser.add_argument("--batch-size", type=int, default=32)
parser.add_argument("--lr", type=float, default=1e-3)
parser.add_argument("--synth-per-epoch", type=int, default=5000, help="synthetic samples generated per 'epoch' pass")
parser.add_argument("--real-data", type=str, default=None, help="path to a RealHTRDataset root, once available")
parser.add_argument("--num-workers", type=int, default=2)
parser.add_argument("--log-every", type=int, default=20)
parser.add_argument("--checkpoint-every", type=int, default=500)
parser.add_argument("--checkpoint-dir", type=str, default="checkpoints")
args = parser.parse_args()
random.seed(0)
torch.manual_seed(0)
train(args)