diff --git a/src/arborist/data/datasets.py b/src/arborist/data/datasets.py index cee8f86..e3fbb77 100644 --- a/src/arborist/data/datasets.py +++ b/src/arborist/data/datasets.py @@ -8,6 +8,7 @@ """ +from collections import defaultdict from copy import deepcopy from torch.utils.data import Dataset, DataLoader, Sampler @@ -16,6 +17,8 @@ import pandas as pd import torch +from arborist.utils.graph_utils import topological_decomposition + # --- Dataset Classes --- class CurveDataset(Dataset): @@ -83,84 +86,207 @@ def __repr__(self): ) -class CurveDatasetCollection(Dataset): +class GraphDataset(Dataset): + """ + Dataset over rooted subgraphs of a single SkeletonGraph. + + Each item corresponds to one root node. __getitem__ extracts the + subgraph within "max_depth" microns, decomposes it into irreducible paths + via topological decomposition, computes first-order finite differences for + each path, and returns a TreeSample with those curves and the line-graph + connectivity between them. + """ - def __init__(self, datasets, is_val=False, n_val_examples=1000, seed=42): + def __init__(self, graph, root_nodes, max_depth=None, transform=None): """ + Instantiates a GraphDataset object. + Parameters ---------- - datasets : List[CurveDataset] - List of PathsDataset instances, one per brain. - is_val : bool, optional - If True, precomputes a fixed set of examples at construction time. - Default is False. - n_val_examples : int, optional - Number of fixed validation examples to precompute. Default is 1000. - seed : int, optional - Random seed for reproducible val set. Default is 42. + graph : SkeletonGraph + The full skeleton graph to sample from. + root_nodes : List[int] + One root node per dataset item. + max_depth : float + Depth in microns for rooted subgraph extraction. + transform : callable, optional + Applied to each raw xyz array before differencing (e.g. + CurveTransforms for augmentation). Default is None. """ + # Call parent class + super().__init__() + # Instance attributes + self.graph = graph + self.root_nodes = root_nodes + self.max_depth = max_depth + self.transform = transform + + def __getitem__(self, i): + # Extract tree sample components + root = self.root_nodes[i] + subgraph = self.graph.rooted_subgraph(root, self.max_depth) + _, paths, topo_edge_index = topological_decomposition(subgraph) + + # Create list of curves + curves = [] + for path in paths: + xyz = subgraph.node_xyz[path].copy() + if self.transform: + xyz = self.transform(xyz) + xyz -= xyz[0] + xyz[1:] -= xyz[:-1].copy() + curves.append(xyz) + + # Create TreeSample + edge_index = _build_line_graph_edge_index(topo_edge_index) + return TreeSample(curves=curves, edge_index=edge_index) + + def __len__(self): + return len(self.root_nodes) + + +class TreeSample: + """ + A rooted subgraph ready to pass through CurveEncoder then GraphTransformer. + + Attributes + ---------- + curves : List[numpy.ndarray] + One array per irreducible path, each of shape (N_i, 3). Values are + first-order finite differences with a leading zero row, matching the + convention expected by CurveEncoder. + edge_index : numpy.ndarray + Shape (2, E), dtype int64. Line-graph adjacency: two curves share an + edge when they meet at a topological node (branch point or leaf), + so message passing over this graph communicates between neighboring + branches. + """ + + def __init__(self, curves, edge_index): + self.curves = curves + self.edge_index = edge_index + + def __repr__(self): + return ( + f"TreeSample(" + f"n_curves={len(self.curves)}, " + f"n_edges={self.edge_index.shape[1]})" + ) + + +def _build_line_graph_edge_index(topo_edge_index): + """ + Converts topological-graph edge pairs into line-graph edge pairs. + + In the topological graph each node is a branching/leaf point and each + edge is an irreducible path (a curve). In the line graph each curve + becomes a node and two curve-nodes are connected when they share a + topological endpoint. + + Parameters + ---------- + topo_edge_index : List[Tuple[int, int]] + Edges of the topological graph as (src_topo_idx, dst_topo_idx) pairs, + parallel to the list of curves. + + Returns + ------- + numpy.ndarray + Shape (2, E), int64. + """ + topo_to_curves = defaultdict(list) + for curve_idx, (u, v) in enumerate(topo_edge_index): + topo_to_curves[u].append(curve_idx) + topo_to_curves[v].append(curve_idx) + + src, dst = [], [] + for neighbors in topo_to_curves.values(): + for i in neighbors: + for j in neighbors: + if i != j: + src.append(i) + dst.append(j) + + if not src: + return np.zeros((2, 0), dtype=np.int64) + return np.array([src, dst], dtype=np.int64) + + +class DatasetCollection(Dataset): + """ + A flat, indexable view over multiple datasets (one per brain/specimen). + + Parameters + ---------- + datasets : List[Dataset] + Constituent datasets to combine. + weight_fn : callable, optional + Maps a dataset to a 1-D array of per-item sampling weights. Used by + samplers for non-uniform drawing (e.g. length-weighted curve sampling). + Defaults to uniform weights when None. Default is None. + is_val : bool, optional + If True, precomputes a fixed set of examples at construction time. + Default is False. + n_val_examples : int, optional + Number of validation examples to precompute. Default is 1000. + seed : int, optional + Random seed for reproducible val set. Default is 42. + """ + + def __init__( + self, + datasets, + weight_fn=None, + is_val=False, + n_val_examples=1000, + seed=42, + ): self.datasets = datasets self.is_val = is_val - self.set_examples_df() - - # Check whether to set validation examples + self._build_index(weight_fn) if is_val: - self.val_examples = self.set_val_examples(n_val_examples, seed) + self.val_examples = self._precompute_val(n_val_examples, seed) - def set_examples_df(self): + def _build_index(self, weight_fn): rows = [] for ds_idx, dataset in enumerate(self.datasets): - ds_idxs = np.full(len(dataset), ds_idx) - p_idxs = np.arange(len(dataset)) - ds_lengths = dataset.curve_lengths() - rows.append( - pd.DataFrame( - { - "ds_idx": ds_idxs, - "path_idx": p_idxs, - "length": ds_lengths, - } - ) - ) - self.examples_df = pd.concat(rows, ignore_index=True) - - def set_val_examples(self, n, seed): - """ - Samples n examples with fixed seed, strips transforms, and caches - the resulting examples. - """ + n = len(dataset) + weights = weight_fn(dataset) if weight_fn is not None else np.ones(n) + rows.append(pd.DataFrame({ + "ds_idx": np.full(n, ds_idx, dtype=int), + "item_idx": np.arange(n), + "weight": weights, + })) + self.index = pd.concat(rows, ignore_index=True) + + def _precompute_val(self, n, seed): rng = np.random.default_rng(seed) - indices = rng.choice(len(self.examples_df), size=n, replace=False) + idxs = rng.choice(len(self.index), size=n, replace=False) examples = [] - for i in indices: - ds_idx = self.examples_df["ds_idx"][i] - path_idx = self.examples_df["path_idx"][i] - dataset = self.datasets[ds_idx] - examples.append(dataset[path_idx]) + for i in idxs: + ds_idx = self.index["ds_idx"][i] + item_idx = self.index["item_idx"][i] + examples.append(self.datasets[ds_idx][item_idx]) return examples - # --- Data Fetching --- def __getitem__(self, i): - # Case 1: validation example if self.is_val: return self.val_examples[i] - - # Case 2: train example - ds_idx = self.examples_df["ds_idx"][i] - path_idx = self.examples_df["path_idx"][i] - return self.datasets[ds_idx][path_idx] + ds_idx = self.index["ds_idx"][i] + item_idx = self.index["item_idx"][i] + return self.datasets[ds_idx][item_idx] def __len__(self): if self.is_val: return len(self.val_examples) - return len(self.examples_df) + return len(self.index) def __repr__(self): return ( - f"CurveDatasetCollection(" - f"num_brains={len(self.datasets)}, " - f"num_curves={len(self.examples_df)}) " + f"DatasetCollection(" + f"num_datasets={len(self.datasets)}, " + f"num_items={len(self.index)})" ) @@ -178,8 +304,8 @@ def __init__(self, dataset, examples_per_epoch): self.examples_per_epoch = examples_per_epoch def __iter__(self): - idxs = self.dataset.examples_df.sample( - self.examples_per_epoch, replace=True, weights="length" + idxs = self.dataset.index.sample( + self.examples_per_epoch, replace=True, weights="weight" ).index return iter(np.array(idxs)) diff --git a/src/arborist/models/arborist.py b/src/arborist/models/arborist.py new file mode 100644 index 0000000..eb289bc --- /dev/null +++ b/src/arborist/models/arborist.py @@ -0,0 +1,158 @@ +""" +Created on Mon Aug 6 17:00:00 2026 + +@author: Anna Grim +@email: anna.grim@alleninstitute.org + +End-to-end ArboristModel for neuron morphology encoding. + +""" + +import json +import torch +import torch.nn as nn + +from arborist.models.curve_transformer import CurveEncoder +from arborist.models.graph_transformer import GraphTransformer + + +class Arborist(nn.Module): + """ + End-to-end neuron morphology encoder. + + Encodes a TreeSample in three stages: + 1. CurveEncoder — each irreducible path → latent vector z_i + 2. GraphTransformer — message-pass over the line graph of the skeleton + so each curve sees its neighboring branches + 3. Mean-pool — average over all curve embeddings → global tree z + + Parameters + ---------- + segment_len : int, optional + Points per curve segment for CurveEncoder. Default is 10. + d_token : int, optional + CurveEncoder token dimension. Default is 128. + curve_n_heads : int, optional + Attention heads in CurveEncoder. Default is 4. + curve_n_layers : int, optional + Transformer layers in CurveEncoder. Default is 4. + d_ff_curve : int, optional + CurveEncoder feed-forward dimension. Default is 256. + latent_dim : int, optional + Shared latent dimension: CurveEncoder output = GraphTransformer input. + Default is 64. + graph_n_heads : int, optional + Attention heads in GraphTransformer. Default is 4. + graph_n_layers : int, optional + Transformer layers in GraphTransformer. Default is 3. + d_ff_graph : int, optional + GraphTransformer feed-forward dimension. Default is 256. + dropout : float, optional + Dropout probability shared across both sub-models. Default is 0.1. + """ + + def __init__( + self, + segment_len=10, + d_token=128, + curve_n_heads=4, + curve_n_layers=4, + d_ff_curve=256, + latent_dim=64, + graph_n_heads=4, + graph_n_layers=3, + d_ff_graph=256, + dropout=0.1, + ): + super().__init__() + self.config = { + "segment_len": segment_len, + "d_token": d_token, + "curve_n_heads": curve_n_heads, + "curve_n_layers": curve_n_layers, + "d_ff_curve": d_ff_curve, + "latent_dim": latent_dim, + "graph_n_heads": graph_n_heads, + "graph_n_layers": graph_n_layers, + "d_ff_graph": d_ff_graph, + "dropout": dropout, + } + self.curve_encoder = CurveEncoder( + segment_len=segment_len, + d_token=d_token, + n_heads=curve_n_heads, + n_layers=curve_n_layers, + d_ff=d_ff_curve, + latent_dim=latent_dim, + dropout=dropout, + ) + self.graph_transformer = GraphTransformer( + d_model=latent_dim, + n_heads=graph_n_heads, + n_layers=graph_n_layers, + d_ff=d_ff_graph, + dropout=dropout, + ) + + def _collate_curves(self, curves): + device = next(self.parameters()).device + lengths = [len(c) for c in curves] + n_max = max(lengths) + B = len(curves) + diffs = torch.zeros(B, n_max, 3, device=device) + mask = torch.ones(B, n_max, dtype=torch.bool, device=device) + for i, (c, l) in enumerate(zip(curves, lengths)): + diffs[i, :l] = torch.tensor(c, dtype=torch.float32, device=device) + mask[i, :l] = False + return diffs, mask + + def encode(self, sample): + """ + Encodes a TreeSample into per-curve and global tree embeddings. + + Parameters + ---------- + sample : TreeSample + A rooted subgraph as returned by GraphDataset.__getitem__. + + Returns + ------- + z_tree : torch.Tensor + Shape (latent_dim,) — global tree embedding, mean-pooled over + all curve embeddings after graph contextualization. + z_curves : torch.Tensor + Shape (n_curves, latent_dim) — per-curve embeddings after the + GraphTransformer contextualizes each curve by its neighbors. + """ + device = next(self.parameters()).device + + if not sample.curves: + empty = torch.zeros(self.config["latent_dim"], device=device) + return empty, empty.unsqueeze(0) + + # Encode all curves in parallel (treat each as an independent batch item) + diffs, mask = self._collate_curves(sample.curves) # (n_curves, N_max, 3) + z, _ = self.curve_encoder(diffs, mask) # (n_curves, latent_dim) + + # Contextualize via graph topology + edge_index = torch.tensor( + sample.edge_index, dtype=torch.long, device=device + ) + z_curves = self.graph_transformer(z, edge_index) # (n_curves, latent_dim) + + z_tree = z_curves.mean(dim=0) # (latent_dim,) + return z_tree, z_curves + + def forward(self, sample): + return self.encode(sample) + + def save_config(self, path): + with open(path, "w") as f: + json.dump(self.config, f) + + @classmethod + def load(cls, path): + checkpoint = torch.load(path) + model = cls(**checkpoint["config"]) + model.load_state_dict(checkpoint["model_state_dict"]) + return model diff --git a/src/arborist/models/graph_transformer.py b/src/arborist/models/graph_transformer.py new file mode 100644 index 0000000..0052f85 --- /dev/null +++ b/src/arborist/models/graph_transformer.py @@ -0,0 +1,131 @@ +""" +Created on Mon Aug 6 17:00:00 2026 + +@author: Anna Grim +@email: anna.grim@alleninstitute.org + +Graph transformer for contextualizing curve embeddings over skeleton topology. + +""" + +import torch +import torch.nn as nn + + +class GraphTransformerLayer(nn.Module): + """ + One layer of graph-masked multi-head self-attention with pre-norm and FFN. + + Each node attends only to itself and its immediate graph neighbors. + """ + + def __init__(self, d_model, n_heads, d_ff, dropout=0.1): + """ + Parameters + ---------- + d_model : int + Node feature dimension. + n_heads : int + Number of attention heads. + d_ff : int + Feed-forward hidden dimension. + dropout : float, optional + Dropout probability. Default is 0.1. + """ + super().__init__() + self.attn = nn.MultiheadAttention( + d_model, n_heads, dropout=dropout, batch_first=True + ) + self.ff = nn.Sequential( + nn.Linear(d_model, d_ff), + nn.GELU(), + nn.Dropout(dropout), + nn.Linear(d_ff, d_model), + ) + self.norm1 = nn.LayerNorm(d_model) + self.norm2 = nn.LayerNorm(d_model) + self.dropout = nn.Dropout(dropout) + + def forward(self, x, attn_mask=None): + """ + Parameters + ---------- + x : torch.Tensor + Shape (n, d_model) — one feature vector per graph node. + attn_mask : torch.Tensor, optional + Shape (n, n) BoolTensor; True at (i, j) blocks attention from + node i to node j. Default is None (full attention). + + Returns + ------- + torch.Tensor + Shape (n, d_model). + """ + h = self.norm1(x).unsqueeze(0) # (1, n, d_model) + h, _ = self.attn(h, h, h, attn_mask=attn_mask) + x = x + self.dropout(h.squeeze(0)) # (n, d_model) + x = x + self.dropout(self.ff(self.norm2(x))) + return x + + +class GraphTransformer(nn.Module): + """ + Stack of graph-masked transformer layers over a set of node features. + + Restricts each node's attention to itself and its 1-hop graph neighbors, + coupling transformer expressiveness with explicit graph topology. + """ + + def __init__(self, d_model, n_heads=4, n_layers=3, d_ff=256, dropout=0.1): + """ + Parameters + ---------- + d_model : int + Node feature dimension (must equal CurveEncoder latent_dim when + used inside ArboristModel). + n_heads : int, optional + Number of attention heads. Default is 4. + n_layers : int, optional + Number of transformer layers. Default is 3. + d_ff : int, optional + Feed-forward hidden dimension. Default is 256. + dropout : float, optional + Dropout probability. Default is 0.1. + """ + super().__init__() + self.layers = nn.ModuleList([ + GraphTransformerLayer(d_model, n_heads, d_ff, dropout) + for _ in range(n_layers) + ]) + self.norm = nn.LayerNorm(d_model) + + def _build_attn_mask(self, n, edge_index, device): + """ + True at (i, j) blocks attention from node i to node j. + Every node attends to itself and all immediate neighbors. + """ + mask = torch.ones(n, n, dtype=torch.bool, device=device) + mask.fill_diagonal_(False) + if edge_index.shape[1] > 0: + mask[edge_index[0], edge_index[1]] = False + return mask + + def forward(self, x, edge_index): + """ + Parameters + ---------- + x : torch.Tensor + Shape (n, d_model) — one feature vector per graph node (curve). + edge_index : torch.Tensor + Shape (2, E), dtype long — bidirectional line-graph adjacency + as returned by GraphDataset. + + Returns + ------- + torch.Tensor + Shape (n, d_model) — enriched node features. + """ + attn_mask = self._build_attn_mask(x.shape[0], edge_index, x.device) + for layer in self.layers: + x = layer(x, attn_mask) + return self.norm(x)