Skip to content

Commit 28872bc

Browse files
committed
Added cryosegnet-style guided denoise filter command with multiprocessing layer.
1 parent 567644d commit 28872bc

3 files changed

Lines changed: 172 additions & 2 deletions

File tree

partinet/__init__.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -61,8 +61,8 @@ def main():
6161
@click.option("--labels", required=True, help="Path to the labels directory")
6262
@click.option("--images", required=True, help="Path to the images directory")
6363
@click.option("--output", required=True, help="Path to the output directory")
64-
def preprocess(labels, images, output):
65-
click.echo("Preprocessing the micrographs...")
64+
def train_split(labels, images, output):
65+
click.echo("Splitting micrographs for training and validation...")
6666
import partinet.split_train
6767
partinet.split_train.main(labels, images, output)
6868

@@ -76,6 +76,14 @@ def star(labels, images, output,conf):
7676
import partinet.star_file
7777
partinet.star_file.main(labels,images,output,conf)
7878

79+
@main.command()
80+
@click.option("--source", required=True, help="Path to Raw micrographs")
81+
@click.option('--project', required=True, help='save denoised micrographs to project/denoised', show_default=True))
82+
def denoise(source, project):
83+
click.echo("Denoising micrographs...")
84+
import partinet.pooled_denoise_proc
85+
partinet.pooled_denoise_proc(source,project)
86+
7987

8088
@main.group()
8189
def train():

partinet/guided_denoiser.py

Lines changed: 97 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,97 @@
1+
# Adapted from CryoSegNet
2+
# https://github.com/jianlin-cheng/CryoSegNet
3+
4+
import numpy as np
5+
import mrcfile
6+
import cv2
7+
from numpy.fft import fft2, ifft2
8+
from scipy.signal import gaussian
9+
10+
def transform(image):
11+
i_min = image.min()
12+
i_max = image.max()
13+
14+
image = ((image - i_min)/(i_max - i_min)) * 255
15+
return image.astype(np.uint8)
16+
17+
18+
def standard_scaler(image):
19+
kernel_size = 9
20+
image = cv2.GaussianBlur(image, (kernel_size, kernel_size), 0)
21+
mu = np.mean(image)
22+
sigma = np.std(image)
23+
image = (image - mu)/sigma
24+
image = transform(image).astype(np.uint8)
25+
return image
26+
27+
def contrast_enhancement(image):
28+
enhanced_image = cv2.fastNlMeansDenoising(image, None, h=10, templateWindowSize=7, searchWindowSize=21)
29+
30+
return enhanced_image
31+
32+
33+
def gaussian_kernel(kernel_size = 3):
34+
h = gaussian(kernel_size, kernel_size / 3).reshape(kernel_size, 1)
35+
h = np.dot(h, h.transpose())
36+
h /= np.sum(h)
37+
return h
38+
39+
def wiener_filter(img, kernel, K):
40+
kernel /= np.sum(kernel)
41+
dummy = np.copy(img)
42+
dummy = fft2(dummy)
43+
kernel = fft2(kernel, s = img.shape)
44+
kernel = np.conj(kernel) / (np.abs(kernel) ** 2 + K)
45+
dummy = dummy * kernel
46+
dummy = np.abs(ifft2(dummy))
47+
return dummy
48+
49+
def clahe(image):
50+
clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(16,16))
51+
52+
# Apply CLAHE to the image
53+
img_equalized = clahe.apply(transform(image))
54+
return img_equalized
55+
56+
def guided_filter(input_image, guidance_image, radius=20, epsilon=0.1):
57+
# Convert images to float32
58+
input_image = input_image.astype(np.float32) / 255.0
59+
guidance_image = guidance_image.astype(np.float32) / 255.0
60+
61+
# Compute mean values of the guidance image and input image
62+
mean_guidance = cv2.boxFilter(guidance_image, -1, (radius, radius))
63+
mean_input = cv2.boxFilter(input_image, -1, (radius, radius))
64+
65+
# Compute correlation and covariance of the guidance and input images
66+
mean_guidance_input = cv2.boxFilter(guidance_image * input_image, -1, (radius, radius))
67+
covariance_guidance_input = mean_guidance_input - mean_guidance * mean_input
68+
69+
# Compute squared mean of the guidance image
70+
mean_guidance_sq = cv2.boxFilter(guidance_image * guidance_image, -1, (radius, radius))
71+
variance_guidance = mean_guidance_sq - mean_guidance * mean_guidance
72+
73+
# Compute weights and mean of the weights
74+
a = covariance_guidance_input / (variance_guidance + epsilon)
75+
b = mean_input - a * mean_guidance
76+
mean_a = cv2.boxFilter(a, -1, (radius, radius))
77+
mean_b = cv2.boxFilter(b, -1, (radius, radius))
78+
79+
# Compute the filtered image
80+
output_image = mean_a * guidance_image + mean_b
81+
82+
return transform(output_image)
83+
84+
85+
def denoise(image_path):
86+
kernel = gaussian_kernel(kernel_size = 9)
87+
image = mrcfile.read(image_path)
88+
image = image.T
89+
image = np.rot90(image)
90+
normalized_image = standard_scaler(np.array(image))
91+
contrast_enhanced_image = contrast_enhancement(normalized_image)
92+
weiner_filtered_image = wiener_filter(contrast_enhanced_image, kernel, K = 30)
93+
clahe_image = clahe(weiner_filtered_image)
94+
guided_filter_image = guided_filter(clahe_image, weiner_filtered_image)
95+
96+
return guided_filter_image
97+

partinet/pooled_denoise_proc.py

Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,65 @@
1+
import os
2+
import cv2
3+
from pathlib import Path
4+
from guided_denoiser import denoise
5+
import logging
6+
import multiprocessing
7+
import argparse
8+
import gc
9+
from concurrent.futures import ProcessPoolExecutor
10+
11+
MAX_WORKERS = max(1, multiprocessing.cpu_count() // 2)
12+
13+
def clahe_denoise(args):
14+
try:
15+
src_path, dest_path = args
16+
denoised = denoise(src_path)
17+
cv2.imwrite(dest_path, denoised)
18+
logging.info(f"Processed image {src_path} to dest. {dest_path}")
19+
del denoised
20+
gc.collect()
21+
except Exception as e:
22+
logging.error(f"Failed to process {src_path}: {str(e)}")
23+
24+
def process_directory(micrographs_dir, clahe_denoised_dir):
25+
os.makedirs(clahe_denoised_dir, exist_ok=True)
26+
logging.info(f"Directory ready: {clahe_denoised_dir}")
27+
28+
tasks = []
29+
for file_name in os.listdir(micrographs_dir):
30+
if file_name.endswith(".mrc"):
31+
src_path = os.path.join(micrographs_dir, file_name)
32+
dest_path = os.path.join(clahe_denoised_dir, file_name.replace(".mrc", ".png"))
33+
if os.path.exists(dest_path):
34+
logging.info(f"{dest_path} already exists!")
35+
else:
36+
tasks.append((src_path, dest_path))
37+
38+
with ProcessPoolExecutor(max_workers=MAX_WORKERS) as executor:
39+
futures = [executor.submit(clahe_denoise, task) for task in tasks]
40+
for future in futures:
41+
future.result()
42+
gc.collect()
43+
44+
def main(source_dir, project_dir):
45+
logging.basicConfig(filename="partinet_denoise.log", level=logging.INFO, format='%(asctime)s - %(message)s')
46+
denoise_dir = os.path.join(project_dir,"denoised")
47+
logger_name = project_dir + "/partinet_denoise.log"
48+
logging.basicConfig(filename=logger_name, level=logging.INFO, format='%(asctime)s - %(message)s')
49+
num_cpus = multiprocessing.cpu_count()
50+
logging.info(f"Number of available CPUs: {num_cpus}")
51+
logging.info(f"Processing raw micrographs in {source_dir}")
52+
logging.info(f"Saving denoised micrographs in {denoise_dir}")
53+
54+
process_directory(source_dir, denoise_dir)
55+
56+
def parse_args():
57+
parser = argparse.ArgumentParser(description="Denoise micrographs with guided CryoSegNet-style filter")
58+
parser.add_argument("--raw", required=True, help="Path to raw micrographs")
59+
parser.add_argument("--project", required=True, help="Denoised micrographs saved in project/denoised")
60+
61+
return parser.parse_args()
62+
63+
if __name__ == "__main__":
64+
args = parse_args()
65+
main(args.raw, args.project)

0 commit comments

Comments
 (0)