Source code for vdjtools.sc.pgen

"""Paired-chain generation probability for single-cell repertoires.

Under chain independence the paired generation probability of a cell is
``Pgen(α) · Pgen(β)`` — the product of each chain's junction Pgen under the native
:mod:`vdjtools.model` engine (bundled per-locus models). This is the single-cell
paired-Pgen residual from Phase 7; it is computed entirely from the native model
(no ``vdjmatch`` dependency).

The paired frame is the :func:`vdjtools.sc.pair.pair_chains` layout — ``alpha_v_call``,
``alpha_j_call``, ``alpha_junction_aa`` and the ``beta_*`` counterparts (α/light and
β/heavy). Each chain's locus is inferred from its V-call prefix (``TRA``/``TRB``, or
``IGK``/``IGL`` + ``IGH`` for BCR) unless given explicitly.

Conditioning on V/J requires the call to match a model **allele** (e.g. ``TRBV20-1*01``);
a gene-level or unmatched call marginalises over all V/J for that chain (still a valid,
if less specific, Pgen). Pass ``condition_vj=False`` to marginalise unconditionally.
"""
from __future__ import annotations

import polars as pl

from ..model import load_bundled, native

ALPHA_V, ALPHA_J, ALPHA_AA = "alpha_v_call", "alpha_j_call", "alpha_junction_aa"
BETA_V, BETA_J, BETA_AA = "beta_v_call", "beta_j_call", "beta_junction_aa"


def _infer_locus(vcalls: pl.Series) -> str | None:
    """Most common three-letter locus prefix among the non-null V calls."""
    pref = vcalls.drop_nulls().str.slice(0, 3)
    if pref.len() == 0:
        return None
    m = pref.mode()
    return m[0] if m.len() else None


def _chain_scores(model, aas: list, vs: list, js: list, condition_vj: bool) -> list:
    """Pgen per row for one chain -- one native batch over the **distinct** clonotype keys.

    Three things this has to get right, and a per-row loop over ``native.pgen_aa`` only got the
    first two:

    * **Deduplicate, do not cache.** Cells of an expanded clone share a clonotype key and the
      native Pgen is deterministic in ``(junction, v, j)``, so scoring the distinct keys within one
      call is exact deduplication -- nothing survives the call.
    * **An unnameable V/J is a null for its own rows only.** The model is keyed by allele and
      raises on anything else, deliberately: mapping an unknown call to "marginalise" returned a
      value 2.38x too high with no error. ``pgen_aa_batch`` resolves names itself and would raise
      for the **whole batch**, so names are resolved here first and the unresolvable keys never
      enter it.
    * **Batch it.** ``pgen_aa_batch`` releases the GIL and threads across the input; the per-row
      loop was serial, and a whole barcoded dataset is tens of thousands of receptors.
    """
    n = len(aas)
    if model is None:
        return [None] * n
    _pm, vi, ji = native.pack(model)
    row_key: list = [None] * n
    nameable: dict = {}
    for i in range(n):
        aa = aas[i]
        if not isinstance(aa, str) or not aa:
            continue
        k = (aa, vs[i], js[i]) if condition_vj else (aa, None, None)
        row_key[i] = k
        if k not in nameable:
            try:
                native._gene_idx(vi, k[1], "V")
                native._gene_idx(ji, k[2], "J")
                nameable[k] = True
            except KeyError:
                nameable[k] = False
    keys = [k for k, good in nameable.items() if good]
    scored = dict(zip(keys, native.pgen_aa_batch(
        model, [k[0] for k in keys], [k[1] for k in keys], [k[2] for k in keys]))) if keys else {}
    return [None if k is None else scored.get(k) for k in row_key]


def _warn_if_all_null(values, model, locus, col) -> None:
    """A whole column of nulls is almost always a naming mismatch -- say so, don't ship it."""
    if model is None or not values or any(v is not None for v in values):
        return
    import warnings

    warnings.warn(
        f"paired_pgen: every {locus} chain scored null ({len(values)} rows). The usual cause "
        f"is a {col} naming the model does not carry; check a value against the model's "
        "alleles, or pass condition_vj=False to marginalise over V/J deliberately.",
        UserWarning, stacklevel=3,
    )


[docs] def paired_pgen( paired: pl.DataFrame, *, source: str = "olga", condition_vj: bool = True, resolve_genes: bool = True, alpha_locus: str | None = None, beta_locus: str | None = None, ) -> pl.DataFrame: """Add ``pgen_alpha``, ``pgen_beta`` and ``pgen_paired`` to a paired single-cell frame. Args: paired: A paired-chain frame (:func:`vdjtools.sc.pair.pair_chains` layout). source: Bundled model set — ``"olga"`` (OLGA-derived) or ``"learned"`` (native EM). condition_vj: Condition each chain's Pgen on its V/J call. ``False`` marginalises over all V/J unconditionally. resolve_genes: Resolve a **gene**-level call (``TRBV10-3``) to a representative model allele (``TRBV10-3*01``) before scoring -- see :func:`vdjtools.model.native.gene_to_allele`. Default ``True``, because CellRanger reports genes and without this every 10x row scores ``None``. Set ``False`` to score only exact allele matches. alpha_locus: Locus of the α/light chain (e.g. ``"TRA"``, ``"IGK"``); inferred from the ``alpha_v_call`` prefix if ``None``. beta_locus: Locus of the β/heavy chain (e.g. ``"TRB"``, ``"IGH"``); inferred from the ``beta_v_call`` prefix if ``None``. Returns: ``paired`` with three added Float64 columns. ``pgen_alpha`` / ``pgen_beta`` are null for a cell missing that chain's junction, or carrying a V/J call the model does not know; ``pgen_paired`` is null unless both are set. Warns: UserWarning: If every chain of a locus scored null -- the usual cause is a V/J naming the model does not recognise, which would otherwise be an entire column of silent nulls. """ a_loc = alpha_locus or (_infer_locus(paired[ALPHA_V]) if ALPHA_V in paired.columns else None) b_loc = beta_locus or (_infer_locus(paired[BETA_V]) if BETA_V in paired.columns else None) ma = load_bundled(a_loc, source) if a_loc else None mb = load_bundled(b_loc, source) if b_loc else None # gene -> representative allele, per model (the two loci have disjoint gene names). aliases: dict[str, str] = {} if resolve_genes and condition_vj: for m in (ma, mb): if m is not None: aliases.update(native.gene_to_allele(m)) def _call(name): return aliases.get(name, name) if name else name def col(name): return (paired[name].to_list() if name in paired.columns else [None] * paired.height) pa = _chain_scores(ma, col(ALPHA_AA), [_call(x) for x in col(ALPHA_V)], [_call(x) for x in col(ALPHA_J)], condition_vj) pb = _chain_scores(mb, col(BETA_AA), [_call(x) for x in col(BETA_V)], [_call(x) for x in col(BETA_J)], condition_vj) pp = [a * b if (a is not None and b is not None) else None for a, b in zip(pa, pb)] _warn_if_all_null(pa, ma, a_loc, ALPHA_V) _warn_if_all_null(pb, mb, b_loc, BETA_V) return paired.with_columns( pl.Series("pgen_alpha", pa, dtype=pl.Float64), pl.Series("pgen_beta", pb, dtype=pl.Float64), pl.Series("pgen_paired", pp, dtype=pl.Float64), )