CoolFace
Apppublic

sandl/private_aggregates_segmentation

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
app.py425 linesDownload Raw Back to root
1import gradio as gr2import tensorflow as tf3import numpy as np4import os5import pickle6 7# import csv8 9# from huggingface_hub import Repository10# from datetime import datetime11import segmentation12from skimage import img_as_ubyte13import pandas as pd14import cv215from image_utils import sort_contours16import matplotlib.pyplot as plt17from image_utils import save_to_zip18 19# from PIL import Image20# import pickle21import torch22 23# Load access tokens24WRITE_TOKEN = os.environ.get("WRITE_PER")  # write25 26# Logs repo path27dataset_url = "https://huggingface.co/datasets/sandl/upload_bond_coat_segmentation"28dataset_path = "logs_bond_coat_segmentation.csv"29 30# Model path31# model_path = 'binary_segmentation_example_low_lr.pth.tar'32# Model with cracks33model_path = "segmentation_with_cracks_example_low_lr_cpu.pth.tar"34model_path = "models/model_best_cracks.pth.tar"35model_path = "../lambdalabs/20230607_multiclass_cracks_patience53_cracks_sam_dataset2.pth.tar"36model_path = "../aws/20230724_multiclass_aggregates_patience50_Data.pth.tar"37model_path = "20230725_multiclass_aggregates_patience80_filab.pth.tar"38num_classes = 139model, preprocessing_fn = segmentation.load_segmentation_model(40    model_path,41    classes=num_classes,42    # device=torch.device("cpu")43)44 45# Number of images to process46NUM_IMAGES = 10  # This number is updated based on the number of images passed47SCALE = float(30 / 260)48 49 50def write_logs(message, log_type="Prediction"):51    """52    Write logs53    """54    # with Repository(local_dir="data", clone_from=dataset_url, use_auth_token=WRITE_TOKEN).commit(commit_message="from private", blocking=False):55    #    with open(dataset_path, "a") as csvfile:56    #            writer = csv.DictWriter(csvfile, fieldnames=["name", "message", "time"])57    #            writer.writerow(58    #                {"name": log_type, "message": message, "time": str(datetime.now())}59    #            )60    print(message)61    return62 63 64def compute_cracks_statistics(input_image, prediction_mask, cracks_index_position=0):65    """66    Compute statistics on the cracks based on the prediction mask67    """68    print("*********")69    print("Computing cracks")70    cracks_mask = prediction_mask[:, :, cracks_index_position]71    cracks_gray = np.where(cracks_mask, 200, 50)  # Values selected to be compatible with the thresholds72    cracks_gray2 = cracks_gray.astype("uint8")73    _, binary = cv2.threshold(cracks_gray2, 150, 255, cv2.THRESH_BINARY)  # Necessary for the findcontours method74    contours, hierarchy = cv2.findContours(binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)75 76    contour_lengths = []77    contour_thicknesses = []78    dist_ls = []79 80    contours_sorted, _ = sort_contours(contours, "left-to-right")81    i = 082    for contour in contours_sorted:83        # Cracks lengths84        perimeter = cv2.arcLength(contour, True)85        length = perimeter / 286        contour_lengths.append(round(SCALE * length, 3))87        # Cracks thickness88        x, y, w, h = cv2.boundingRect(contour)89        contour_thicknesses.append(round(SCALE * h, 3))90        # Distances91        if i < len(contours_sorted) - 1:92            cnt1 = contours_sorted[i][0][0]93            cnt2 = contours_sorted[i + 1][0][0]94            dist = np.linalg.norm(cnt1 - cnt2) * SCALE95            print(cnt1, cnt2, dist)96            dist_ls.append(dist)97        i += 198    df_stats_per_crack = pd.DataFrame(99        np.transpose([contour_lengths, contour_thicknesses]),100        columns=["Crack length (micrometers)", "Crack thickness (micrometers)"],101    )102    print(df_stats_per_crack.shape)103    cracks_density_length = (np.sum(contour_lengths)) / np.shape(input_image)[0]104    cracks_density_surface = np.sum(cracks_mask == 1) / (np.shape(input_image)[0] * np.max(contour_lengths))105    cracks_stats_single = {106        "mean_distance_cracks": round(np.mean(dist_ls), 3),107        "std_distance_cracks": round(np.std(dist_ls), 3),108        "contour lengths (micrometers)": contour_lengths,109        "average cracks length (micrometers)": round(np.mean(contour_lengths), 3),110        "cracks thicknesses (micrometers)": contour_thicknesses,111        "average cracks thicknesses (micrometers)": round(np.mean(contour_thicknesses), 3),112        "cracks density (surface %)": cracks_density_surface,113        "cracks density (length ratio)": cracks_density_length,114        "cracks_detailed_dataframe": df_stats_per_crack,115    }116    return cracks_stats_single117 118 119def plot_distribution(distrib, title="test"):120    """121    Returns a matlplotlib object122    """123    fig, ax = plt.subplots(figsize=(4.5, 2.5))  # Figsize selected manually to fit the interface124    ax.hist(distrib)125    fig.suptitle(f"Distribution of {title}")126    return fig127 128 129def predict_multiple_image(input_files, request: gr.Request):130    if type(input_files) != list:131        input_files = [input_files]132    n_images = len(input_files)133    print(n_images)134    i = 0135    file_names = []136    images_batched = []137    for file in input_files:138        content = cv2.imread(file.name)139        content_expanded = np.expand_dims(img_as_ubyte(content), axis=0)  # To be able to do batch processing140        images_batched.append(content_expanded)141        file_names.append(file.name.split("/")[-1])142        i += 1143    images_batched = np.concatenate(images_batched, axis=0)  # Concatenate for batch processing144    print(file_names)145 146    # Temporary load the predictions from a static file147    # predictions_batched = segmentation.segmentation_models_inference(148    #     images_batched,149    #     model,150    #     preprocessing_fn,151    #     batch_size=4,152    #     patch_size=512,153    #     num_classes=num_classes,154    #     # device=torch.device("cpu"),155    # )156    with open("aggregates_prediction.pkl", "rb") as file:157        predictions_batched = pickle.load(file)158        # pred = predictions_batched[0]159        # predictions_batched = [pred] * n_images160 161    # Put the mask over the raw images162    output_image_list = []163    for i in range(n_images):164        output_image = images_batched[i].copy()165        pred_single = predictions_batched[i]166        output_image[pred_single[:, :, 0], 0] = 255167        # output_image[pred_single[:, :, 0], 1] = 0168        # output_image[pred_single[:, :, 0], 2] = 0169        output_image_list.append(output_image)170 171    cracks_stats_list = []172    # Compute the statistics per image173    for i in range(n_images):174        cracks_stats_single = compute_cracks_statistics(images_batched[i], predictions_batched[i])175        cracks_stats_list.append(cracks_stats_single)176 177    # Add cracks length histogram over all images178    cracks_length = [179        list(cracks_stats_list[i]["cracks_detailed_dataframe"]["Crack length (micrometers)"]) for i in range(n_images)180    ]181    cracks_length = [item for sublist in cracks_length for item in sublist]182    cracks_length_plot = plot_distribution(cracks_length, "particles length accross all images")183 184    # Compute statistics at the batch level185    (186        mean_cracks,187        std_cracks,188        cracks_density_surface,189        cracks_density_length,190        mean_cracks_lengths,191        mean_cracks_thicknesses,192    ) = compute_batch_statistics(cracks_stats_list)193    print("aggregated crack density", cracks_density_length)194 195    if request is not None:196        message = f"{request.username}_{request.client.host}"197        # print("Logs are disabled")198        write_logs(message)199        print(message)200 201    print("Prediction done")202    detailed_stats_df = pd.DataFrame(cracks_stats_list)203    detailed_stats_df.drop(columns=["cracks_detailed_dataframe"], inplace=True)204    detailed_stats_df.rename(205        columns={206            "mean_distance_cracks": "mean_distance_cracks (in micrometers)",207            "std_distance_cracks": "std_distance_cracks (in micrometers)",208            "cracks density (surface %)": "cracks density (surface %)",209            "cracks density (length ratio)": "cracks density (length ratio)",210            "average cracks length": "average cracks length (micometer)",211        },212        inplace=True,213    )214    detailed_stats_df.index = file_names215    aggregated_stats_df = pd.DataFrame(216        [217            [218                round(mean_cracks, 3),219                round(std_cracks, 3),220                round(cracks_density_surface * 100, 3),221                round(cracks_density_length, 3),222                round(mean_cracks_lengths, 3),223                round(mean_cracks_thicknesses, 3),224            ]225        ],226        columns=[227            "Average distance between cracks (in micrometers)",228            "Standard deviation of distance between cracks (in micrometers)",229            "Cracks density (surface %)",230            "Cracks density (length ratio)",231            "Average cracks length (micrometers)",232            "Average cracks thickness (micrometers)",233        ],234    )235    # Prepare the dataframes for the detailed statistics of cracks per image236    cracks_detailed_list = [cracks_stats_list[i]["cracks_detailed_dataframe"] for i in range(n_images)]237 238    # Generate the zip folder containing all the statistics239    zip_file = save_to_zip(240        detailed_stats_df, aggregated_stats_df, cracks_length_plot, output_image_list, file_names, cracks_detailed_list241    )242    input_images_list = [images_batched[i, :, :, :] for i in range(n_images)]243    none_images = [None] * (NUM_IMAGES - n_images)244    output_list = (245        [zip_file, detailed_stats_df, aggregated_stats_df, cracks_length_plot]246        + input_images_list247        + none_images248        + output_image_list249        + none_images250    )251    print("Prediction final done ---------")252    # output_list = [zip_file, aggregated_stats_df] + input_images_list + none_images + output_image_list + none_images253    return tuple(output_list)254 255 256def compute_batch_statistics(stats_list):257    """258    Compute the statistics of the input image batch259    """260 261    df_stats = pd.DataFrame(stats_list)262    return (263        df_stats["mean_distance_cracks"].median(),264        df_stats["std_distance_cracks"].median(),265        df_stats["cracks density (surface %)"].mean(),266        df_stats["cracks density (length ratio)"].mean(),267        df_stats["average cracks length (micrometers)"].mean(),268        df_stats["average cracks thicknesses (micrometers)"].mean(),269    )270 271 272osium_theme_colors = gr.themes.Color(273        c50="#e4f3fa",  # Dataframe background cell content - light mode only274        c100="#e4f3fa",  # Top corner of clear button in light mode + markdown text in dark mode275        c200="#a1c6db",  # Component borders276        c300="#FFFFFF",  #277        c400="#e4f3fa",  # Footer text278        c500="#0c1538",  # Text of component headers in light mode only279        c600="#a1c6db",  # Top corner of button in dark mode280        c700="#475383",  # Button text in light mode + component borders in dark mode281        c800="#0c1538",  # Markdown text in light mode282        c900="#a1c6db",  # Background of dataframe - dark mode283        c950="#0c1538",284)  # Background in dark mode only285# secondary color used for highlight box content when typing in light mode, and download option in dark mode286# primary color used for login button in dark mode287osium_theme = gr.themes.Default(primary_hue="cyan", secondary_hue="cyan", neutral_hue=osium_theme_colors)288 289css_styling = """#submit {background: #1eccd8} 290    #submit:hover {background: #a2f1f6} 291    .output-image, .input-image, .image-preview {height: 250px !important}292    .output-plot {height: 250px !important}293    #interpretation {height: 250px !important}"""294 295osium_theme = gr.themes.Default(primary_hue="cyan", secondary_hue="cyan", neutral_hue=osium_theme_colors)296page_title = "Particles aggregates segmentation"297favicon_path = "osiumai_favicon.ico"298logo_path = "osiumai_logo.jpg"299html = f"""<html> <link rel="icon" type="image/x-icon" href="file={favicon_path}">300<img src='file={logo_path}' alt='Osium AI logo' width='200' height='100'> </html>"""301 302 303with gr.Blocks(css=css_styling, page_title=page_title, theme=osium_theme) as demo:304    #gr.HTML(html)305    gr.Markdown("# <p style='text-align: center;'>Identify the particles aggregates in your microscopy images</p>")306    gr.Markdown("This AI model computes the statistics of the particles aggregates in the images provided")307    with gr.Row():308        clear_button = gr.Button("Clear")309        prediction_button = gr.Button("Predict", elem_id="submit")310    with gr.Row():311        with gr.Column():312            gr.Markdown("### Your input files")313            input_file = gr.File(label="Your input files", file_count="multiple", elem_id="input_files")314            gr.Examples(["aggregates1.png"], input_file)315    with gr.Row():316        with gr.Column():317            gr.Markdown("### Your input images")318            input_image1 = gr.Image(elem_classes="input-image")319            input_image2 = gr.Image(elem_classes="input-image")320            input_image3 = gr.Image(elem_classes="input-image")321            input_image4 = gr.Image(elem_classes="input-image")322            input_image5 = gr.Image(elem_classes="input-image")323            input_image6 = gr.Image(elem_classes="input-image")324            input_image7 = gr.Image(elem_classes="input-image")325            input_image8 = gr.Image(elem_classes="input-image")326            input_image9 = gr.Image(elem_classes="input-image")327            input_image10 = gr.Image(elem_classes="input-image")328        with gr.Column():329            gr.Markdown("### The predicted masks")330            output_image1 = gr.Image(elem_classes="output-image")331            output_image2 = gr.Image(elem_classes="output-image")332            output_image3 = gr.Image(elem_classes="output-image")333            output_image4 = gr.Image(elem_classes="output-image")334            output_image5 = gr.Image(elem_classes="output-image")335            output_image6 = gr.Image(elem_classes="output-image")336            output_image7 = gr.Image(elem_classes="output-image")337            output_image8 = gr.Image(elem_classes="output-image")338            output_image9 = gr.Image(elem_classes="output-image")339            output_image10 = gr.Image(elem_classes="output-image")340        with gr.Column():341            gr.Markdown("### The images' statistics")342            with gr.Row():343                output_file = gr.File(label="Download all your results and statistics")344            with gr.Row():345                aggregated_stats = gr.DataFrame(label="Aggregated statistics")346            with gr.Row():347                detailed_stats = gr.DataFrame(label="Per image statistics")348            with gr.Row():349                length_plot = gr.Plot(elem_classes="output-plot", type="matplotlib")350 351    prediction_button.click(352        fn=predict_multiple_image,353        inputs=[input_file],354        outputs=[355            output_file,356            detailed_stats,357            aggregated_stats,358            length_plot,359            input_image1,360            input_image2,361            input_image3,362            input_image4,363            input_image5,364            input_image6,365            input_image7,366            input_image8,367            input_image9,368            input_image10,369            output_image1,370            output_image2,371            output_image3,372            output_image4,373            output_image5,374            output_image6,375            output_image7,376            output_image8,377            output_image9,378            output_image10,379        ],380        show_progress=True,381    )382 383    clear_button.click(384        lambda x: [gr.update(value=None)] * (2 * NUM_IMAGES + 3),385        [],386        [387            input_file,388            input_image1,389            input_image2,390            input_image3,391            input_image4,392            input_image5,393            input_image6,394            input_image7,395            input_image8,396            input_image9,397            input_image10,398            detailed_stats,399            aggregated_stats,400            output_image1,401            output_image2,402            output_image3,403            output_image4,404            output_image5,405            output_image6,406            output_image7,407            output_image8,408            output_image9,409            output_image10,410            output_file,411        ],412    )413 414 415def authenticate(username, password):416    return (username == "snolf") & (password == "matterai23")417 418 419if __name__ == "__main__":420    demo.queue(concurrency_count=2)421    demo.launch(422        server_port=7860, 423        # server_name="0.0.0.0"424        )425