ChantaroNtw's picture
Update app.py
c9c92c7 verified
Raw
History Blame Contribute Delete
13.1 kB
import os
import io
import gc
import json
import logging
import cv2
import gradio as gr
import numpy as np
import tensorflow as tf
from PIL import Image
from huggingface_hub import (
hf_hub_download,
snapshot_download,
)
# -------------------------------------------------
# TensorFlow Memory Optimization
# -------------------------------------------------
os.environ["TF_CPP_MIN_LOG_LEVEL"] = "2"
os.environ["TF_ENABLE_ONEDNN_OPTS"] = "0"
os.environ["TF_NUM_INTRAOP_THREADS"] = "1"
os.environ["TF_NUM_INTEROP_THREADS"] = "1"
os.environ["OMP_NUM_THREADS"] = "1"
tf.config.threading.set_inter_op_parallelism_threads(1)
tf.config.threading.set_intra_op_parallelism_threads(1)
# -------------------------------------------------
# Logging
# -------------------------------------------------
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("SkinClassifier")
# -------------------------------------------------
# Config
# -------------------------------------------------
MODEL_REPO = "ChantaroNtw/Skin-model"
DERM_MODEL_ID = "google/derm-foundation"
HF_TOKEN = (
os.getenv("HF_TOKEN")
or os.getenv("HUGGINGFACE_HUB_TOKEN")
)
CACHE_DIR = "./hf_cache"
DERM_DIR = "./derm-foundation"
IMAGE_SIZE = (448, 448)
TOPK = 5
BATCH_SIZE = 8
# -------------------------------------------------
# Create directories
# -------------------------------------------------
os.makedirs(CACHE_DIR, exist_ok=True)
os.makedirs(DERM_DIR, exist_ok=True)
# -------------------------------------------------
# Global variables
# -------------------------------------------------
head = None
infer = None
mu = None
sd = None
best_threshold = None
CLASS_NAMES = None
# ============================================================
# Download model files
# ============================================================
HEAD_PATH = hf_hub_download(
repo_id=MODEL_REPO,
filename="mlp_best.keras",
cache_dir=CACHE_DIR,
)
MU_PATH = hf_hub_download(
repo_id=MODEL_REPO,
filename="mu.npy",
cache_dir=CACHE_DIR,
)
SD_PATH = hf_hub_download(
repo_id=MODEL_REPO,
filename="sd.npy",
cache_dir=CACHE_DIR,
)
THRESHOLD_PATH = hf_hub_download(
repo_id=MODEL_REPO,
filename="mlp_thresholds.npy",
cache_dir=CACHE_DIR,
)
LABEL_PATH = hf_hub_download(
repo_id=MODEL_REPO,
filename="class_names.json",
cache_dir=CACHE_DIR,
)
# ============================================================
# Lazy Load Resources
# ============================================================
def load_resources():
"""
โหลดโมเดลทั้งหมดเพียงครั้งเดียว
"""
global head
global infer
global mu
global sd
global best_threshold
global CLASS_NAMES
if infer is not None:
return
logger.info("Loading MLP head...")
import keras
head = keras.saving.load_model(
HEAD_PATH,
compile=False,
)
logger.info("Loading normalization parameters...")
mu = np.load(MU_PATH).astype(np.float32)
sd = np.load(SD_PATH).astype(np.float32)
best_threshold = np.load(
THRESHOLD_PATH
).astype(np.float32)
with open(LABEL_PATH, "r", encoding="utf-8") as f:
CLASS_NAMES = json.load(f)
logger.info("Downloading Derm Foundation...")
print("MODEL_REPO =", MODEL_REPO, type(MODEL_REPO))
print("DERM_MODEL_ID =", DERM_MODEL_ID, type(DERM_MODEL_ID))
derm_path = snapshot_download(
repo_id=DERM_MODEL_ID,
repo_type="model",
allow_patterns=[
"saved_model.pb",
"variables/*",
],
token=HF_TOKEN,
cache_dir=CACHE_DIR,
local_dir=DERM_DIR,
)
if os.path.exists(
os.path.join(DERM_DIR, "saved_model.pb")
):
derm_path = DERM_DIR
else:
derm_path = snapshot_download(...)
logger.info("Loading Derm Foundation...")
derm = tf.saved_model.load(derm_path)
infer = derm.signatures["serving_default"]
logger.info("Models loaded successfully.")
# ============================================================
# Image Utilities
# ============================================================
def preprocess_image(image: Image.Image) -> np.ndarray:
"""
Resize image and normalize to [0,1]
"""
image = image.convert("RGB")
image = image.resize(IMAGE_SIZE)
image = np.asarray(image).astype(np.float32)
image /= 255.0
return image
# ============================================================
# TF Example
# ============================================================
def create_tf_example(image: np.ndarray) -> bytes:
"""
Convert RGB image -> TF Example
"""
image_uint8 = (image * 255).astype(np.uint8)
encoded = tf.io.encode_jpeg(image_uint8).numpy()
example = tf.train.Example(
features=tf.train.Features(
feature={
"image/encoded": tf.train.Feature(
bytes_list=tf.train.BytesList(
value=[encoded]
)
)
}
)
)
return example.SerializeToString()
# ============================================================
# Embedding
# ============================================================
def get_embedding(image: np.ndarray) -> np.ndarray:
"""
Derm Foundation embedding
"""
load_resources()
example = create_tf_example(image)
outputs = infer(
inputs=tf.constant([example])
)
embedding = outputs["embedding"].numpy()[0]
return embedding.astype(np.float32)
# ============================================================
# Normalize Embedding
# ============================================================
def normalize_embedding(
embedding: np.ndarray,
) -> np.ndarray:
embedding = embedding.reshape(1, -1)
embedding = (
embedding - mu
) / (sd + 1e-6)
return embedding.astype(np.float32)
# ============================================================
# Prediction
# ============================================================
def predict_probs(image: Image.Image) -> np.ndarray:
"""
Return probability of every class
"""
load_resources()
image = preprocess_image(image)
embedding = get_embedding(image)
embedding = normalize_embedding(embedding)
probs = head.predict(
embedding,
verbose=0,
)[0]
probs = probs.astype(np.float32)
del image
del embedding
gc.collect()
return probs
# ============================================================
# Top-K Prediction
# ============================================================
def predict_topk(
probs: np.ndarray,
topk: int = TOPK,
):
"""
Return Top-K predictions
"""
idx = np.argsort(probs)[::-1][:topk]
results = []
for i in idx:
results.append(
{
"label": CLASS_NAMES[i],
"prob": float(probs[i]),
}
)
return results
# ============================================================
# Multi-label Prediction
# ============================================================
def predict_multilabel(
probs: np.ndarray,
):
"""
Return labels above threshold
"""
results = []
for i in range(len(CLASS_NAMES)):
if probs[i] >= best_threshold[i]:
results.append(
{
"label": CLASS_NAMES[i],
"prob": float(probs[i]),
}
)
return results
# ============================================================
# Complete Prediction
# ============================================================
def predict(image: Image.Image):
"""
Main prediction function
"""
probs = predict_probs(image)
topk = predict_topk(probs)
multilabel = predict_multilabel(probs)
return {
"topk": topk,
"positives": multilabel,
"probs": {
CLASS_NAMES[i]: float(probs[i])
for i in range(len(CLASS_NAMES))
}
}
# ============================================================
# Heatmap
# ============================================================
PATCH_SIZE = 64
STRIDE = 48
def make_patch_heatmap(image: Image.Image):
load_resources()
image = image.resize((224, 224))
image = np.asarray(image).astype(np.float32) / 255.0
base_embedding = get_embedding(image)
coords = []
tf_examples = []
heatmap = np.zeros((224, 224), dtype=np.float32)
for y in range(0, 224, STRIDE):
for x in range(0, 224, STRIDE):
occluded = image.copy()
occluded[
y:y+PATCH_SIZE,
x:x+PATCH_SIZE
] = 0
tf_examples.append(
create_tf_example(occluded)
)
coords.append((y, x))
embeddings = []
for i in range(
0,
len(tf_examples),
BATCH_SIZE
):
batch = tf.constant(
tf_examples[i:i+BATCH_SIZE]
)
output = infer(
inputs=batch
)["embedding"].numpy()
embeddings.extend(output)
del batch
del output
gc.collect()
for emb, (y, x) in zip(
embeddings,
coords
):
score = np.linalg.norm(
base_embedding - emb
)
heatmap[
y:y+PATCH_SIZE,
x:x+PATCH_SIZE
] = score
heatmap = cv2.normalize(
heatmap,
None,
0,
1,
cv2.NORM_MINMAX
)
return heatmap
# ============================================================
# Overlay Heatmap
# ============================================================
def overlay_heatmap(
image: Image.Image,
heatmap: np.ndarray,
):
image = image.resize((224, 224))
image = np.asarray(image)
heatmap = cv2.GaussianBlur(
heatmap,
(21, 21),
0
)
heatmap = np.power(
heatmap,
1.5
)
heatmap = cv2.normalize(
heatmap,
None,
0,
1,
cv2.NORM_MINMAX
)
heatmap_uint8 = np.uint8(
heatmap * 255
)
heatmap_color = cv2.applyColorMap(
heatmap_uint8,
cv2.COLORMAP_JET
)
mask = heatmap > (
np.mean(heatmap)
+ np.std(heatmap)
)
overlay = image.copy()
overlay[mask] = (
0.6 * overlay[mask]
+ 0.4 * heatmap_color[mask]
).astype(np.uint8)
return overlay
# ============================================================
# Gradio Prediction
# ============================================================
def gradio_predict(
image: Image.Image,
generate_heatmap: bool,
):
if image is None:
return (
{},
None,
"❌ Please upload an image."
)
try:
result = predict(image)
label_result = {
item["label"]: item["prob"]
for item in result["topk"]
}
overlay = None
if generate_heatmap:
heatmap = make_patch_heatmap(image)
overlay = overlay_heatmap(
image,
heatmap,
)
del heatmap
gc.collect()
return (
label_result,
overlay,
"✅ Prediction completed."
)
except Exception as e:
gc.collect()
return (
{},
None,
f"❌ {str(e)}"
)
# ============================================================
# UI
# ============================================================
with gr.Blocks(
title="Skin Disease Classifier"
) as demo:
gr.Markdown(
"# 🧠 Skin Disease Classifier"
)
with gr.Row():
with gr.Column():
image_input = gr.Image(
type="pil",
label="Skin Image",
)
generate_heatmap = gr.Checkbox(
value=False,
label="Generate Heatmap (Slower)"
)
analyze_btn = gr.Button(
"Analyze",
variant="primary",
)
with gr.Column():
prediction_output = gr.Label(
label="Prediction",
num_top_classes=5,
)
heatmap_output = gr.Image(
label="Heatmap"
)
status_output = gr.Markdown()
analyze_btn.click(
fn=gradio_predict,
inputs=[
image_input,
generate_heatmap,
],
outputs=[
prediction_output,
heatmap_output,
status_output,
],
)
if __name__ == "__main__":
demo.launch(
server_name="0.0.0.0",
server_port=7860,
show_error=True,
ssr_mode=False,
)