Files
pfm-ocr/backend/config/classify_ocr_server.py
T
fhanyuh caf8e98378 chore: normalize line endings (CRLF -> LF)
No content changes: git diff --ignore-all-space over these files is empty.
The churn came from editing on Windows against a repo checked out with LF.
2026-08-27 10:40:49 +07:00

607 lines
26 KiB
Python

import base64
import io
import math
import os
import re
import traceback
import pickle
from datetime import date
from pathlib import Path
import torch
from torchvision import transforms
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=["*"],
)
# DINOv2 Image preprocessing
DINOV2_TRANSFORMS = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
# Load DINOv2 Vector Similarity Search Model on startup
print("Loading DINOv2 for similarity search...")
dinov2_model = None
dinov2_index = None
models_dir = resolve_classifier_models_dir()
dinov2_index_path = models_dir / "dinov2_index.pkl" if models_dir else None
if dinov2_index_path and dinov2_index_path.is_file():
try:
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using DINOv2 on device: {device}")
# Load the index
print(f"Loading DINOv2 index from {dinov2_index_path}...")
with open(dinov2_index_path, "rb") as f:
dinov2_index = pickle.load(f)
print(f"DINOv2 index loaded with {len(dinov2_index['metadata'])} reference images.")
# Load DINOv2 model
dinov2_model = torch.hub.load("facebookresearch/dinov2", "dinov2_vits14").to(device)
dinov2_model.eval()
print("DINOv2 model loaded successfully.")
except Exception as e:
print(f"Error loading DINOv2 model or index: {e}")
dinov2_model = None
dinov2_index = None
else:
print(f"DINOv2 index not found at {dinov2_index_path}. DINOv2 search disabled.")
# 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()
# Expiry-date extraction cascade lives in date_extract.py (same dir) so it
# can be offline-tested without loading models.
from date_extract import (
clean_date_line,
extract_expired_date,
find_expired_crop_index,
line_has_exp_keyword,
)
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")
# Classification always sees the original upright orientation - the
# 90-degree expiry-date search below may rotate `image` to a
# sideways/upside-down orientation that DINOv2/YOLO were never
# trained on (their reference photos are all shot upright), so using
# a rotated frame there would hurt classification, not help it.
classification_image = image
# Multi-orientation expiry-date search: some photos are captured with
# the whole frame rotated ~90 degrees from upright (e.g. staff held
# the phone in portrait for a package whose printed date runs
# horizontally), so the expiry stamp - and the product framing -
# ends up sideways. Try 0/90/180/270 degree rotations in order and
# stop at the first one where PaddleOCR actually finds an expiry
# date; if none of the four find one, fall back to the 0-degree
# result so behaviour for genuinely-undetectable photos is unchanged.
# This costs extra OCR passes (up to 4x) only on images where the
# first pass found nothing - already-working images stay on the fast
# single-pass path below.
rotated_image_used = False
res_list = []
text_lines = []
text_polys = []
expired_date = None
expired_idx = None
expired_source_line = None
if ocr:
base_image = image
for step_angle in (0, 90, 180, 270):
try:
candidate_image = (
base_image.rotate(step_angle, resample=Image.BICUBIC, expand=True)
if step_angle else base_image
)
img_arr = np.array(candidate_image)
candidate_res_list = list(ocr.predict(img_arr))
candidate_res_entry = candidate_res_list[0] if candidate_res_list else {}
candidate_text_lines = candidate_res_entry.get("rec_texts", [])
candidate_text_polys = ocr_text_polys(candidate_res_entry)
candidate_expired_date, candidate_expired_idx, candidate_expired_source_line = (
extract_expired_date(candidate_text_lines)
)
if step_angle == 0:
# Always keep the 0-degree pass as the fallback result.
image, res_list, text_lines, text_polys = (
candidate_image, candidate_res_list, candidate_text_lines, candidate_text_polys
)
expired_date, expired_idx, expired_source_line = (
candidate_expired_date, candidate_expired_idx, candidate_expired_source_line
)
if candidate_expired_date is not None:
if step_angle != 0:
print(f"[Auto-Rotate-90] Expiry date found after rotating {step_angle} degrees.")
image, res_list, text_lines, text_polys = (
candidate_image, candidate_res_list, candidate_text_lines, candidate_text_polys
)
rotated_image_used = True
expired_date, expired_idx, expired_source_line = (
candidate_expired_date, candidate_expired_idx, candidate_expired_source_line
)
break
except Exception as rot_err:
print(f"Error during {step_angle}-degree OCR pass: {rot_err}")
traceback.print_exc()
# Tiled full-resolution pass: PaddleOCR downscales anything over
# its 4000px max_side_limit, which is exactly what kills small
# inkjet date stamps on these ~3200x5700 phone photos. Split the
# original image into overlapping tiles that each fit under the
# limit (so the date region is OCR'd at native resolution) and
# run the cascade per tile. Failure-path only, keyword-anchored
# acceptance like the VL fallback below.
if expired_date is None and max(base_image.size) > 2600:
TILE, OVERLAP = 2400, 400
W, H = base_image.size
step = TILE - OVERLAP
try:
found = False
for y0 in range(0, H, step):
if found:
break
for x0 in range(0, W, step):
tile = base_image.crop((x0, y0, min(x0 + TILE, W), min(y0 + TILE, H)))
if tile.width < 300 or tile.height < 300:
continue
tile_res = list(ocr.predict(np.array(tile)))
tile_lines = tile_res[0].get("rec_texts", []) if tile_res else []
if not tile_lines:
continue
t_date, _t_idx, t_source = extract_expired_date(tile_lines)
if t_date is not None and t_source and line_has_exp_keyword(
clean_date_line(t_source)
):
print(f"[Tile-Pass] Expiry date {t_date} found in full-res tile ({x0},{y0}) (line: {t_source!r})")
expired_date = t_date
expired_idx = None # tile polys don't map to the full image
expired_source_line = t_source
found = True
break
except Exception as tile_err:
print(f"[Tile-Pass] failed: {tile_err}")
traceback.print_exc()
# VL fallback: the lightweight PP-OCRv6 detector missed the date
# at every orientation. The vLLM-backed PaddleOCR-VL pipeline
# (:8090, same container) is a much stronger reader of small,
# low-contrast inkjet codes - ask it to read the whole package
# and run the same date cascade over its text output. Only fires
# on already-failed images, so the happy path stays single-pass.
# Acceptance is stricter than the local cascade: the matched
# line must carry an expiry keyword (BB/EXP/Baik digunakan...),
# so a bare number elsewhere on the package can't be
# hallucinated into a date on photos where none is visible.
# Even when no date is found, the VL's (much cleaner) text lines
# are kept and appended to text_lines below - they feed the
# gateway's OCR-evidence classification re-ranking.
vl_text_lines = []
if expired_date is None:
try:
vl_url = os.environ.get(
"VL_PIPELINE_URL", "http://localhost:8090/layout-parsing"
)
buffered = io.BytesIO()
base_image.save(buffered, format="JPEG")
vl_payload = {
"file": base64.b64encode(buffered.getvalue()).decode("utf-8"),
"matchHistoryJob": False,
"useLayoutDetection": True,
"fileType": 1,
"useDocUnwarping": False,
"useDocOrientationClassify": True,
}
vl_resp = requests.post(vl_url, json=vl_payload, timeout=120)
if vl_resp.status_code == 200:
vl_data = vl_resp.json()
if vl_data.get("errorCode") == 0:
layout_results = vl_data.get("result", {}).get("layoutParsingResults", [])
md_text = ""
if layout_results:
md_text = (layout_results[0].get("markdown") or {}).get("text", "") or ""
vl_lines = [ln.strip() for ln in md_text.splitlines() if ln.strip()]
vl_text_lines = vl_lines
if vl_lines:
vl_date, vl_idx, vl_source_line = extract_expired_date(vl_lines)
if vl_date is not None and vl_source_line and line_has_exp_keyword(
clean_date_line(vl_source_line)
):
print(f"[VL-Fallback] Expiry date {vl_date} found by VL pipeline (line: {vl_source_line!r})")
expired_date = vl_date
expired_idx = None # no OCR polys for VL text; skip crop
expired_source_line = vl_source_line
else:
print(f"[VL-Fallback] pipeline error: {vl_resp.status_code} {vl_resp.text[:200]}")
except Exception as vl_err:
print(f"[VL-Fallback] failed: {vl_err}")
traceback.print_exc()
# Fine tilt-straighten correction (<90 degrees), applied on top of
# whichever 90-degree orientation the search above landed on.
try:
if expired_idx is not None and expired_idx < len(text_polys):
poly = text_polys[expired_idx]
if len(poly) >= 2:
p0 = poly[0]
p1 = poly[1]
dx = float(p1[0]) - float(p0[0])
dy = float(p1[1]) - float(p0[1])
angle_rad = math.atan2(dy, dx)
angle_deg = math.degrees(angle_rad)
if abs(angle_deg) > 3.0:
print(f"[Auto-Rotate] Detected Expiry Date text line angle: {angle_deg:.2f} degrees. Rotating image...")
image = image.rotate(angle_deg, resample=Image.BICUBIC, expand=True)
rotated_image_used = True
except Exception as pre_ocr_err:
print(f"Error in fine tilt-straighten pass: {pre_ocr_err}")
traceback.print_exc()
# 1. Run DINOv2 Similarity Search or YOLO Classification
classification_result = {}
top1_name = None
# Try DINOv2 first if index exists
if dinov2_model and dinov2_index:
try:
device = "cuda" if torch.cuda.is_available() else "cpu"
# Preprocess image
preprocessed = DINOV2_TRANSFORMS(classification_image).unsqueeze(0).to(device)
# Extract query embedding
with torch.no_grad():
query_emb = dinov2_model(preprocessed)
query_emb = query_emb / query_emb.norm(dim=-1, keepdim=True)
query_emb = query_emb.squeeze(0).cpu().numpy()
ref_embeddings = dinov2_index["embeddings"] # (N, 384)
ref_metadata = dinov2_index["metadata"] # List of dicts
# Compute Cosine Similarity
similarities = np.dot(ref_embeddings, query_emb)
# Aggregate similarities per class (max similarity of any reference image in the class)
class_sims = {}
for idx, meta in enumerate(ref_metadata):
c_name = meta["class_name"]
sim = float(similarities[idx])
if c_name not in class_sims or sim > class_sims[c_name]:
class_sims[c_name] = sim
# Sort classes by similarity
sorted_classes = sorted(class_sims.items(), key=lambda x: x[1], reverse=True)
all_probs = []
for c_name, sim in sorted_classes:
all_probs.append({
"name": c_name,
"confidence": sim
})
if all_probs:
top1_name = all_probs[0]["name"]
top1_conf = all_probs[0]["confidence"]
classification_result = {
"top1_name": top1_name,
"top1_confidence": top1_conf,
"all_probabilities": all_probs,
"method": "dinov2_similarity"
}
print(f"[DINOv2] Best match: {top1_name} ({top1_conf:.4f})")
except Exception as dinov2_err:
print(f"[DINOv2 Error] Similarity search failed, falling back to YOLO: {dinov2_err}")
traceback.print_exc()
top1_name = None
# Fallback to YOLO if DINOv2 was not run or failed
if not top1_name:
if yolo_model:
results = yolo_model(classification_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,
"method": "yolo_classifier"
}
print(f"[YOLO Fallback] Best match: {top1_name} ({top1_conf:.4f})")
else:
classification_result = {
"error": "Both DINOv2 index and YOLO model are unavailable"
}
# 2. Run PaddleOCR
ocr_result = {}
if ocr:
if rotated_image_used:
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)
)
else:
sku = extract_sku(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)
crop_idx = find_expired_crop_index(
text_lines, expired_idx, expired_date, len(text_polys)
)
# Merge the VL pipeline's text lines (when its fallback ran) into
# the returned text_lines: the gateway's classification re-ranking
# feeds on them, and they're much cleaner than local OCR on hard
# photos. Appended after all poly-aligned work above, so rec_polys
# indexing is unaffected. Also retry SKU extraction over them -
# a VL-read 8-digit SKU enables the gateway's exact-match pin.
if vl_text_lines:
text_lines = list(text_lines) + vl_text_lines
if not sku:
sku = extract_sku(vl_text_lines)
if sku:
print(f"[VL-Fallback] SKU {sku} extracted from VL text lines.")
# 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()
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
}
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)