From e2b48a6bd5e1dd00e15a7bf60a47f5a95c81750b Mon Sep 17 00:00:00 2001 From: anna-grim Date: Tue, 14 Jul 2026 21:06:08 +0000 Subject: [PATCH 1/4] refactor: updates for merge eval --- .../machine_learning/image_dataloader.py | 4 +- .../merge_proofreading/merge_detection.py | 76 +++++++++---------- .../merge_proofreading/search_datasets.py | 22 ++---- src/neuron_proofreader/utils/swc_util.py | 4 +- src/neuron_proofreader/utils/util.py | 59 +++++--------- 5 files changed, 68 insertions(+), 97 deletions(-) diff --git a/src/neuron_proofreader/machine_learning/image_dataloader.py b/src/neuron_proofreader/machine_learning/image_dataloader.py index 81c526fc..78a727d9 100644 --- a/src/neuron_proofreader/machine_learning/image_dataloader.py +++ b/src/neuron_proofreader/machine_learning/image_dataloader.py @@ -246,7 +246,7 @@ def __call__(self, node): # Check whether to apply image augmentation if self.transform: patches = self.transform(patches) - return patches + return node, patches def compute_patch_specs(self, node): voxel = self.graph.node_voxel(node) @@ -264,7 +264,7 @@ def __call__(self, nodes): # Load patches img = self.read_image(center, shape) mask = self.create_mask(center, shape, nodes[len(nodes) // 2]) - return self.stack(img, mask), offset + return nodes, self.stack(img, mask), offset def compute_patch_specs(self, nodes): # Compute bounding box diff --git a/src/neuron_proofreader/merge_proofreading/merge_detection.py b/src/neuron_proofreader/merge_proofreading/merge_detection.py index 81e4fb0a..a1844a13 100644 --- a/src/neuron_proofreader/merge_proofreading/merge_detection.py +++ b/src/neuron_proofreader/merge_proofreading/merge_detection.py @@ -9,6 +9,7 @@ """ +from copy import deepcopy from torch.nn.functional import sigmoid from torch.utils.data import DataLoader from time import time @@ -28,10 +29,8 @@ def __init__( self, dataset, model, - model_path, batch_size=16, device="cuda", - remove_detected_sites=False, threshold=0.5, ): # Instance attributes @@ -40,12 +39,11 @@ def __init__( self.device = device self.node_preds = np.zeros((len(dataset.node_xyz))) self.patch_shape = dataset.patch_shape - self.remove_detected_sites = remove_detected_sites + self.visited_sites = list() self.threshold = threshold # Load model self.model = model - ml_util.load_model(model, model_path, device=self.device) # --- Core routines --- def search_graph(self): @@ -55,6 +53,7 @@ def search_graph(self): pbar = tqdm(total=self.dataset.estimate_iterations()) for nodes, x_nodes in dataloader: self.node_preds[np.array(nodes)] = self.predict(x_nodes) + self.visited_sites.extend(nodes.tolist()) pbar.update(len(nodes)) # Non-maximum suppression of detected sites @@ -63,14 +62,9 @@ def search_graph(self): merge_sites = self.filter_with_nms(merge_sites, likelihoods) # Report results - rate = self.dataset.distance_traversed / (time() - t0) - print("\n# Detected Merge Sites:", len(merge_sites)) - print(f"Distance Traversed: {self.dataset.distance_traversed:.2f}μm") - print(f"Merge Proofreading Rate: {rate:.2f}μm/s") - - # Remove merge mistakes (optional) - if self.remove_detected_sites: - pass + rate = len(self.visited_sites) / (time() - t0) + print("\n# Detected Merges:", len(merge_sites)) + print(f"Proofreading Rate: {rate:.2f} site/s") return merge_sites def predict(self, x): @@ -153,10 +147,22 @@ def remove_merge_sites(self, merge_site_nodes, max_depth=10): self.dataset.remove_nodes(rm_nodes) print("# Nodes Deleted:", len(rm_nodes)) - # --- Helpers --- - def get_detected_sites(self, threshold): - nodes = np.where(self.node_preds >= threshold)[0] - return [self.dataset.node_xyz[i] for i in nodes] + # --- Report Results --- + def save(self, output_dir, inplace=True): + self.save_fragment_predictions(output_dir, inplace=inplace) + self.save_parameters(output_dir) + self.save_predictions(output_dir) + self.save_sites(output_dir) + + def save_fragment_predictions(self, output_dir, inplace=True): + fragments_path = os.path.join(output_dir, "fragment_preds.zip") + if inplace: + self.dataset.node_radius = 10 * np.maximum(self.node_preds, 0.1) + self.dataset.to_zipped_swcs(fragments_path, use_radius=True) + else: + graph = deepcopy(self.dataset.graph) + graph.node_radius = 10 * np.maximum(self.node_preds, 0.1) + graph.to_zipped_swcs(fragments_path, use_radius=True) def save_parameters(self, output_dir): json_path = os.path.join(output_dir, "detection_parameters.json") @@ -171,35 +177,18 @@ def save_parameters(self, output_dir): } util.write_json(json_path, parameters) - def save_results( - self, output_dir, output_prefix_s3=None, save_fragments=True - ): - self.save_sites(output_dir) - if save_fragments: - self.dataset.graph.node_radius = 10 * np.maximum( - self.node_preds, 0.1 - ) - fragments_path = os.path.join(output_dir, "fragments.zip") - self.dataset.to_zipped_swcs(fragments_path, use_radius=True) - - # Upload results to S3 (if applicable) - if output_prefix_s3: - bucket_name, prefix = util.parse_cloud_path(output_prefix_s3) - util.upload_dir_to_s3(output_dir, bucket_name, prefix) - - def save_sites(self, output_dir): - # Save model predictions + def save_predictions(self, output_dir): + nodes = np.array(self.visited_sites, dtype=int) df = pd.DataFrame( - columns=["World", "Segment_ID", "Prediction", "Degree"] + columns=["xyz", "Segment_ID", "Prediction", "Degree"] ) - df["World"] = list(map(tuple, self.dataset.node_xyz)) - df["Prediction"] = self.node_preds - df["Segment_ID"] = [ - self.dataset.node_segment_id(i) for i in self.dataset.nodes - ] - df["Degree"] = [self.dataset.degree[i] for i in self.dataset.nodes] + df["xyz"] = list(map(tuple, self.dataset.node_xyz[nodes])) + df["Prediction"] = self.node_preds[nodes] + df["Segment_ID"] = [self.dataset.node_segment_id(i) for i in nodes] + df["Degree"] = [self.dataset.degree[i] for i in nodes] df.to_csv(os.path.join(output_dir, "model_predictions.csv")) + def save_sites(self, output_dir): # Get predicted merge sites nodes = np.where(self.node_preds >= self.threshold)[0] detected_sites = [self.dataset.node_xyz[i] for i in nodes] @@ -230,3 +219,8 @@ def save_train_dataset(self, output_dir): self.dataset._batch_to_zipped_swcs(roots, zip_path, False) self.save_sites(output_dir) print("# Fragments Saved:", len(roots)) + + # --- Helpers --- + def get_detected_sites(self, threshold): + nodes = np.where(self.node_preds >= threshold)[0] + return [self.dataset.node_xyz[i] for i in nodes] diff --git a/src/neuron_proofreader/merge_proofreading/search_datasets.py b/src/neuron_proofreader/merge_proofreading/search_datasets.py index d2666dac..133df540 100644 --- a/src/neuron_proofreader/merge_proofreading/search_datasets.py +++ b/src/neuron_proofreader/merge_proofreading/search_datasets.py @@ -42,7 +42,6 @@ def __init__( super().__init__() # Instance attributes - self.distance_traversed = 0 self.graph = graph self.is_multimodal = is_multimodal self.min_size = min_search_size @@ -64,33 +63,33 @@ def __iter__(self): def producer(): sites = self._all_sites() with ThreadPoolExecutor(max_workers=self.prefetch) as executor: - futures = {} + futures = set() def fill(): while len(futures) < self.prefetch: try: site = next(sites) - futures[ - executor.submit(self.patch_loader, site) - ] = site + futures.add(executor.submit(self.patch_loader, site)) except StopIteration: break - fill() + fill() while futures: done, _ = wait(futures, return_when=FIRST_COMPLETED) for f in done: - patch_queue.put((futures.pop(f), f.result())) + futures.remove(f) + patch_queue.put(f.result()) + fill() patch_queue.put(sentinel) Thread(target=producer, daemon=True).start() - while True: item = patch_queue.get() if item is sentinel: break + yield from self.get_input(*item) def _all_sites(self): @@ -208,7 +207,6 @@ def generate_component_sites(self, root): nodes = list() for i, j in nx.dfs_edges(self.graph, source=root): # Check if starting new batch - self.distance_traversed += self.dist(i, j) if len(nodes) == 0: if self.is_node_valid(i): root = i @@ -271,10 +269,6 @@ def estimate_iterations(self): class SparseSearchDataset(SearchDataset): - pass - - -class BranchingSearchDataset(SearchDataset): def __init__( self, @@ -298,7 +292,7 @@ def __init__( # Instance attributes self.patch_loader = DetectionPatchLoader(self.graph, img_config) - self.search_mode = "branching_nodes" + self.search_mode = "sparse" def estimate_iterations(self): return len(self.branching_nodes()) diff --git a/src/neuron_proofreader/utils/swc_util.py b/src/neuron_proofreader/utils/swc_util.py index 70379e94..c71cf2ce 100644 --- a/src/neuron_proofreader/utils/swc_util.py +++ b/src/neuron_proofreader/utils/swc_util.py @@ -313,8 +313,8 @@ def read_from_cloud(self, path): use_s3 = util.is_s3_path(path) # List paths - swc_paths = util.list_cloud_paths(path, ".swc") - zip_paths = util.list_cloud_paths(path, ".zip") + swc_paths = util.list_paths(path, extension=".swc") + zip_paths = util.list_paths(path, extension=".zip") # Call reader if swc_paths: diff --git a/src/neuron_proofreader/utils/util.py b/src/neuron_proofreader/utils/util.py index 06842732..f2a92ccc 100644 --- a/src/neuron_proofreader/utils/util.py +++ b/src/neuron_proofreader/utils/util.py @@ -68,7 +68,7 @@ def list_files_in_zip(zip_content): return zip_file.namelist() -def list_paths(dir_path, extension=None): +def list_paths(dir_path, extension=""): """ Lists all paths within "directory" that end with "extension" if provided. @@ -78,15 +78,20 @@ def list_paths(dir_path, extension=None): Path to directory to be searched. extension : str, optional If provided, only paths of files with the extension are returned. - Default is None. + Default is an empty string. Returns ------- paths : List[str] List of all paths within "directory". """ - filenames = listdir(dir_path, extension=extension) - return [os.path.join(dir_path, f) for f in filenames] + if is_gcs_path(dir_path): + return list_gcs_paths(dir_path, extension=extension) + elif is_s3_path(dir_path): + return list_s3_paths(dir_path, extension) + else: + filenames = listdir(dir_path, extension=extension) + return [os.path.join(dir_path, f) for f in filenames] def list_subdirs(path, keyword=None, return_paths=False): @@ -362,29 +367,6 @@ def get_google_swcs_dirname(prefix): return "swcs" -def list_cloud_paths(path, extension=""): - """ - Lists all files in a GCS/S3 bucket with the given extension. - - Parameters - ---------- - path : str - Path to cloud prefix to be searched, must be in the format: - f"{scheme}://{bucket_name}/{prefix}". - extension : str, optional - File extension of filenames to be listed. Default is an empty string. - - Returns - ------- - List[str] - Filenames stored at the GCS path with the given extension. - """ - assert is_gcs_path(path) or is_s3_path(path) - bucket_name, prefix = parse_cloud_path(path) - list_fn = list_gcs_paths if is_gcs_path(path) else list_s3_paths - return list_fn(bucket_name, prefix, extension=extension) - - def parse_cloud_path(path): """ Parses a cloud storage path into its bucket name and key/prefix. Supports @@ -458,16 +440,14 @@ def is_gcs_path(path): return path.startswith("gs://") -def list_gcs_paths(bucket_name, prefix, extension=""): +def list_gcs_paths(path, extension=""): """ Lists paths at a GCS prefix with the given extension. Parameters ---------- - bucket_name : str - Name of bucket containing prefix. - prefix : str - Path to location within bucket to be searched. + path : str + Path to location in a GCS bucket. extension : str, optional File extension of filenames to be listed. Default is an empty string. @@ -476,12 +456,16 @@ def list_gcs_paths(bucket_name, prefix, extension=""): List[str] Paths under the GCS prefix with the given extension. """ + # Create bucket + bucket_name, prefix = parse_cloud_path(path) bucket = storage.Client().bucket(bucket_name) + + # List paths paths = list() for name in [b.name for b in bucket.list_blobs(prefix=prefix)]: if extension in name: paths.append(os.path.join(f"gs://{bucket_name}", name)) - return paths + return sorted(paths) def list_gcs_subprefixes(path): @@ -556,17 +540,15 @@ def is_s3_path(path): return path.startswith("s3://") -def list_s3_paths(bucket_name, prefix, extension=""): +def list_s3_paths(path, extension=""): """ Lists all object keys in a public S3 bucket under a given prefix, optionally filters by file extension. Parameters ---------- - bucket_name : str - Name of the S3 bucket. - prefix : str - Prefix to search under. + path : str + Path to location in an S3 bucket. extension : str, optional File extension to filter by. Default is an empty string. @@ -576,6 +558,7 @@ def list_s3_paths(bucket_name, prefix, extension=""): S3 object keys that match the prefix and extension filter. """ # Create an anonymous client for public buckets + bucket_name, prefix = parse_cloud_path(path) s3 = boto3.client("s3", config=Config(signature_version=UNSIGNED)) response = s3.list_objects_v2(Bucket=bucket_name, Prefix=prefix) From 158f052e33194cd6372f0b4b1c5e362c9ca68b13 Mon Sep 17 00:00:00 2001 From: anna-grim Date: Wed, 15 Jul 2026 04:05:30 +0000 Subject: [PATCH 2/4] feat: compute eval batch size --- .../merge_proofreading/merge_detection.py | 2 +- src/neuron_proofreader/utils/ml_util.py | 62 ++++++++++++++++++- 2 files changed, 62 insertions(+), 2 deletions(-) diff --git a/src/neuron_proofreader/merge_proofreading/merge_detection.py b/src/neuron_proofreader/merge_proofreading/merge_detection.py index a1844a13..b4f2229c 100644 --- a/src/neuron_proofreader/merge_proofreading/merge_detection.py +++ b/src/neuron_proofreader/merge_proofreading/merge_detection.py @@ -55,6 +55,7 @@ def search_graph(self): self.node_preds[np.array(nodes)] = self.predict(x_nodes) self.visited_sites.extend(nodes.tolist()) pbar.update(len(nodes)) + pbar.close() # Non-maximum suppression of detected sites merge_sites = np.where(self.node_preds > self.threshold)[0] @@ -171,7 +172,6 @@ def save_parameters(self, output_dir): "is_multimodal": self.dataset.is_multimodal, "min_search_size": self.dataset.min_size, "patch_shape": self.patch_shape, - "remove_detected_sites": self.remove_detected_sites, "search_mode": self.dataset.search_mode, "subgraph_radius": self.dataset.subgraph_radius, } diff --git a/src/neuron_proofreader/utils/ml_util.py b/src/neuron_proofreader/utils/ml_util.py index 562aa6e3..a99a43b6 100644 --- a/src/neuron_proofreader/utils/ml_util.py +++ b/src/neuron_proofreader/utils/ml_util.py @@ -215,7 +215,67 @@ def move(self, v, device): # --- Miscellaneous --- -def find_max_batch_size( +def find_max_eval_batch_size( + model, input_shape, device="cuda", start_batch_size=1, max_batch_size=128 +): + """ + Finds the largest batch size that fits in GPU memory for a forward + pass of "model" on inputs of "input_shape", via binary search. + + Parameters + ---------- + model : torch.nn.Module + Model to run inference with. Assumed to already be on "device" and + in eval mode. + input_shape : Tuple[int] + Shape of a single input sample, excluding the batch dimension. + device : str, optional + Device to run inference on. Default is "cuda". + start_batch_size : int, optional + Initial batch size to test. Default is 1. + max_batch_size : int, optional + Upper bound on batch size to consider. Default is 4096. + + Returns + ------- + max_batch_size : int + Largest batch size that ran successfully. + """ + def fits(batch_size): + try: + x = torch.zeros((batch_size, *input_shape), device=device) + with torch.no_grad(): + model(x) + del x + torch.cuda.empty_cache() + return True + except RuntimeError as e: + if "out of memory" not in str(e).lower(): + raise + torch.cuda.empty_cache() + return False + + # Exponential search for an upper bound that fails + lo, hi = 0, start_batch_size + while hi <= max_batch_size and fits(hi): + lo = hi + hi *= 2 + hi = min(hi, max_batch_size + 1) + + if lo == 0: + return 0 + + # Binary search between last success (lo) and first failure (hi) + while hi - lo > 1: + mid = (lo + hi) // 2 + if fits(mid): + lo = mid + else: + hi = mid + return lo + + +def find_max_train_batch_size( model, input_shape, optimizer_cls, device="cuda", start=1, max_bs=32 ): model.to(device) From fc7ad41e3d734130a448960b0b496f1bbd0765de Mon Sep 17 00:00:00 2001 From: anna-grim Date: Wed, 15 Jul 2026 22:12:43 +0000 Subject: [PATCH 3/4] refactor: upd merge inference --- src/neuron_proofreader/merge_proofreading/merge_detection.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/neuron_proofreader/merge_proofreading/merge_detection.py b/src/neuron_proofreader/merge_proofreading/merge_detection.py index b4f2229c..7014ee8c 100644 --- a/src/neuron_proofreader/merge_proofreading/merge_detection.py +++ b/src/neuron_proofreader/merge_proofreading/merge_detection.py @@ -83,6 +83,7 @@ def predict(self, x): numpy.ndarray Predicted merge site likelihoods. """ + self.model.eval() with torch.inference_mode(): x = x.to(self.device) y = sigmoid(self.model(x)) @@ -116,7 +117,7 @@ def filter_with_nms(self, merge_sites, likelihoods): iou = img_util.compute_iou3d( xyz_i, xyz_root, self.patch_shape, self.patch_shape ) - if iou > 0.35: + if iou > 0.3 and self.dataset.degree[i] == 2: merge_sites_set.remove(i) self.node_preds[i] = 0 From 98b41652f01ad011eee5bc83abfab3e98b505da3 Mon Sep 17 00:00:00 2001 From: anna-grim Date: Mon, 20 Jul 2026 20:38:33 +0000 Subject: [PATCH 4/4] refactor: bug fixes --- .../merge_proofreading/merge_detection.py | 80 +++++++++++++++++-- 1 file changed, 72 insertions(+), 8 deletions(-) diff --git a/src/neuron_proofreader/merge_proofreading/merge_detection.py b/src/neuron_proofreader/merge_proofreading/merge_detection.py index 7014ee8c..31821164 100644 --- a/src/neuron_proofreader/merge_proofreading/merge_detection.py +++ b/src/neuron_proofreader/merge_proofreading/merge_detection.py @@ -10,11 +10,13 @@ """ from copy import deepcopy +from scipy.spatial import KDTree from torch.nn.functional import sigmoid from torch.utils.data import DataLoader from time import time from tqdm import tqdm +import networkx as nx import numpy as np import os import pandas as pd @@ -60,7 +62,14 @@ def search_graph(self): # Non-maximum suppression of detected sites merge_sites = np.where(self.node_preds > self.threshold)[0] likelihoods = self.node_preds[merge_sites] - merge_sites = self.filter_with_nms(merge_sites, likelihoods) + merge_sites = self.apply_graph_nms(merge_sites, likelihoods) + + # Iteratively average nearby sites + while True: + before = len(merge_sites) + merge_sites = self.avg_nearby_sites(merge_sites) + if before == len(merge_sites): + break # Report results rate = len(self.visited_sites) / (time() - t0) @@ -83,22 +92,22 @@ def predict(self, x): numpy.ndarray Predicted merge site likelihoods. """ - self.model.eval() with torch.inference_mode(): x = x.to(self.device) - y = sigmoid(self.model(x)) - return np.squeeze(ml_util.to_cpu(y, to_numpy=True), axis=1) + y = y.detach().cpu().numpy() + return np.squeeze(y, axis=1) - def filter_with_nms(self, merge_sites, likelihoods): + def apply_graph_nms(self, merge_sites, likelihoods): # Sort by confidence - merge_sites = [merge_sites[i] for i in np.argsort(likelihoods)] + merge_sites = [merge_sites[i] for i in np.argsort(likelihoods)[::-1]] + merge_sites = deque(merge_sites) # NMS merge_sites_set = set(merge_sites) filtered_merge_sites = set() while merge_sites: # Local max - root = merge_sites.pop() + root = merge_sites.popleft() xyz_root = self.dataset.node_xyz[root] if root in merge_sites_set: filtered_merge_sites.add(root) @@ -129,6 +138,61 @@ def filter_with_nms(self, merge_sites, likelihoods): visited.add(j) return filtered_merge_sites + def avg_nearby_sites(self, merge_sites, max_dist=24): + # Sort sites by likelihood + merge_sites = list(merge_sites) + likelihoods = [self.node_preds[i] for i in merge_sites] + merge_sites = [merge_sites[i] for i in np.argsort(likelihoods)[::-1]] + + # Search for spatially nearby sites + visited = set() + new_merge_sites = list() + sites_kdtree = KDTree([self.dataset.node_xyz[i] for i in merge_sites]) + for root in merge_sites: + # Check whether to skip + if root in visited: + continue + else: + visited.add(root) + + # Get nearby sites + xyz_query = self.dataset.node_xyz[root] + idxs = sites_kdtree.query_ball_point(xyz_query, max_dist) + nodes = [merge_sites[i] for i in idxs] + + # Check whether to combine sites + if len(nodes) > 1: + hits = list() + for node in nodes: + try: + path = nx.shortest_path( + self.dataset.graph, source=root, target=node + ) + if self.dataset.path_length(path) < max_dist + 4: + hits.append(node) + visited.add(node) + except nx.exception.NetworkXNoPath: + pass + + # Add site to list + xyz_arr = np.array([self.dataset.node_xyz[i] for i in hits]) + xyz_avg = xyz_arr.mean(axis=0) + best_node = min( + hits, + key=lambda n: np.linalg.norm(self.dataset.node_xyz[n] - xyz_avg) + ) + new_merge_sites.append(best_node) + + # Update node predictions + likelihood = self.node_preds[root] + for node in hits: + if node != best_node: + self.node_preds[node] = 0 + self.node_preds[best_node] = likelihood + else: + new_merge_sites.append(root) + return new_merge_sites + def remove_merge_sites(self, merge_site_nodes, max_depth=10): rm_nodes = set() for root in tqdm(merge_site_nodes, desc="Remove Merge Sites"): @@ -149,7 +213,7 @@ def remove_merge_sites(self, merge_site_nodes, max_depth=10): self.dataset.remove_nodes(rm_nodes) print("# Nodes Deleted:", len(rm_nodes)) - # --- Report Results --- + # --- Save Results --- def save(self, output_dir, inplace=True): self.save_fragment_predictions(output_dir, inplace=inplace) self.save_parameters(output_dir)