segtraq 0.0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
segtraq/__init__.py ADDED
@@ -0,0 +1,9 @@
1
+ """Top-level package for SegTraQ."""
2
+
3
+ __author__ = """Daria Lazic, Matthias Meyer-Bender, Martin Emons"""
4
+ __email__ = "daria.lazic@embl.de, matthias.meyerbender@embl.de, martin.emons@uzh.ch"
5
+
6
+ from . import bl, cs, io, nc, pl, segtraq_metrics, sp
7
+ from .utils import run_label_transfer # TODO - maybe move to baseline metrics later?
8
+
9
+ __all__ = ["bl", "cs", "nc", "io", "sp", "pl", "segtraq_metrics", "run_label_transfer"]
segtraq/bl/__init__.py ADDED
@@ -0,0 +1,21 @@
1
+ from .baseline import (
2
+ genes_per_cell,
3
+ morphological_features,
4
+ num_cells,
5
+ num_genes,
6
+ num_transcripts,
7
+ perc_unassigned_transcripts,
8
+ transcript_density,
9
+ transcripts_per_cell,
10
+ )
11
+
12
+ __all__ = [
13
+ "num_cells",
14
+ "num_transcripts",
15
+ "num_genes",
16
+ "perc_unassigned_transcripts",
17
+ "transcripts_per_cell",
18
+ "genes_per_cell",
19
+ "transcript_density",
20
+ "morphological_features",
21
+ ]
segtraq/bl/baseline.py ADDED
@@ -0,0 +1,372 @@
1
+ import geopandas as gpd
2
+ import numpy as np
3
+ import pandas as pd
4
+ import spatialdata as sd
5
+ from joblib import Parallel, delayed
6
+ from shapely.geometry import Polygon
7
+
8
+
9
+ def num_cells(sdata: sd.SpatialData, table_key: str = "table") -> int:
10
+ """
11
+ Counts the number of cells in the given SpatialData object based on the specified table key.
12
+
13
+ Parameters
14
+ ----------
15
+ sdata : sd.SpatialData
16
+ The SpatialData object containing spatial information and a table.
17
+ cell_id : str, optional
18
+ The key in the `tables` attribute of `sdata` that corresponds to table.
19
+ Default is "table".
20
+
21
+ Returns
22
+ -------
23
+ int
24
+ The number of cells found under the specified table key.
25
+ """
26
+ return len(sdata.tables[table_key])
27
+
28
+
29
+ def num_transcripts(sdata: sd.SpatialData, transcript_key: str = "transcripts"):
30
+ """
31
+ Counts the total number of transcripts in the given SpatialData object.
32
+
33
+ Parameters
34
+ ----------
35
+ sdata : sd.SpatialData
36
+ The SpatialData object containing transcript information.
37
+ transcript_key : str, optional
38
+ The key to access transcript data within the spatial data object. Default is "transcripts".
39
+
40
+ Returns
41
+ -------
42
+ int
43
+ The total number of transcripts in the specified SpatialData object.
44
+ """
45
+ return sdata.points[transcript_key].shape[0].compute()
46
+
47
+
48
+ def num_genes(
49
+ sdata: sd.SpatialData,
50
+ transcript_key: str = "transcripts",
51
+ gene_key: str = "feature_name",
52
+ ) -> int:
53
+ """
54
+ Counts the number of unique genes in the given SpatialData object.
55
+
56
+ Parameters
57
+ ----------
58
+ sdata : sd.SpatialData
59
+ The SpatialData object containing gene information.
60
+ transcript_key : str, optional
61
+ The key to access transcript data within the spatial data object. Default is "transcripts".
62
+ gene_key : str, optional
63
+ The key to access gene names within the transcript data. Default is "feature_name".
64
+
65
+ Returns
66
+ -------
67
+ int
68
+ The number of unique genes found in the specified SpatialData object.
69
+ """
70
+ # converting from np.int64 to int for consistency
71
+ return int(sdata.points[transcript_key][gene_key].nunique().compute())
72
+
73
+
74
+ def perc_unassigned_transcripts(
75
+ sdata: sd.SpatialData,
76
+ transcript_key: str = "transcripts",
77
+ cell_key: str = "cell_id",
78
+ unassigned_key: int = -1,
79
+ ) -> float:
80
+ """
81
+ Calculates the proportion of unassigned transcripts in a SpatialData object.
82
+
83
+ Parameters
84
+ ----------
85
+ sdata : sd.SpatialData
86
+ The spatial data object containing transcript information.
87
+ transcript_key : str, optional
88
+ The key to access transcript data within the spatial data object. Default is "transcripts".
89
+ cell_key : str, optional
90
+ The key to access cell assignment information within the transcript data. Default is "cell_id".
91
+ unassigned_key : int, optional
92
+ The value indicating an unassigned transcript. Default is -1.
93
+
94
+ Returns
95
+ -------
96
+ float
97
+ The fraction of transcripts that are unassigned.
98
+ """
99
+ counts = sdata.points[transcript_key][cell_key].compute().value_counts()
100
+ num_unassigned = counts.get(unassigned_key, 0)
101
+ # converting from np.float64 to float for consistency
102
+ return float(num_unassigned / counts.sum())
103
+
104
+
105
+ def transcripts_per_cell(
106
+ sdata: sd.SpatialData,
107
+ transcript_key: str = "transcripts",
108
+ cell_key: str = "cell_id",
109
+ ) -> pd.DataFrame:
110
+ """
111
+ Counts the number of transcripts assigned to each cell.
112
+
113
+ Parameters
114
+ ----------
115
+ sdata : sd.SpatialData
116
+ A SpatialData object containing transcript and cell assignment information.
117
+ transcript_key : str, optional
118
+ The key in `sdata.points` corresponding to transcript data. Default is "transcripts".
119
+ cell_key : str, optional
120
+ The column name in the transcript data that contains cell assignment information. Default is "cell_id".
121
+
122
+ Returns
123
+ -------
124
+ pd.DataFrame
125
+ A DataFrame with two columns: the cell identifier (`cell_key`) and the
126
+ corresponding transcript count ("transcript_count").
127
+ """
128
+ counts = sdata.points[transcript_key][cell_key].compute().value_counts().astype("int64")
129
+ counts_df = counts.reset_index()
130
+ counts_df.columns = [cell_key, "transcript_count"]
131
+ return counts_df
132
+
133
+
134
+ def genes_per_cell(sdata, transcript_key="transcripts", cell_key="cell_id", gene_key="feature_name"):
135
+ """
136
+ Calculates the number of unique genes detected per cell.
137
+
138
+ Parameters
139
+ ----------
140
+ sdata : object
141
+ An object containing spatial transcriptomics data with a `points` attribute.
142
+ transcript_key : str, optional
143
+ The key to access the transcript data within `sdata.points` (default is "transcripts").
144
+ cell_key : str, optional
145
+ The column name in the transcript data representing cell identifiers (default is "cell_id").
146
+ gene_key : str, optional
147
+ The column name in the transcript data representing gene names (default is "feature_name").
148
+
149
+ Returns
150
+ -------
151
+ pandas.DataFrame
152
+ A DataFrame with one row per cell, containing the cell identifier and
153
+ the count of unique genes detected in that cell.
154
+ """
155
+ df = sdata.points[transcript_key].compute()
156
+ # Group by cell and count unique genes
157
+ gene_counts = df.groupby(cell_key)[gene_key].nunique().reset_index()
158
+ gene_counts.columns = [cell_key, "gene_count"]
159
+ return gene_counts
160
+
161
+
162
+ def transcript_density(
163
+ sdata: sd.SpatialData,
164
+ table_key: str = "table",
165
+ transcript_key: str = "transcripts",
166
+ cell_key: str = "cell_id",
167
+ area_key: str = "cell_area",
168
+ ) -> pd.DataFrame:
169
+ """
170
+ Calculates the transcript density for each cell in a SpatialData object.
171
+ Transcript density is defined as the number of transcripts per unit area for each cell.
172
+
173
+ Parameters
174
+ ----------
175
+ sdata : sd.SpatialData
176
+ The SpatialData object containing spatial transcriptomics data.
177
+ table_key : str, optional
178
+ The key to access the AnnData table from `sdata.tables`. Default is "table".
179
+ transcript_key : str, optional
180
+ The key in the transcript table indicating transcript identifiers. Default is "transcripts".
181
+ cell_key : str, optional
182
+ The key in the table indicating cell identifiers. Default is "cell_id".
183
+ area_key: str, optional
184
+ The key in the table indicating the cell area/volume. Default is "cell_area".
185
+
186
+ Returns
187
+ -------
188
+ pd.DataFrame
189
+ A DataFrame with columns `[cell_key, "transcript_density"]`,
190
+ where "transcript_density" is the number of transcripts per unit area for
191
+ each cell. Rows with missing values are dropped.
192
+ """
193
+ adata = sdata.tables[table_key]
194
+ counts_df = transcripts_per_cell(sdata, transcript_key, cell_key)
195
+ area_df = adata.obs[[cell_key, area_key]]
196
+
197
+ merged = counts_df.merge(area_df, on=cell_key, how="left")
198
+ merged["transcript_density"] = merged["transcript_count"] / merged[area_key]
199
+
200
+ return merged[[cell_key, "transcript_density"]].dropna()
201
+
202
+
203
+ def morphological_features(
204
+ sdata,
205
+ shape_key: str = "cell_boundaries",
206
+ id_key: str = "cell_id",
207
+ features_to_compute: list = None,
208
+ n_jobs: int = -1, # number of parallel jobs, -1 uses all CPUs
209
+ ):
210
+ """
211
+ Compute morphological features for cell shapes in a spatial transcriptomics dataset.
212
+
213
+ Parameters
214
+ ----------
215
+ sdata : object
216
+ Spatial data object containing cell shape information. Must have a `.shapes` attribute with geometries.
217
+ shape_key : str, optional
218
+ Key in `sdata.shapes` specifying the geometry column (default is "cell_boundaries").
219
+ id_key : str, optional
220
+ Key in `sdata.shapes` specifying the unique cell identifier column (default is "cell_id").
221
+ features_to_compute : list of str, optional
222
+ List of morphological features to compute. If None, all available features are computed.
223
+ Available features: "cell_area", "perimeter", "circularity", "bbox_width", "bbox_height",
224
+ "extent", "solidity", "convexity", "elongation", "eccentricity", "compactness", "sphericity".
225
+ n_jobs : int, optional
226
+ Number of parallel jobs to use for computation. -1 uses all available CPUs (default is -1).
227
+
228
+ Returns
229
+ -------
230
+ features : pandas.DataFrame
231
+ DataFrame containing the computed morphological features for each cell, indexed by `id_key`.
232
+
233
+ Raises
234
+ ------
235
+ ValueError
236
+ If any requested feature in `features_to_compute` is not recognized.
237
+
238
+ Notes
239
+ -----
240
+ - Requires `geopandas`, `shapely`, `numpy`, `pandas`, and `joblib`.
241
+ - Some features are proxies or approximations (e.g., "sphericity" uses "circularity").
242
+ - Invalid or null geometries are filtered out before computation.
243
+ """
244
+ # Define all possible features
245
+ all_features = [
246
+ "cell_area",
247
+ "perimeter",
248
+ "circularity",
249
+ "bbox_width",
250
+ "bbox_height",
251
+ "extent",
252
+ "solidity",
253
+ "convexity",
254
+ "elongation",
255
+ "eccentricity",
256
+ "compactness",
257
+ ]
258
+
259
+ # If no features specified, compute all
260
+ if features_to_compute is None:
261
+ features_to_compute = all_features
262
+ else:
263
+ # Validate features requested
264
+ invalid_feats = set(features_to_compute) - set(all_features)
265
+ if invalid_feats:
266
+ raise ValueError(f"Unknown features requested: {invalid_feats}")
267
+
268
+ cells = sdata.shapes[shape_key]
269
+ if not isinstance(cells, gpd.GeoDataFrame):
270
+ cells = cells.to_gdf()
271
+
272
+ # Filter valid geometries
273
+ cells = cells[cells.geometry.notnull() & cells.geometry.is_valid].copy().reset_index()
274
+
275
+ features = pd.DataFrame()
276
+ features[id_key] = cells[id_key].values
277
+ geom = cells.geometry
278
+
279
+ # Compute features conditionally
280
+ if "cell_area" in features_to_compute or any(
281
+ f in features_to_compute for f in ["circularity", "extent", "solidity", "compactness", "sphericity"]
282
+ ):
283
+ areas = geom.area
284
+ if "cell_area" in features_to_compute:
285
+ features["cell_area"] = areas
286
+ else:
287
+ areas = None
288
+
289
+ if "perimeter" in features_to_compute or any(
290
+ f in features_to_compute
291
+ for f in [
292
+ "circularity",
293
+ "compactness",
294
+ "convexity",
295
+ "compactness",
296
+ "sphericity",
297
+ ]
298
+ ):
299
+ perimeters = geom.length
300
+ if "perimeter" in features_to_compute:
301
+ features["perimeter"] = perimeters
302
+ else:
303
+ perimeters = None
304
+
305
+ if "circularity" in features_to_compute:
306
+ if areas is None:
307
+ areas = geom.area
308
+ if perimeters is None:
309
+ perimeters = geom.length
310
+ features["circularity"] = 4 * np.pi * areas / (perimeters**2 + 1e-6)
311
+
312
+ if any(f in features_to_compute for f in ["bbox_width", "bbox_height", "extent"]):
313
+ bounds = geom.bounds
314
+ if "bbox_width" in features_to_compute:
315
+ features["bbox_width"] = bounds["maxx"] - bounds["minx"]
316
+ if "bbox_height" in features_to_compute:
317
+ features["bbox_height"] = bounds["maxy"] - bounds["miny"]
318
+ if "extent" in features_to_compute:
319
+ width = bounds["maxx"] - bounds["minx"]
320
+ height = bounds["maxy"] - bounds["miny"]
321
+ if areas is None:
322
+ areas = geom.area
323
+ features["extent"] = areas / (width * height + 1e-6)
324
+
325
+ if "solidity" in features_to_compute or "convexity" in features_to_compute:
326
+ convex_hull = geom.convex_hull
327
+ if "solidity" in features_to_compute:
328
+ convex_areas = convex_hull.area
329
+ if areas is None:
330
+ areas = geom.area
331
+ features["solidity"] = areas / (convex_areas + 1e-6)
332
+ if "convexity" in features_to_compute:
333
+ convex_perimeters = convex_hull.length
334
+ if perimeters is None:
335
+ perimeters = geom.length
336
+ features["convexity"] = convex_perimeters / (perimeters + 1e-6)
337
+
338
+ # Parallelized elongation and eccentricity calculation
339
+ def compute_elong_ecc(poly):
340
+ if not isinstance(poly, Polygon) or poly.is_empty:
341
+ return np.nan, np.nan
342
+
343
+ min_rect = poly.minimum_rotated_rectangle
344
+ coords = list(min_rect.exterior.coords)
345
+ edges = [np.linalg.norm(np.array(coords[i]) - np.array(coords[i + 1])) for i in range(4)]
346
+ edges = sorted(edges)
347
+ if len(edges) < 2 or edges[1] == 0:
348
+ return np.nan, np.nan
349
+
350
+ elongation = edges[2] / edges[1]
351
+ a = edges[2] / 2
352
+ b = edges[1] / 2
353
+ eccentricity = np.sqrt(a**2 - b**2) / a if a > 0 else np.nan
354
+
355
+ return elongation, eccentricity
356
+
357
+ if "elongation" in features_to_compute or "eccentricity" in features_to_compute:
358
+ results = Parallel(n_jobs=n_jobs)(delayed(compute_elong_ecc)(poly) for poly in geom)
359
+ elongations, eccentricities = zip(*results, strict=False)
360
+ if "elongation" in features_to_compute:
361
+ features["elongation"] = elongations
362
+ if "eccentricity" in features_to_compute:
363
+ features["eccentricity"] = eccentricities
364
+
365
+ if "compactness" in features_to_compute:
366
+ if perimeters is None:
367
+ perimeters = geom.length
368
+ if areas is None:
369
+ areas = geom.area
370
+ features["compactness"] = (perimeters**2) / (areas + 1e-6)
371
+
372
+ return features
segtraq/cs/__init__.py ADDED
@@ -0,0 +1,17 @@
1
+ from .clustering_stability import (
2
+ compute_ari,
3
+ compute_mean_cosine_distance,
4
+ compute_purity,
5
+ compute_rmsd,
6
+ compute_silhouette_score,
7
+ compute_z_plane_correlation,
8
+ )
9
+
10
+ __all__ = [
11
+ "compute_ari",
12
+ "compute_silhouette_score",
13
+ "compute_purity",
14
+ "compute_rmsd",
15
+ "compute_mean_cosine_distance",
16
+ "compute_z_plane_correlation",
17
+ ]