File size: 3,578 Bytes
2c7a090
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
"""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()