pet-classifier / app_hf.py
xin.zeng
fix
09c4cbb
Raw
History Blame Contribute Delete
21.6 kB
"""
Hugging Face Spaces 版本 - 宠物图片分类器
使用 Gradio Blocks 自定义界面,与 Flask 前端样式一致
"""
import os
import json
import sqlite3
import base64
from datetime import datetime
from pathlib import Path
from io import BytesIO
import torch
import torch.nn.functional as F
from PIL import Image
import gradio as gr
import timm
from torchvision import transforms
# ==================== 配置 ====================
MODEL_NAME = "resnet50"
NUM_CLASSES = 37
IMAGE_SIZE = 224
MEAN = [0.485, 0.456, 0.406]
STD = [0.229, 0.224, 0.225]
# 数据库配置 - 使用绝对路径确保在 Hugging Face Spaces 上正确工作
# 优先使用应用根目录,如果不可写则回退到 /tmp
APP_DIR = Path(__file__).parent.resolve()
def get_writable_base_dir():
"""获取可写的基础目录
优先级:
1. /data - Hugging Face Persistent Storage(付费功能,数据持久化)
2. APP_DIR - 应用根目录(临时,容器重启后丢失)
3. /tmp - 临时目录(临时,容器重启后丢失)
"""
test_dirs = [
(Path("/data"), "Hugging Face Persistent Storage (持久化)"),
(APP_DIR, "应用根目录 (临时)"),
(Path("/tmp"), "临时目录 (临时)"),
]
for base_dir, desc in test_dirs:
test_path = base_dir / ".write_test"
try:
base_dir.mkdir(parents=True, exist_ok=True)
test_path.touch()
test_path.unlink()
print(f"✅ 使用目录: {base_dir} - {desc}")
if "临时" in desc:
print(f"⚠️ 警告: 数据将在 Space 重启后丢失!")
print(f" 如需持久化,请在 Settings 中启用 Persistent Storage")
return base_dir
except Exception as e:
print(f"⚠️ 目录不可用 {base_dir}: {e}")
continue
# 最后回退到 /tmp
print("⚠️ 回退到 /tmp 目录,数据不会持久化")
return Path("/tmp")
WRITABLE_BASE = get_writable_base_dir()
DATABASE_DIR = WRITABLE_BASE / "database"
UPLOADS_DIR = WRITABLE_BASE / "uploads"
DATABASE_PATH = DATABASE_DIR / "predictions_hf.db"
MAX_RECORDS = 10 # 只保存最近10次记录
print(f"📂 数据库目录: {DATABASE_DIR}")
print(f"📂 上传目录: {UPLOADS_DIR}")
# 37个类别名称
CLASS_NAMES = [
"Abyssinian", "american_bulldog", "american_pit_bull_terrier",
"basset_hound", "beagle", "Bengal", "Birman", "Bombay", "boxer",
"British_Shorthair", "chihuahua", "Egyptian_Mau", "english_cocker_spaniel",
"english_setter", "german_shorthaired", "great_pyrenees", "havanese",
"japanese_chin", "keeshond", "leonberger", "Maine_Coon", "miniature_pinscher",
"newfoundland", "Persian", "pomeranian", "pug", "Ragdoll", "Russian_Blue",
"saint_bernard", "samoyed", "scottish_terrier", "shiba_inu", "Siamese",
"Sphynx", "staffordshire_bull_terrier", "wheaten_terrier", "yorkshire_terrier",
]
# 类别中英文映射
CLASS_NAMES_CN = {
"Abyssinian": "阿比西尼亚猫", "american_bulldog": "美国斗牛犬",
"american_pit_bull_terrier": "美国比特犬", "basset_hound": "巴吉度猎犬",
"beagle": "比格犬", "Bengal": "孟加拉猫", "Birman": "伯曼猫",
"Bombay": "孟买猫", "boxer": "拳师犬", "British_Shorthair": "英国短毛猫",
"chihuahua": "吉娃娃", "Egyptian_Mau": "埃及猫",
"english_cocker_spaniel": "英国可卡犬", "english_setter": "英国塞特犬",
"german_shorthaired": "德国短毛猎犬", "great_pyrenees": "大白熊犬",
"havanese": "哈瓦那犬", "japanese_chin": "日本狆", "keeshond": "荷兰毛狮犬",
"leonberger": "莱昂伯格犬", "Maine_Coon": "缅因猫",
"miniature_pinscher": "迷你杜宾犬", "newfoundland": "纽芬兰犬",
"Persian": "波斯猫", "pomeranian": "博美犬", "pug": "巴哥犬",
"Ragdoll": "布偶猫", "Russian_Blue": "俄罗斯蓝猫",
"saint_bernard": "圣伯纳犬", "samoyed": "萨摩耶",
"scottish_terrier": "苏格兰梗", "shiba_inu": "柴犬", "Siamese": "暹罗猫",
"Sphynx": "斯芬克斯猫", "staffordshire_bull_terrier": "斯塔福郡斗牛梗",
"wheaten_terrier": "软毛麦色梗", "yorkshire_terrier": "约克夏梗",
}
# ==================== 数据库操作 ====================
def init_db():
"""初始化数据库和上传目录"""
print("=" * 50)
print("🚀 初始化数据库和目录...")
print(f"📂 DATABASE_DIR: {DATABASE_DIR}")
print(f"📂 UPLOADS_DIR: {UPLOADS_DIR}")
print(f"📂 DATABASE_PATH: {DATABASE_PATH}")
# 创建目录
try:
DATABASE_DIR.mkdir(parents=True, exist_ok=True)
print(f"✅ 数据库目录已创建: {DATABASE_DIR}")
print(f" 目录存在: {DATABASE_DIR.exists()}")
print(f" 是否可写: {os.access(DATABASE_DIR, os.W_OK)}")
except Exception as e:
print(f"❌ 创建数据库目录失败: {e}")
import traceback
traceback.print_exc()
try:
UPLOADS_DIR.mkdir(parents=True, exist_ok=True)
print(f"✅ 上传目录已创建: {UPLOADS_DIR}")
print(f" 目录存在: {UPLOADS_DIR.exists()}")
print(f" 是否可写: {os.access(UPLOADS_DIR, os.W_OK)}")
except Exception as e:
print(f"❌ 创建上传目录失败: {e}")
import traceback
traceback.print_exc()
# 创建数据库
try:
conn = sqlite3.connect(str(DATABASE_PATH))
cursor = conn.cursor()
cursor.execute("""
CREATE TABLE IF NOT EXISTS predictions (
id INTEGER PRIMARY KEY AUTOINCREMENT,
predicted_class TEXT NOT NULL,
predicted_class_cn TEXT,
confidence REAL NOT NULL,
top5_json TEXT,
image_path TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)
""")
conn.commit()
conn.close()
print(f"✅ 数据库已初始化: {DATABASE_PATH}")
print(f" 数据库文件存在: {DATABASE_PATH.exists()}")
except Exception as e:
print(f"❌ 数据库初始化失败: {e}")
import traceback
traceback.print_exc()
print("=" * 50)
# 启动时清理旧文件(保持最近10次)
cleanup_old_images()
def save_uploaded_image(image, record_id):
"""保存上传的图片"""
try:
# 确保上传目录存在
if not UPLOADS_DIR.exists():
print(f"⚠️ 上传目录不存在,尝试创建: {UPLOADS_DIR}")
UPLOADS_DIR.mkdir(parents=True, exist_ok=True)
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
image_filename = f"{record_id}_{timestamp}.jpg"
image_path = UPLOADS_DIR / image_filename
# 确保是 RGB 模式
if image.mode != "RGB":
image = image.convert("RGB")
image.save(image_path, "JPEG", quality=85)
print(f"✅ 图片已保存: {image_path}")
print(f" 文件存在: {image_path.exists()}")
return str(image_path)
except Exception as e:
print(f"❌ 保存图片失败: {e}")
import traceback
traceback.print_exc()
return None
def add_prediction_to_db(predicted_class, predicted_class_cn, confidence, top5_json, image_path=None):
"""添加预测记录到数据库"""
conn = sqlite3.connect(str(DATABASE_PATH))
cursor = conn.cursor()
cursor.execute("""
INSERT INTO predictions (predicted_class, predicted_class_cn, confidence, top5_json, image_path)
VALUES (?, ?, ?, ?, ?)
""", (predicted_class, predicted_class_cn, confidence, top5_json, image_path))
# 获取刚插入的记录ID
record_id = cursor.lastrowid
# 删除超过限制的旧记录
cursor.execute("""
DELETE FROM predictions WHERE id NOT IN (
SELECT id FROM predictions ORDER BY created_at DESC LIMIT ?
)
""", (MAX_RECORDS,))
conn.commit()
conn.close()
return record_id
def cleanup_old_images():
"""清理旧的图片文件,只保留最近10次对应的图片"""
try:
conn = sqlite3.connect(str(DATABASE_PATH))
cursor = conn.cursor()
# 获取所有保留记录的图片路径
cursor.execute("SELECT image_path FROM predictions ORDER BY created_at DESC LIMIT ?", (MAX_RECORDS,))
keep_paths = {row[0] for row in cursor.fetchall() if row[0]}
conn.close()
# 删除不在保留列表中的图片文件
if UPLOADS_DIR.exists():
deleted_count = 0
for img_file in UPLOADS_DIR.glob("*.jpg"):
img_path_str = str(img_file)
if img_path_str not in keep_paths:
try:
img_file.unlink()
deleted_count += 1
except Exception as e:
print(f"⚠️ 删除图片失败 {img_file.name}: {e}")
if deleted_count > 0:
print(f"🗑️ 清理了 {deleted_count} 张旧图片")
except Exception as e:
print(f"⚠️ 清理图片失败: {e}")
def get_recent_predictions(limit=10):
"""获取最近的预测记录,包括图片路径"""
conn = sqlite3.connect(str(DATABASE_PATH))
cursor = conn.cursor()
cursor.execute("""
SELECT predicted_class, predicted_class_cn, confidence, image_path, created_at
FROM predictions ORDER BY created_at DESC LIMIT ?
""", (limit,))
rows = cursor.fetchall()
conn.close()
return rows
# ==================== 模型 ====================
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = None
transform = transforms.Compose([
transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
transforms.ToTensor(),
transforms.Normalize(mean=MEAN, std=STD),
])
def init_model():
global model
if model is None:
model_path = "models/best_model.pth"
if os.path.exists(model_path):
print(f"📦 加载模型: {model_path}")
m = timm.create_model(MODEL_NAME, pretrained=False, num_classes=NUM_CLASSES)
checkpoint = torch.load(model_path, map_location=device, weights_only=False)
if "model_state_dict" in checkpoint:
m.load_state_dict(checkpoint["model_state_dict"])
else:
m.load_state_dict(checkpoint)
m = m.to(device)
m.eval()
model = m
print(f"✅ 模型加载完成,使用设备: {device}")
else:
raise FileNotFoundError(f"模型文件不存在: {model_path}")
return model
# ==================== 预测函数 ====================
def predict(image):
if image is None:
return "", "", ""
m = init_model()
image_rgb = image.convert("RGB")
image_tensor = transform(image_rgb).unsqueeze(0).to(device)
with torch.no_grad():
output = m(image_tensor)
probs = F.softmax(output, dim=1)
top5_probs, top5_indices = probs.topk(5, dim=1)
top5_list = []
for i in range(5):
idx = top5_indices[0][i].item()
prob = top5_probs[0][i].item()
class_name = CLASS_NAMES[idx]
cn_name = CLASS_NAMES_CN.get(class_name, class_name)
top5_list.append({
"class": class_name,
"class_cn": cn_name,
"probability": round(prob * 100, 2)
})
top1_class = CLASS_NAMES[top5_indices[0][0].item()]
top1_cn = CLASS_NAMES_CN.get(top1_class, top1_class)
top1_confidence = round(top5_probs[0][0].item() * 100, 2)
# 保存预测结果和图片到数据库
image_path = None
try:
print(f"📝 保存预测结果到数据库...")
print(f" 数据库路径: {DATABASE_PATH}")
print(f" 数据库存在: {DATABASE_PATH.exists()}")
# 确保数据库目录存在
if not DATABASE_DIR.exists():
print(f"⚠️ 数据库目录不存在,尝试创建: {DATABASE_DIR}")
DATABASE_DIR.mkdir(parents=True, exist_ok=True)
# 先插入记录获取ID(不清理,等保存图片后再清理)
conn = sqlite3.connect(str(DATABASE_PATH))
cursor = conn.cursor()
cursor.execute("""
INSERT INTO predictions (predicted_class, predicted_class_cn, confidence, top5_json, image_path)
VALUES (?, ?, ?, ?, ?)
""", (top1_class, top1_cn, top1_confidence, json.dumps(top5_list, ensure_ascii=False), None))
record_id = cursor.lastrowid
conn.commit()
print(f"✅ 预测记录已插入,ID: {record_id}")
# 保存图片
image_path = save_uploaded_image(image_rgb, record_id)
# 更新数据库记录中的图片路径
if image_path:
cursor.execute("UPDATE predictions SET image_path = ? WHERE id = ?", (image_path, record_id))
conn.commit()
print(f"✅ 图片路径已更新到数据库")
# 删除超过限制的旧记录(在同一个事务中)
cursor.execute("""
DELETE FROM predictions WHERE id NOT IN (
SELECT id FROM predictions ORDER BY created_at DESC LIMIT ?
)
""", (MAX_RECORDS,))
conn.commit()
conn.close()
# 清理旧的图片文件
cleanup_old_images()
except Exception as e:
print(f"❌ 保存失败: {e}")
import traceback
traceback.print_exc()
# 构建结果显示
result_class = f"{top1_class} ({top1_cn})"
result_confidence = f"{top1_confidence}%"
# 构建 Top5 HTML
top5_html = ""
for i, item in enumerate(top5_list):
top5_html += f'''
<div class="top5-item">
<span class="top5-rank">{i+1}</span>
<span class="top5-name">{item["class"]} ({item["class_cn"]})</span>
<span class="top5-prob">{item["probability"]}%</span>
</div>
<div class="top5-bar">
<div class="top5-bar-fill" style="width: {item["probability"]}%"></div>
</div>
'''
return result_class, result_confidence, top5_html
def image_to_base64(image_path):
"""将图片转换为 base64 编码"""
try:
if image_path and Path(image_path).exists():
with open(image_path, "rb") as img_file:
img_data = img_file.read()
img_base64 = base64.b64encode(img_data).decode("utf-8")
return f"data:image/jpeg;base64,{img_base64}"
except Exception as e:
print(f"⚠️ 转换图片失败: {e}")
return None
def get_history_html():
"""生成历史记录HTML,包含图片"""
rows = get_recent_predictions(MAX_RECORDS)
if not rows:
return '''<div class="empty-state"><div class="empty-icon">📭</div><p>暂无历史记录</p></div>'''
html = ""
for row in rows:
pred_class, pred_cn, confidence, image_path, created_at = row
# 处理图片 - 转换为 base64 编码
image_html = ""
if image_path:
img_base64 = image_to_base64(image_path)
if img_base64:
image_html = f'<img src="{img_base64}" class="history-image" alt="预测图片" />'
else:
image_html = '<div class="history-image-placeholder">📷</div>'
else:
image_html = '<div class="history-image-placeholder">📷</div>'
html += f'''
<div class="history-item">
{image_html}
<div class="history-info">
<div class="history-class">{pred_class} ({pred_cn})</div>
<div class="history-meta">{created_at}</div>
</div>
<div class="history-confidence">{confidence}%</div>
</div>
'''
return html
# ==================== 自定义 CSS ====================
custom_css = """
.gradio-container {
background: linear-gradient(135deg, #667eea 0%, #764ba2 100%) !important;
min-height: 100vh;
}
.main-header {
text-align: center;
color: white;
padding: 20px 0;
}
.main-header h1 {
font-size: 2.5rem;
margin-bottom: 10px;
text-shadow: 0 2px 10px rgba(0,0,0,0.2);
}
.main-header p {
font-size: 1.1rem;
opacity: 0.9;
}
.card {
background: white !important;
border-radius: 20px !important;
padding: 25px !important;
box-shadow: 0 15px 35px rgba(0,0,0,0.1) !important;
}
.card h2 {
color: #333;
margin-bottom: 15px;
font-size: 1.4rem;
}
.prediction-main {
background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);
border-radius: 16px;
padding: 25px;
text-align: center;
color: white;
margin: 15px 0;
}
.prediction-class {
font-size: 1.6rem;
font-weight: bold;
margin-bottom: 10px;
}
.prediction-confidence {
font-size: 2.5rem;
font-weight: bold;
text-shadow: 0 2px 10px rgba(0,0,0,0.2);
}
.prediction-label {
font-size: 0.9rem;
opacity: 0.9;
margin-top: 5px;
}
.top5-item {
display: flex;
justify-content: space-between;
align-items: center;
padding: 10px 12px;
border-radius: 10px;
margin-bottom: 8px;
background: #f8f9fa;
}
.top5-rank {
width: 26px;
height: 26px;
background: #667eea;
color: white;
border-radius: 50%;
display: inline-flex;
align-items: center;
justify-content: center;
font-weight: bold;
font-size: 0.85rem;
margin-right: 10px;
}
.top5-name {
flex: 1;
font-weight: 500;
color: #333;
}
.top5-prob {
font-weight: bold;
color: #667eea;
}
.top5-bar {
width: 100%;
height: 4px;
background: #eee;
border-radius: 2px;
margin-bottom: 10px;
overflow: hidden;
}
.top5-bar-fill {
height: 100%;
background: linear-gradient(90deg, #667eea, #764ba2);
border-radius: 2px;
}
.history-item {
display: flex;
align-items: center;
padding: 12px;
border-radius: 10px;
margin-bottom: 8px;
background: #f8f9fa;
gap: 12px;
}
.history-image {
width: 60px;
height: 60px;
object-fit: cover;
border-radius: 8px;
border: 2px solid #e0e0e0;
}
.history-image-placeholder {
width: 60px;
height: 60px;
display: flex;
align-items: center;
justify-content: center;
background: #e0e0e0;
border-radius: 8px;
font-size: 24px;
}
.history-info {
flex: 1;
}
.history-class {
font-weight: bold;
color: #333;
margin-bottom: 3px;
}
.history-meta {
font-size: 0.8rem;
color: #999;
}
.history-confidence {
font-weight: bold;
color: #667eea;
font-size: 1.1rem;
}
.empty-state {
text-align: center;
padding: 30px;
color: #999;
}
.empty-icon {
font-size: 40px;
margin-bottom: 10px;
}
"""
# ==================== Gradio 界面 ====================
# 初始化数据库和目录
init_db()
with gr.Blocks(title="🐾 宠物图片分类器", css=custom_css) as demo:
gr.HTML("""
<div class="main-header">
<h1>🐾 宠物图片分类器</h1>
<p>基于 Oxford-IIIT Pet 数据集训练,支持 37 种宠物分类</p>
</div>
""")
with gr.Row():
# 左侧:上传和结果
with gr.Column(scale=1):
with gr.Group(elem_classes="card"):
gr.HTML("<h2>📷 上传图片</h2>")
image_input = gr.Image(type="pil", height=300)
predict_btn = gr.Button("🔍 开始识别", variant="primary", size="lg")
with gr.Column(elem_classes="card", visible=False) as result_container:
gr.HTML("<h2>🎯 识别结果</h2>")
result_html = gr.HTML()
# 右侧:历史记录
with gr.Column(scale=1):
with gr.Group(elem_classes="card"):
gr.HTML("<h2>📜 历史记录</h2>")
history_html = gr.HTML(value=get_history_html())
refresh_btn = gr.Button("🔄 刷新", size="sm")
def on_predict(image):
if image is None:
return gr.update(visible=False), "", get_history_html()
result_class, result_confidence, top5_html = predict(image)
result_display = f'''
<div class="prediction-main">
<div class="prediction-class">{result_class}</div>
<div class="prediction-confidence">{result_confidence}</div>
<div class="prediction-label">置信度</div>
</div>
<h3 style="margin: 15px 0 10px; color: #333;">Top 5 预测</h3>
{top5_html}
'''
return gr.update(visible=True), result_display, get_history_html()
def on_refresh():
return get_history_html()
predict_btn.click(
fn=on_predict,
inputs=[image_input],
outputs=[result_container, result_html, history_html]
)
refresh_btn.click(fn=on_refresh, outputs=[history_html])
# ==================== 启动 ====================
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0", server_port=7860)