achouffe's picture
fix: typo in docstring
e33cf6e verified
Raw
History Blame
6.54 kB
"""
Gradio app to showcase the bear face recognition model.
"""
import uuid
from pathlib import Path
from typing import Any, Tuple
import gradio as gr
from PIL import Image
from ui import bearid_ui
from utils import load_models, run_pipeline, setup
def prediction_to_str(bear_ids: list[str], indexed_k_nearest_individuals) -> str:
"""
Turn the prediction into a human friendly string.
"""
nearest_individuals = []
for j in range(len(bear_ids)):
bear_id = bear_ids[j]
nearest_individual = indexed_k_nearest_individuals[bear_id][0]
distance = nearest_individual["distance"]
nearest_individuals.append(
{"bear_id": bear_id, "distance": distance, "data": nearest_individual}
)
distance_str = "\n".join(
[
f"- {n['bear_id']} at distance {n['distance']:.2f} in the embedding space"
for n in nearest_individuals
]
)
return f"The model found that the closest individual is {bear_ids[0]}:\n\n{distance_str}"
def interface_fn(
loaded_models: dict[str, Any],
pil_image: Image.Image,
param_square_dim: int,
param_k: int,
param_n_samples_per_individual: int,
) -> Tuple[Image.Image, Image.Image, Image.Image, str, list[Image.Image], str]:
"""
Main interface function that runs the model on the provided pil_image and
returns the exepected tuple to populate the gradio interface.
Args:
loaded_models (dict[str, Any]): Loaded ML models for the pipeline to
run.
pil_image (PIL): image to run inference on.
param_square_dim (int): square dimension for the cropped chips.
param_k (int): how many individuals to compare it to?
param_n_samples_per_individual (int): how many images for each
individual to sample.
Returns:
pil_image_with_prediction (PIL): image with prediction from the model.
raw_prediction_str (str): string representing the raw prediction from the
model.
"""
result = run_pipeline(
loaded_models=loaded_models,
pil_image=pil_image,
param_square_dim=param_square_dim,
param_k=param_k,
param_n_samples_per_individual=param_n_samples_per_individual,
knn_index_filepath=METRIC_LEARNING_KNN_INDEX_FILEPATH,
)
pil_image_segmented_head = result["stages"]["segmentation"]["output"]["pil_image"]
pil_image_cropped_head = result["stages"]["crop"]["output"]["pil_images"]["resized"]
bear_ids = result["stages"]["identification"]["output"]["bear_ids"]
indexed_samples = result["stages"]["identification"]["output"]["indexed_samples"]
indexed_k_nearest_individuals = result["stages"]["identification"]["output"][
"indexed_k_nearest_individuals"
]
bear_id = bear_ids[0]
output_filepath = OUTPUT_DIR_PREDICTION / f"prediction_{uuid.uuid4()}.png"
output_filepath.parent.mkdir(parents=True, exist_ok=True)
if output_filepath.exists():
output_filepath.unlink()
bearid_ui(
pil_image_chip=pil_image_cropped_head,
indexed_k_nearest_individuals=indexed_k_nearest_individuals,
indexed_samples=indexed_samples,
save_filepath=output_filepath,
k_closest_neighbors=param_k,
)
pil_image_identification_prediction = Image.open(output_filepath)
raw_prediction_str = prediction_to_str(
bear_ids=bear_ids,
indexed_k_nearest_individuals=indexed_k_nearest_individuals,
)
return (
pil_image_segmented_head,
pil_image_cropped_head,
pil_image_identification_prediction,
bear_ids[0],
[Image.open(fp) for fp in indexed_samples[bear_id]],
raw_prediction_str,
)
def examples(dir_examples: Path) -> list[Path]:
"""
List the images from the dir_examples directory.
Returns:
filepaths (list[Path]): list of image filepaths.
"""
return list(dir_examples.glob("*.jpg"))
# Setting up the model artifacts
INPUT_PACKAGED_PIPELINE = Path("./data/09_external/artifacts/packaged_pipeline.zip")
PIPELINE_INSTALL_PATH = Path("./data/06_models/pipeline/metriclearning/")
METRIC_LEARNING_MODEL_FILEPATH = Path(
"./data/06_models/pipeline/metriclearning/bearidentification/model.pt"
)
METRIC_LEARNING_KNN_INDEX_FILEPATH = Path(
"./data/06_models/pipeline/metriclearning/bearidentification/knn.index"
)
INSTANCE_SEGMENTATION_WEIGHTS_FILEPATH = Path(
"./data/06_models/pipeline/metriclearning/bearfacesegmentation/model.pt"
)
setup(
input_packaged_pipeline=INPUT_PACKAGED_PIPELINE,
install_path=PIPELINE_INSTALL_PATH,
)
# Main Gradio interface
DIR_EXAMPLES = Path("data/images/")
DEFAULT_IMAGE_INDEX = 0
PARAM_K = 5
PARAM_N_SAMPLES_PER_INDIVIDUAL = 4
PARAM_SQUARE_DIM = 300
OUTPUT_DIR_PREDICTION = Path("./data/07_model_output/predict/")
with gr.Blocks() as demo:
loaded_models = load_models(
filepath_metric_learning_weights=METRIC_LEARNING_MODEL_FILEPATH,
filepath_segmentation_weights=INSTANCE_SEGMENTATION_WEIGHTS_FILEPATH,
)
image_filepaths = examples(dir_examples=DIR_EXAMPLES)
default_value_input = Image.open(image_filepaths[DEFAULT_IMAGE_INDEX])
input = gr.Image(
value=default_value_input,
type="pil",
label="input image",
sources=["upload", "clipboard"],
)
output_segmentation_image = gr.Image(type="pil", label="model prediction")
output_cropped_image = gr.Image(type="pil", label="cropped bear face")
output_bear_id = gr.Text(label="Most likely BearID")
output_bear_id_samples = gr.Gallery(
label="Similar faces of the predicted BearID", preview=True
)
output_identification_prediction_image = gr.Image(
type="pil",
label="prediction",
)
output_raw = gr.Text(label="raw prediction")
fn = lambda pil_image: interface_fn(
loaded_models=loaded_models,
pil_image=pil_image,
param_square_dim=PARAM_SQUARE_DIM,
param_k=PARAM_K,
param_n_samples_per_individual=PARAM_N_SAMPLES_PER_INDIVIDUAL,
)
gr.Interface(
title="ML pipeline for identifying bears from their faces 🐻",
fn=fn,
inputs=input,
outputs=[
output_segmentation_image,
output_cropped_image,
output_identification_prediction_image,
output_bear_id,
output_bear_id_samples,
output_raw,
],
examples=image_filepaths,
flagging_mode="never",
)
demo.launch()