"""Recognize the Hangul syllable and font of a single glyph image using the trained `FontRecognitionModel` checkpoint (model/hfc.pt). Usage: python example/demo.py python example/demo.py --image path/to/glyph.png --topk 10 """ from __future__ import annotations import argparse import csv import sys from pathlib import Path import torch from PIL import Image from model import CHAR_SIZE, FontRecognitionModel, decode_open, decode_restricted EXAMPLE_DIR = Path(__file__).resolve().parent PROJECT_ROOT = EXAMPLE_DIR.parent DEFAULT_CHECKPOINT = PROJECT_ROOT / "model" / "hfc.pt" DEFAULT_INDEX = PROJECT_ROOT / "model" / "index.csv" DEFAULT_IMAGE = EXAMPLE_DIR / "0408-0000.png" def load_font_index(index_csv: Path) -> dict[int, str]: """Load the font-id -> font-name mapping produced from model/index.json.""" with index_csv.open(encoding="utf-8-sig", newline="") as f: return {int(row["id"]): row["font_name"] for row in csv.DictReader(f)} def load_model(checkpoint_path: Path, num_font_classes: int, device: torch.device) -> FontRecognitionModel: checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) model = FontRecognitionModel(num_font_classes) model.load_state_dict(checkpoint["model"]) model.to(device) model.eval() return model def load_glyph_tensor(image_path: Path, device: torch.device) -> torch.Tensor: """Load a glyph image as a (1, 1, CHAR_SIZE, CHAR_SIZE) float tensor in [0, 1], matching the training data's black-glyph-on-white-background, unnormalized convention.""" image = Image.open(image_path).convert("L") if image.size != (CHAR_SIZE, CHAR_SIZE): image = image.resize((CHAR_SIZE, CHAR_SIZE), Image.BILINEAR) tensor = torch.frombuffer(bytearray(image.tobytes()), dtype=torch.uint8) tensor = tensor.reshape(1, 1, CHAR_SIZE, CHAR_SIZE).to(torch.float32) / 255.0 return tensor.to(device) def main() -> None: sys.stdout.reconfigure(encoding="utf-8") parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--checkpoint", type=Path, default=DEFAULT_CHECKPOINT) parser.add_argument("--index", type=Path, default=DEFAULT_INDEX) parser.add_argument("--image", type=Path, default=DEFAULT_IMAGE) parser.add_argument("--topk", type=int, default=5) args = parser.parse_args() device = torch.device("cuda" if torch.cuda.is_available() else "cpu") font_index = load_font_index(args.index) model = load_model(args.checkpoint, len(font_index), device) x = load_glyph_tensor(args.image, device) with torch.no_grad(): output = model.encode(x) restricted_char = decode_restricted(output.cho_logits, output.jung_logits, output.jong_logits)[0] open_char = decode_open(output.cho_logits, output.jung_logits, output.jong_logits)[0] font_probs = torch.softmax(output.font_logits, dim=-1)[0] topk_probs, topk_ids = font_probs.topk(min(args.topk, font_probs.numel())) print(f"Image: {args.image}") print() print("=== Hangul character recognition ===") print(f"Restricted decode (KS X 1001, 2,350 chars): {restricted_char} (U+{ord(restricted_char):04X})") print(f"Open decode (all 11,172 syllables): {open_char} (U+{ord(open_char):04X})") print() print("=== Font recognition ===") for rank, (font_id, prob) in enumerate(zip(topk_ids.tolist(), topk_probs.tolist()), start=1): print(f"{rank}. id={font_id:<5} prob={prob * 100:6.2f}% font_name={font_index[font_id]}") if __name__ == "__main__": main()