Spaces:
Sleeping
Sleeping
Commit ·
7e03ee9
1
Parent(s): 4ff776d
Initial deploy: KIE Receipt Florence-2 demo
Browse files- README.md +13 -13
- kie/__init__.py +0 -0
- kie/config.py +235 -0
- kie/postprocessing.py +366 -0
- requirements.txt +8 -3
README.md
CHANGED
|
@@ -1,19 +1,19 @@
|
|
| 1 |
---
|
| 2 |
-
title:
|
| 3 |
-
emoji:
|
| 4 |
-
colorFrom:
|
| 5 |
-
colorTo:
|
| 6 |
-
sdk:
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
- streamlit
|
| 10 |
pinned: false
|
| 11 |
-
|
| 12 |
---
|
| 13 |
|
| 14 |
-
#
|
| 15 |
|
| 16 |
-
|
| 17 |
|
| 18 |
-
|
| 19 |
-
|
|
|
|
|
|
| 1 |
---
|
| 2 |
+
title: KIE Receipt Florence-2
|
| 3 |
+
emoji: 🧾
|
| 4 |
+
colorFrom: indigo
|
| 5 |
+
colorTo: blue
|
| 6 |
+
sdk: streamlit
|
| 7 |
+
sdk_version: 1.39.0
|
| 8 |
+
app_file: app/app.py
|
|
|
|
| 9 |
pinned: false
|
| 10 |
+
license: mit
|
| 11 |
---
|
| 12 |
|
| 13 |
+
# KIE Receipt — Florence-2 + LoRA Demo
|
| 14 |
|
| 15 |
+
Fine-tuned **Florence-2** (LoRA, 1.87% params) cho multi-domain receipt KIE · Macro ANLS **88.21%** · Giải quyết catastrophic forgetting bằng **Weighted Replay** (BWT +6.21%)
|
| 16 |
|
| 17 |
+
Upload ảnh hóa đơn (EN hoặc VI) → nhận JSON gồm company, date, address, total.
|
| 18 |
+
|
| 19 |
+
> Source code: [github.com/thaidinh1206/kie-receipt-florence2](https://github.com/thaidinh1206/kie-receipt-florence2)
|
kie/__init__.py
ADDED
|
File without changes
|
kie/config.py
ADDED
|
@@ -0,0 +1,235 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
config.py — Cấu hình trung tâm cho demo sản phẩm KIE Receipt
|
| 3 |
+
Số liệu thực nghiệm được load từ output/ của Kaggle (đã verify).
|
| 4 |
+
"""
|
| 5 |
+
|
| 6 |
+
from pathlib import Path
|
| 7 |
+
|
| 8 |
+
# ── Root của toàn bộ project (folder chứa kie/ và output/)
|
| 9 |
+
_HERE = Path(__file__).resolve().parent # …/project_kie/kie
|
| 10 |
+
_OUTROOT = _HERE.parent / "output" # …/project_kie/output
|
| 11 |
+
|
| 12 |
+
# ============================================================
|
| 13 |
+
# CHECKPOINT — E5 stage3/best (Full Pipeline, best model)
|
| 14 |
+
# Dùng absolute path để Streamlit không bị lỗi relative path
|
| 15 |
+
# ============================================================
|
| 16 |
+
CHECKPOINT_DIR = _OUTROOT / "checkpoint/E5/E05_full_pipeline/checkpoints/stage3/best"
|
| 17 |
+
|
| 18 |
+
# ── Checkpoint cho từng thí nghiệm (dùng khi muốn so sánh real inference)
|
| 19 |
+
CHECKPOINT_MAP = {
|
| 20 |
+
"E1": _OUTROOT / "checkpoint/E1/E01_no_replay/checkpoints/stage3/best",
|
| 21 |
+
"E2": _OUTROOT / "checkpoint/E2/E02_uniform_replay/checkpoints/stage3/best",
|
| 22 |
+
"E3": _OUTROOT / "checkpoint/E3/E03_weighted_replay/checkpoints/stage3/best",
|
| 23 |
+
"E4": _OUTROOT / "checkpoint/E3/E03_weighted_replay/checkpoints/stage3/best", # E4 = E3 ckpt + postproc
|
| 24 |
+
"E5": _OUTROOT / "checkpoint/E5/E05_full_pipeline/checkpoints/stage3/best",
|
| 25 |
+
}
|
| 26 |
+
|
| 27 |
+
# ── Florence-2 base model
|
| 28 |
+
MODEL_ID = "microsoft/Florence-2-base"
|
| 29 |
+
MODEL_REVISION = "refs/pr/6"
|
| 30 |
+
|
| 31 |
+
# ── HF Hub — checkpoint E5 (fallback khi không có local output/)
|
| 32 |
+
# Upload bằng: huggingface-cli upload <HF_CHECKPOINT_REPO> <local_ckpt_dir> --repo-type model
|
| 33 |
+
HF_CHECKPOINT_REPO = "thaidinhz1/kie-receipt-florence2-lora"
|
| 34 |
+
|
| 35 |
+
# ── KIE prompt — KHÔNG THAY ĐỔI (phải giống lúc train)
|
| 36 |
+
PROMPT = (
|
| 37 |
+
"<KIE_RECEIPT> Extract receipt information. "
|
| 38 |
+
"Return exactly JSON with keys: company, date, address, total."
|
| 39 |
+
)
|
| 40 |
+
|
| 41 |
+
# ── Inference
|
| 42 |
+
MAX_NEW_TOKENS = 128
|
| 43 |
+
NUM_BEAMS = 3
|
| 44 |
+
|
| 45 |
+
# ── 4 trường thông tin
|
| 46 |
+
FIELDS = ["company", "date", "address", "total"]
|
| 47 |
+
FIELD_LABELS = {
|
| 48 |
+
"company" : "🏪 Tên cơ sở",
|
| 49 |
+
"date" : "📅 Ngày GD",
|
| 50 |
+
"address" : "📍 Địa chỉ",
|
| 51 |
+
"total" : "💰 Tổng tiền",
|
| 52 |
+
}
|
| 53 |
+
|
| 54 |
+
# ── LoRA (chỉ dùng nếu load từ .pt state dict)
|
| 55 |
+
LORA_R = 16
|
| 56 |
+
LORA_ALPHA = 32
|
| 57 |
+
LORA_DROPOUT = 0.05
|
| 58 |
+
LORA_TARGETS = ["q_proj", "k_proj", "v_proj", "out_proj", "fc1", "fc2"]
|
| 59 |
+
|
| 60 |
+
# ============================================================
|
| 61 |
+
# KẾT QUẢ THỰC NGHIỆM E1–E5 — SỐ LIỆU THỰC TỪ KAGGLE OUTPUT
|
| 62 |
+
# Nguồn:
|
| 63 |
+
# E*X*_final_test_summary.json — test ANLS per field + macro
|
| 64 |
+
# E*X*_bwt_summary.json — BWT (Backward Transfer Weight)
|
| 65 |
+
#
|
| 66 |
+
# Ghi chú E4: cùng checkpoint với E3, chỉ khác inference postprocessing
|
| 67 |
+
# → BWT giống hệt E3; VN ANLS giống E3 (postprocess không đổi VN)
|
| 68 |
+
# ============================================================
|
| 69 |
+
|
| 70 |
+
EXPERIMENT_RESULTS = {
|
| 71 |
+
# ── E1: STL thuần (No Replay) ─────────────────────────────
|
| 72 |
+
"E1 — No Replay": {
|
| 73 |
+
"description": "STL thuần — không replay, không augmentation, không postprocessing nâng cao",
|
| 74 |
+
# Macro ANLS test (%) — Bảng 4.10
|
| 75 |
+
"sroie" : 72.21,
|
| 76 |
+
"mcocr" : 85.56,
|
| 77 |
+
"vn" : 83.66, # 5-fold CV mean (165 ảnh)
|
| 78 |
+
"macro" : 80.48, # (72.21+85.56+83.66)/3
|
| 79 |
+
# Per-field ANLS test — SROIE (Bảng C.1 — cần xác minh)
|
| 80 |
+
"sroie_company": 68.68, "sroie_date": 87.66,
|
| 81 |
+
"sroie_address": 60.37, "sroie_total": 66.51,
|
| 82 |
+
# Per-field ANLS test — MC-OCR (Bảng C.2 — cần xác minh)
|
| 83 |
+
"mcocr_company": 89.87, "mcocr_date": 96.55,
|
| 84 |
+
"mcocr_address": 50.38, "mcocr_total": 91.99,
|
| 85 |
+
# Per-field ANLS — VN Supermarket (Bảng 4.11, 5-fold CV mean)
|
| 86 |
+
"vn_company": 87.7, "vn_date": 94.2,
|
| 87 |
+
"vn_address": 74.4, "vn_total": 78.3,
|
| 88 |
+
# BWT (%) — Bảng 4.17, đo trên test set
|
| 89 |
+
"bwt_sroie" : -14.11,
|
| 90 |
+
"bwt_mcocr" : -5.67,
|
| 91 |
+
"macro_bwt" : -9.89,
|
| 92 |
+
# Validation ANLS tại snapshot sau từng stage (chưa có số liệu chính thức)
|
| 93 |
+
"r_1_1_sroie": 80.42, # SROIE val sau Stage 1
|
| 94 |
+
"r_2_2_mcocr": 88.78, # MC-OCR val sau Stage 2
|
| 95 |
+
"r_3_1_sroie": 67.80, # SROIE val sau Stage 3
|
| 96 |
+
"r_3_2_mcocr": 80.94, # MC-OCR val sau Stage 3
|
| 97 |
+
},
|
| 98 |
+
|
| 99 |
+
# ── E2: Uniform Replay [50/50, 34/33/33] ──────────────────
|
| 100 |
+
"E2 — Uniform Replay": {
|
| 101 |
+
"description": "Replay đồng đều: [MC-OCR 50% / SROIE 50%] và [VN 34% / MC-OCR 33% / SROIE 33%]",
|
| 102 |
+
# Macro ANLS test (%) — Bảng 4.10
|
| 103 |
+
"sroie" : 82.67,
|
| 104 |
+
"mcocr" : 91.73,
|
| 105 |
+
"vn" : 80.07, # 5-fold CV mean
|
| 106 |
+
"macro" : 84.82, # (82.67+91.73+80.07)/3
|
| 107 |
+
# Per-field ANLS test — SROIE (Bảng C.1 — cần xác minh)
|
| 108 |
+
"sroie_company": 65.58, "sroie_date": 69.53,
|
| 109 |
+
"sroie_address": 63.11, "sroie_total": 59.51,
|
| 110 |
+
# Per-field ANLS test — MC-OCR (Bảng C.2 — cần xác minh)
|
| 111 |
+
"mcocr_company": 95.60, "mcocr_date": 97.13,
|
| 112 |
+
"mcocr_address": 83.28, "mcocr_total": 94.70,
|
| 113 |
+
# Per-field ANLS — VN Supermarket (Bảng 4.11, 5-fold CV mean)
|
| 114 |
+
"vn_company": 85.4, "vn_date": 93.0,
|
| 115 |
+
"vn_address": 62.8, "vn_total": 79.0,
|
| 116 |
+
# BWT (%) — Bảng 4.17, đo trên test set
|
| 117 |
+
"bwt_sroie" : -3.02,
|
| 118 |
+
"bwt_mcocr" : +1.81,
|
| 119 |
+
"macro_bwt" : -0.61,
|
| 120 |
+
# Validation ANLS tại snapshot sau từng stage (chưa có số liệu chính thức)
|
| 121 |
+
"r_1_1_sroie": 81.30,
|
| 122 |
+
"r_2_2_mcocr": 87.99,
|
| 123 |
+
"r_3_1_sroie": 68.04,
|
| 124 |
+
"r_3_2_mcocr": 92.77,
|
| 125 |
+
},
|
| 126 |
+
|
| 127 |
+
# ── E3: Weighted Replay [85/15, 75/20/5] ──────────────────
|
| 128 |
+
"E3 — Weighted Replay": {
|
| 129 |
+
"description": "Replay có trọng số: [MC-OCR 85% / SROIE 15%] và [VN 75% / MC-OCR 20% / SROIE 5%]",
|
| 130 |
+
# Macro ANLS test (%) — Bảng 4.10
|
| 131 |
+
"sroie" : 87.12,
|
| 132 |
+
"mcocr" : 92.70,
|
| 133 |
+
"vn" : 83.23, # 5-fold CV mean
|
| 134 |
+
"macro" : 87.68, # (87.12+92.70+83.23)/3
|
| 135 |
+
# Per-field ANLS test — SROIE (Bảng C.1 — cần xác minh)
|
| 136 |
+
"sroie_company": 86.68, "sroie_date": 92.34,
|
| 137 |
+
"sroie_address": 83.34, "sroie_total": 80.49,
|
| 138 |
+
# Per-field ANLS test — MC-OCR (Bảng C.2 — cần xác minh)
|
| 139 |
+
"mcocr_company": 94.33, "mcocr_date": 95.63,
|
| 140 |
+
"mcocr_address": 80.56, "mcocr_total": 92.46,
|
| 141 |
+
# Per-field ANLS — VN Supermarket (Bảng 4.11, 5-fold CV mean)
|
| 142 |
+
"vn_company": 87.7, "vn_date": 94.1,
|
| 143 |
+
"vn_address": 72.4, "vn_total": 78.7,
|
| 144 |
+
# BWT (%) — Bảng 4.17, đo trên test set
|
| 145 |
+
"bwt_sroie" : +9.13,
|
| 146 |
+
"bwt_mcocr" : +3.28,
|
| 147 |
+
"macro_bwt" : +6.21,
|
| 148 |
+
# Validation ANLS tại snapshot sau từng stage (chưa có số liệu chính thức)
|
| 149 |
+
"r_1_1_sroie": 81.94,
|
| 150 |
+
"r_2_2_mcocr": 88.58,
|
| 151 |
+
"r_3_1_sroie": 86.69,
|
| 152 |
+
"r_3_2_mcocr": 90.99,
|
| 153 |
+
},
|
| 154 |
+
|
| 155 |
+
# ── E4: E3 + Tolerant PostProcess ─────────────────────────
|
| 156 |
+
# Cùng checkpoint với E3; chỉ khác: tolerant JSON parser + chuẩn hóa date/total
|
| 157 |
+
"E4 — +PostProcess": {
|
| 158 |
+
"description": "E3 + Tolerant JSON parser + chuẩn hóa date (DD/MM/YYYY) + chuẩn hóa total (bỏ VND, dấu phẩy)",
|
| 159 |
+
# Macro ANLS test (%) — Bảng 4.10
|
| 160 |
+
"sroie" : 86.52,
|
| 161 |
+
"mcocr" : 90.60,
|
| 162 |
+
"vn" : 84.31, # 5-fold CV mean
|
| 163 |
+
"macro" : 87.14, # (86.52+90.60+84.31)/3
|
| 164 |
+
# Per-field ANLS test — SROIE (Bảng 4.12)
|
| 165 |
+
"sroie_company": 88.4, "sroie_date": 96.6,
|
| 166 |
+
"sroie_address": 82.3, "sroie_total": 78.7,
|
| 167 |
+
# Per-field ANLS test — MC-OCR (Bảng 4.12)
|
| 168 |
+
"mcocr_company": 95.8, "mcocr_date": 97.8,
|
| 169 |
+
"mcocr_address": 74.8, "mcocr_total": 93.9,
|
| 170 |
+
# Per-field ANLS — VN Supermarket (Bảng 4.11, 5-fold CV mean)
|
| 171 |
+
"vn_company": 88.7, "vn_date": 94.5,
|
| 172 |
+
"vn_address": 73.8, "vn_total": 80.3,
|
| 173 |
+
# BWT (%) — cùng checkpoint E3, giống E3
|
| 174 |
+
"bwt_sroie" : +9.13,
|
| 175 |
+
"bwt_mcocr" : +3.28,
|
| 176 |
+
"macro_bwt" : +6.21,
|
| 177 |
+
# Validation ANLS tại snapshot sau từng stage (giống E3)
|
| 178 |
+
"r_1_1_sroie": 81.94,
|
| 179 |
+
"r_2_2_mcocr": 88.58,
|
| 180 |
+
"r_3_1_sroie": 86.69,
|
| 181 |
+
"r_3_2_mcocr": 90.99,
|
| 182 |
+
},
|
| 183 |
+
|
| 184 |
+
# ── E5: Full Pipeline (E4 + Augmentation Stage 3) ─────────
|
| 185 |
+
"E5 — Full Pipeline ✅": {
|
| 186 |
+
"description": "E4 + Augmentation cho VN train (Rotate/Perspective/Brightness/GaussNoise/Blur) — Best overall",
|
| 187 |
+
|
| 188 |
+
# ------------------------------------------------------------------
|
| 189 |
+
# KẾT QUẢ CHÍNH — SROIE & MC-OCR: holdout cố định; VN: 5-fold CV
|
| 190 |
+
# ------------------------------------------------------------------
|
| 191 |
+
# Macro ANLS test (%) — Bảng 4.10
|
| 192 |
+
# VN dùng 5-fold CV (165 ảnh, seed=42) vì holdout 17 ảnh quá nhỏ
|
| 193 |
+
"sroie" : 86.50,
|
| 194 |
+
"mcocr" : 92.89,
|
| 195 |
+
"vn" : 85.23, # 5-fold CV mean ± 4.44 pp — Bảng 4.10/4.11
|
| 196 |
+
"macro" : 88.21, # (86.50+92.89+85.23)/3
|
| 197 |
+
|
| 198 |
+
# Per-field ANLS — SROIE test holdout 64 ảnh (Bảng 4.12)
|
| 199 |
+
"sroie_company": 89.4, "sroie_date": 96.1,
|
| 200 |
+
"sroie_address": 85.6, "sroie_total": 74.8,
|
| 201 |
+
|
| 202 |
+
# Per-field ANLS — MC-OCR test holdout 87 ảnh (Bảng 4.12)
|
| 203 |
+
"mcocr_company": 95.2, "mcocr_date": 99.1,
|
| 204 |
+
"mcocr_address": 82.3, "mcocr_total": 95.0,
|
| 205 |
+
|
| 206 |
+
# Per-field ANLS — VN Supermarket, 5-fold CV mean (Bảng 4.11, 165 ảnh)
|
| 207 |
+
"vn_company": 87.7, "vn_date": 95.9,
|
| 208 |
+
"vn_address": 75.3, "vn_total": 82.0,
|
| 209 |
+
|
| 210 |
+
# ------------------------------------------------------------------
|
| 211 |
+
# KẾT QUẢ PHỤ — VN holdout cố định 17 ảnh
|
| 212 |
+
# Dùng cho ablation (Bảng 4.12/4.14/4.18/4.19); KHÔNG d��ng làm kết quả chính
|
| 213 |
+
# vì n=17 gây phương sai cao; k-fold trên 165 ảnh đáng tin cậy hơn
|
| 214 |
+
# ------------------------------------------------------------------
|
| 215 |
+
"vn_holdout" : 90.30, # Bảng 4.14 (macro, holdout 17 ảnh)
|
| 216 |
+
"vn_holdout_company" : 93.5, # Bảng 4.12
|
| 217 |
+
"vn_holdout_date" : 95.9,
|
| 218 |
+
"vn_holdout_address" : 83.4,
|
| 219 |
+
"vn_holdout_total" : 88.4,
|
| 220 |
+
|
| 221 |
+
# ------------------------------------------------------------------
|
| 222 |
+
# BWT — đo trên validation set (Bảng 4.13)
|
| 223 |
+
# Lưu ý: E1-E3 BWT đo trên test set (Bảng 4.17); E5 BWT đo trên val set
|
| 224 |
+
# ------------------------------------------------------------------
|
| 225 |
+
"bwt_sroie" : -4.78,
|
| 226 |
+
"bwt_mcocr" : +2.52,
|
| 227 |
+
"macro_bwt" : -1.13,
|
| 228 |
+
|
| 229 |
+
# Validation ANLS tại snapshot sau từng stage — Bảng 4.13
|
| 230 |
+
"r_1_1_sroie": 89.50, # SROIE val sau Stage 1
|
| 231 |
+
"r_2_2_mcocr": 89.36, # MC-OCR val sau Stage 2
|
| 232 |
+
"r_3_1_sroie": 84.72, # SROIE val sau Stage 3
|
| 233 |
+
"r_3_2_mcocr": 91.88, # MC-OCR val sau Stage 3
|
| 234 |
+
},
|
| 235 |
+
}
|
kie/postprocessing.py
ADDED
|
@@ -0,0 +1,366 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
postprocessing.py — Tolerant JSON parser + chuẩn hóa đầu ra mô hình
|
| 3 |
+
Bám sát 100% logic trong E05_full_pipeline.py (postprocess / tolerant_json_parse)
|
| 4 |
+
"""
|
| 5 |
+
|
| 6 |
+
import re
|
| 7 |
+
import json
|
| 8 |
+
from .config import FIELDS
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
# ─────────────────────────────────────────────────────────────
|
| 12 |
+
# 1. TOLERANT JSON PARSER
|
| 13 |
+
# ─────────────────────────────────────────────────────────────
|
| 14 |
+
|
| 15 |
+
def _try_parse(text: str):
|
| 16 |
+
"""Thử json.loads đơn thuần, trả None nếu lỗi."""
|
| 17 |
+
try:
|
| 18 |
+
return json.loads(text)
|
| 19 |
+
except Exception:
|
| 20 |
+
return None
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def _regex_extract(raw: str) -> dict:
|
| 24 |
+
"""Fallback: dùng regex tìm từng key='value' khi JSON hoàn toàn sai cú pháp."""
|
| 25 |
+
result = {k: None for k in FIELDS}
|
| 26 |
+
for key in FIELDS:
|
| 27 |
+
m = re.search(
|
| 28 |
+
rf'["\']?{key}["\']?\s*:\s*["\']([^"\'{{}}]+)["\']',
|
| 29 |
+
raw, re.IGNORECASE
|
| 30 |
+
)
|
| 31 |
+
if m:
|
| 32 |
+
result[key] = m.group(1).strip()
|
| 33 |
+
return result
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def tolerant_json_parse(raw: str) -> dict:
|
| 37 |
+
"""
|
| 38 |
+
Phân tích chuỗi đầu ra mô hình thành dict {company, date, address, total}.
|
| 39 |
+
6 bước xử lý tích lũy (khớp Config 1–6 trong ablation_postproc.py):
|
| 40 |
+
1. json.loads thuần
|
| 41 |
+
2. Tách substring {...}
|
| 42 |
+
3. Tự động thêm dấu '}' thiếu
|
| 43 |
+
4. Sửa key không có nháy kép
|
| 44 |
+
5. Regex fallback
|
| 45 |
+
6. Trả về dict rỗng
|
| 46 |
+
"""
|
| 47 |
+
if not raw:
|
| 48 |
+
return {k: None for k in FIELDS}
|
| 49 |
+
raw = raw.strip()
|
| 50 |
+
|
| 51 |
+
# Bước 1 — thử parse toàn bộ chuỗi
|
| 52 |
+
obj = _try_parse(raw)
|
| 53 |
+
if obj is not None:
|
| 54 |
+
return {k: obj.get(k) for k in FIELDS}
|
| 55 |
+
|
| 56 |
+
# Bước 2 — tìm cặp ngoặc {...}
|
| 57 |
+
s, e = raw.find("{"), raw.rfind("}")
|
| 58 |
+
if s < 0 or e <= s:
|
| 59 |
+
return _regex_extract(raw) # không có ngoặc → regex
|
| 60 |
+
|
| 61 |
+
text = raw[s:e + 1]
|
| 62 |
+
obj = _try_parse(text)
|
| 63 |
+
if obj is not None:
|
| 64 |
+
return {k: obj.get(k) for k in FIELDS}
|
| 65 |
+
|
| 66 |
+
# Bước 3 — sửa ngoặc đóng thiếu
|
| 67 |
+
open_c = text.count("{")
|
| 68 |
+
close_c = text.count("}")
|
| 69 |
+
if open_c > close_c:
|
| 70 |
+
text += "}" * (open_c - close_c)
|
| 71 |
+
obj = _try_parse(text)
|
| 72 |
+
if obj is not None:
|
| 73 |
+
return {k: obj.get(k) for k in FIELDS}
|
| 74 |
+
|
| 75 |
+
# Bước 4 — sửa key không có nháy kép (company: "..." → "company": "...")
|
| 76 |
+
fixed = re.sub(r'([a-zA-Z_]+)":', r'"\1":', text)
|
| 77 |
+
obj = _try_parse(fixed)
|
| 78 |
+
if obj is not None:
|
| 79 |
+
return {k: obj.get(k) for k in FIELDS}
|
| 80 |
+
|
| 81 |
+
# Bước 5 — regex fallback (trích từng trường độc lập)
|
| 82 |
+
result = _regex_extract(raw)
|
| 83 |
+
|
| 84 |
+
# Bước 6 — trả về dict rỗng nếu không tìm được gì
|
| 85 |
+
return result if any(v for v in result.values()) else {k: None for k in FIELDS}
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
# ─────────────────────────────────────────────────────────────
|
| 89 |
+
# 2. TỰ NHẬN DIỆN NGÔN NGỮ HÓA ĐƠN
|
| 90 |
+
# ─────────────────────────────────────────────────────────────
|
| 91 |
+
|
| 92 |
+
# Ký tự đặc trưng tiếng Việt (không có trong tiếng Anh/Latin cơ bản)
|
| 93 |
+
_VN_CHARS = re.compile(
|
| 94 |
+
r'[àáảãạăắằẳẵặâấầẩẫậèéẻẽẹêếềểễệìíỉĩịòóỏõọôốồổỗộơớờởỡợùúủũụưứừửữựỳýỷỹỵđ'
|
| 95 |
+
r'ÀÁẢÃẠĂẮẰẲẴẶÂẤẦẨẪẬÈÉẺẼẸÊẾỀỂỄỆÌÍỈĨỊÒÓỎÕỌÔỐỒỔỖỘƠỚỜỞỠỢÙÚỦŨỤƯỨỪỬỮỰỲÝỶỸỴĐ]',
|
| 96 |
+
re.UNICODE
|
| 97 |
+
)
|
| 98 |
+
|
| 99 |
+
# Từ khóa xuất hiện trên hóa đơn VN kể cả khi không có dấu
|
| 100 |
+
_VN_KEYWORDS = re.compile(
|
| 101 |
+
r'\b(tong\s*tien|phieu\s*tinh\s*tien|hoa\s*don|don\s*gia|thanh\s*tien'
|
| 102 |
+
r'|so\s*luong|mat\s*hang|hang\s*hoa|giam\s*gia|tien\s*mat|tien\s*thua'
|
| 103 |
+
r'|tien\s*tra\s*lai|khach\s*hang|nhan\s*vien|thu\s*ngan'
|
| 104 |
+
r'|winmart|vincommerce|co\.?op|bach\s*hoa\s*xanh|circle\s*k'
|
| 105 |
+
r'|fujimart|big\s*c|lotte\s*mart|emart|coopmart|coopfood)\b',
|
| 106 |
+
re.IGNORECASE
|
| 107 |
+
)
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
def _is_vnd_amount(total_str) -> bool:
|
| 111 |
+
"""Kiểm tra xem giá trị total có khả năng là VND không.
|
| 112 |
+
VND thường >= 5000 và là số nguyên (không có phần thập phân .xx như USD)."""
|
| 113 |
+
if not total_str:
|
| 114 |
+
return False
|
| 115 |
+
# Bỏ ký hiệu tiền tệ đã có sẵn
|
| 116 |
+
cleaned = re.sub(r'[,\.\s đ₫VNDvnd]', '', str(total_str))
|
| 117 |
+
try:
|
| 118 |
+
val = float(cleaned)
|
| 119 |
+
# VND >= 5000 và không có phần thập phân cents như 12.50
|
| 120 |
+
return val >= 5000 and val == int(val)
|
| 121 |
+
except (ValueError, TypeError):
|
| 122 |
+
return False
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
def auto_detect_domain(parsed_dict: dict) -> str:
|
| 126 |
+
"""
|
| 127 |
+
Tự động nhận diện ngôn ngữ hóa đơn dựa trên nội dung các trường đã parse.
|
| 128 |
+
|
| 129 |
+
Hệ thống điểm đa tín hiệu (để xử lý hóa đơn VN không có dấu):
|
| 130 |
+
+3 — có ký tự tiếng Việt có dấu (tín hiệu chắc chắn nhất)
|
| 131 |
+
+2 — có từ khóa VN không dấu (TONG TIEN, HOA DON, WINMART, v.v.)
|
| 132 |
+
+1 — total >= 5000 và là số nguyên (đặc trưng VND)
|
| 133 |
+
Ngưỡng quyết định: tổng điểm >= 1 → "vn", ngược lại → "sroie"
|
| 134 |
+
|
| 135 |
+
Parameters
|
| 136 |
+
----------
|
| 137 |
+
parsed_dict : dict — kết quả từ tolerant_json_parse
|
| 138 |
+
|
| 139 |
+
Returns
|
| 140 |
+
-------
|
| 141 |
+
"vn" hoặc "sroie"
|
| 142 |
+
"""
|
| 143 |
+
text_to_check = " ".join(
|
| 144 |
+
str(parsed_dict.get(k) or "")
|
| 145 |
+
for k in ("company", "address", "total", "date")
|
| 146 |
+
)
|
| 147 |
+
|
| 148 |
+
score = 0
|
| 149 |
+
|
| 150 |
+
# Tín hiệu 1: ký tự có dấu tiếng Việt (chắc chắn)
|
| 151 |
+
if _VN_CHARS.search(text_to_check):
|
| 152 |
+
score += 3
|
| 153 |
+
|
| 154 |
+
# Tín hiệu 2: từ khóa VN không dấu (tên chuỗi, thuật ngữ hóa đơn)
|
| 155 |
+
if _VN_KEYWORDS.search(text_to_check):
|
| 156 |
+
score += 2
|
| 157 |
+
|
| 158 |
+
# Tín hiệu 3: giá trị total đặc trưng VND (>= 5000, không có cents)
|
| 159 |
+
if _is_vnd_amount(parsed_dict.get("total")):
|
| 160 |
+
score += 1
|
| 161 |
+
|
| 162 |
+
return "vn" if score >= 1 else "sroie"
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
# ─────────────────────────────────────────────────────────────
|
| 166 |
+
# 3. CHUẨN HÓA TỪNG TRƯỜNG
|
| 167 |
+
# ─────────────────────────────────────────────────────────────
|
| 168 |
+
|
| 169 |
+
_MONTH_MAP = {
|
| 170 |
+
"jan": "01", "feb": "02", "mar": "03", "apr": "04",
|
| 171 |
+
"may": "05", "jun": "06", "jul": "07", "aug": "08",
|
| 172 |
+
"sep": "09", "oct": "10", "nov": "11", "dec": "12"
|
| 173 |
+
}
|
| 174 |
+
|
| 175 |
+
_DATE_PATTERNS = [
|
| 176 |
+
# DD/MM/YYYY hoặc DD-MM-YYYY hoặc DD.MM.YYYY (swap m[0] và m[1] để tạo đúng YYYY-MM-DD)
|
| 177 |
+
(r'(\d{1,2})[/\-\.](\d{1,2})[/\-\.](\d{4})',
|
| 178 |
+
lambda m: f"{m[2]}-{int(m[1]):02d}-{int(m[0]):02d}"),
|
| 179 |
+
# ngày DD tháng MM năm YYYY (tiếng Việt có hoặc không dấu)
|
| 180 |
+
(r'(?i)ng[àa]y\s+(\d{1,2})\s+th[áa]ng\s+(\d{1,2})\s+n[ăa]m\s+(\d{4})',
|
| 181 |
+
lambda m: f"{m[2]}-{int(m[1]):02d}-{int(m[0]):02d}"),
|
| 182 |
+
# YYYY/MM/DD
|
| 183 |
+
(r'(\d{4})[/\-\.](\d{1,2})[/\-\.](\d{1,2})',
|
| 184 |
+
lambda m: f"{m[0]}-{int(m[1]):02d}-{int(m[2]):02d}"),
|
| 185 |
+
# DD Mon YYYY (e.g. 25 Mar 2023)
|
| 186 |
+
(r'(\d{1,2})\s+([A-Za-z]{3})\s+(\d{4})',
|
| 187 |
+
lambda m: (f"{m[2]}-{_MONTH_MAP.get(m[1].lower()[:3], '??')}-{int(m[0]):02d}"
|
| 188 |
+
if _MONTH_MAP.get(m[1].lower()[:3]) else None)),
|
| 189 |
+
# YYYYMMDD
|
| 190 |
+
(r'^(\d{4})(\d{2})(\d{2})$',
|
| 191 |
+
lambda m: f"{m[0]}-{m[1]}-{m[2]}"),
|
| 192 |
+
]
|
| 193 |
+
|
| 194 |
+
|
| 195 |
+
def normalize_date(val) -> str:
|
| 196 |
+
"""Chuẩn hóa ngày về định dạng DD/MM/YYYY."""
|
| 197 |
+
if not val:
|
| 198 |
+
return val
|
| 199 |
+
val = str(val).strip()
|
| 200 |
+
|
| 201 |
+
# Đã đúng định dạng DD/MM/YYYY → giữ nguyên
|
| 202 |
+
if re.match(r'^\d{2}/\d{2}/\d{4}$', val):
|
| 203 |
+
return val
|
| 204 |
+
|
| 205 |
+
# Đã là YYYY-MM-DD (ISO) → chuyển sang DD/MM/YYYY
|
| 206 |
+
m = re.match(r'^(\d{4})-(\d{2})-(\d{2})$', val)
|
| 207 |
+
if m:
|
| 208 |
+
return f"{m.group(3)}/{m.group(2)}/{m.group(1)}"
|
| 209 |
+
|
| 210 |
+
# Thử các pattern khác, parse về YYYY-MM-DD rồi chuyển sang DD/MM/YYYY
|
| 211 |
+
for pattern, formatter in _DATE_PATTERNS:
|
| 212 |
+
m = re.search(pattern, val)
|
| 213 |
+
if m:
|
| 214 |
+
try:
|
| 215 |
+
iso = formatter(m.groups())
|
| 216 |
+
if iso and "??" not in iso:
|
| 217 |
+
# iso = "YYYY-MM-DD"
|
| 218 |
+
parts = iso.split("-")
|
| 219 |
+
if len(parts) == 3:
|
| 220 |
+
return f"{parts[2]}/{parts[1]}/{parts[0]}"
|
| 221 |
+
except Exception:
|
| 222 |
+
pass
|
| 223 |
+
|
| 224 |
+
# Thử dateutil nếu cài
|
| 225 |
+
try:
|
| 226 |
+
from dateutil import parser as du_parser
|
| 227 |
+
return du_parser.parse(val, dayfirst=True).strftime("%d/%m/%Y")
|
| 228 |
+
except Exception:
|
| 229 |
+
pass
|
| 230 |
+
|
| 231 |
+
return val
|
| 232 |
+
|
| 233 |
+
|
| 234 |
+
def normalize_total(val, domain: str = "sroie") -> str:
|
| 235 |
+
"""
|
| 236 |
+
Chuẩn hóa tổng tiền:
|
| 237 |
+
- Bỏ ký tự tiền tệ gốc (VND, VNĐ, đồng, ₫, đ)
|
| 238 |
+
- Bỏ dấu phân cách hàng nghìn để lấy số thuần
|
| 239 |
+
- Nếu domain tiếng Việt (mcocr / vn):
|
| 240 |
+
→ dùng dấu chấm phân cách nghìn: 5000 → "5.000"
|
| 241 |
+
→ thêm ký hiệu " đ" ở cuối: "5.000 đ"
|
| 242 |
+
- Nếu domain tiếng Anh (sroie): trả về số thuần (không đơn vị)
|
| 243 |
+
"""
|
| 244 |
+
if not val:
|
| 245 |
+
return val
|
| 246 |
+
val = str(val).strip()
|
| 247 |
+
|
| 248 |
+
# Bỏ ký hiệu tiền tệ gốc (thêm RM, MYR, $, USD)
|
| 249 |
+
val = re.sub(r'(?i)(vnd|vnđ|đồng|dong|₫|đ|rm|myr|\$|usd)\s*', '', val).strip()
|
| 250 |
+
|
| 251 |
+
# Bỏ dấu phân cách hàng nghìn (dấu , hoặc . trước đúng 3 chữ số)
|
| 252 |
+
val_clean = re.sub(r'[,\.](?=\d{3}(?:[,\.]|$))', '', val)
|
| 253 |
+
val_clean = val_clean.replace(',', '.').strip()
|
| 254 |
+
|
| 255 |
+
is_vn = str(domain).lower() in ("mcocr", "vn", "mc-ocr", "vn supermarket",
|
| 256 |
+
"vn_supermarket", "vietnamese")
|
| 257 |
+
try:
|
| 258 |
+
num = float(val_clean)
|
| 259 |
+
if 1e2 <= num <= 1e9:
|
| 260 |
+
int_num = int(num) if num == int(num) else num
|
| 261 |
+
if is_vn:
|
| 262 |
+
# ── Tiếng Việt: dấu phẩy nghìn + ký hiệu đ ──────────
|
| 263 |
+
# VD: 125000 → "125,000 đ"
|
| 264 |
+
if isinstance(int_num, int):
|
| 265 |
+
formatted = f"{int_num:,}"
|
| 266 |
+
else:
|
| 267 |
+
formatted = f"{int(int_num):,}"
|
| 268 |
+
return f"{formatted} đ"
|
| 269 |
+
else:
|
| 270 |
+
# ── Tiếng Anh: format số thực 2 chữ số thập phân ─────
|
| 271 |
+
if isinstance(int_num, int):
|
| 272 |
+
# Số nguyên >= 100: chèn dấu chấm thập phân 2 từ phải
|
| 273 |
+
# VD: 2345 → 23.45 | 500 → 5.00
|
| 274 |
+
dollars = int_num / 100
|
| 275 |
+
return f"{dollars:.2f}"
|
| 276 |
+
else:
|
| 277 |
+
# Đã có thập phân → chỉ làm tròn 2 chữ số
|
| 278 |
+
return f"{num:.2f}"
|
| 279 |
+
elif num < 1e2 and num >= 0:
|
| 280 |
+
# Số nhỏ (< 100) kể cả tiếng Anh: hiển thị 2 chữ số thập phân
|
| 281 |
+
# VD: 5 → "5.00" | 9.5 → "9.50"
|
| 282 |
+
if not is_vn:
|
| 283 |
+
return f"{num:.2f}"
|
| 284 |
+
else:
|
| 285 |
+
return f"{int(num):,} đ"
|
| 286 |
+
except ValueError:
|
| 287 |
+
pass
|
| 288 |
+
|
| 289 |
+
# Không parse được số — với VN vẫn thêm " đ" nếu chưa có
|
| 290 |
+
raw = val.strip()
|
| 291 |
+
if is_vn and raw and not re.search(r'[đ₫]', raw, re.IGNORECASE):
|
| 292 |
+
return f"{raw} đ"
|
| 293 |
+
return raw
|
| 294 |
+
|
| 295 |
+
|
| 296 |
+
def normalize_address(val) -> str:
|
| 297 |
+
"""Chuẩn hóa địa chỉ: bỏ ký tự điều khiển, chuẩn hóa khoảng trắng."""
|
| 298 |
+
if not val:
|
| 299 |
+
return val
|
| 300 |
+
# Lọc các chuỗi đại diện cho giá trị null/rỗng
|
| 301 |
+
if str(val).strip().lower() in ("null", "none", "n/a", "nan"):
|
| 302 |
+
return None
|
| 303 |
+
val = re.sub(r'[\x00-\x1f\x7f]', ' ', str(val))
|
| 304 |
+
return re.sub(r'\s+', ' ', val).strip()
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
# ─────────────────────────────────────────────────────────────
|
| 308 |
+
# 3. HÀM POSTPROCESS TỔNG HỢP
|
| 309 |
+
# ─────────────────────────────────────────────────────────────
|
| 310 |
+
|
| 311 |
+
def postprocess(pred_obj: dict, domain: str = "sroie") -> dict:
|
| 312 |
+
"""
|
| 313 |
+
Áp dụng chuẩn hóa đầy đủ lên dict đã parse.
|
| 314 |
+
Tương đương Config 6 trong ablation_postproc.py.
|
| 315 |
+
|
| 316 |
+
Parameters
|
| 317 |
+
----------
|
| 318 |
+
pred_obj : dict — kết quả từ tolerant_json_parse
|
| 319 |
+
domain : str — "sroie" | "mcocr" | "vn"
|
| 320 |
+
Ảnh hưởng cách format trường total:
|
| 321 |
+
• mcocr / vn → dấu chấm nghìn + ký hiệu "đ" (5000 → "5.000 đ")
|
| 322 |
+
• sroie → số thuần (5000 → "5000")
|
| 323 |
+
"""
|
| 324 |
+
return {
|
| 325 |
+
"company": (pred_obj.get("company") or "").strip() or None,
|
| 326 |
+
"date" : normalize_date(pred_obj.get("date")),
|
| 327 |
+
"total" : normalize_total(pred_obj.get("total"), domain=domain),
|
| 328 |
+
"address": normalize_address(pred_obj.get("address")),
|
| 329 |
+
}
|
| 330 |
+
|
| 331 |
+
|
| 332 |
+
def full_postprocess(raw_text: str, domain: str = "auto") -> tuple[dict, dict, str]:
|
| 333 |
+
"""
|
| 334 |
+
Pipeline hoàn chỉnh từ raw text → dict đã chuẩn hóa.
|
| 335 |
+
Trả về (parsed_raw, parsed_normalized, detected_domain).
|
| 336 |
+
|
| 337 |
+
Parameters
|
| 338 |
+
----------
|
| 339 |
+
raw_text : str — chuỗi JSON thô từ mô hình
|
| 340 |
+
domain : str — "auto" (mặc định) | "sroie" | "mcocr" | "vn"
|
| 341 |
+
"auto" → tự nhận diện ngôn ngữ từ nội dung parsed
|
| 342 |
+
"""
|
| 343 |
+
parsed_raw = tolerant_json_parse(raw_text)
|
| 344 |
+
|
| 345 |
+
# Tự động nhận diện domain nếu không chỉ định
|
| 346 |
+
if domain == "auto":
|
| 347 |
+
detected = auto_detect_domain(parsed_raw)
|
| 348 |
+
else:
|
| 349 |
+
detected = domain
|
| 350 |
+
|
| 351 |
+
parsed_norm = postprocess(parsed_raw, domain=detected)
|
| 352 |
+
return parsed_raw, parsed_norm, detected
|
| 353 |
+
|
| 354 |
+
|
| 355 |
+
def is_valid_json(text: str) -> bool:
|
| 356 |
+
"""Kiểm tra xem mô hình có sinh ra JSON hợp lệ không."""
|
| 357 |
+
if not text:
|
| 358 |
+
return False
|
| 359 |
+
s, e = text.find("{"), text.rfind("}")
|
| 360 |
+
if s < 0 or e <= s:
|
| 361 |
+
return False
|
| 362 |
+
try:
|
| 363 |
+
json.loads(text[s:e + 1])
|
| 364 |
+
return True
|
| 365 |
+
except Exception:
|
| 366 |
+
return False
|
requirements.txt
CHANGED
|
@@ -1,3 +1,8 @@
|
|
| 1 |
-
|
| 2 |
-
|
| 3 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
torch
|
| 2 |
+
transformers==4.47.1
|
| 3 |
+
peft==0.13.2
|
| 4 |
+
streamlit>=1.39.0
|
| 5 |
+
pillow
|
| 6 |
+
plotly
|
| 7 |
+
huggingface_hub
|
| 8 |
+
numpy
|