CellSamCells / nellie2.py
champ7's picture
Create nellie2.py
f680f54 verified
Raw
History Blame
14.3 kB
import os
import glob
import logging
import random
import cv2
import numpy as np
import matplotlib.pyplot as plt
import tifffile as tiff
from scipy.sparse import csr_matrix
from tabulate import tabulate
from tqdm import tqdm
from huggingface_hub import snapshot_download
from nellie.im_info.verifier import FileInfo, ImInfo
from nellie.segmentation.filtering import Filter
from nellie.segmentation.labelling import Label
from skimage.filters import rank, sobel
from skimage import exposure
from skimage.morphology import disk, binary_dilation
from skimage.segmentation import find_boundaries
from skimage.util import img_as_ubyte
# Logging setup
logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s", datefmt="%H:%M:%S")
logger = logging.getLogger(__name__)
# =====================================================================
# Preprocessing Logic
# =====================================================================
def preprocess_method_1(image_obj: np.ndarray) -> np.ndarray:
""" W2: Organelles (Red Fluorescence) -> CLAHE & Top-Hat preprocessing """
if image_obj.ndim == 3:
image_obj = image_obj[0]
p1, p99 = np.percentile(image_obj, (1, 99.9))
robust_img = exposure.rescale_intensity(image_obj, in_range=(p1, p99), out_range=(0, 255)).astype(np.uint8)
clahe = cv2.createCLAHE(clipLimit=5.0, tileGridSize=(8, 8))
clahe_img = clahe.apply(robust_img)
return clahe_img
def preprocess_method_2(img_input: np.ndarray) -> np.ndarray:
img_min, img_max = img_input.min(), img_input.max()
img_normalized = (img_input - img_min) / (img_max - img_min if img_max > img_min else 1.0)
img_8bit = img_as_ubyte(img_normalized)
nellie_input = rank.entropy(img_8bit, disk(25))
edge_map = sobel(img_normalized)
texture_features = nellie_input * (1.0 + edge_map * 5.0)
tex_min, tex_max = texture_features.min(), texture_features.max()
tex_normalized = (texture_features - tex_min) / (tex_max - tex_min if tex_max > tex_min else 1.0)
return (tex_normalized * 255).astype(np.uint8)
def preprocess_combined(img_raw: np.ndarray) -> np.ndarray:
return preprocess_method_2(preprocess_method_1(img_raw))
# =====================================================================
# Vectorized Metric Evaluation
# =====================================================================
def evaluate_predictions(pred_mask: np.ndarray, gt_mask: np.ndarray, iou_threshold: float = 0.5):
"""Vectorized instance-level matching using strict IoU >= 0.5"""
pred_mask, gt_mask = pred_mask.astype(np.int32), gt_mask.astype(np.int32)
pred_labels = np.unique(pred_mask)[np.unique(pred_mask) != 0]
gt_labels = np.unique(gt_mask)[np.unique(gt_mask) != 0]
num_pred, num_gt = len(pred_labels), len(gt_labels)
pred_bg, gt_bg = pred_mask > 0, gt_mask > 0
union_pixel = np.logical_or(pred_bg, gt_bg).sum()
global_pixel_iou = np.logical_and(pred_bg, gt_bg).sum() / union_pixel if union_pixel > 0 else 0.0
if num_gt == 0 and num_pred == 0: return global_pixel_iou, 1.0, 1.0, 1.0, num_pred, num_gt
if num_gt == 0 or num_pred == 0: return global_pixel_iou, 0.0, 0.0, 0.0, num_pred, num_gt
overlapping = gt_bg & pred_bg
if not np.any(overlapping): return global_pixel_iou, 0.0, 0.0, 0.0, num_pred, num_gt
gt_id_map = {id_: i for i, id_ in enumerate(gt_labels)}
pred_id_map = {id_: j for j, id_ in enumerate(pred_labels)}
gt_indices = np.array([gt_id_map[x] for x in gt_mask[overlapping]])
pred_indices = np.array([pred_id_map[x] for x in pred_mask[overlapping]])
intersection = csr_matrix((np.ones(len(gt_indices), dtype=np.int32), (gt_indices, pred_indices)), shape=(num_gt, num_pred)).toarray()
gt_vec = np.array([np.bincount(gt_mask.ravel())[id_] for id_ in gt_labels])[:, None]
pred_vec = np.array([np.bincount(pred_mask.ravel())[id_] for id_ in pred_labels])[None, :]
union = gt_vec + pred_vec - intersection
iou_matrix = np.divide(intersection, union, out=np.zeros_like(intersection, dtype=float), where=union != 0)
matches = iou_matrix > iou_threshold
tp = np.sum(np.any(matches, axis=0))
precision = tp / num_pred if num_pred > 0 else 0.0
recall = tp / num_gt if num_gt > 0 else 0.0
instance_f1 = (2 * precision * recall) / (precision + recall) if (precision + recall) > 0 else 0.0
return global_pixel_iou, precision, recall, instance_f1, num_pred, num_gt
# =====================================================================
# Dataset Prep & Core Pipeline
# =====================================================================
def prepare_datasets(raw_tif_files, base_dir="generated_datasets"):
dirs = {
"Raw": os.path.join(base_dir, "dataset_raw"),
"Prep 1": os.path.join(base_dir, "dataset_prep1"),
"Prep 2": os.path.join(base_dir, "dataset_prep2"),
"Prep 1->2": os.path.join(base_dir, "dataset_prep1_2")
}
for d in dirs.values(): os.makedirs(d, exist_ok=True)
for filepath in tqdm(raw_tif_files, desc="Building Datasets"):
fname = os.path.basename(filepath)
if os.path.exists(os.path.join(dirs["Prep 1->2"], fname)):
continue
img_raw = np.squeeze(tiff.imread(filepath))
tiff.imwrite(os.path.join(dirs["Raw"], fname), img_raw)
tiff.imwrite(os.path.join(dirs["Prep 1"], fname), preprocess_method_1(img_raw))
tiff.imwrite(os.path.join(dirs["Prep 2"], fname), preprocess_method_2(img_raw))
tiff.imwrite(os.path.join(dirs["Prep 1->2"], fname), preprocess_combined(img_raw))
return dirs
def run_nellie_experiment(dataset_dir, combined_data_path, run_native_filter, target_filenames):
files = [os.path.join(dataset_dir, fname) for fname in target_filenames]
results, saved_masks = [], {}
for filepath in tqdm(files, desc=f"Evaluating", leave=False):
if not os.path.exists(filepath): continue
filename = os.path.basename(filepath)
gt_mask_path = os.path.join(combined_data_path, filename.replace("_w2.TIF", "_w3_seg.npy"))
if not os.path.exists(gt_mask_path): continue
f_info = FileInfo(filepath=filepath)
f_info.find_metadata()
f_info.load_metadata()
img_array = np.squeeze(tiff.imread(filepath))
f_info.axes = "ZYX" if img_array.ndim == 3 else "YX"
f_info.good_axes, f_info.good_dims = True, True
im_info = ImInfo(file_info=f_info)
if run_native_filter:
frangi_filter = Filter(im_info=im_info, device="auto")
frangi_filter.run()
else:
prep_path = im_info.create_output_path(pipeline_path="im_preprocessed", ext=".ome.tif")
im_info.allocate_memory(output_path=prep_path, data=img_array, dtype=str(img_array.dtype))
label_processor = Label(im_info=im_info, otsu_thresh_intensity=False, min_radius_um=10.0, device="auto")
label_processor.run()
pred_mask = np.squeeze(tiff.imread(im_info.pipeline_paths["im_instance_label"]))
gt_data = np.load(gt_mask_path, allow_pickle=True)
if gt_data.ndim == 0 and isinstance(gt_data.item(), dict):
gt_data = gt_data.item().get("masks", gt_data.item())
gt_to_compare = np.squeeze(gt_data)
metrics = evaluate_predictions(pred_mask, gt_to_compare, iou_threshold=0.5)
results.append(metrics)
saved_masks[filename] = (pred_mask, metrics[4])
return results, saved_masks
# =====================================================================
# Plot Single Experiment
# =====================================================================
def plot_single_experiment_visual(sample_filename, combined_data_path, data_src_dir, raw_data_dir, config_name, pred_mask, pred_count):
logger.info(f"Generating visual plot for: {config_name}")
# --- Load Data ---
gt_data = np.load(os.path.join(combined_data_path, sample_filename.replace("_w2.TIF", "_w3_seg.npy")), allow_pickle=True)
if gt_data.ndim == 0 and isinstance(gt_data.item(), dict):
gt_data = gt_data.item().get("masks", gt_data.item())
gt_mask = np.squeeze(gt_data)
gt_count = len(np.unique(gt_mask)[np.unique(gt_mask) != 0])
feed_img = tiff.imread(os.path.join(data_src_dir, sample_filename))
raw_img = tiff.imread(os.path.join(raw_data_dir, sample_filename))
# --- Ground Truth High-Def Assets ---
gt_binary = gt_mask > 0
gt_boundaries = find_boundaries(gt_binary, mode="outer")
gt_thick_boundaries = binary_dilation(gt_boundaries, disk(1))
gt_masked_fill = np.ma.masked_where(~gt_binary, gt_binary)
gt_masked_edges = np.ma.masked_where(~gt_thick_boundaries, gt_thick_boundaries)
# --- Prediction High-Def Assets ---
pred_binary = pred_mask > 0
pred_boundaries = find_boundaries(pred_binary, mode="outer")
pred_thick_boundaries = binary_dilation(pred_boundaries, disk(1))
pred_masked_fill = np.ma.masked_where(~pred_binary, pred_binary)
pred_masked_edges = np.ma.masked_where(~pred_thick_boundaries, pred_thick_boundaries)
cmap_choice = "gray" if "Prep 1" in data_src_dir or "Raw" in data_src_dir else "magma"
fig, axes = plt.subplots(1, 3, figsize=(18, 6))
# 1. Ground Truth Overlay (on Raw Image)
axes[0].imshow(raw_img, cmap="gray")
axes[0].imshow(gt_masked_fill, cmap="summer", alpha=0.25)
axes[0].imshow(gt_masked_edges, cmap="Greens", alpha=0.9)
axes[0].set_title(f"GT Overlay on Raw (Count: {gt_count})", fontsize=12, fontweight="bold")
axes[0].axis("off")
# 2. Input Image Alone (The preprocessed image actually fed to Nellie)
axes[1].imshow(feed_img, cmap=cmap_choice)
axes[1].set_title(f"Data Fed into Nellie", fontsize=12, fontweight="bold")
axes[1].axis("off")
# 3. Predicted Mask Overlay (on Raw Image)
axes[2].imshow(raw_img, cmap="gray")
axes[2].imshow(pred_masked_fill, cmap="autumn", alpha=0.25)
axes[2].imshow(pred_masked_edges, cmap="Wistia", alpha=0.9)
axes[2].set_title(f"Predicted Mask on Raw (Count: {pred_count})", fontsize=12, fontweight="bold")
axes[2].axis("off")
plt.suptitle(f"EXPERIMENT: {config_name}\nFile: {sample_filename}", fontsize=16, fontweight="bold")
plt.tight_layout()
plt.show()
# =====================================================================
# Main Execution Flow
# =====================================================================
def main():
n_samples = 10 # Evaluate 10 images
logger.info("Downloading dataset...")
dataset_dir = snapshot_download(repo_id="champ7/celldatamag", repo_type="dataset", allow_patterns="combined_data/*")
combined_data_path = os.path.join(dataset_dir, "combined_data")
# Extract complete list and sample N files
raw_tif_files = sorted(glob.glob(os.path.join(combined_data_path, "*.TIF")))
if n_samples and n_samples < len(raw_tif_files):
random.seed(42) # Ensure the same random subset is used across runs
raw_tif_files = random.sample(raw_tif_files, n_samples)
logger.info(f"Subsampled dataset to {n_samples} random images.")
target_filenames = [os.path.basename(f) for f in raw_tif_files]
sample_filename = target_filenames[0]
dirs = prepare_datasets(raw_tif_files)
configs = [
{"name": "Prep 1 Only (Bypassing Nellie Filter)", "data_src": "Prep 1", "run_filter": False},
{"name": "Prep 2 Only (Bypassing Nellie Filter)", "data_src": "Prep 2", "run_filter": False},
{"name": "Prep 1->2 Only (Bypassing Nellie Filter)", "data_src": "Prep 1->2", "run_filter": False},
{"name": "Nellie Native Only (Raw -> Nellie Filter)", "data_src": "Raw", "run_filter": True},
{"name": "Prep 1 + Nellie Native Filter", "data_src": "Prep 1", "run_filter": True},
{"name": "Prep 2 + Nellie Native Filter", "data_src": "Prep 2", "run_filter": True},
{"name": "Prep 1->2 + Nellie Native Filter", "data_src": "Prep 1->2", "run_filter": True},
]
aggregate_table = []
headers = ["Strategy Configuration", "Total GT Cells", "Total Pred Cells", "Mean Pixel IoU", "Mean Precision", "Mean Recall", "Mean F1"]
for config in configs:
logger.info(f"Running strategy: {config['name']}")
metrics, saved_masks = run_nellie_experiment(
dataset_dir=dirs[config["data_src"]],
combined_data_path=combined_data_path,
run_native_filter=config["run_filter"],
target_filenames=target_filenames
)
ious, precs, recs, f1s = [m[0] for m in metrics], [m[1] for m in metrics], [m[2] for m in metrics], [m[3] for m in metrics]
total_pred, total_gt = sum([m[4] for m in metrics]), sum([m[5] for m in metrics])
current_result_row = [
config["name"],
total_gt, total_pred,
f"{np.mean(ious):.4f}", f"{np.mean(precs):.4f}", f"{np.mean(recs):.4f}", f"{np.mean(f1s):.4f}"
]
aggregate_table.append(current_result_row)
# Immediate output for this run
print(f"\n" + "="*80)
print(f" EXPERIMENT RESULTS: {config['name']}")
print(f"="*80)
print(tabulate([current_result_row], headers=headers, tablefmt="fancy_grid"))
print("\n")
# Visual Plot
pred_mask_sample, pred_count_sample = saved_masks[sample_filename]
plot_single_experiment_visual(
sample_filename=sample_filename,
combined_data_path=combined_data_path,
data_src_dir=dirs[config["data_src"]],
raw_data_dir=dirs["Raw"],
config_name=config["name"],
pred_mask=pred_mask_sample,
pred_count=pred_count_sample
)
# Final Summary Table
print("\n" + "=" * 115)
print(" " * 35 + f"FINAL 7-WAY EVALUATION REPORT (N={len(target_filenames)} images)")
print("=" * 115)
print(tabulate(aggregate_table, headers=headers, tablefmt="fancy_grid"))
if __name__ == "__main__":
main()