Skip to content

Commit ecb4990

Browse files
committed
update: downsample
1 parent 4f3c224 commit ecb4990

2 files changed

Lines changed: 32 additions & 8 deletions

File tree

‎src/scloop/data/types.py‎

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,5 +8,11 @@
88

99
Index_t = Annotated[int, Field(ge=0)]
1010
# 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[list[Index_t], Field(min_length=2)]
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+
]
1218
# TODO: make a type for boundary matrix. Restrict matrix size for efficient computation
Lines changed: 25 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,33 @@
11
# Copyright 2025 Zhiyuan Yu (Heemskerk's lab, University of Michigan)
22
from numba import jit
33
from anndata import AnnData
4-
from ..data.types import IndexListDownSample
4+
from ..data.types import IndexListDownSample, SizeDownSample, EmbeddingMethod
55
from pydantic import validate_call
6-
from typing import TypeAlias
6+
from scipy.spatial.distance import directed_hausdorff
77
import numpy as np
88

9-
@jit
10-
def _sample_impl(data: np.ndarray, n: int) -> np.ndarray:
11-
return np.array([1, 1])
9+
@jit(nopython=True)
10+
def _sample_impl(data: np.ndarray, seed_idx: int, n: int) -> np.ndarray:
11+
indices = np.zeros(n, dtype=np.int64)
12+
indices[0] = seed_idx
13+
min_dists = np.full(len(data), np.inf)
14+
15+
for i in range(1, n):
16+
last_point = data[indices[i-1]]
17+
dists = np.sum((data - last_point) ** 2, axis=1)
18+
min_dists = np.minimum(min_dists, dists)
19+
min_dists[indices[:i]] = -1
20+
indices[i] = np.argmax(min_dists)
21+
22+
return indices.tolist()
1223

1324
@validate_call(config={"arbitrary_types_allowed": True})
14-
def sample(adata: AnnData, n: int) -> IndexListDownSample:
15-
return _sample_impl(np.array(), n).tolist()
25+
def sample(adata: AnnData, embedding_method: EmbeddingMethod, n: SizeDownSample) -> IndexListDownSample:
26+
"""
27+
topology-aware downsampling
28+
"""
29+
assert n <= adata.shape[0]
30+
assert f"X_{embedding_method}" in adata.obsm
31+
downsample_embedding = adata.obsm[f"X_{embedding_method}"]
32+
assert type(downsample_embedding) is np.ndarray
33+
return _sample_impl(downsample_embedding, 0, n)

0 commit comments

Comments
 (0)