phuongsuga commited on
Commit
dc37c6d
·
verified ·
1 Parent(s): 3f1030d

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +9 -8
app.py CHANGED
@@ -15,8 +15,8 @@ from flask_cors import CORS
15
  # =====================
16
  # CONFIG
17
  # =====================
18
- TEXT_MODEL_REPO = "phuongsuga/PBL6_AI_Model_Text_Image" # repo chứa mô hình text
19
- IMAGE_MODEL_REPO = "phuongsuga/PBL6_AI_Model_Image" # repo chứa mô hình image
20
  DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
21
  THRESHOLD = 0.65
22
 
@@ -24,12 +24,13 @@ THRESHOLD = 0.65
24
  # LOAD TEXT MODEL
25
  # =====================
26
  print("🔹 Downloading text model...")
27
- tokenizer = AutoTokenizer.from_pretrained(TEXT_MODEL_REPO, use_fast=False)
28
- text_model = AutoModelForSequenceClassification.from_pretrained(TEXT_MODEL_REPO)
29
  text_model.to(DEVICE).eval()
30
 
31
- # Load label2id.json
32
- label_url = f"https://huggingface.co/{TEXT_MODEL_REPO}/resolve/main/label2id.json"
 
33
  label2id = requests.get(label_url).json()
34
  id2label_text = {i: l for l, i in label2id.items()}
35
 
@@ -43,13 +44,13 @@ def build_model(num_classes=4):
43
  return timm.create_model("efficientnet_b3", pretrained=False, num_classes=num_classes)
44
 
45
  image_model = build_model()
 
46
 
47
- # tải file weight từ repo Hugging Face
48
- image_model_path = f"https://huggingface.co/{IMAGE_MODEL_REPO}/resolve/main/efficientnet_b3.pth"
49
  torch.hub.download_url_to_file(image_model_path, "efficientnet_b3.pth")
50
  image_model.load_state_dict(torch.load("efficientnet_b3.pth", map_location=DEVICE))
51
  image_model.to(DEVICE).eval()
52
 
 
53
  # Chuẩn hoá ảnh đầu vào
54
  val_transforms = transforms.Compose([
55
  transforms.Resize((224, 224)),
 
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"
21
  THRESHOLD = 0.65
22
 
 
24
  # LOAD TEXT MODEL
25
  # =====================
26
  print("🔹 Downloading text model...")
27
+ tokenizer = AutoTokenizer.from_pretrained(f"{TEXT_MODEL_REPO}/text_model", use_fast=False)
28
+ text_model = AutoModelForSequenceClassification.from_pretrained(f"{TEXT_MODEL_REPO}/text_model/checkpoint-3390")
29
  text_model.to(DEVICE).eval()
30
 
31
+ # Load label2id.json từ repo
32
+ import requests
33
+ label_url = f"https://huggingface.co/{TEXT_MODEL_REPO}/resolve/main/text_model/label2id.json"
34
  label2id = requests.get(label_url).json()
35
  id2label_text = {i: l for l, i in label2id.items()}
36
 
 
44
  return timm.create_model("efficientnet_b3", pretrained=False, num_classes=num_classes)
45
 
46
  image_model = build_model()
47
+ image_model_path = f"https://huggingface.co/{IMAGE_MODEL_REPO}/resolve/main/image_model/efficientnet_b3.pth"
48
 
 
 
49
  torch.hub.download_url_to_file(image_model_path, "efficientnet_b3.pth")
50
  image_model.load_state_dict(torch.load("efficientnet_b3.pth", map_location=DEVICE))
51
  image_model.to(DEVICE).eval()
52
 
53
+
54
  # Chuẩn hoá ảnh đầu vào
55
  val_transforms = transforms.Compose([
56
  transforms.Resize((224, 224)),