"""Deterministic scientific scoring embedding for scCS v0.8."""
from __future__ import annotations
from dataclasses import asdict, dataclass
from typing import Optional, Sequence, Union
import numpy as np
from .furcation import Furcation
from .geometry import SimplexStarGeometry
from .ordering import FurcationOrderingResult, FurcationOrderingScaler
@dataclass(frozen=True)
class ScoringEmbeddingResult:
"""Scientific coordinates for one annotation-defined furcation.
Root cells are ordered along the incoming arm. Every annotated terminal
cell is placed at its equal-radius simplex vertex. Display-only terminal
clouds are generated by :meth:`SingleScorer.plot_star` and are not stored
in the scientific embedding.
"""
coordinates: np.ndarray
selected_indices: np.ndarray
selected_cell_ids: np.ndarray
selected_labels: np.ndarray
selected_ordering_values: np.ndarray
root_mask: np.ndarray
terminal_mask: np.ndarray
terminal_names: np.ndarray
geometry: SimplexStarGeometry
ordering: FurcationOrderingResult
arm_scale: float
ordering_key: Optional[str]
@property
def n_selected(self) -> int:
return len(self.selected_indices)
@property
def dimension(self) -> int:
return self.coordinates.shape[1]
def write_to_adata(
self,
adata,
*,
coordinate_key: str = "X_sccs_score",
metadata_key: str = "sccs_v08",
) -> None:
"""Write full-length coordinates and metadata without subsetting AnnData."""
full = np.full(
(adata.n_obs, self.dimension),
np.nan,
dtype=float,
)
full[self.selected_indices] = self.coordinates
adata.obsm[coordinate_key] = full
metadata = dict(adata.uns.get(metadata_key, {}))
metadata["fate_names"] = list(self.geometry.fate_names)
metadata["root_direction"] = self.geometry.root_direction.tolist()
metadata["terminal_directions"] = self.geometry.terminal_directions.tolist()
metadata["selected_indices"] = self.selected_indices.tolist()
metadata["arm_scale"] = float(self.arm_scale)
metadata["ordering_key"] = self.ordering_key
metadata["terminal_coordinate_mode"] = "fixed_vertex"
metadata["terminal_scientific_radius"] = float(self.arm_scale)
metadata["ordering_diagnostics"] = asdict(self.ordering.diagnostics)
adata.uns[metadata_key] = metadata
def _resolve_ordering(
adata,
ordering: Union[str, Sequence[float], np.ndarray],
) -> tuple[np.ndarray, Optional[str]]:
if isinstance(ordering, str):
if ordering not in adata.obs:
raise KeyError(f"Ordering column {ordering!r} is missing from adata.obs.")
values = adata.obs[ordering].to_numpy(dtype=float)
key = ordering
else:
values = np.asarray(ordering, dtype=float)
key = None
if values.ndim != 1 or len(values) != adata.n_obs:
raise ValueError("Array ordering must be one-dimensional and match adata.n_obs.")
# Finiteness is validated after the supervised furcation mask is applied
# by FurcationOrderingScaler. Missing values on unrelated cells are not
# relevant to this furcation.
return values, key
[docs]
def build_scoring_embedding(
adata,
furcation: Furcation,
*,
ordering: Union[str, Sequence[float], np.ndarray],
ordering_scaler: Optional[FurcationOrderingScaler] = None,
arm_scale: float = 1.0,
write_to_adata: bool = False,
coordinate_key: str = "X_sccs_score",
) -> ScoringEmbeddingResult:
"""Build deterministic root-plus-simplex scientific coordinates.
The function performs no topology inference and adds no jitter. Only the
manually annotated root and terminal populations are selected. Root cells
are ordered along the incoming arm; terminal cells are fixed at equal
simplex vertices independent of terminal pseudotime or abundance.
"""
if not np.isfinite(arm_scale) or arm_scale <= 0:
raise ValueError("arm_scale must be positive and finite.")
validation = furcation.validate_adata(adata)
labels_full = adata.obs[furcation.obs_key].astype(str).to_numpy()
ordering_values, ordering_key = _resolve_ordering(adata, ordering)
selected_indices = np.flatnonzero(validation.selected_mask)
selected_labels = labels_full[selected_indices]
selected_cell_ids = np.asarray(adata.obs_names[selected_indices]).astype(str)
scaler = ordering_scaler or FurcationOrderingScaler()
selected_ordering_values = np.asarray(ordering_values[selected_indices], dtype=float)
if not scaler.higher_is_later:
selected_ordering_values = -selected_ordering_values
ordering_result = scaler.fit_transform(
ordering_values,
labels_full,
furcation,
)
geometry = SimplexStarGeometry(furcation.terminal_names)
coordinates = np.empty(
(len(selected_indices), geometry.dimension),
dtype=float,
)
root_mask = ordering_result.root_mask
terminal_mask = ordering_result.terminal_mask
coordinates[root_mask] = geometry.root_coordinates(
ordering_result.root_progress[root_mask],
arm_scale=arm_scale,
)
coordinates[terminal_mask] = geometry.terminal_coordinates(
ordering_result.terminal_names[terminal_mask],
arm_scale=arm_scale,
)
if not np.all(np.isfinite(coordinates)):
raise RuntimeError("Scientific scoring coordinates contain non-finite values.")
result = ScoringEmbeddingResult(
coordinates=coordinates,
selected_indices=selected_indices,
selected_cell_ids=selected_cell_ids,
selected_labels=selected_labels,
selected_ordering_values=selected_ordering_values,
root_mask=root_mask,
terminal_mask=terminal_mask,
terminal_names=ordering_result.terminal_names,
geometry=geometry,
ordering=ordering_result,
arm_scale=float(arm_scale),
ordering_key=ordering_key,
)
if write_to_adata:
result.write_to_adata(adata, coordinate_key=coordinate_key)
return result