Spaces:
Running
Running
File size: 8,660 Bytes
cdf3b6a 0dcad07 cdf3b6a 0dcad07 cdf3b6a 0dcad07 cdf3b6a 0dcad07 cdf3b6a 0dcad07 cdf3b6a 0dcad07 cdf3b6a 0dcad07 cdf3b6a 0dcad07 cdf3b6a 0dcad07 cdf3b6a 0dcad07 cdf3b6a b149d41 cdf3b6a b149d41 cdf3b6a b149d41 cdf3b6a 0dcad07 b149d41 cdf3b6a 0dcad07 cdf3b6a b149d41 cdf3b6a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 | import os
import io
import json
import base64
import datetime
import logging
from functools import lru_cache
from fastapi import FastAPI, File, UploadFile, Form, HTTPException
from fastapi.responses import HTMLResponse, JSONResponse
from PIL import Image, UnidentifiedImageError
from ultralytics import YOLO
from huggingface_hub import HfApi
# ──────────────────────────────────────────────
# Config & logging
# ──────────────────────────────────────────────
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("afb")
MODEL_REPO = os.environ.get("MODEL_REPO", "") # เช่น "Teeradon/AFB-Detect-YOLOv11"
MODEL_FILE = os.environ.get("MODEL_FILE", "best.pt")
LOCAL_MODEL = os.environ.get("LOCAL_MODEL", "best.pt")
CONF_THRESHOLD = float(os.environ.get("CONF_THRESHOLD", "0.15"))
IMGSZ = int(os.environ.get("IMGSZ", "1536")) # ขนาดภาพตอน inference (ปรับได้ผ่าน env)
MAX_FILE_MB = 10
ALLOWED_TYPES = {"image/jpeg", "image/png", "image/jpg", "image/webp"}
# ── ที่เก็บภาพ: Hugging Face Dataset ──
HF_TOKEN = os.environ.get("HF_TOKEN")
DATASET_REPO = os.environ.get("DATASET_REPO", "") # เช่น "Teeradon/afb-research-images"
app = FastAPI(title="AFB Detection")
# ──────────────────────────────────────────────
# Startup validation
# ──────────────────────────────────────────────
@app.on_event("startup")
def validate_config():
problems = []
if not HF_TOKEN:
problems.append("HF_TOKEN ไม่ได้ตั้งค่า")
if not DATASET_REPO:
problems.append("DATASET_REPO ไม่ได้ตั้งค่า (ที่เก็บภาพ)")
if problems:
for p in problems:
logger.warning("CONFIG WARNING: %s", p)
else:
logger.info("Config validated OK")
# ──────────────────────────────────────────────
# Lazy-loaded singletons
# ──────────────────────────────────────────────
@lru_cache(maxsize=1)
def get_model():
"""โหลด YOLO model ครั้งเดียว แล้ว cache ไว้"""
if MODEL_REPO:
from huggingface_hub import hf_hub_download
path = hf_hub_download(repo_id=MODEL_REPO, filename=MODEL_FILE, token=HF_TOKEN)
else:
path = LOCAL_MODEL
logger.info("Loading YOLO model from %s", path)
return YOLO(path)
@lru_cache(maxsize=1)
def get_hf_api():
"""สร้าง HfApi client ครั้งเดียว แล้ว cache ไว้"""
if not HF_TOKEN:
raise RuntimeError("HF_TOKEN ไม่ได้ตั้งค่า")
return HfApi(token=HF_TOKEN)
def upload_to_dataset(img_bytes: bytes, path_in_repo: str):
"""อัปโหลดภาพขึ้น HF Dataset repo"""
api = get_hf_api()
api.upload_file(
path_or_fileobj=io.BytesIO(img_bytes),
path_in_repo=path_in_repo,
repo_id=DATASET_REPO,
repo_type="dataset",
)
# ──────────────────────────────────────────────
# Routes — pages
# ──────────────────────────────────────────────
def read_html(name: str) -> str:
with open(name, "r", encoding="utf-8") as f:
return f.read()
@app.get("/", response_class=HTMLResponse)
async def index():
return read_html("index.html")
@app.get("/result", response_class=HTMLResponse)
async def result_page():
return read_html("result.html")
@app.get("/health")
async def health():
return {"status": "ok"}
# ──────────────────────────────────────────────
# Route — analyze
# ──────────────────────────────────────────────
@app.post("/analyze")
async def analyze(
file: UploadFile = File(...),
consent: str = Form("true"),
sample_ref: str = Form(""),
):
# 1) ตรวจ consent
if consent.lower() != "true":
raise HTTPException(status_code=400, detail="ต้องยินยอมก่อนจึงจะวิเคราะห์ได้")
# 2) ตรวจชนิดไฟล์
if file.content_type not in ALLOWED_TYPES:
raise HTTPException(
status_code=415,
detail=f"ชนิดไฟล์ไม่รองรับ: {file.content_type} (รองรับ JPG, PNG, WEBP)",
)
# 3) อ่าน + ตรวจขนาด
img_bytes = await file.read()
if len(img_bytes) > MAX_FILE_MB * 1024 * 1024:
raise HTTPException(status_code=413, detail=f"ไฟล์ใหญ่เกิน {MAX_FILE_MB}MB")
# 4) เปิดภาพ
try:
img = Image.open(io.BytesIO(img_bytes)).convert("RGB")
except (UnidentifiedImageError, OSError):
raise HTTPException(status_code=422, detail="ไม่สามารถเปิดภาพได้ (ไฟล์อาจเสียหาย)")
# 5) รัน inference
try:
model = get_model()
results = model(img, conf=CONF_THRESHOLD, imgsz=IMGSZ, verbose=False)
result = results[0]
except Exception:
logger.exception("Inference failed")
raise HTTPException(status_code=500, detail="การวิเคราะห์ล้มเหลว กรุณาลองใหม่")
# 6) สรุปผล + เก็บ raw boxes ทุกอันไว้ให้ frontend กรองเอง
img_w, img_h = img.size
all_boxes = []
if len(result.boxes) > 0:
xyxy = result.boxes.xyxy.tolist()
confs_raw = result.boxes.conf.tolist()
for (x1, y1, x2, y2), c in zip(xyxy, confs_raw):
all_boxes.append({
"x1": round(x1 / img_w, 4),
"y1": round(y1 / img_h, 4),
"x2": round(x2 / img_w, 4),
"y2": round(y2 / img_h, 4),
"conf": round(c, 4),
})
all_boxes.sort(key=lambda b: b["conf"], reverse=True)
boxes_at_default = [b for b in all_boxes if b["conf"] >= CONF_THRESHOLD]
afb_count = len(boxes_at_default)
avg_conf = round(sum(b["conf"] for b in boxes_at_default) / afb_count * 100, 1) if afb_count > 0 else 0.0
# 7) ส่งภาพต้นฉบับ (ไม่มี box) — frontend วาด box เองจาก all_boxes
buf = io.BytesIO()
img.save(buf, format="JPEG", quality=90)
orig_jpeg = buf.getvalue()
orig_b64 = base64.b64encode(orig_jpeg).decode()
# 8) บันทึกขึ้น HF Dataset — error ไม่ทำให้ทั้ง request ล้ม
saved = False
save_error = None
try:
now = datetime.datetime.now()
ts = now.strftime("%Y%m%d_%H%M%S")
date_folder = now.strftime("%Y-%m-%d")
ref = (sample_ref or "noref").replace("/", "-").replace(" ", "_")[:40]
base_name = f"images/{date_folder}/afb_{ts}_{ref}_n{afb_count}_c{avg_conf}"
upload_to_dataset(orig_jpeg, f"{base_name}_original.jpg")
try:
det_img = Image.fromarray(result.plot()[..., ::-1])
det_buf = io.BytesIO()
det_img.save(det_buf, format="JPEG", quality=90)
upload_to_dataset(det_buf.getvalue(), f"{base_name}_detection.jpg")
except Exception:
pass
saved = True
except Exception:
logger.exception("Dataset upload failed")
save_error = "บันทึกภาพไม่สำเร็จ (ผลการวิเคราะห์ยังแสดงได้ปกติ)"
return JSONResponse({
"afb_count": afb_count,
"avg_conf": avg_conf,
"orig_b64": orig_b64,
"all_boxes": all_boxes,
"conf_threshold": CONF_THRESHOLD,
"saved": saved,
"save_error": save_error,
"sample_ref": sample_ref or "",
}) |