phuongsuga commited on
Commit
8a4f04c
·
verified ·
1 Parent(s): d991985

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +17 -11
app.py CHANGED
@@ -15,6 +15,7 @@ from flask_cors import CORS
15
  # =====================
16
  # CONFIG
17
  # =====================
 
18
  TEXT_MODEL_REPO = "phuongsuga/PBL6_AI_Model_Text_Image"
19
  IMAGE_MODEL_REPO = "phuongsuga/PBL6_AI_Model_Image"
20
  DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
@@ -35,33 +36,35 @@ text_model = AutoModelForSequenceClassification.from_pretrained(
35
  TEXT_MODEL_REPO,
36
  subfolder="text_model/checkpoint-3390"
37
  )
38
-
39
  text_model.to(DEVICE).eval()
40
 
41
- # Load label2id.json từ repo
42
  label_url = f"https://huggingface.co/{TEXT_MODEL_REPO}/resolve/main/text_model/label2id.json"
43
  label2id = requests.get(label_url).json()
44
  id2label_text = {i: l for l, i in label2id.items()}
45
 
46
-
47
  # =====================
48
  # LOAD IMAGE MODEL
49
  # =====================
50
  print("🔹 Downloading image model...")
 
51
  class_names = ["an_toan", "bao_luc", "khieu_dam_doi_truy", "nhay_cam_chinh_tri"]
52
 
53
  def build_model(num_classes=4):
54
  return timm.create_model("efficientnet_b3", pretrained=False, num_classes=num_classes)
55
 
56
  image_model = build_model()
57
- image_model_path = f"https://huggingface.co/{IMAGE_MODEL_REPO}/resolve/main/image_model/efficientnet_b3.pth"
58
 
59
- torch.hub.download_url_to_file(image_model_path, "efficientnet_b3.pth")
60
- image_model.load_state_dict(torch.load("efficientnet_b3.pth", map_location=DEVICE))
61
- image_model.to(DEVICE).eval()
 
 
62
 
 
 
63
 
64
- # Chuẩn hoá ảnh đầu vào
65
  val_transforms = transforms.Compose([
66
  transforms.Resize((224, 224)),
67
  transforms.ToTensor(),
@@ -76,7 +79,7 @@ def seg_pyvi(text: str) -> str:
76
  try:
77
  seg = ViTokenizer.tokenize(text)
78
  seg = seg.replace(" ", "_")
79
- except:
80
  seg = text
81
  return seg
82
 
@@ -129,6 +132,10 @@ def predict_image(pil_image: Image.Image):
129
  app = Flask(__name__)
130
  CORS(app)
131
 
 
 
 
 
132
  @app.route("/analyze", methods=["POST"])
133
  def analyze():
134
  result = {"text_result": [], "image_result": []}
@@ -152,6 +159,5 @@ def analyze():
152
  # RUN APP
153
  # =====================
154
  if __name__ == "__main__":
155
- import os
156
- port = int(os.environ.get("PORT", 7860)) # Hugging Face sẽ truyền PORT vào đây
157
  app.run(host="0.0.0.0", port=port)
 
15
  # =====================
16
  # CONFIG
17
  # =====================
18
+ os.environ["TRANSFORMERS_CACHE"] = "/tmp/hf_cache" # tránh vượt storage limit
19
  TEXT_MODEL_REPO = "phuongsuga/PBL6_AI_Model_Text_Image"
20
  IMAGE_MODEL_REPO = "phuongsuga/PBL6_AI_Model_Image"
21
  DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
 
36
  TEXT_MODEL_REPO,
37
  subfolder="text_model/checkpoint-3390"
38
  )
 
39
  text_model.to(DEVICE).eval()
40
 
41
+ # load label2id.json
42
  label_url = f"https://huggingface.co/{TEXT_MODEL_REPO}/resolve/main/text_model/label2id.json"
43
  label2id = requests.get(label_url).json()
44
  id2label_text = {i: l for l, i in label2id.items()}
45
 
 
46
  # =====================
47
  # LOAD IMAGE MODEL
48
  # =====================
49
  print("🔹 Downloading image model...")
50
+
51
  class_names = ["an_toan", "bao_luc", "khieu_dam_doi_truy", "nhay_cam_chinh_tri"]
52
 
53
  def build_model(num_classes=4):
54
  return timm.create_model("efficientnet_b3", pretrained=False, num_classes=num_classes)
55
 
56
  image_model = build_model()
 
57
 
58
+ # tải model tạm trong /tmp để không chiếm storage
59
+ image_model_path = "/tmp/efficientnet_b3.pth"
60
+ if not os.path.exists(image_model_path):
61
+ url = f"https://huggingface.co/{IMAGE_MODEL_REPO}/resolve/main/image_model/efficientnet_b3.pth"
62
+ torch.hub.download_url_to_file(url, image_model_path)
63
 
64
+ image_model.load_state_dict(torch.load(image_model_path, map_location=DEVICE))
65
+ image_model.to(DEVICE).eval()
66
 
67
+ # chuẩn hóa ảnh
68
  val_transforms = transforms.Compose([
69
  transforms.Resize((224, 224)),
70
  transforms.ToTensor(),
 
79
  try:
80
  seg = ViTokenizer.tokenize(text)
81
  seg = seg.replace(" ", "_")
82
+ except Exception:
83
  seg = text
84
  return seg
85
 
 
132
  app = Flask(__name__)
133
  CORS(app)
134
 
135
+ @app.route("/")
136
+ def home():
137
+ return jsonify({"message": "✅ AI moderation API is running!"})
138
+
139
  @app.route("/analyze", methods=["POST"])
140
  def analyze():
141
  result = {"text_result": [], "image_result": []}
 
159
  # RUN APP
160
  # =====================
161
  if __name__ == "__main__":
162
+ port = int(os.environ.get("PORT", 7860)) # Hugging Face Space truyền PORT vào
 
163
  app.run(host="0.0.0.0", port=port)