sandl/private_aggregates_segmentation
0
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 