Skip to content

Commit 5fa95da

Browse files
authored
Merge pull request #18 from stanfish06/downsample
update: downsample (sign)
2 parents 60fd724 + 19ff00f commit 5fa95da

3 files changed

Lines changed: 104 additions & 2 deletions

File tree

‎src/scloop/data/types.py‎

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,18 @@
11
# Copyright 2025 Zhiyuan Yu (Heemskerk's lab, University of Michigan)
2-
from typing import Literal
2+
from typing import Literal, Annotated, TypeAlias
3+
from pydantic import Field
34

45
FeatureSelectionMethod = Literal["hvg", "delve", "none"]
56
EmbeddingMethod = Literal["pca", "diffmap", "scvi"]
67
EmbeddingNeighbors = Literal["pca", "scvi"]
8+
9+
Index_t = Annotated[int, Field(ge=0)]
10+
# need at least 2 points to compute PH. Maybe also set an upper bound later as it is not feasible to compute PH on a lot of points
11+
IndexListDownSample: TypeAlias = Annotated[
12+
list[Index_t],
13+
Field(min_length=2, description="Downsampled indices for PH computation"),
14+
]
15+
SizeDownSample = Annotated[
16+
int, Field(ge=2, description="Sample to this number of cells")
17+
]
18+
# TODO: make a type for boundary matrix. Restrict matrix size for efficient computation
Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,2 +1,3 @@
11
from .prepare import prepare_adata
2+
from .downsample import sample
23
from .delve import delve_fs
Lines changed: 90 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1,90 @@
1-
# TODO: use numba
1+
# Copyright 2025 Zhiyuan Yu (Heemskerk's lab, University of Michigan)
2+
from numba import jit
3+
from anndata import AnnData
4+
from ..data.types import IndexListDownSample, SizeDownSample, EmbeddingMethod
5+
from pydantic import validate_call
6+
import pandas as pd
7+
import numpy as np
8+
9+
__all__ = ["sample"]
10+
11+
12+
@jit(nopython=True)
13+
def _sample_impl(
14+
data: np.ndarray,
15+
class_labels: np.ndarray,
16+
classes: np.ndarray,
17+
seed_indices: np.ndarray,
18+
n: int,
19+
) -> np.ndarray:
20+
# selected observation indices
21+
indices = np.zeros(n, dtype=np.int64)
22+
min_dists = np.full(len(data), np.inf)
23+
24+
num_seeds = len(seed_indices)
25+
num_classes = len(classes)
26+
indices[0] = seed_indices[0]
27+
28+
# for each newly added points, recompute and select out-of-bag hausdorff point
29+
for i in range(1, n):
30+
last_point = data[indices[i - 1]]
31+
dists = np.sum((data - last_point) ** 2, axis=1)
32+
min_dists = np.minimum(min_dists, dists)
33+
min_dists[indices[:i]] = -1
34+
if i >= num_seeds:
35+
class_indicies = np.where(class_labels == classes[i % num_classes])[0]
36+
next_idx = class_indicies[np.argmax(min_dists[class_indicies])]
37+
# if this class is exhausted, fall back to reguler sampling
38+
if min_dists[next_idx] == -1:
39+
next_idx = np.argmax(min_dists)
40+
indices[i] = next_idx
41+
else:
42+
indices[i] = seed_indices[i]
43+
44+
return indices
45+
46+
47+
@validate_call(config={"arbitrary_types_allowed": True})
48+
def sample(
49+
adata: AnnData,
50+
groupby: str | None,
51+
embedding_method: EmbeddingMethod,
52+
n: SizeDownSample,
53+
random_state: int = 0,
54+
) -> IndexListDownSample:
55+
"""
56+
Topology-preserving downsampling using greedy farthest-point sampling.
57+
58+
Args:
59+
adata: AnnData object containing the data
60+
groupby: column in adata.obs for class-balanced sampling, or None
61+
embedding_method: which embedding to use from adata.obsm
62+
n: number of points to sample
63+
random_state: random seed for reproducibility
64+
65+
Returns:
66+
list of indices into adata.obs for the downsampled points
67+
"""
68+
assert n <= adata.shape[0]
69+
assert f"X_{embedding_method}" in adata.obsm
70+
downsample_embedding = adata.obsm[f"X_{embedding_method}"]
71+
assert type(downsample_embedding) is np.ndarray
72+
73+
if groupby is None:
74+
class_labels = np.zeros(adata.shape[0], dtype=np.int64)
75+
classes = np.array([0])
76+
seed_indices = np.array([np.random.randint(adata.shape[0])])
77+
else:
78+
assert type(adata.obs) is pd.DataFrame
79+
assert groupby in adata.obs.columns
80+
class_labels, classes = pd.factorize(adata.obs.loc[:, groupby])
81+
classes = np.arange(len(classes), dtype=np.int64)
82+
seed_indices = []
83+
np.random.seed(random_state)
84+
for c in classes:
85+
seed_indices.append(np.random.choice(np.where(class_labels == c)[0]))
86+
seed_indices = np.array(seed_indices)
87+
88+
return _sample_impl(
89+
downsample_embedding, class_labels, classes, seed_indices, n
90+
).tolist()

0 commit comments

Comments
 (0)