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)