GLINER_DEMO / app.py
TienVu2204's picture
Update app.py
6db9175 verified
Raw History Blame Contribute Delete
8.5 kB
import gradio as gr
import json
import torch
import time
import psutil
import gc
from gliner import GLiNER
# ---------------------------------------------------------
# CẤU HÌNH CHUNG
# ---------------------------------------------------------
device = "cuda" if torch.cuda.is_available() else "cpu"
MODEL_CHOICES = [
"gliner-community/gliner_large-v2.5",
"gliner-community/gliner_small-v2.5",
"knowledgator/gliner-x-base",
"knowledgator/gliner-relex-base-v1.0",
"knowledgator/gliner-bi-base-v1.0",
"knowledgator/gliner-multitask-v1.0"
]
# Schema chung: 5 trường cần bóc tách
SCHEMA_FIELDS = ["Name", "Location", "Date or Time", "Pet or Pet type", "Family member"]
# Biến toàn cục để cache model đang nạp trong VRAM
current_model_name = None
current_model = None
# ---------------------------------------------------------
# HÀM ĐO LƯỜNG VÀ QUẢN LÝ TÀI NGUYÊN
# ---------------------------------------------------------
def get_memory_stats():
ram_info = psutil.virtual_memory()
ram_used_gb = ram_info.used / (1024**3)
vram_used_gb = 0
if torch.cuda.is_available():
vram_used_gb = torch.cuda.memory_allocated() / (1024**3)
return ram_used_gb, vram_used_gb
def free_current_model():
"""Dọn dẹp model hiện tại ra khỏi VRAM."""
global current_model, current_model_name
if current_model is not None:
del current_model
current_model = None
current_model_name = None
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
# ---------------------------------------------------------
# LOADER VÀ EXTRACTOR CHO GLINER MODELS
# ---------------------------------------------------------
def load_model(model_name):
"""Tự động unload model cũ rồi load model mới."""
global current_model, current_model_name
if current_model_name == model_name:
return # Đã load sẵn
print(f"Switching model: {current_model_name} -> {model_name}")
free_current_model()
# Load standard GLiNER model
current_model = GLiNER.from_pretrained(model_name)
if torch.cuda.is_available():
current_model = current_model.to("cuda")
current_model_name = model_name
print(f"Loaded: {model_name}")
def extract_with_gliner(text):
"""Dùng GLiNER tiêu chuẩn để trích xuất, có xử lý ngoại lệ cho các format dị biệt."""
model = current_model
labels = SCHEMA_FIELDS
# Lấy output thô từ model
raw_entities = model.predict_entities(text, labels, threshold=0.4)
result_dict = {label: [] for label in labels}
# 1. Bẻ bọc (Unwrap): Xử lý trường hợp model trả về mảng 2 chiều do lỗi batching ngầm
# VD: raw_entities = [ [ {'text': 'Bob', 'label': 'Name'} ] ]
if isinstance(raw_entities, list) and len(raw_entities) > 0 and isinstance(raw_entities[0], list):
entities = raw_entities[0]
else:
entities = raw_entities
# 2. Xử lý an toàn từng entity
for ent in entities:
try:
ent_label = None
ent_text = None
# Format A: Chuẩn (Dictionary) -> Mở bình thường
if isinstance(ent, dict):
ent_label = ent.get("label")
ent_text = ent.get("text")
# Format B: Tuple/List -> Quét tìm xem phần tử nào là label, phần tử nào là text
elif isinstance(ent, (list, tuple)):
# Lọc ra giá trị nào trùng với danh sách SCHEMA_FIELDS thì nó chính là nhãn
ent_label = next((item for item in ent if item in labels), None)
# Text là giá trị chuỗi (string) còn lại không phải là nhãn
ent_text = next((item for item in ent if isinstance(item, str) and item not in labels), None)
# Đưa vào kết quả nếu tìm thấy hợp lệ
if ent_label in result_dict and ent_text:
result_dict[ent_label].append(ent_text)
except Exception as e:
# Nếu có entity nào cực kỳ dị biệt, log ra thay vì làm sập cả app
print(f"Lỗi đọc entity: {ent} - Chi tiết: {e}")
continue
return json.dumps(result_dict, indent=4, ensure_ascii=False)
# ---------------------------------------------------------
# HÀM XỬ LÝ CHÍNH TRÊN UI
# ---------------------------------------------------------
def extract_information(text, model_choice):
if not text or not text.strip():
return json.dumps({"error": "Vui lòng nhập câu đầu vào."}, indent=4), "Lỗi: Không có đầu vào."
# Reset peak memory tracker trước khi load
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
# 1. Load model (sẽ unload model cũ nếu có)
load_start = time.time()
try:
load_model(model_choice)
except Exception as e:
return f"Lỗi khi tải mô hình: {str(e)}", "Model Load Error"
load_time = time.time() - load_start
# 2. Inference
infer_start = time.time()
try:
extracted_data = extract_with_gliner(text)
except Exception as e:
return f"Lỗi khi inference: {str(e)}", f"Inference Error trên {model_choice}"
infer_time = time.time() - infer_start
# 3. Đo lường resource
ram_end, _ = get_memory_stats()
vram_peak = torch.cuda.max_memory_allocated() / (1024**3) if torch.cuda.is_available() else 0
log_output = (
f"⚙️ Model: {model_choice}\n"
f"⏱️ Load model: {load_time:.3f}s (0.000s nếu đã cache)\n"
f"⏱️ Inference: {infer_time:.3f}s\n"
f"💻 Device: {device.upper()}\n"
f"📊 RAM Hệ thống: ~{ram_end:.2f} GB\n"
f"🎮 VRAM peak: {vram_peak:.2f} GB\n"
)
return extracted_data, log_output
# ---------------------------------------------------------
# GIAO DIỆN GRADIO
# ---------------------------------------------------------
EXAMPLE_INPUTS = [
"Hey, my husband John and I are planning to bring our golden retriever Max to the clinic in downtown Boston next Tuesday around 3 PM.",
"Can you update the delivery address for Sarah to 1204 Pine Street, Seattle? Also, my sister will be there to receive it tomorrow morning.",
"I need to book a flight for myself and my daughter Emma from NYC to London on the 24th of October. We can't leave our two siamese cats behind.",
"We recently adopted a rescue iguana named Lizzy. My brother Mike is taking her to the vet on 5th Ave tomorrow at noon for a quick checkup.",
]
with gr.Blocks(theme=gr.themes.Soft()) as demo:
gr.Markdown("## 🕵️‍♂️ GLiNER Models Arena")
gr.Markdown(
"So sánh hiệu năng và độ chính xác của các phiên bản GLiNER khác nhau:\n"
"- **GLiNER v2.5** (Small/Large): Phiên bản cập nhật của community\n"
"- **Knowledgator Models** (x-base, relex, bi-base, multitask): Các phiên bản fine-tune/multi-task đặc biệt.\n"
"Model sẽ tự động giải phóng khỏi VRAM khi bạn chuyển đổi."
)
with gr.Row():
with gr.Column(scale=1):
input_text = gr.Textbox(
lines=5,
label="Input Sentence (Tiếng Anh)",
placeholder="VD: My husband John and I are bringing our golden retriever Max to Boston next Tuesday at 3 PM...",
)
model_dropdown = gr.Radio(
choices=MODEL_CHOICES,
value="gliner-community/gliner_small-v2.5",
label="Chọn Model để test",
)
submit_btn = gr.Button("Extract Data", variant="primary")
gr.Examples(
examples=[[ex] for ex in EXAMPLE_INPUTS],
inputs=[input_text],
label="📚 Test cases có sẵn",
)
with gr.Column(scale=1):
output_json = gr.Code(language="json", label="Extracted Information")
output_logs = gr.Textbox(
lines=7,
label="⚙️ System Logs & Resource Monitor",
interactive=False,
)
submit_btn.click(
fn=extract_information,
inputs=[input_text, model_dropdown],
outputs=[output_json, output_logs],
)
if __name__ == "__main__":
demo.launch()