kritaphatson's picture
Upload folder using huggingface_hub
7d27854 verified
Raw
History Blame Contribute Delete
7.14 kB
import os
import io
import json
import base64
import torch
import torch.nn as nn
import torch.nn.functional as F
from PIL import Image
import cv2
import numpy as np
import gradio as gr
from transformers import AutoModelForImageClassification, AutoImageProcessor
# --- Configuration ---
# Use an environment variable for the token to stay secure
HF_TOKEN = os.environ.get("HF_TOKEN")
USE_HF_TRANSFORMERS = True
HF_MODEL_ID = "kritaphatson/thai_snake_image_classifier" # Use repo ID for Space deployment
# Determine model path (prioritize local if exists, else Hub)
LOCAL_MODEL_PATH = r'C:\Users\ADMIN\PycharmProjects\snake_research\transformer_experiment_results_swin_exact\Swin-Base_augnone_fineT_thres500'
model_path = LOCAL_MODEL_PATH if os.path.exists(LOCAL_MODEL_PATH) else HF_MODEL_ID
# Grad-CAM Settings
STAGE_INDEX = 2
AGG_METHOD = "layercam"
ALPHA = 0.40
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
def load_inference_engine():
try:
print(f"Attempting to load model from: {model_path}")
# Use token for both local and Hub to be safe
model = AutoModelForImageClassification.from_pretrained(model_path, token=HF_TOKEN)
processor = AutoImageProcessor.from_pretrained(model_path, token=HF_TOKEN)
model.to(device)
model.eval()
print("Model loaded successfully.")
except Exception as e:
print(f"Error loading model: {e}")
return None, None, {}
# Try to load id2label
idx_to_class = {}
if hasattr(model.config, 'id2label') and model.config.id2label:
idx_to_class = model.config.id2label
print("Using model.config.id2label")
# Check for local file in Space dir
if not idx_to_class:
local_id2label = "id2label.json"
if os.path.exists(local_id2label):
with open(local_id2label, 'r') as f:
idx_to_class = json.load(f)
print(f"Loaded id2label from local {local_id2label}")
if not idx_to_class:
print("Warning: No id2label found, using defaults.")
idx_to_class = {str(i): f"Class {i}" for i in range(model.config.num_labels)}
return model, processor, idx_to_class
model, processor, idx_to_class = load_inference_engine()
str_idx_to_class = {str(k): v for k, v in idx_to_class.items()} if idx_to_class else {}
# --- Grad-CAM Logic ---
# (Same logic as before, but with safety)
def get_gradcam(model, pixel_values, target_layer):
if model is None: return None
holder = {"acts": None, "grads": None}
def fwd_hook(_m, _inp, out):
hs = out[0] if isinstance(out, (tuple, list)) else out
hs.requires_grad_(True)
hs.retain_grad()
holder["acts"] = hs
hs.register_hook(lambda g: holder.__setitem__("grads", g))
hook = target_layer.register_forward_hook(fwd_hook)
model.zero_grad()
outputs = model(pixel_values=pixel_values)
logits = outputs.logits
idx = torch.argmax(logits, dim=-1)
score = logits[0, idx]
score.backward()
hook.remove()
acts = holder["acts"][0].detach().cpu()
grads = holder["grads"][0].detach().cpu()
if AGG_METHOD == "layercam":
cam = torch.relu(grads) * torch.relu(acts)
cam = cam.sum(dim=-1)
else:
weights = grads.mean(dim=0)
cam = torch.relu(acts @ weights)
cam = cam - cam.min()
if cam.max() > 0:
cam = cam / cam.max()
s = int(np.sqrt(cam.shape[0]))
return cam.reshape(s, s).numpy()
def apply_overlay(orig_img, heatmap, alpha=ALPHA):
heatmap = cv2.resize(heatmap, (orig_img.size[0], orig_img.size[1]))
heatmap = np.uint8(255 * heatmap)
heatmap = cv2.applyColorMap(heatmap, cv2.COLORMAP_JET)
heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)
overlayed = cv2.addWeighted(heatmap, alpha, np.array(orig_img), 1 - alpha, 0)
return Image.fromarray(overlayed)
def predict(image):
if image is None: return None, {"Error": "No image uploaded"}
if model is None or processor is None:
return None, {"Error": "Model failed to load. Please check Space application logs."}
try:
raw_img = Image.fromarray(image).convert('RGB')
inputs = processor(images=raw_img, return_tensors="pt").to(device)
with torch.no_grad():
outputs = model(**inputs)
probabilities = F.softmax(outputs.logits, dim=1)[0]
target_layer = model.swin.encoder.layers[STAGE_INDEX]
heatmap = get_gradcam(model, inputs["pixel_values"], target_layer)
cam_pil = apply_overlay(raw_img, heatmap)
top_prob, top_idx = torch.topk(probabilities, k=min(5, len(str_idx_to_class)))
confidences = {str_idx_to_class.get(str(idx.item()), f"Class {idx.item()}"): float(prob.item())
for prob, idx in zip(top_prob, top_idx)}
return cam_pil, confidences
except Exception as e:
return None, {"Error": str(e)}
# --- Gradio Blocks Interface ---
with gr.Blocks(title="Thai Snake Classifier", theme="soft") as demo:
gr.Markdown("# Thai Snake Classifier")
gr.Markdown("Classify Thai snakes with Swin-Base and Grad-CAM visualization.")
with gr.Row():
with gr.Column():
input_img = gr.Image(label="Upload Snake Image", type="numpy")
btn = gr.Button("Classify", variant="primary")
with gr.Column():
output_cam = gr.Image(label="Grad-CAM Visualization", type="pil")
output_labels = gr.Label(num_top_classes=5, label="Confidences")
# Add Examples with explicit labels matching model id2label
gr.Examples(
examples=[
["examples/Trimeresurus_albolabris_1.png", "Trimeresurus_albolabris"],
["examples/Trimeresurus_albolabris_2.png", "Trimeresurus_albolabris"],
["examples/Trimeresurus_albolabris_3.png", "Trimeresurus_albolabris"],
["examples/Trimeresurus_macrops_1.png", "Trimeresurus_macrops"],
["examples/Trimeresurus_macrops_2.png", "Trimeresurus_macrops"],
["examples/Trimeresurus_macrops_3.png", "Trimeresurus_macrops"],
["examples/Ptyas_mucosa_1.png", "Ptyas_mucosa"],
["examples/Ptyas_mucosa_2.png", "Ptyas_mucosa"],
["examples/Oligodon_taeniatus_1.png", "Oligodon_taeniatus"],
["examples/Oligodon_taeniatus_2.png", "Oligodon_taeniatus"],
["examples/Oligodon_taeniatus_3.png", "Oligodon_taeniatus"]
],
inputs=[input_img, gr.Textbox(visible=False, label="Species Label")], # Dummy input for label
outputs=[output_cam, output_labels],
fn=predict,
cache_examples=False,
)
btn.click(fn=predict, inputs=input_img, outputs=[output_cam, output_labels])
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0", server_port=7860)