class PrioritizationEngine:
"""Score and rank variants by pathogenicity.
Processes an iterator of annotated variants through pathogenicity scoring:
normalizes CADD/REVEL scores, computes a composite rank, and sorts each
batch in descending order by composite rank (nulls last).
Frequency-based exclusion is handled downstream by the ACMG classifier
via BA1/BS1 benign evidence tags, not by this engine.
Variants are processed in configurable batches to bound memory usage.
On MemoryError, the engine retries with a reduced chunk size (capped at
500,000 variants per chunk).
Parameters
----------
config : PrioritizationConfig, optional
Configuration containing ``batch_size`` and optional score file
paths. When None, defaults are used (batch_size=10,000).
Raises
------
ValueError
If ``config.batch_size`` is outside [1,000, 100,000]. Enforced at
config construction time via ``PrioritizationConfig.__post_init__``.
"""
def __init__(
self,
config: PrioritizationConfig | None = None,
remote_config: RemoteTabixConfig | None = None,
) -> None:
if config is None:
config = PrioritizationConfig()
self._config = config
self._batch_size = config.batch_size
self._score_loader = ScoreLoader()
self._cadd_scores: dict[CoordinateKey, float] = {}
self._revel_scores: dict[CoordinateKey, float] = {}
self._spliceai_scores: dict[CoordinateKey, float] = {}
self._spliceai_db: SpliceAISQLiteLoader | None = None
self._remote_cadd: RemoteTabixCADD | None = None
if config.cadd_scores_path is not None:
self._cadd_scores = self._score_loader.load_cadd(config.cadd_scores_path)
elif remote_config is not None and remote_config.is_cadd_active:
# Remote tabix CADD: no local file, use remote backend
from vartriage.remote.cadd import RemoteTabixCADD as _RemoteCADD
self._remote_cadd = _RemoteCADD(remote_config)
logger.info(
"Remote CADD tabix backend active: %s", remote_config.cadd_remote_url
)
if config.revel_scores_path is not None:
self._revel_scores = self._score_loader.load_revel(config.revel_scores_path)
# SpliceAI: SQLite backend takes precedence over TSV
if config.spliceai_db_path is not None:
from vartriage.prioritization.spliceai_db import (
SpliceAISQLiteLoader as _SpliceAILoader,
)
self._spliceai_db = _SpliceAILoader(config.spliceai_db_path)
logger.info("SpliceAI SQLite backend active: %s", config.spliceai_db_path)
elif config.spliceai_scores_path is not None:
self._spliceai_scores = self._score_loader.load_spliceai(
config.spliceai_scores_path
)
def close(self) -> None:
"""Release resources held by the engine."""
if self._spliceai_db is not None:
self._spliceai_db.close()
self._spliceai_db = None
def prioritize(
self, variants: Iterator[AnnotatedVariant]
) -> Iterator[ScoredVariant]:
"""Score variants for pathogenicity ranking.
All variants pass through scoring regardless of allele frequency.
Frequency-based evidence (BA1/BS1) is applied downstream by the
ACMG classifier, not by this method.
Parameters
----------
variants : Iterator[AnnotatedVariant]
Input stream of annotated variant records. May be empty.
Yields
------
ScoredVariant
Scored variants sorted in descending order by composite
pathogenicity rank within each batch. Variants with null
composite rank appear last.
"""
yield from self._process_in_batches(variants)
def _process_in_batches(
self, variants: Iterator[AnnotatedVariant]
) -> Iterator[ScoredVariant]:
"""Score variants in configurable batches.
Parameters
----------
variants : Iterator[AnnotatedVariant]
Filtered variant stream.
Yields
------
ScoredVariant
Scored variants from each batch, sorted within batch.
"""
batch_size = self._batch_size
while True:
batch = list(islice(variants, batch_size))
if not batch:
break
try:
scored = self._score_batch(batch)
except MemoryError:
logger.warning(
"MemoryError during scoring of batch (size=%d). "
"Falling back to chunked processing.",
len(batch),
)
scored = self._chunked_fallback(batch)
yield from scored
def _score_batch(self, batch: list[AnnotatedVariant]) -> list[ScoredVariant]:
"""Score a single batch of variants.
Extracts coordinate keys from each variant and performs lookups
against pre-loaded CADD and REVEL score dictionaries. When no local
CADD scores exist and a remote tabix backend is active, fetches
CADD scores via HTTP byte-range queries.
Parameters
----------
batch : list[AnnotatedVariant]
Batch of annotated variants to score.
Returns
-------
list[ScoredVariant]
Scored variants sorted descending by composite rank, nulls last.
"""
keys: list[CoordinateKey] = [
(v.variant.chrom, v.variant.pos, v.variant.ref, v.variant.alt)
for v in batch
]
# CADD: prefer local dict, fall back to remote tabix
if self._cadd_scores:
cadd_scores = self._score_loader.lookup_batch(keys, self._cadd_scores)
elif self._remote_cadd is not None:
remote_dict = self._remote_cadd.lookup_batch(keys)
cadd_scores = [remote_dict.get(k) for k in keys]
else:
cadd_scores = [None] * len(keys)
revel_scores = self._score_loader.lookup_batch(keys, self._revel_scores)
spliceai_scores: list[float | None] | None = None
if self._spliceai_db is not None:
spliceai_scores = self._spliceai_db.lookup_batch(keys)
elif self._spliceai_scores:
spliceai_scores = self._score_loader.lookup_batch(
keys, self._spliceai_scores
)
return score_variants(batch, cadd_scores, revel_scores, spliceai_scores)
def _chunked_fallback(self, batch: list[AnnotatedVariant]) -> list[ScoredVariant]:
"""Process a batch in smaller chunks after a MemoryError.
Splits the batch into chunks of at most ``_MAX_CHUNK_SIZE``
(500,000) variants and scores each chunk independently. Results
from all chunks are merged and re-sorted by composite rank.
Parameters
----------
batch : list[AnnotatedVariant]
The full batch that triggered MemoryError.
Returns
-------
list[ScoredVariant]
All scored variants merged and sorted.
"""
chunk_size = min(_MAX_CHUNK_SIZE, max(1, len(batch) // 2))
all_scored: list[ScoredVariant] = []
for start in range(0, len(batch), chunk_size):
chunk = batch[start : start + chunk_size]
try:
scored_chunk = self._score_batch(chunk)
all_scored.extend(scored_chunk)
except MemoryError:
logger.error(
"MemoryError persists even with chunk_size=%d. Reducing further.",
chunk_size,
)
smaller_size = max(1, chunk_size // 2)
for sub_start in range(0, len(chunk), smaller_size):
sub_chunk = chunk[sub_start : sub_start + smaller_size]
scored_sub = self._score_batch(sub_chunk)
all_scored.extend(scored_sub)
return _merge_sort_scored(all_scored)