File size: 5,134 Bytes
c71be75
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""โหลดโมเดลและให้บริการ inference — เวอร์ชัน Hugging Face Space

ต่างจาก backend/inference.py (ตัวสำหรับรันบนเครื่อง/Render) ตรงที่:
- ไม่มีไฟล์โมเดลติดมากับ image (Space repo เบา ไม่เก็บไฟล์ 400MB+)
- ดาวน์โหลดโมเดลทั้ง 3 ตัวจาก Hugging Face Model Hub ตอน container start
  (repo กำหนดผ่าน env HF_MODEL_REPO) แล้วเก็บไว้ที่ /app/models
"""

import json
import os
import time
from pathlib import Path

from preprocess import clean_text

MODELS_DIR = Path(__file__).resolve().parent / "models"
HF_MODEL_REPO = os.environ.get("HF_MODEL_REPO", "")


def ensure_models_downloaded() -> None:
    """ดาวน์โหลดโมเดลจาก HF Hub ถ้ายังไม่มีในเครื่อง (รันครั้งเดียวตอน start)"""
    if (MODELS_DIR / "metrics.json").exists():
        return
    if not HF_MODEL_REPO:
        raise RuntimeError(
            "ไม่พบโมเดลในเครื่องและไม่ได้ตั้ง env HF_MODEL_REPO ให้ดาวน์โหลด"
        )
    from huggingface_hub import snapshot_download

    print(f"กำลังดาวน์โหลดโมเดลจาก {HF_MODEL_REPO} ...")
    t0 = time.time()
    snapshot_download(
        repo_id=HF_MODEL_REPO, repo_type="model", local_dir=MODELS_DIR
    )
    print(f"ดาวน์โหลดเสร็จใน {time.time() - t0:.0f} วินาที")


class ModelRegistry:
    def __init__(self) -> None:
        ensure_models_downloaded()
        self.metrics = json.loads((MODELS_DIR / "metrics.json").read_text())
        self.available = {
            n for n in self.metrics["models"] if self._artifact_exists(n)
        }
        for name, m in self.metrics["models"].items():
            m["available"] = name in self.available
        best = self.metrics["best_model"]
        self.best = (
            best
            if best in self.available
            else max(self.available, key=lambda n: self.metrics["models"][n]["f1"])
        )
        self._baselines: dict = {}
        self._berta = None  # (tokenizer, model, device)

    @staticmethod
    def _artifact_exists(name: str) -> bool:
        if name == "wangchanberta":
            return (MODELS_DIR / "wangchanberta").is_dir()
        return (MODELS_DIR / f"{name}.joblib").exists()

    def _load_baseline(self, name: str):
        if name not in self._baselines:
            import joblib

            self._baselines[name] = joblib.load(MODELS_DIR / f"{name}.joblib")
        return self._baselines[name]

    def _load_berta(self):
        if self._berta is None:
            import torch
            from transformers import (
                AutoModelForSequenceClassification,
                AutoTokenizer,
            )

            path = MODELS_DIR / "wangchanberta"
            device = torch.device("cpu")  # HF Space free tier = CPU basic
            tokenizer = AutoTokenizer.from_pretrained(path)
            model = AutoModelForSequenceClassification.from_pretrained(path)
            model.to(device).eval()
            self._berta = (tokenizer, model, device)
        return self._berta

    def warmup(self) -> None:
        for name in self.available:
            if name == "wangchanberta":
                self._load_berta()
            else:
                self._load_baseline(name)

    def predict(self, text: str, model_name: str | None = None) -> dict:
        name = model_name or self.best
        if name not in self.metrics["models"]:
            raise KeyError(f"ไม่รู้จักโมเดล: {name}")
        if name not in self.available:
            raise KeyError(f"โมเดล {name} ไม่พร้อมใช้งานบนเซิร์ฟเวอร์นี้")
        t0 = time.time()

        if name == "wangchanberta":
            import torch

            tokenizer, model, device = self._load_berta()
            enc = tokenizer(
                clean_text(text), truncation=True, max_length=256,
                return_tensors="pt",
            ).to(device)
            with torch.no_grad():
                probs = model(**enc).logits.softmax(dim=-1)[0].cpu().tolist()
        else:
            pipe = self._load_baseline(name)
            probs = pipe.predict_proba([text])[0].tolist()

        label_idx = int(probs[1] >= 0.5)
        return {
            "label": "fake" if label_idx == 1 else "real",
            "label_th": "ข่าวปลอม" if label_idx == 1 else "ข่าวจริง",
            "probability": round(probs[label_idx], 4),
            "prob_fake": round(probs[1], 4),
            "model_used": name,
            "display_name": self.metrics["models"][name]["display_name"],
            "latency_ms": round((time.time() - t0) * 1000, 1),
        }