|
1 | 1 | # Copyright 2025 Zhiyuan Yu (Heemskerk's lab, University of Michigan) |
2 | 2 | from numba import jit |
3 | 3 | from anndata import AnnData |
4 | | -from ..data.types import IndexListDownSample |
| 4 | +from ..data.types import IndexListDownSample, SizeDownSample, EmbeddingMethod |
5 | 5 | from pydantic import validate_call |
6 | | -from typing import TypeAlias |
| 6 | +from scipy.spatial.distance import directed_hausdorff |
7 | 7 | import numpy as np |
8 | 8 |
|
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() |
12 | 23 |
|
13 | 24 | @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