Source code for scCS.scoring_embedding

"""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