Commit 059a13d4 authored by Vũ Hoàng Anh's avatar Vũ Hoàng Anh

Merge feature/ocr-vision into main

parents bbb548b6 451352c1
...@@ -137,8 +137,39 @@ class ImageSearchGraph: ...@@ -137,8 +137,39 @@ class ImageSearchGraph:
if analysis.get("success"): if analysis.get("success"):
feats = analysis["features"] feats = analysis["features"]
extracted_features_text += f"Ảnh {idx+1} [Local Vision AI]: Danh mục: {feats.get('category')}, Màu: {feats.get('color')}, Phong cách: {feats.get('style')}\n" ocr_text = feats.get("ocr_text", "")
ocr_info = f", Chữ in trên áo (OCR): {ocr_text}" if ocr_text else ""
extracted_features_text += f"Ảnh {idx+1} [Local Vision AI + OCR]: Danh mục: {feats.get('category')}, Màu: {feats.get('color')}, Phong cách: {feats.get('style')}{ocr_info}\n"
extracted_features_text += f"Raw tags: {', '.join(feats.get('raw_labels', []))}\n" extracted_features_text += f"Raw tags: {', '.join(feats.get('raw_labels', []))}\n"
if ocr_text:
import sqlite3
from common.constants import SQLITE_DB_PATH
try:
conn = sqlite3.connect(SQLITE_DB_PATH)
conn.row_factory = sqlite3.Row
cursor = conn.cursor()
keywords = [k for k in ocr_text.split() if len(k) > 2]
if keywords:
like_clauses = " OR ".join(["clean_description LIKE ? OR product_name LIKE ?"] * len(keywords))
params = []
for k in keywords:
params.extend([f"%{k}%", f"%{k}%"])
rows = cursor.execute(
f"SELECT product_name, internal_ref_code, product_image_url "
f"FROM pg__dashboard_canifa__ultra_descriptions "
f"WHERE {like_clauses} LIMIT 5",
params
).fetchall()
if rows:
extracted_features_text += f"→ Đã tìm thấy {len(rows)} sản phẩm khớp chữ '{ocr_text}' trong SQLite:\n"
for r in rows:
extracted_features_text += f" - [{r['internal_ref_code']}] {r['product_name']}\n"
conn.close()
except Exception as e:
logger.error(f"Lỗi truy vấn SQLite OCR: {e}")
except Exception as e: except Exception as e:
logger.error(f"Lỗi đọc ảnh cho Local Vision Model: {e}") logger.error(f"Lỗi đọc ảnh cho Local Vision Model: {e}")
......
...@@ -129,41 +129,71 @@ class LocalVisionModel: ...@@ -129,41 +129,71 @@ class LocalVisionModel:
logger.error("Lỗi khởi tạo Vision Model: %s", e) logger.error("Lỗi khởi tạo Vision Model: %s", e)
self._classifier = None self._classifier = None
# Khởi tạo RapidOCR siêu nhẹ
try:
from rapidocr_onnxruntime import RapidOCR
logger.info("Khởi tạo RapidOCR (ONNX) siêu nhẹ cho nhận diện chữ...")
self._ocr = RapidOCR()
logger.info("Đã khởi tạo thành công RapidOCR.")
except ImportError:
logger.warning("Chưa cài đặt RapidOCR. Chạy: uv pip install rapidocr-onnxruntime")
self._ocr = None
def analyze_image(self, image_path: str) -> dict: def analyze_image(self, image_path: str) -> dict:
""" """
Analyze image at given path and extract fashion-specific features. Analyze image at given path and extract fashion-specific features.
Returns: Returns:
dict with keys: success, features, confidence dict with keys: success, features, confidence
features contains: raw_labels, category, color, style, all_categories features contains: raw_labels, category, color, style, all_categories, ocr_text
""" """
if not self._classifier:
return {"error": "Vision model chưa sẵn sàng. Vui lòng kiểm tra logs."}
try: try:
# 1. Load image # 1. Load image
img = Image.open(image_path).convert("RGB") img = Image.open(image_path).convert("RGB")
# 2. Classify via model category = "unknown"
predictions = self._classifier(img) color = "unknown"
style = "casual"
all_categories = []
top_tags = []
confidence = 1.0
# 3. Filter by confidence threshold # 2. Classify via model (if available)
strong_predictions = [ if self._classifier:
p for p in predictions if p["score"] >= MIN_CONFIDENCE try:
] predictions = self._classifier(img)
strong_predictions = [p for p in predictions if p["score"] >= MIN_CONFIDENCE]
if not strong_predictions: if not strong_predictions:
strong_predictions = predictions[:1] # Keep at least top-1 strong_predictions = predictions[:1]
top_tags = [p["label"] for p in strong_predictions[:5]] top_tags = [p["label"] for p in strong_predictions[:5]]
tags_str = " ".join(top_tags).lower() tags_str = " ".join(top_tags).lower()
# 4. Extract features using curated rules
category = self._match_first(tags_str, CATEGORY_RULES, "unknown") category = self._match_first(tags_str, CATEGORY_RULES, "unknown")
color = self._match_first(tags_str, COLOR_RULES, "unknown") color = self._match_first(tags_str, COLOR_RULES, "unknown")
style = self._match_first(tags_str, STYLE_RULES, "casual") style = self._match_first(tags_str, STYLE_RULES, "casual")
# Collect ALL matching categories (for multi-tag search)
all_categories = self._match_all(tags_str, CATEGORY_RULES) all_categories = self._match_all(tags_str, CATEGORY_RULES)
confidence = strong_predictions[0]["score"]
except Exception as e:
logger.error("Lỗi khi phân loại ảnh: %s", e)
# 5. Extract Text via RapidOCR (ONNX CPU)
extracted_text = []
if self._ocr:
try:
import numpy as np
# Convert PIL RGB to OpenCV BGR format
img_cv = np.array(img)
img_cv = img_cv[:, :, ::-1].copy()
ocr_result, _ = self._ocr(img_cv)
if ocr_result:
# ocr_result format: [[box, text, confidence], ...]
for box, text, score in ocr_result:
if score > 0.3: # Độ tin cậy > 30% để lấy chữ in
extracted_text.append(text)
except Exception as e:
logger.error("Lỗi khi chạy RapidOCR: %s", e)
extracted_features = { extracted_features = {
"raw_labels": top_tags, "raw_labels": top_tags,
...@@ -171,12 +201,13 @@ class LocalVisionModel: ...@@ -171,12 +201,13 @@ class LocalVisionModel:
"color": color, "color": color,
"style": style, "style": style,
"all_categories": all_categories, "all_categories": all_categories,
"ocr_text": " ".join(extracted_text)
} }
return { return {
"success": True, "success": True,
"features": extracted_features, "features": extracted_features,
"confidence": strong_predictions[0]["score"], "confidence": confidence,
} }
except Exception as e: except Exception as e:
logger.error("Lỗi khi phân tích ảnh: %s", e) logger.error("Lỗi khi phân tích ảnh: %s", e)
......
import asyncio
from PIL import Image, ImageDraw
import base64
import io
def create_dummy_image_with_text():
# Create a blank black image
img = Image.new('RGB', (300, 300), color = (0, 0, 0))
d = ImageDraw.Draw(img)
# Add text "CANIFA" to the image
d.text((50, 150), "CANIFA", fill=(255, 255, 255))
# Save to base64
buffered = io.BytesIO()
img.save(buffered, format="JPEG")
img_str = base64.b64encode(buffered.getvalue()).decode()
return img_str
async def test():
print("Testing Image Vision OCR with SQLite...")
img_b64 = create_dummy_image_with_text()
from agent.image_search_agent.image_search_graph import get_image_search_agent
agent = get_image_search_agent()
print("Running Image Search Agent...")
result = await agent.chat(user_message="Tìm cho tôi chiếc áo này", images=[img_b64])
print("\n=== AI RESPONSE ===")
print(result.get("response"))
print("===================\n")
print("--- Pipeline Trace ---")
for diag in result.get("pipeline", []):
print(f"Step: {diag.get('label')}")
if diag.get("step") == "user":
print(diag.get("content"))
if __name__ == "__main__":
asyncio.run(test())
\ No newline at end of file
version = 1
revision = 3
requires-python = ">=3.14"
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment