Hali5's picture
revert model back
982ac6c
Raw
History Blame Contribute Delete
5.06 kB
import spaces
import torch
from torchvision.transforms import v2
import gradio
from PIL import Image
from huggingface_hub import hf_hub_download
from models.linear_predictor import Predictor
import numpy
import os
device = "cuda" if torch.cuda.is_available() else "cpu"
LABELS = [
"Adipose",
"Background",
"Debris",
"Lymphocytes",
"Mucus",
"Smooth Muscle",
"Normal Colon Mucosa",
"Cancer-associated Stroma",
"Colorectal Adenocarcinoma Epithelium",
]
model = Predictor(n_labels=len(LABELS))
model_file = hf_hub_download(
repo_id="Hali5/Mae-Model-MedMNIST-Predictor",
filename="checkpoints/model_linear_v2_epoch_100.pt"
)
model.load_state_dict(torch.load(model_file,map_location=device))
model.to(device)
model.eval()
tf = v2.Compose([
v2.ToImage(),
v2.Resize((64, 64), antialias=True),
v2.ToDtype(torch.float32, scale=True),
])
dataset = numpy.load("test_samples.npz")
images = dataset["images"]
labels = dataset["labels"]
val_dataset = numpy.load("val_samples.npz")
val_images = val_dataset["images"]
val_labels = val_dataset["labels"]
example_rows_test = []
example_rows_val = []
number_of_examples = len(labels)
os.makedirs("ui_examples", exist_ok=True)
for i in range(number_of_examples):
img_array = images[i]
val_img_arr = val_images[i]
label_index = int(labels[i].item() if hasattr(labels[i], 'item') else labels[i])
val_label_index = int(val_labels[i].item() if hasattr(val_labels[i], 'item') else val_labels[i])
if img_array.max() <= 1.0:
img_array = (img_array * 255).astype(numpy.uint8)
val_img_arr = (val_img_arr * 255).astype(numpy.uint8)
else:
img_array = img_array.astype(numpy.uint8)
val_img_arr = val_img_arr.astype(numpy.uint8)
truth_label_text = LABELS[label_index] if label_index < len(LABELS) else f"Class {label_index}"
val_truth_label_text = LABELS[val_label_index] if val_label_index < len(LABELS) else f"Class {val_label_index}"
file_path = f"ui_examples/sample_{i}.jpg"
val_file_path = f"ui_examples/val_sample_{i}.jpg"
Image.fromarray(img_array, "RGB").save(file_path)
Image.fromarray(val_img_arr, "RGB").save(val_file_path)
example_rows_test.append([file_path, truth_label_text])
example_rows_val.append([val_file_path, val_truth_label_text])
@spaces.GPU
def predict(image,truth_labels=True):
if image is None:
return None
# uplouded image
img_tensor = tf(image).unsqueeze(0).to(device)
with torch.no_grad():
outputs = model(img_tensor)
print(outputs.shape)
probabilities = torch.nn.functional.softmax(outputs.squeeze(0), dim=0)
return {LABELS[i]: float(probabilities[i]) for i in range(len(LABELS))}
custom_css = """
.tab img{
object-fit: fill !important;
width: 100% !important;
height: 100% !important;
image-rendering: pixelated !important; /* Forces crisp pixel lines */
}
"""
with gradio.Blocks(css=custom_css) as demo:
gradio.Markdown("# PathMNIST Image Classification")
with gradio.Tab("Predict", elem_classes="tab"):
gradio.Markdown("## Upload a tissue image for classification")
with gradio.Row():
input_img = gradio.Image(height=512, width=512)
with gradio.Column():
output_lbl = gradio.Label(num_top_classes=9)
btn = gradio.Button("Predict")
btn.click(fn=predict, inputs=input_img, outputs=output_lbl)
with gradio.Tab("Examples (Validation)", elem_classes="tab"):
gradio.Markdown("## Select an example below to test the model against the PathMNIST validation dataset")
with gradio.Row():
with gradio.Column():
input_img_val_ex = gradio.Image(label="Selected Test Image", height=512, width=512)
truth_box_val = gradio.Textbox(label="Ground Truth Label", interactive=False)
output_lbl_val_ex = gradio.Label(num_top_classes=9, label="Model Prediction")
gradio.Examples(
examples=example_rows_val,
inputs=[input_img_val_ex, truth_box_val],
outputs=output_lbl_val_ex,
fn=predict,
cache_examples=True,
)
with gradio.Tab("Examples (Test)", elem_classes="tab"):
gradio.Markdown("## Select an example below to test the model against the PathMNIST test dataset")
with gradio.Row():
with gradio.Column():
input_img_test_ex = gradio.Image(label="Selected Test Image", height=512, width=512)
truth_box_test = gradio.Textbox(label="Ground Truth Label", interactive=False)
output_lbl_test_ex = gradio.Label(num_top_classes=9, label="Model Prediction")
gradio.Examples(
examples=example_rows_test,
inputs=[input_img_test_ex, truth_box_test],
outputs=output_lbl_test_ex,
fn=predict,
cache_examples=True,
)
if __name__ == "__main__":
demo.launch()