Files
pfm-ocr/backend/config/classify_ocr_server.py
T

593 lines
22 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import base64
import io
import os
import re
import traceback
from datetime import date
from pathlib import Path
import requests
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from PIL import Image
import numpy as np
from ultralytics import YOLO
from paddleocr import PaddleOCR
app = FastAPI(title="PFM Product Classifier and OCR API")
CLASSIFIER_WEIGHTS_GLOB = "produk-pfm-classifier-26n-*e-*.pt"
CLASSIFIER_DATE_IN_NAME = re.compile(
r"produk-pfm-classifier-26n-\d+e-(\d{4}-\d{2}-\d{2})\.pt$"
)
def _classifier_date_from_name(path: Path) -> date | None:
match = CLASSIFIER_DATE_IN_NAME.match(path.name)
if not match:
return None
year, month, day = (int(part) for part in match.group(1).split("-"))
return date(year, month, day)
def latest_classifier_weights(models_dir: Path) -> Path | None:
"""Return the newest produk-pfm-classifier .pt weights in models/."""
if not models_dir.is_dir():
return None
candidates = list(models_dir.glob(CLASSIFIER_WEIGHTS_GLOB))
if not candidates:
return None
def sort_key(path: Path) -> tuple[date, float]:
name_date = _classifier_date_from_name(path) or date.min
return (name_date, path.stat().st_mtime)
return max(candidates, key=sort_key)
def resolve_classifier_models_dir() -> Path | None:
"""Locate produk-pfm/models (env override, repo path, or Docker mount)."""
env_dir = os.environ.get("CLASSIFIER_MODELS_DIR")
if env_dir:
path = Path(env_dir)
if path.is_dir():
return path
repo_root = Path(__file__).resolve().parent.parent
for candidate in (
repo_root / "pfm-web-app/public/produk-pfm/models",
Path("/app/pfm-web-app/public/produk-pfm/models"),
Path(__file__).resolve().parent,
):
if candidate.is_dir():
return candidate
return None
def resolve_classifier_weights_path() -> Path | None:
"""Resolve YOLO weights: CLASSIFIER_MODEL_PATH or newest file in models/."""
explicit = os.environ.get("CLASSIFIER_MODEL_PATH")
if explicit:
path = Path(explicit)
if path.is_file():
return path
print(f"CLASSIFIER_MODEL_PATH not found: {path}")
models_dir = resolve_classifier_models_dir()
if models_dir is None:
return None
weights = latest_classifier_weights(models_dir)
if weights is None:
print(f"No classifier weights matching {CLASSIFIER_WEIGHTS_GLOB} in {models_dir}")
return weights
# Enable CORS
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# Load models on startup
print("Loading YOLO model...")
yolo_model = None
try:
classifier_weights = resolve_classifier_weights_path()
if classifier_weights is None:
raise FileNotFoundError(
"No produk-pfm classifier weights found. Train with train_classifier.py "
"or set CLASSIFIER_MODEL_PATH / CLASSIFIER_MODELS_DIR."
)
print(f"Using classifier weights: {classifier_weights}")
yolo_model = YOLO(str(classifier_weights))
print("YOLO model loaded successfully.")
except Exception as e:
print(f"Error loading YOLO model: {e}")
yolo_model = None
print("Loading PaddleOCR...")
try:
# Use standard textline orientation detection for PaddleOCR 3.x
ocr = PaddleOCR(use_textline_orientation=True, lang='en')
print("PaddleOCR loaded successfully.")
except Exception as e:
print(f"Error loading PaddleOCR: {e}")
ocr = None
class ScanRequest(BaseModel):
image_base64: str
def clean_ocr_text(text: str) -> str:
return re.sub(r'^[^\w\s./-]+|[^\w\s./-]+$', '', text).strip()
EXP_KEYWORD_RE = re.compile(
r'(?:exp(?:\.|ired)?|tgl(?:\s*exp)?|expiry|bbd|best\s*before|before|best|baik\s*digunakan|\bbb\b)',
re.IGNORECASE,
)
DD_MM_YYYY_RE = re.compile(
r'(?<!\d)(0[1-9]|[12]\d|3[01]).*?(0[1-9]|1[0-2]).*?(20\d{2})(?!\d)'
)
DDMMYYYY_RE = re.compile(
r'(?<!\d)(0[1-9]|[12]\d|3[01])(0[1-9]|1[0-2])(20\d{2})(?!\d)'
)
BB_ATTACHED_DATE_RE = re.compile(
r'\b(?:bb|bestbefore)\s*[:.-]?\s*(0[1-9]|[12]\d|3[01])(0[1-9]|1[0-2])(20\d{2})(?!\d)',
re.IGNORECASE,
)
KEYWORD_DIGITS_RE = re.compile(
r'(?:exp|expired|tgl|expiry|bbd|before|best|bb|baik|digunakan)\s*[:.-]?\s*(\d{6,8})\b',
re.IGNORECASE,
)
LENIENT_DATE_RE = re.compile(
r'(?<!\d)(\d{1,2}).*?(\d{1,2}).*?((?:20)?\d{2})(?!\d)'
)
def is_valid_ddmmyyyy_digits(val: str) -> bool:
if len(val) != 8 or not val.isdigit():
return False
day, month, year = int(val[0:2]), int(val[2:4]), int(val[4:8])
return 1 <= day <= 31 and 1 <= month <= 12 and 2000 <= year <= 2099
def format_ddmmyyyy(val: str) -> str:
if is_valid_ddmmyyyy_digits(val):
return f"{val[0:2]}/{val[2:4]}/{val[4:8]}"
return val.upper()
def format_ddmmyy(val: str) -> str:
if len(val) == 6 and val.isdigit():
day, month = int(val[0:2]), int(val[2:4])
if 1 <= day <= 31 and 1 <= month <= 12:
return f"{val[0:2]}/{val[2:4]}/{val[4:6]}"
return val.upper()
def line_has_exp_keyword(line: str) -> bool:
if EXP_KEYWORD_RE.search(line):
return True
# BB05032027 — keyword directly followed by digits
return bool(re.search(r'(?i)\b(?:bb|bestbefore)(?:\s*[:.-]?\s*)?\d', line))
def clean_date_line(line: str) -> str:
# 1) Replace "1)" with "0"
cleaned = line.replace("1)", "0")
# 2) Replace "()" with "0"
cleaned = cleaned.replace("()", "0")
# Clean BB misrecognitions (convert B8, 8B, 88 to BB when followed by digits)
cleaned = re.sub(r'\b(?:B8|8B|88)(?=\d)', 'BB', cleaned)
cleaned = re.sub(r'^(?:B8|8B|88)(?=\d)', 'BB', cleaned)
# Clean 012 month misrecognition (e.g. 020122027 -> 02022027)
cleaned = re.sub(r'(?<!\d)(\d{1,2})012(20\d{2})(?!\d)', r'\g<1>02\g<2>', cleaned)
cleaned = re.sub(r'(?<!\d)(\d{1,2})([-./\s]+)012([-./\s]+)(20\d{2})(?!\d)', r'\g<1>\g<2>02\g<3>\g<4>', cleaned)
# Clean 112 month misrecognition (e.g. 021122027 -> 02022027)
cleaned = re.sub(r'(?<!\d)(\d{1,2})112(20\d{2})(?!\d)', r'\g<1>02\g<2>', cleaned)
cleaned = re.sub(r'(?<!\d)(\d{1,2})([-./\s]+)112([-./\s]+)(20\d{2})(?!\d)', r'\g<1>\g<2>02\g<3>\g<4>', cleaned)
# Run contextual replacements
for _ in range(3):
# letter o/O flanked by digits or boundary -> 0
cleaned = re.compile(r'(\d)[oO](\d|\b)').sub(r'\g<1>0\g<2>', cleaned)
cleaned = re.compile(r'(\b|\d)[oO](\d)').sub(r'\g<1>0\g<2>', cleaned)
# letter I/i/l/| flanked by digits -> 1
cleaned = re.compile(r'(\d)[Ii|l](\d|\b)').sub(r'\g<1>1\g<2>', cleaned)
cleaned = re.compile(r'(\b|\d)[Ii|l](\d)').sub(r'\g<1>1\g<2>', cleaned)
# letter S/s flanked by digits -> 5
cleaned = re.compile(r'(\d)[Ss](\d|\b)').sub(r'\g<1>5\g<2>', cleaned)
cleaned = re.compile(r'(\b|\d)[Ss](\d)').sub(r'\g<1>5\g<2>', cleaned)
# letter Z/z flanked by digits -> 2
cleaned = re.compile(r'(\d)[Zz](\d|\b)').sub(r'\g<1>2\g<2>', cleaned)
cleaned = re.compile(r'(\b|\d)[Zz](\d)').sub(r'\g<1>2\g<2>', cleaned)
# letter B flanked by digits -> 8
cleaned = re.compile(r'(\d)B(\d|\b)').sub(r'\g<1>8\g<2>', cleaned)
cleaned = re.compile(r'(\b|\d)B(\d)').sub(r'\g<1>8\g<2>', cleaned)
return cleaned
def extract_expired_date(text_lines):
"""Return (formatted_date, line_index, source_line). Prioritises BB/EXP + DDMMYYYY or DD MM YYYY."""
if not text_lines:
return None, None, None
cleaned_lines = [clean_date_line(line) for line in text_lines]
def pick(match, idx, cleaned_line, formatter=None):
raw = match.group(0)
original_line = text_lines[idx].strip()
if match.lastindex and match.lastindex >= 3:
formatted = f"{match.group(1)}/{match.group(2)}/{match.group(3)}"
elif match.lastindex and match.lastindex >= 1 and match.group(1).isdigit():
digits = match.group(1)
if len(digits) == 8:
formatted = format_ddmmyyyy(digits)
elif len(digits) == 6:
formatted = format_ddmmyy(digits)
else:
formatted = digits
elif formatter:
formatted = formatter(raw)
else:
formatted = raw.strip().upper()
return formatted, idx, original_line
# 1) BB/EXP keyword lines — compact DDMMYYYY (e.g. BB05032027, EXP 05032027)
for idx, line in enumerate(cleaned_lines):
if not line_has_exp_keyword(line):
continue
match = BB_ATTACHED_DATE_RE.search(line) or DDMMYYYY_RE.search(line)
if match:
return pick(match, idx, line)
# 2) BB/EXP keyword lines — spaced DD MM YYYY (e.g. BB 05 03 2027)
for idx, line in enumerate(cleaned_lines):
if not line_has_exp_keyword(line):
continue
match = DD_MM_YYYY_RE.search(line)
if match:
return pick(match, idx, line)
# 3) Keyword + 6–8 digit run (BB05032027 via keyword_digits)
for idx, line in enumerate(cleaned_lines):
match = KEYWORD_DIGITS_RE.search(line)
if match:
digits = match.group(1)
if len(digits) == 8 and is_valid_ddmmyyyy_digits(digits):
return format_ddmmyyyy(digits), idx, text_lines[idx].strip()
if len(digits) == 6:
return format_ddmmyy(digits), idx, text_lines[idx].strip()
# 3.5) BB/EXP keyword lines — lenient check for unclear/noisy date formats (e.g. BB 02J 132027)
for idx, line in enumerate(cleaned_lines):
if not line_has_exp_keyword(line):
continue
match = LENIENT_DATE_RE.search(line)
if match:
return pick(match, idx, line)
# 4) Any line — spaced DD MM YYYY
for idx, line in enumerate(cleaned_lines):
match = DD_MM_YYYY_RE.search(line)
if match:
return pick(match, idx, line)
# 5) Any line — compact DDMMYYYY (skip likely SKU: same line has 8-digit product code context)
for idx, line in enumerate(cleaned_lines):
for match in DDMMYYYY_RE.finditer(line):
digits = f"{match.group(1)}{match.group(2)}{match.group(3)}"
if is_valid_ddmmyyyy_digits(digits):
# Skip if this 8-digit block is the only digits and looks like SKU on label top
if re.search(r'\b\d{8}\b', line) and not line_has_exp_keyword(line):
if re.search(r'(?:nugget|chicken|fiesta|champ|okey|akumo|frozen|gr)', line, re.I):
continue
return format_ddmmyyyy(digits), idx, text_lines[idx].strip()
# 6) Legacy patterns (slashes, month names, etc.)
date_patterns = [
r'\b\d{2}[-./]\d{2}[-./]\d{2,4}\b',
r'\b\d{4}[-./]\d{2}[-./]\d{2}\b',
r'\b\d{2}\s+(?:JAN|FEB|MAR|APR|MAY|JUN|JUL|AUG|SEP|OCT|NOV|DEC)[a-zA-Z]*\s+\d{2,4}\b',
]
for idx, line in enumerate(cleaned_lines):
if not line_has_exp_keyword(line):
continue
for pat in date_patterns:
match = re.search(pat, line, re.IGNORECASE)
if match:
return match.group(0).upper(), idx, text_lines[idx].strip()
return None, None, None
def line_contains_expired_date(line: str, expired_date: str) -> bool:
if not line or not expired_date:
return False
digits_only = re.sub(r"\D", "", expired_date)
line_digits = re.sub(r"\D", "", line)
if len(digits_only) >= 6 and digits_only in line_digits:
return True
compact = expired_date.replace("/", "")
return compact in line.replace(" ", "") or expired_date in line
def find_expired_crop_index(text_lines, expired_idx, expired_date, polys_len):
"""Pick OCR box index for cropping; prefer the line that actually contains the date."""
if not expired_date or polys_len <= 0:
return None
if (
expired_idx is not None
and expired_idx < polys_len
and expired_idx < len(text_lines)
and line_contains_expired_date(text_lines[expired_idx], expired_date)
):
return expired_idx
keyword_match = None
for idx, line in enumerate(text_lines):
if idx >= polys_len:
break
if not line_contains_expired_date(line, expired_date):
continue
if line_has_exp_keyword(line):
return idx
if keyword_match is None:
keyword_match = idx
if keyword_match is not None:
return keyword_match
if expired_idx is not None and expired_idx < polys_len:
return expired_idx
return None
def ocr_coordinate_image(res_entry, fallback_image: Image.Image) -> Image.Image:
"""Image in the same pixel space as rec_polys (after doc orientation + unwarping)."""
dpr = res_entry.get("doc_preprocessor_res") or {}
output_arr = dpr.get("output_img")
if output_arr is not None:
return Image.fromarray(np.asarray(output_arr)).convert("RGB")
return fallback_image
def ocr_text_polys(res_entry):
"""Recognition polygons — 1:1 aligned with rec_texts."""
return res_entry.get("rec_polys") or res_entry.get("dt_polys") or []
def extract_sku(text_lines):
# SKU is usually an 8-digit number (e.g. 12010111)
for line in text_lines:
match = re.search(r'\b(\d{8})\b', line)
if match:
return match.group(1)
# Try finding 7-9 digit numbers
for line in text_lines:
match = re.search(r'\b(\d{7,9})\b', line)
if match:
return match.group(1)
return None
def crop_poly_region(image, poly, padding=8):
x_coords = [float(p[0]) for p in poly]
y_coords = [float(p[1]) for p in poly]
x_min = max(0, int(min(x_coords)))
y_min = max(0, int(min(y_coords)))
x_max = min(image.width, int(max(x_coords)))
y_max = min(image.height, int(max(y_coords)))
crop_left = max(0, x_min - padding)
crop_top = max(0, y_min - padding)
crop_right = min(image.width, x_max + padding)
crop_bottom = min(image.height, y_max + padding)
if crop_right <= crop_left or crop_bottom <= crop_top:
return None
cropped = image.crop((crop_left, crop_top, crop_right, crop_bottom))
buffered = io.BytesIO()
cropped.save(buffered, format="JPEG")
return "data:image/jpeg;base64," + base64.b64encode(buffered.getvalue()).decode("utf-8")
def extract_product_name(text_lines, classified_name=None):
keywords = ['NUGGET', 'CHICKEN', 'CHAMP', 'FIESTA', 'AKUMO', 'ASIMO', 'OKEY', 'FRIES', 'BURGER', 'SAUSAGE', 'SOSIS', 'KARAGE', 'SIOMAY', 'BUMBU', 'RACIK']
matches = []
for line in text_lines:
line_upper = line.upper()
if any(kw in line_upper for kw in keywords):
cleaned = re.sub(r'\b\d{8}\b', '', line)
cleaned = re.sub(r'(?:exp|expired|tgl|expiry|bbd|before)[^ \n]*', '', cleaned, flags=re.IGNORECASE)
cleaned = re.sub(r'\b\d{2}[-./]\d{2}[-./]\d{2,4}\b', '', cleaned)
cleaned = cleaned.strip()
if len(cleaned) > 3:
matches.append(cleaned)
if matches:
return max(matches, key=len)
if classified_name:
return classified_name
candidate_lines = [l for l in text_lines if not re.match(r'^\d+$', l) and len(l) > 3]
if candidate_lines:
return max(candidate_lines, key=len)
return "Unknown Product"
@app.post("/classify-ocr")
async def classify_ocr(payload: ScanRequest):
try:
# Decode image
from PIL import ImageOps
img_data = base64.b64decode(payload.image_base64.split(",")[-1])
raw_image = Image.open(io.BytesIO(img_data))
image = ImageOps.exif_transpose(raw_image).convert("RGB")
# 1. Run YOLO Classification
classification_result = {}
top1_name = None
if yolo_model:
results = yolo_model(image)
probs = results[0].probs
top1_idx = probs.top1
top1_conf = float(probs.top1conf)
top1_name = results[0].names[top1_idx]
all_probs = []
for idx, val in enumerate(probs.data):
all_probs.append({
"name": results[0].names[idx],
"confidence": float(val)
})
all_probs.sort(key=lambda x: x["confidence"], reverse=True)
classification_result = {
"top1_name": top1_name,
"top1_confidence": top1_conf,
"all_probabilities": all_probs
}
else:
classification_result = {
"error": "YOLO model not loaded"
}
# 2. Run PaddleOCR
ocr_result = {}
if ocr:
img_arr = np.array(image)
# Use predict method and convert generator to list
res_list = list(ocr.predict(img_arr))
text_lines = []
if res_list and len(res_list) > 0:
text_lines = res_list[0].get("rec_texts", [])
sku = extract_sku(text_lines)
expired_date, expired_idx, expired_source_line = extract_expired_date(text_lines)
product_name = extract_product_name(text_lines, top1_name)
res_entry = res_list[0] if res_list else {}
coord_image = ocr_coordinate_image(res_entry, image)
text_polys = ocr_text_polys(res_entry)
crop_idx = find_expired_crop_index(
text_lines, expired_idx, expired_date, len(text_polys)
)
# Create visual OCR image with bounding boxes
vis_image_b64 = None
try:
vis_image = coord_image.copy()
from PIL import ImageDraw, ImageFont
draw = ImageDraw.Draw(vis_image)
try:
font = ImageFont.load_default()
except:
font = None
for idx, poly in enumerate(text_polys):
is_expired = (crop_idx is not None and idx == crop_idx)
pts = [(float(p[0]), float(p[1])) for p in poly]
if is_expired:
color = (245, 158, 11) # Amber
label = "EXP"
else:
color = (13, 148, 136) # Teal
label = "TEXT"
draw.polygon(pts, outline=color, width=3)
x0, y0 = pts[0]
label_w = 32 if label == "EXP" else 38
draw.rectangle([x0, y0 - 15, x0 + label_w, y0], fill=color)
if font:
draw.text((x0 + 4, y0 - 14), label, fill=(255, 255, 255), font=font)
else:
draw.text((x0 + 4, y0 - 14), label, fill=(255, 255, 255))
buffered = io.BytesIO()
vis_image.save(buffered, format="JPEG")
vis_image_b64 = "data:image/jpeg;base64," + base64.b64encode(buffered.getvalue()).decode("utf-8")
except Exception as draw_err:
print(f"Error drawing visual OCR: {draw_err}")
traceback.print_exc()
# Crop expired date OCR region for summary verification
expired_date_crop_b64 = None
try:
if crop_idx is not None and crop_idx < len(text_polys):
expired_date_crop_b64 = crop_poly_region(coord_image, text_polys[crop_idx])
except Exception as crop_err:
print(f"Error cropping expired date image: {crop_err}")
traceback.print_exc()
# 3. Call Spotting API
spotting_image_b64 = None
try:
img_b64_only = payload.image_base64.split(",")[-1]
spotting_payload = {
"file": img_b64_only,
"matchHistoryJob": False,
"useLayoutDetection": False,
"fileType": 1,
"useDocUnwarping": False,
"useDocOrientationClassify": False,
"promptLabel": "spotting"
}
spotting_url = "http://localhost:8090/layout-parsing"
spotting_resp = requests.post(spotting_url, json=spotting_payload, timeout=60)
if spotting_resp.status_code == 200:
spotting_data = spotting_resp.json()
if spotting_data.get("errorCode") == 0:
layout_results = spotting_data.get("result", {}).get("layoutParsingResults", [])
if layout_results:
page0 = layout_results[0]
out_imgs = page0.get("outputImages", {})
spotting_img = out_imgs.get("spotting_res_img")
if spotting_img:
spotting_image_b64 = "data:image/jpeg;base64," + spotting_img
else:
print(f"Spotting API error: {spotting_resp.text}")
except Exception as spotting_err:
print(f"Error calling spotting API: {spotting_err}")
traceback.print_exc()
ocr_result = {
"text_lines": text_lines,
"extracted_product_name": product_name,
"extracted_sku": sku,
"extracted_expired_date": expired_date,
"expired_line_index": crop_idx,
"expired_source_line": expired_source_line,
"expired_date_crop_base64": expired_date_crop_b64,
"vis_image_base64": vis_image_b64,
"spotting_image_base64": spotting_image_b64
}
else:
ocr_result = {
"error": "PaddleOCR not loaded"
}
return {
"classification": classification_result,
"ocr": ocr_result
}
except Exception as e:
traceback.print_exc()
raise HTTPException(status_code=500, detail=str(e))
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8120)