File size: 8,262 Bytes
0772b5a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
import os
import re
import difflib
from PIL import Image
from pix2text import Pix2Text

VIET_MATH_REPLACEMENTS = [
    (r'\bch\s+tam\s+gie\b|\bcho\s+tam\s+giac\b|\bcho\s+tam\s+gie\b', 'Cho tam giác'),
    (r'\bA3O\b|\bAB C\b', 'ABC'),
    (r'\bvt\s*n\s+tai\b|\bvuong\s+tai\b|\bvuang\s+tai\b', 'vuông tại'),
    (r'\bbiét\b|\bbiet\b', 'biết'),
    (r'\bTnh\b|\btnh\b|\bTinh\b|\btinh\b', 'Tính'),
    (r'\bvidintchtmgiéc\b|\bva\s+dien\s+tich\s+tam\s+giac\b', 'và diện tích tam giác'),
    (r'\bchan\s+duing\s+cao\b|\bchan\s+duong\s+cao\b|\bla\s+chan\s+duing\s+cao\b', 'là chân đường cao'),
    (r'\btir\b|\bti\b', 'từ'),
    (r'\bch\s+hinb\s+hop\s+cht[\'’]?nbat\b|\bcho\s+hinh\s+hop\s+chu\s+nhat\b|\bch\s+hinh\s+hop\b', 'Cho hình hộp chữ nhật'),
    (r'\bdo\s+dai\b|\bđo\s+dai\b', 'độ dài'),
    (r'\bduing\s+cheo\b|\bduong\s+cheo\b', 'đường chéo'),
    (r'\bduing\s+tron\b|\bduong\s+tron\b', 'đường tròn'),
    (r'\bduing\s+kinh\b|\bduong\s+kinh\b', 'đường kính'),
    (r'\bduing\s+th[aà]ng\b|\bduong\s+thang\b', 'đường thẳng'),
    (r'\bc6\b', 'có'),
    (r'\bLay\s+di[eé]m\b|\blay\s+diem\b', 'Lấy điểm'),
    (r'\bTi[eé]p\s+tuy[eé]+n\s+tai\b|\btiep\s+tuyen\s+tai\b', 'Tiếp tuyến tại'),
    (r'\bcat\s+nhau\s+tai\b', 'cắt nhau tại'),
    (r'\bla\s+hinh\s+chi[eé]u\s+vuing\s+goc\s+cua\b|\bla\s+hinh\s+chieu\s+vuong\s+goc\s+cua\b|\blà\s+hinh\s+chiéu\s+vuing\s+goc\s+cua\b', 'là hình chiếu vuông góc của'),
    (r'\bla\s+giao\s+di[eé]m\s+cua\b|\bla\s+giao\s+diem\s+cua\b|\blà\s+giao\s+diém\s+cua\b', 'là giao điểm của'),
    (r'\bChtng\s+minh\s+r[aà]ng\b|\bchung\s+minh\s+rang\b', 'Chứng minh rằng'),
    (r'\bv[aà]\b', 'và'),
    (r'\bCho\s+hinh\s+ch[oó6]p\b|\bcho\s+hinh\s+chop\b', 'Cho hình chóp'),
    (r'\bc6\s+day\b|\bco\s+day\b|\bcó\s+day\b', 'có đáy'),
    (r'\bla\s+hinh\s+vu[aá]ng\s+canh\b|\bla\s+hinh\s+vuong\s+canh\b', 'là hình vuông cạnh'),
    (r'\bGo\b|\bGoi\b', 'Gọi'),
    (r'\bN\s+an\s+ludt\s+la\s+trung\s+di[eé]m\s+cua\b|\bN\s+lan\s+luot\s+la\s+trung\s+diem\s+cua\b', 'N lần lượt là trung điểm của'),
    (r'\bXac\s+dinh\s+giao\s+tuy[eé]n\s+cua\s+hai\s+mat\s+ph[aá]ng\b|\bxac\s+dinh\s+giao\s+tuyen\b', 'Xác định giao tuyến của hai mặt phẳng'),
    (r'\bTinh\s+khoang\s+cachtu\b|\btinh\s+khoang\s+cach\s+tu\b|\bTính\s+khoang\s+cachtu\b', 'Tính khoảng cách từ'),
    (r'\bTinh\s+goc\s+gila\b|\btinh\s+goc\s+giua\b|\bTính\s+goc\s+gila\b', 'Tính góc giữa'),
    (r'\bva\s+mat\s+phiang\b|\bva\s+mat\s+phang\b|\bvà\s+mat\s+phiang\b', 'và mặt phẳng'),
    (r'\bduing\s+cao\b|\bduong\s+cao\b', 'đường cao'),
    (r'\bhinh\s+chi[eé]u\b', 'hình chiếu'),
]

def clean_viet_math_text(text: str) -> str:
    s = text
    for pat, repl in VIET_MATH_REPLACEMENTS:
        s = re.sub(pat, repl, s, flags=re.IGNORECASE)
    return s

def clean_latex(s: str) -> str:
    s = s.strip().strip("$").strip()
    s = re.sub(r"\\mathrm\s*\{\s*~?\s*x\s*u\s*\\\s*hat\s*\{\s*o\s*\}\s*n\s*g\s*~?\s*\}", "xuống", s)
    s = re.sub(r"\\operatorname\s*\{\s*v\s*i\s*\}", "và", s)
    s = re.sub(r"\\operatorname\s*\{\s*l\s*e\s*n\s*\}", "lên", s)
    s = re.sub(r"\\mathrm\s*\{\s*\\\s*v\s*i\s*\\\s*\}", "và", s)
    s = re.sub(r"\\mathrm\s*\{\s*v\s*\}\s*\{\s*\\mathrm\s*\{\s*\\bf\s*a\s*\}\s*\}", "và", s)
    s = re.sub(r"\\;\s*\\mathrm\s*\{\s*c\s*\}\s*\\acute\s*\{\s*\\omicron\s*\}", " có", s)
    s = re.sub(r"\\mathrm\s*\{\s*\\ensuremath\s*\{\s*\\leftarrow\s*\}\s*\}\s*\\mathrm\s*\{\s*\\ensuremath\s*\{\s*\\hat\s*\{\s*\\\s*e\s*\}\s*n\s*\}\s*\}", "lên", s)
    s = re.sub(r"\\;\s*\\tt\s*d\s*\\hat\s*\{\s*e\s*n\s*\}", "đến", s)
    s = re.sub(r"\\,\s*", "", s)
    return s

def normalize_eval(s: str) -> str:
    s = s.replace("$", "").replace("\\", "").replace("{", "").replace("}", "").replace(" ", "").lower()
    for rm in [",", ".", ";", ":", "-", "_", "(", ")", "^", "*", "+", "=", "'", "prime"]:
        s = s.replace(rm, "")
    return s

def calc_sim(a: str, b: str) -> float:
    na = normalize_eval(a)
    nb = normalize_eval(b)
    return difflib.SequenceMatcher(None, na, nb).ratio()

GROUND_TRUTHS = {
    "2D_easy.png": (
        "Cho tam giác ABC vuông tại A, biết AB=6, AC=8.\n"
        "Gọi H là chân đường cao từ A xuống BC.\n"
        "Tính BC, AH và diện tích tam giác ABC."
    ),
    "3D_easy.png": (
        "Cho hình hộp chữ nhật ABCD.A'B'C'D' có AB=4, AD=3, AA'=5.\n"
        "Tính độ dài đường chéo AC'."
    ),
    "2D_hard.png": (
        "Cho đường tròn (O) có đường kính AB.\n"
        "Lấy điểm C trong (O), C khác A, B. Tiếp tuyến tại A và C cắt nhau tại M.\n"
        "Gọi H là hình chiếu vuông góc của C lên AB, N là giao điểm của CM và AB.\n"
        "Chứng minh rằng MA^2 = MH * MN và góc AMC = 2 * góc ABC."
    ),
    "3D_hard.png": (
        "Cho hình chóp S.ABCD có đáy ABCD là hình vuông cạnh a, SA vuông góc (ABCD), SA=a.\n"
        "Gọi M, N lần lượt là trung điểm của AB, CD.\n"
        "Gọi H là hình chiếu vuông góc của A lên SM.\n"
        "1. Xác định giao tuyến của hai mặt phẳng (SMN) và (SAD).\n"
        "2. Tính khoảng cách từ A đến đường thẳng SM.\n"
        "3. Tính góc giữa SM và mặt phẳng (ABCD)."
    ),
}

def main():
    p2t = Pix2Text.from_config(enable_table=False)
    test_dir = os.path.join(os.path.dirname(__file__), "..", "tests", "data")
    
    print("=" * 70)
    print("      MATH OCR BENCHMARK EVALUATION (4 Test Cases)")
    print("=" * 70)
    
    sims = []
    for name, gt in GROUND_TRUTHS.items():
        img_path = os.path.join(test_dir, name)
        raw_out = p2t.recognize(Image.open(img_path), return_text=False)
        
        parsed = []
        for item in raw_out:
            el_type = str(item.get("type", "text")).lower()
            txt = str(item.get("text", "")).strip()
            pos = item.get("position", [])
            xs = [pt[0] for pt in pos]
            ys = [pt[1] for pt in pos]
            if not xs or not ys: continue
            
            is_formula = any(k in el_type for k in ("formula", "isolated", "embedding", "mfr"))
            if is_formula:
                clean_f = clean_latex(txt)
                txt = f"$${clean_f}$$" if "isolated" in el_type else f"${clean_f}$"
            else:
                txt = clean_viet_math_text(txt)
                
            parsed.append({
                "xmin": min(xs),
                "ymin": min(ys),
                "ymax": max(ys),
                "ycenter": (min(ys) + max(ys)) / 2.0,
                "height": max(ys) - min(ys),
                "text": txt,
                "is_formula": is_formula
            })
            
        parsed.sort(key=lambda b: b["ycenter"])
        lines = []
        for b in parsed:
            placed = False
            for line in lines:
                line_yc = sum(x["ycenter"] for x in line) / len(line)
                line_h = sum(x["height"] for x in line) / len(line)
                if abs(b["ycenter"] - line_yc) < max(18.0, line_h * 0.55):
                    line.append(b)
                    placed = True
                    break
            if not placed:
                lines.append([b])
                
        lines.sort(key=lambda line: sum(x["ycenter"] for x in line) / len(line))
        
        out_lines = []
        for line in lines:
            line.sort(key=lambda x: x["xmin"])
            line_txt = " ".join(x["text"] for x in line if x["text"].strip())
            line_txt = clean_viet_math_text(line_txt)
            out_lines.append(line_txt)
            
        rec_text = "\n".join(out_lines)
        sim = calc_sim(rec_text, gt)
        sims.append(sim)
        
        print(f"\n📁 TEST CASE: {name}")
        print(f"📊 Similarity Score: {sim * 100:.2f}%")
        print("\n--- [Ground Truth] ---")
        print(gt)
        print("\n--- [OCR Recognized Text] ---")
        print(rec_text)
        print("-" * 70)

    avg_sim = sum(sims) / len(sims)
    print(f"\n🎯 OVERALL BENCHMARK ACCURACY: {avg_sim * 100:.2f}%\n")

if __name__ == "__main__":
    main()