Skip to content

Cohort Analysis

Multi-sample cohort analysis classes for cross-sample variant aggregation and reporting.

CohortPipeline

vartriage.cohort.pipeline.CohortPipeline

Orchestrate multi-sample cohort analysis.

Processes each sample VCF through the standard vartriage pipeline, collects classified variants, then merges them via CohortAggregator for cross-sample analysis.

Parameters

cohort_config : CohortConfig Cohort-level settings (sample list, thresholds, output). pipeline_config : PipelineConfig | None Base pipeline configuration applied to each sample. When None, a minimal config is constructed per-sample using only the cohort_config's sample VCF paths with default quality/prioritization settings. annotation_config : AnnotationConfig | None Shared annotation config for all samples. Overrides the pipeline_config's annotation setting when provided. prioritization_config : PrioritizationConfig | None Shared prioritization config. Overrides pipeline_config when provided.

Source code in vartriage/cohort/pipeline.py
class CohortPipeline:
    """Orchestrate multi-sample cohort analysis.

    Processes each sample VCF through the standard vartriage pipeline,
    collects classified variants, then merges them via CohortAggregator
    for cross-sample analysis.

    Parameters
    ----------
    cohort_config : CohortConfig
        Cohort-level settings (sample list, thresholds, output).
    pipeline_config : PipelineConfig | None
        Base pipeline configuration applied to each sample. When None,
        a minimal config is constructed per-sample using only the
        cohort_config's sample VCF paths with default quality/prioritization
        settings.
    annotation_config : AnnotationConfig | None
        Shared annotation config for all samples. Overrides the
        pipeline_config's annotation setting when provided.
    prioritization_config : PrioritizationConfig | None
        Shared prioritization config. Overrides pipeline_config when provided.
    """

    def __init__(
        self,
        cohort_config: CohortConfig,
        pipeline_config: PipelineConfig | None = None,
        annotation_config: AnnotationConfig | None = None,
        prioritization_config: PrioritizationConfig | None = None,
    ) -> None:
        self._cohort_config = cohort_config
        self._base_pipeline_config = pipeline_config
        self._annotation_config = annotation_config
        self._prioritization_config = prioritization_config
        self._aggregator = CohortAggregator(cohort_config)

        # Results populated after run()
        self._variants: list[CohortVariant] = []
        self._gene_burdens: list[GeneBurden] = []
        self._summary: CohortSummary | None = None
        self._samples_processed: list[str] = []

    @property
    def variants(self) -> list[CohortVariant]:
        """Aggregated cohort variants (populated after run())."""
        return self._variants

    @property
    def gene_burdens(self) -> list[GeneBurden]:
        """Per-gene burden records (populated after run())."""
        return self._gene_burdens

    @property
    def summary(self) -> CohortSummary | None:
        """Cohort summary statistics (populated after run())."""
        return self._summary

    def run(self) -> list[Path]:
        """Execute the full cohort analysis pipeline.

        Sequence:
        1. Process each sample VCF through the standard pipeline
        2. Aggregate variants across samples
        3. Compute cohort statistics
        4. Generate reports

        Returns
        -------
        list[Path]
            Paths to generated report files.

        Raises
        ------
        FileNotFoundError
            If any sample VCF file does not exist.
        """
        logger.info(
            "Starting cohort analysis '%s' with %d samples",
            self._cohort_config.cohort_name,
            self._cohort_config.sample_count,
        )

        # Validate all VCF files exist before processing
        for vcf_path in self._cohort_config.sample_vcfs:
            if not vcf_path.exists():
                raise FileNotFoundError(f"Sample VCF not found: {vcf_path}")

        # TemporaryDirectory context manager guarantees cleanup on
        # success, failure, or keyboard interrupt.
        with tempfile.TemporaryDirectory(prefix="vartriage_cohort_") as tmp_dir:
            tmp_path = Path(tmp_dir)

            if self._cohort_config.parallel:
                self._process_parallel(tmp_path)
            else:
                self._process_sequential(tmp_path)

        # Aggregate
        logger.info(
            "Aggregating variants across %d samples",
            len(self._samples_processed),
        )
        self._variants = self._aggregator.aggregate()

        # Statistics
        stats = CohortStatistics(self._cohort_config, self._variants)
        self._gene_burdens = stats.compute_gene_burden()
        self._summary = stats.compute_summary(self._samples_processed)

        # Report generation
        reporter = CohortReportGenerator(self._cohort_config)
        report_paths = reporter.generate(
            self._variants, self._gene_burdens, self._summary
        )

        logger.info(
            "Cohort analysis complete: %d variants, %d shared, %d genes",
            self._summary.total_variants,
            self._summary.shared_variants,
            self._summary.genes_affected,
        )

        return report_paths

    def _process_sequential(self, tmp_output_dir: Path) -> None:
        """Process each sample VCF sequentially."""
        for vcf_path in self._cohort_config.sample_vcfs:
            sample_id = self._cohort_config.label_for(vcf_path)
            classified = self._run_single_sample(vcf_path, sample_id, tmp_output_dir)
            self._aggregator.add_sample(sample_id, vcf_path, classified)
            self._samples_processed.append(sample_id)

    def _process_parallel(self, tmp_output_dir: Path) -> None:
        """Process sample VCFs concurrently using a thread pool.

        Thread-based parallelism works here because the per-sample
        pipeline is I/O-bound (VCF parsing, reference file reads via
        pysam which releases the GIL during C-level I/O).
        """
        max_workers = self._cohort_config.max_workers
        futures_map: dict[Future[list[ClassifiedVariant]], tuple[str, Path]] = {}

        with ThreadPoolExecutor(max_workers=max_workers) as executor:
            for vcf_path in self._cohort_config.sample_vcfs:
                sample_id = self._cohort_config.label_for(vcf_path)
                future = executor.submit(
                    self._run_single_sample,
                    vcf_path,
                    sample_id,
                    tmp_output_dir,
                )
                futures_map[future] = (sample_id, vcf_path)

            for future in as_completed(futures_map):
                sample_id, vcf_path = futures_map[future]
                try:
                    classified: list[ClassifiedVariant] = future.result()
                    self._aggregator.add_sample(sample_id, vcf_path, classified)
                    self._samples_processed.append(sample_id)
                except Exception:
                    logger.exception(
                        "Failed to process sample '%s' (%s)",
                        sample_id,
                        vcf_path,
                    )
                    raise

    def _run_single_sample(
        self, vcf_path: Path, sample_id: str, tmp_output_dir: Path
    ) -> list[ClassifiedVariant]:
        """Run the standard pipeline on a single VCF and collect results.

        Delegates to Pipeline.run_to_classification() which executes
        all stages up through ACMG classification without writing a
        report file.

        Parameters
        ----------
        vcf_path : Path
            Path to the sample's VCF file.
        sample_id : str
            Sample identifier for logging.
        tmp_output_dir : Path
            Temporary directory for placeholder output paths.

        Returns
        -------
        list[ClassifiedVariant]
            All classified variants from this sample.
        """
        from vartriage.pipeline import Pipeline

        logger.info("Processing sample '%s': %s", sample_id, vcf_path)

        config = self._build_sample_config(vcf_path, tmp_output_dir)
        pipeline = Pipeline(config)
        classified = list(pipeline.run_to_classification(vcf_path))

        logger.info(
            "Sample '%s' yielded %d classified variants",
            sample_id,
            len(classified),
        )
        return classified

    def _build_sample_config(
        self, vcf_path: Path, tmp_output_dir: Path
    ) -> PipelineConfig:
        """Build a PipelineConfig for a single sample.

        If a base pipeline_config was provided, clones it with the
        sample's VCF path. Otherwise builds a minimal config.
        """
        tmp_output = tmp_output_dir / f"{vcf_path.stem}_output.json"

        if self._base_pipeline_config is not None:
            return self._clone_base_config(
                self._base_pipeline_config, vcf_path, tmp_output
            )

        # Minimal config when no base is provided
        return PipelineConfig(
            vcf_path=vcf_path,
            output_path=tmp_output,
            annotation=self._annotation_config,
            prioritization=(self._prioritization_config or PrioritizationConfig()),
            report=ReportConfig(output_format="json"),
        )

    def _clone_base_config(
        self, base: PipelineConfig, vcf_path: Path, output_path: Path
    ) -> PipelineConfig:
        """Clone a base PipelineConfig with per-sample overrides.

        Propagates all base config fields so cohort samples behave
        consistently with the single-sample pipeline. Only vcf_path,
        output_path, and report format are overridden per-sample.

        Note: sample, inheritance, and clinical_report are intentionally
        passed through from the base config. run_to_classification()
        does not apply SampleExtractor or InheritanceFilter (those are
        single-sample concerns handled by the full run() path), but
        passing them preserves config integrity for future use.
        """
        return PipelineConfig(
            vcf_path=vcf_path,
            output_path=output_path,
            quality_filter=base.quality_filter,
            annotation=self._annotation_config or base.annotation,
            prioritization=(self._prioritization_config or base.prioritization),
            report=ReportConfig(output_format="json"),
            missing_data=base.missing_data,
            gene_filter=base.gene_filter,
            region_filter=base.region_filter,
            sample=base.sample,
            inheritance=base.inheritance,
            clinical_report=None,
            use_bundles=base.use_bundles,
            genome_build=base.genome_build,
            api=base.api,
        )

gene_burdens property

Per-gene burden records (populated after run()).

summary property

Cohort summary statistics (populated after run()).

variants property

Aggregated cohort variants (populated after run()).

run()

Execute the full cohort analysis pipeline.

Sequence: 1. Process each sample VCF through the standard pipeline 2. Aggregate variants across samples 3. Compute cohort statistics 4. Generate reports

Returns

list[Path] Paths to generated report files.

Raises

FileNotFoundError If any sample VCF file does not exist.

Source code in vartriage/cohort/pipeline.py
def run(self) -> list[Path]:
    """Execute the full cohort analysis pipeline.

    Sequence:
    1. Process each sample VCF through the standard pipeline
    2. Aggregate variants across samples
    3. Compute cohort statistics
    4. Generate reports

    Returns
    -------
    list[Path]
        Paths to generated report files.

    Raises
    ------
    FileNotFoundError
        If any sample VCF file does not exist.
    """
    logger.info(
        "Starting cohort analysis '%s' with %d samples",
        self._cohort_config.cohort_name,
        self._cohort_config.sample_count,
    )

    # Validate all VCF files exist before processing
    for vcf_path in self._cohort_config.sample_vcfs:
        if not vcf_path.exists():
            raise FileNotFoundError(f"Sample VCF not found: {vcf_path}")

    # TemporaryDirectory context manager guarantees cleanup on
    # success, failure, or keyboard interrupt.
    with tempfile.TemporaryDirectory(prefix="vartriage_cohort_") as tmp_dir:
        tmp_path = Path(tmp_dir)

        if self._cohort_config.parallel:
            self._process_parallel(tmp_path)
        else:
            self._process_sequential(tmp_path)

    # Aggregate
    logger.info(
        "Aggregating variants across %d samples",
        len(self._samples_processed),
    )
    self._variants = self._aggregator.aggregate()

    # Statistics
    stats = CohortStatistics(self._cohort_config, self._variants)
    self._gene_burdens = stats.compute_gene_burden()
    self._summary = stats.compute_summary(self._samples_processed)

    # Report generation
    reporter = CohortReportGenerator(self._cohort_config)
    report_paths = reporter.generate(
        self._variants, self._gene_burdens, self._summary
    )

    logger.info(
        "Cohort analysis complete: %d variants, %d shared, %d genes",
        self._summary.total_variants,
        self._summary.shared_variants,
        self._summary.genes_affected,
    )

    return report_paths

Orchestrates per-sample Pipeline execution, aggregation, statistics computation, and report generation.

from vartriage import CohortPipeline, CohortConfig

config = CohortConfig(
    sample_vcfs=[Path("a.vcf.gz"), Path("b.vcf.gz")],
    output_path=Path("output/"),
)
pipeline = CohortPipeline(cohort_config=config)
report_paths = pipeline.run()

# Access results after run()
pipeline.variants  # list[CohortVariant]
pipeline.gene_burdens  # list[GeneBurden]
pipeline.summary  # CohortSummary

Parameters:

Name Type Description
cohort_config CohortConfig Cohort-level settings
pipeline_config PipelineConfig, optional Base config applied to each sample
annotation_config AnnotationConfig, optional Shared annotation config (overrides pipeline_config)
prioritization_config PrioritizationConfig, optional Shared scoring config (overrides pipeline_config)

Methods:

  • run() -> list[Path] - Execute the full cohort analysis, returns paths to generated reports.

Properties (populated after run):

  • variants: list[CohortVariant]
  • gene_burdens: list[GeneBurden]
  • summary: CohortSummary | None

CohortAggregator

vartriage.cohort.aggregator.CohortAggregator

Merges per-sample classified variants into cohort-level records.

Groups variants by genomic coordinate across all samples, then produces CohortVariant records with cross-sample frequency and merged evidence. Respects the config's min_recurrence threshold and max_af_threshold for filtering.

Parameters

config : CohortConfig Cohort analysis configuration.

Source code in vartriage/cohort/aggregator.py
class CohortAggregator:
    """Merges per-sample classified variants into cohort-level records.

    Groups variants by genomic coordinate across all samples, then
    produces CohortVariant records with cross-sample frequency and
    merged evidence. Respects the config's min_recurrence threshold
    and max_af_threshold for filtering.

    Parameters
    ----------
    config : CohortConfig
        Cohort analysis configuration.
    """

    def __init__(self, config: CohortConfig) -> None:
        self._config = config
        self._variant_map: dict[VariantKey, list[SampleOccurrence]] = defaultdict(list)
        self._samples_added: int = 0

    @property
    def samples_added(self) -> int:
        """Number of samples that have been ingested so far."""
        return self._samples_added

    @property
    def total_distinct_variants(self) -> int:
        """Number of distinct variant coordinates seen across all samples."""
        return len(self._variant_map)

    def add_sample(
        self,
        sample_id: str,
        vcf_path: Path,
        variants: list[ClassifiedVariant],
    ) -> int:
        """Ingest classified variants from a single sample.

        Parameters
        ----------
        sample_id : str
            Human-readable sample identifier.
        vcf_path : Path
            Source VCF path for traceability.
        variants : list[ClassifiedVariant]
            All classified variants from this sample's pipeline run.

        Returns
        -------
        int
            Number of variants added from this sample (after AF filtering).
        """
        added = 0
        max_af = self._config.max_af_threshold

        for classified in variants:
            af = classified.scored.annotated.allele_frequency

            # Skip common variants that exceed cohort AF threshold
            if af is not None and af > max_af:
                continue

            key = self._variant_key(classified)
            occurrence = SampleOccurrence(
                sample_id=sample_id,
                vcf_path=vcf_path,
                classified=classified,
            )
            self._variant_map[key].append(occurrence)
            added += 1

        self._samples_added += 1
        logger.info(
            "Added sample '%s': %d variants (of %d total, %d filtered by AF)",
            sample_id,
            added,
            len(variants),
            len(variants) - added,
        )
        return added

    def aggregate(self) -> list[CohortVariant]:
        """Produce merged cohort variants from all ingested samples.

        Filtering logic:
        - Variants with sample_count >= min_recurrence are included.
        - Singletons (sample_count == 1) are included only when
          include_singletons is True.
        - Variants with 1 < sample_count < min_recurrence are excluded.

        This means min_recurrence acts as a hard inclusion threshold,
        not just a highlight marker. Set min_recurrence=1 with
        include_singletons=True to get all variants regardless of
        recurrence.

        Results are sorted by sample_count descending, then by genomic
        coordinate.

        Returns
        -------
        list[CohortVariant]
            Cohort-level variant records meeting inclusion criteria.
        """
        total_samples = self._config.sample_count
        min_rec = self._config.min_recurrence
        include_singletons = self._config.include_singletons

        results: list[CohortVariant] = []

        for key, occurrences in self._variant_map.items():
            sample_count = len(occurrences)

            # Singletons are excluded unless explicitly included
            if sample_count == 1 and not include_singletons:
                continue

            # Variants below min_recurrence are excluded (unless they
            # are singletons that passed the check above)
            if sample_count < min_rec and sample_count > 1:
                continue

            cohort_variant = self._build_cohort_variant(key, occurrences, total_samples)
            results.append(cohort_variant)

        # Sort: highest recurrence first, then by coordinate for stability
        results.sort(key=lambda v: (-v.sample_count, v.chrom, v.pos, v.ref, v.alt))

        logger.info(
            "Aggregation complete: %d variants (%d shared, %d singletons)",
            len(results),
            sum(1 for v in results if v.sample_count >= 2),
            sum(1 for v in results if v.is_singleton),
        )
        return results

    def get_recurrent_variants(self, min_count: int = 2) -> list[CohortVariant]:
        """Return only variants appearing in >= min_count samples.

        Convenience method for extracting shared variants without
        changing the config's min_recurrence permanently.

        Parameters
        ----------
        min_count : int
            Minimum sample count threshold. Default is 2.

        Returns
        -------
        list[CohortVariant]
            Filtered and sorted cohort variants.
        """
        all_variants = self.aggregate()
        return [v for v in all_variants if v.sample_count >= min_count]

    def reset(self) -> None:
        """Clear all ingested data for reuse with a fresh cohort."""
        self._variant_map.clear()
        self._samples_added = 0

    def _build_cohort_variant(
        self,
        key: VariantKey,
        occurrences: list[SampleOccurrence],
        total_samples: int,
    ) -> CohortVariant:
        """Construct a CohortVariant from merged sample occurrences."""
        chrom, pos, ref, alt = key

        # Collect classifications and consequences across samples
        classifications = [occ.classified.classification for occ in occurrences]
        consequences = [
            occ.classified.scored.annotated.consequence for occ in occurrences
        ]

        # Union of all evidence tags
        all_tags: set[EvidenceTag] = set()
        for occ in occurrences:
            all_tags.update(occ.classified.evidence_tags)

        # Gene name: take the first non-None value
        gene_name: str | None = None
        for occ in occurrences:
            gn = occ.classified.scored.annotated.gene_name
            if gn is not None:
                gene_name = gn
                break

        # Allele frequency: consensus (use first non-None)
        allele_frequency: float | None = None
        for occ in occurrences:
            af = occ.classified.scored.annotated.allele_frequency
            if af is not None:
                allele_frequency = af
                break

        return CohortVariant(
            chrom=chrom,
            pos=pos,
            ref=ref,
            alt=alt,
            gene_name=gene_name,
            consequence=_most_severe_consequence(consequences),
            sample_count=len(occurrences),
            total_samples=total_samples,
            occurrences=tuple(occurrences),
            max_classification=_most_severe_classification(classifications),
            all_evidence_tags=frozenset(all_tags),
            allele_frequency=allele_frequency,
        )

    @staticmethod
    def _variant_key(classified: ClassifiedVariant) -> VariantKey:
        """Extract the canonical merge key from a classified variant."""
        v = classified.scored.annotated.variant
        return (v.chrom, v.pos, v.ref, v.alt)

samples_added property

Number of samples that have been ingested so far.

total_distinct_variants property

Number of distinct variant coordinates seen across all samples.

add_sample(sample_id, vcf_path, variants)

Ingest classified variants from a single sample.

Parameters

sample_id : str Human-readable sample identifier. vcf_path : Path Source VCF path for traceability. variants : list[ClassifiedVariant] All classified variants from this sample's pipeline run.

Returns

int Number of variants added from this sample (after AF filtering).

Source code in vartriage/cohort/aggregator.py
def add_sample(
    self,
    sample_id: str,
    vcf_path: Path,
    variants: list[ClassifiedVariant],
) -> int:
    """Ingest classified variants from a single sample.

    Parameters
    ----------
    sample_id : str
        Human-readable sample identifier.
    vcf_path : Path
        Source VCF path for traceability.
    variants : list[ClassifiedVariant]
        All classified variants from this sample's pipeline run.

    Returns
    -------
    int
        Number of variants added from this sample (after AF filtering).
    """
    added = 0
    max_af = self._config.max_af_threshold

    for classified in variants:
        af = classified.scored.annotated.allele_frequency

        # Skip common variants that exceed cohort AF threshold
        if af is not None and af > max_af:
            continue

        key = self._variant_key(classified)
        occurrence = SampleOccurrence(
            sample_id=sample_id,
            vcf_path=vcf_path,
            classified=classified,
        )
        self._variant_map[key].append(occurrence)
        added += 1

    self._samples_added += 1
    logger.info(
        "Added sample '%s': %d variants (of %d total, %d filtered by AF)",
        sample_id,
        added,
        len(variants),
        len(variants) - added,
    )
    return added

aggregate()

Produce merged cohort variants from all ingested samples.

Filtering logic: - Variants with sample_count >= min_recurrence are included. - Singletons (sample_count == 1) are included only when include_singletons is True. - Variants with 1 < sample_count < min_recurrence are excluded.

This means min_recurrence acts as a hard inclusion threshold, not just a highlight marker. Set min_recurrence=1 with include_singletons=True to get all variants regardless of recurrence.

Results are sorted by sample_count descending, then by genomic coordinate.

Returns

list[CohortVariant] Cohort-level variant records meeting inclusion criteria.

Source code in vartriage/cohort/aggregator.py
def aggregate(self) -> list[CohortVariant]:
    """Produce merged cohort variants from all ingested samples.

    Filtering logic:
    - Variants with sample_count >= min_recurrence are included.
    - Singletons (sample_count == 1) are included only when
      include_singletons is True.
    - Variants with 1 < sample_count < min_recurrence are excluded.

    This means min_recurrence acts as a hard inclusion threshold,
    not just a highlight marker. Set min_recurrence=1 with
    include_singletons=True to get all variants regardless of
    recurrence.

    Results are sorted by sample_count descending, then by genomic
    coordinate.

    Returns
    -------
    list[CohortVariant]
        Cohort-level variant records meeting inclusion criteria.
    """
    total_samples = self._config.sample_count
    min_rec = self._config.min_recurrence
    include_singletons = self._config.include_singletons

    results: list[CohortVariant] = []

    for key, occurrences in self._variant_map.items():
        sample_count = len(occurrences)

        # Singletons are excluded unless explicitly included
        if sample_count == 1 and not include_singletons:
            continue

        # Variants below min_recurrence are excluded (unless they
        # are singletons that passed the check above)
        if sample_count < min_rec and sample_count > 1:
            continue

        cohort_variant = self._build_cohort_variant(key, occurrences, total_samples)
        results.append(cohort_variant)

    # Sort: highest recurrence first, then by coordinate for stability
    results.sort(key=lambda v: (-v.sample_count, v.chrom, v.pos, v.ref, v.alt))

    logger.info(
        "Aggregation complete: %d variants (%d shared, %d singletons)",
        len(results),
        sum(1 for v in results if v.sample_count >= 2),
        sum(1 for v in results if v.is_singleton),
    )
    return results

get_recurrent_variants(min_count=2)

Return only variants appearing in >= min_count samples.

Convenience method for extracting shared variants without changing the config's min_recurrence permanently.

Parameters

min_count : int Minimum sample count threshold. Default is 2.

Returns

list[CohortVariant] Filtered and sorted cohort variants.

Source code in vartriage/cohort/aggregator.py
def get_recurrent_variants(self, min_count: int = 2) -> list[CohortVariant]:
    """Return only variants appearing in >= min_count samples.

    Convenience method for extracting shared variants without
    changing the config's min_recurrence permanently.

    Parameters
    ----------
    min_count : int
        Minimum sample count threshold. Default is 2.

    Returns
    -------
    list[CohortVariant]
        Filtered and sorted cohort variants.
    """
    all_variants = self.aggregate()
    return [v for v in all_variants if v.sample_count >= min_count]

reset()

Clear all ingested data for reuse with a fresh cohort.

Source code in vartriage/cohort/aggregator.py
def reset(self) -> None:
    """Clear all ingested data for reuse with a fresh cohort."""
    self._variant_map.clear()
    self._samples_added = 0

Merges classified variants from multiple samples by genomic coordinate.

Methods:

  • add_sample(sample_id, vcf_path, variants) -> int - Ingest one sample's results. Returns count after AF filtering.
  • aggregate() -> list[CohortVariant] - Produce merged cohort variants respecting config thresholds.
  • get_recurrent_variants(min_count=2) -> list[CohortVariant] - Convenience filter for shared variants.
  • reset() - Clear all data for reuse.

Properties:

  • samples_added: int
  • total_distinct_variants: int

CohortStatistics

vartriage.cohort.statistics.CohortStatistics

Compute summary statistics from aggregated cohort variants.

Takes a list of CohortVariant records (output of CohortAggregator) and produces per-gene burden tables, recurrence distributions, and the overall CohortSummary dataclass.

Parameters

config : CohortConfig Cohort configuration (used for cohort_name and sample metadata). variants : list[CohortVariant] Aggregated cohort variants to analyze.

Source code in vartriage/cohort/statistics.py
class CohortStatistics:
    """Compute summary statistics from aggregated cohort variants.

    Takes a list of CohortVariant records (output of CohortAggregator)
    and produces per-gene burden tables, recurrence distributions, and
    the overall CohortSummary dataclass.

    Parameters
    ----------
    config : CohortConfig
        Cohort configuration (used for cohort_name and sample metadata).
    variants : list[CohortVariant]
        Aggregated cohort variants to analyze.
    """

    def __init__(self, config: CohortConfig, variants: list[CohortVariant]) -> None:
        self._config = config
        self._variants = variants

    @property
    def variant_count(self) -> int:
        """Total distinct variants in the cohort."""
        return len(self._variants)

    def compute_summary(self, samples_processed: list[str]) -> CohortSummary:
        """Produce the top-level cohort summary.

        Parameters
        ----------
        samples_processed : list[str]
            Ordered list of sample identifiers that were analyzed.

        Returns
        -------
        CohortSummary
            Aggregate statistics for the cohort run.
        """
        total = len(self._variants)
        shared = sum(1 for v in self._variants if v.sample_count >= 2)
        singletons = sum(1 for v in self._variants if v.is_singleton)
        universal = sum(1 for v in self._variants if v.is_universal)

        pathogenic = sum(
            1
            for v in self._variants
            if v.max_classification == ACMGClassification.PATHOGENIC
        )
        likely_pathogenic = sum(
            1
            for v in self._variants
            if v.max_classification == ACMGClassification.LIKELY_PATHOGENIC
        )

        genes: set[str] = set()
        for v in self._variants:
            if v.gene_name is not None:
                genes.add(v.gene_name)

        top_genes = self._top_recurrent_genes(limit=10)

        return CohortSummary(
            cohort_name=self._config.cohort_name,
            total_samples=self._config.sample_count,
            total_variants=total,
            shared_variants=shared,
            singleton_variants=singletons,
            universal_variants=universal,
            pathogenic_variants=pathogenic,
            likely_pathogenic_variants=likely_pathogenic,
            genes_affected=len(genes),
            top_recurrent_genes=tuple(top_genes),
            samples_processed=tuple(samples_processed),
        )

    def compute_gene_burden(self) -> list[GeneBurden]:
        """Compute per-gene variant burden across the cohort.

        Groups variants by gene, counts pathogenic/likely_pathogenic
        hits, and tracks how many samples are affected per gene.
        Results are sorted by pathogenic_count descending, then by
        samples_affected descending.

        Returns
        -------
        list[GeneBurden]
            Per-gene burden records, sorted by severity.
        """
        # gene -> list of cohort variants
        gene_variants: dict[str, list[CohortVariant]] = defaultdict(list)

        for v in self._variants:
            if v.gene_name is None:
                continue
            gene_variants[v.gene_name].append(v)

        total_samples = self._config.sample_count
        burdens: list[GeneBurden] = []

        for gene_name, variants in gene_variants.items():
            pathogenic_count = sum(
                1
                for v in variants
                if v.max_classification in _PATHOGENIC_CLASSIFICATIONS
            )

            # Unique samples affected in this gene
            samples_in_gene: set[str] = set()
            for v in variants:
                samples_in_gene.update(v.sample_ids)

            # Most severe classification in this gene
            most_severe = self._gene_most_severe(variants)

            burdens.append(
                GeneBurden(
                    gene_name=gene_name,
                    total_variants=len(variants),
                    pathogenic_count=pathogenic_count,
                    samples_affected=len(samples_in_gene),
                    total_samples=total_samples,
                    most_severe=most_severe,
                )
            )

        burdens.sort(
            key=lambda b: (-b.pathogenic_count, -b.samples_affected, b.gene_name)
        )

        logger.info(
            "Gene burden computed: %d genes, %d with pathogenic variants",
            len(burdens),
            sum(1 for b in burdens if b.pathogenic_count > 0),
        )
        return burdens

    def recurrence_distribution(self) -> dict[int, int]:
        """Count how many variants appear in exactly N samples.

        Returns
        -------
        dict[int, int]
            Mapping of sample_count -> number of variants with that count.
            Sorted by key ascending.
        """
        counter: Counter[int] = Counter()
        for v in self._variants:
            counter[v.sample_count] += 1
        return dict(sorted(counter.items()))

    def per_sample_counts(self) -> dict[str, int]:
        """Count total variants per sample across the cohort.

        Returns
        -------
        dict[str, int]
            Mapping of sample_id -> number of cohort variants that
            include that sample. Sorted by count descending.
        """
        counts: Counter[str] = Counter()
        for v in self._variants:
            for sample_id in v.sample_ids:
                counts[sample_id] += 1
        return dict(counts.most_common())

    def classification_distribution(self) -> dict[str, int]:
        """Count variants by their most severe ACMG classification.

        Returns
        -------
        dict[str, int]
            Mapping of classification name -> count.
        """
        counter: Counter[str] = Counter()
        for v in self._variants:
            counter[v.max_classification.value] += 1
        return dict(counter.most_common())

    def consequence_distribution(self) -> dict[str, int]:
        """Count variants by functional consequence type.

        Returns
        -------
        dict[str, int]
            Mapping of consequence name -> count.
        """
        counter: Counter[str] = Counter()
        for v in self._variants:
            counter[v.consequence.value] += 1
        return dict(counter.most_common())

    def _top_recurrent_genes(self, limit: int = 10) -> list[str]:
        """Get genes with the most recurrent variants.

        Ranks genes by the sum of sample_count across all their
        variants, which captures both how many variants a gene has
        and how widely shared they are.
        """
        gene_recurrence: Counter[str] = Counter()
        for v in self._variants:
            if v.gene_name is not None:
                gene_recurrence[v.gene_name] += v.sample_count
        return [gene for gene, _ in gene_recurrence.most_common(limit)]

    @staticmethod
    def _gene_most_severe(variants: list[CohortVariant]) -> ACMGClassification:
        """Find the most severe classification across gene variants."""
        best_idx = len(CLASSIFICATION_SEVERITY_ORDER) - 1
        for v in variants:
            try:
                idx = CLASSIFICATION_SEVERITY_ORDER.index(v.max_classification)
            except ValueError:
                continue
            if idx < best_idx:
                best_idx = idx
        return CLASSIFICATION_SEVERITY_ORDER[best_idx]

variant_count property

Total distinct variants in the cohort.

classification_distribution()

Count variants by their most severe ACMG classification.

Returns

dict[str, int] Mapping of classification name -> count.

Source code in vartriage/cohort/statistics.py
def classification_distribution(self) -> dict[str, int]:
    """Count variants by their most severe ACMG classification.

    Returns
    -------
    dict[str, int]
        Mapping of classification name -> count.
    """
    counter: Counter[str] = Counter()
    for v in self._variants:
        counter[v.max_classification.value] += 1
    return dict(counter.most_common())

compute_gene_burden()

Compute per-gene variant burden across the cohort.

Groups variants by gene, counts pathogenic/likely_pathogenic hits, and tracks how many samples are affected per gene. Results are sorted by pathogenic_count descending, then by samples_affected descending.

Returns

list[GeneBurden] Per-gene burden records, sorted by severity.

Source code in vartriage/cohort/statistics.py
def compute_gene_burden(self) -> list[GeneBurden]:
    """Compute per-gene variant burden across the cohort.

    Groups variants by gene, counts pathogenic/likely_pathogenic
    hits, and tracks how many samples are affected per gene.
    Results are sorted by pathogenic_count descending, then by
    samples_affected descending.

    Returns
    -------
    list[GeneBurden]
        Per-gene burden records, sorted by severity.
    """
    # gene -> list of cohort variants
    gene_variants: dict[str, list[CohortVariant]] = defaultdict(list)

    for v in self._variants:
        if v.gene_name is None:
            continue
        gene_variants[v.gene_name].append(v)

    total_samples = self._config.sample_count
    burdens: list[GeneBurden] = []

    for gene_name, variants in gene_variants.items():
        pathogenic_count = sum(
            1
            for v in variants
            if v.max_classification in _PATHOGENIC_CLASSIFICATIONS
        )

        # Unique samples affected in this gene
        samples_in_gene: set[str] = set()
        for v in variants:
            samples_in_gene.update(v.sample_ids)

        # Most severe classification in this gene
        most_severe = self._gene_most_severe(variants)

        burdens.append(
            GeneBurden(
                gene_name=gene_name,
                total_variants=len(variants),
                pathogenic_count=pathogenic_count,
                samples_affected=len(samples_in_gene),
                total_samples=total_samples,
                most_severe=most_severe,
            )
        )

    burdens.sort(
        key=lambda b: (-b.pathogenic_count, -b.samples_affected, b.gene_name)
    )

    logger.info(
        "Gene burden computed: %d genes, %d with pathogenic variants",
        len(burdens),
        sum(1 for b in burdens if b.pathogenic_count > 0),
    )
    return burdens

compute_summary(samples_processed)

Produce the top-level cohort summary.

Parameters

samples_processed : list[str] Ordered list of sample identifiers that were analyzed.

Returns

CohortSummary Aggregate statistics for the cohort run.

Source code in vartriage/cohort/statistics.py
def compute_summary(self, samples_processed: list[str]) -> CohortSummary:
    """Produce the top-level cohort summary.

    Parameters
    ----------
    samples_processed : list[str]
        Ordered list of sample identifiers that were analyzed.

    Returns
    -------
    CohortSummary
        Aggregate statistics for the cohort run.
    """
    total = len(self._variants)
    shared = sum(1 for v in self._variants if v.sample_count >= 2)
    singletons = sum(1 for v in self._variants if v.is_singleton)
    universal = sum(1 for v in self._variants if v.is_universal)

    pathogenic = sum(
        1
        for v in self._variants
        if v.max_classification == ACMGClassification.PATHOGENIC
    )
    likely_pathogenic = sum(
        1
        for v in self._variants
        if v.max_classification == ACMGClassification.LIKELY_PATHOGENIC
    )

    genes: set[str] = set()
    for v in self._variants:
        if v.gene_name is not None:
            genes.add(v.gene_name)

    top_genes = self._top_recurrent_genes(limit=10)

    return CohortSummary(
        cohort_name=self._config.cohort_name,
        total_samples=self._config.sample_count,
        total_variants=total,
        shared_variants=shared,
        singleton_variants=singletons,
        universal_variants=universal,
        pathogenic_variants=pathogenic,
        likely_pathogenic_variants=likely_pathogenic,
        genes_affected=len(genes),
        top_recurrent_genes=tuple(top_genes),
        samples_processed=tuple(samples_processed),
    )

consequence_distribution()

Count variants by functional consequence type.

Returns

dict[str, int] Mapping of consequence name -> count.

Source code in vartriage/cohort/statistics.py
def consequence_distribution(self) -> dict[str, int]:
    """Count variants by functional consequence type.

    Returns
    -------
    dict[str, int]
        Mapping of consequence name -> count.
    """
    counter: Counter[str] = Counter()
    for v in self._variants:
        counter[v.consequence.value] += 1
    return dict(counter.most_common())

per_sample_counts()

Count total variants per sample across the cohort.

Returns

dict[str, int] Mapping of sample_id -> number of cohort variants that include that sample. Sorted by count descending.

Source code in vartriage/cohort/statistics.py
def per_sample_counts(self) -> dict[str, int]:
    """Count total variants per sample across the cohort.

    Returns
    -------
    dict[str, int]
        Mapping of sample_id -> number of cohort variants that
        include that sample. Sorted by count descending.
    """
    counts: Counter[str] = Counter()
    for v in self._variants:
        for sample_id in v.sample_ids:
            counts[sample_id] += 1
    return dict(counts.most_common())

recurrence_distribution()

Count how many variants appear in exactly N samples.

Returns

dict[int, int] Mapping of sample_count -> number of variants with that count. Sorted by key ascending.

Source code in vartriage/cohort/statistics.py
def recurrence_distribution(self) -> dict[int, int]:
    """Count how many variants appear in exactly N samples.

    Returns
    -------
    dict[int, int]
        Mapping of sample_count -> number of variants with that count.
        Sorted by key ascending.
    """
    counter: Counter[int] = Counter()
    for v in self._variants:
        counter[v.sample_count] += 1
    return dict(sorted(counter.items()))

Computes summary metrics from aggregated cohort variants.

Methods:

  • compute_summary(samples_processed) -> CohortSummary - Top-level metrics.
  • compute_gene_burden() -> list[GeneBurden] - Per-gene mutation burden sorted by severity.
  • recurrence_distribution() -> dict[int, int] - Sample count histogram.
  • per_sample_counts() -> dict[str, int] - Variants per sample.
  • classification_distribution() -> dict[str, int] - Count by ACMG class.
  • consequence_distribution() -> dict[str, int] - Count by consequence type.

CohortReportGenerator

vartriage.cohort.report.CohortReportGenerator

Generate cohort analysis reports in JSON or CSV format.

Writes three output files per cohort run: - variants report: all cohort variants with recurrence data - gene burden report: per-gene statistics - summary report: top-level cohort metrics

Parameters

config : CohortConfig Cohort configuration with output_path and output_format.

Source code in vartriage/cohort/report.py
class CohortReportGenerator:
    """Generate cohort analysis reports in JSON or CSV format.

    Writes three output files per cohort run:
    - variants report: all cohort variants with recurrence data
    - gene burden report: per-gene statistics
    - summary report: top-level cohort metrics

    Parameters
    ----------
    config : CohortConfig
        Cohort configuration with output_path and output_format.
    """

    def __init__(self, config: CohortConfig) -> None:
        self._config = config
        self._output_dir = config.output_path

    def generate(
        self,
        variants: list[CohortVariant],
        gene_burdens: list[GeneBurden],
        summary: CohortSummary,
    ) -> list[Path]:
        """Write all cohort report files.

        Parameters
        ----------
        variants : list[CohortVariant]
            Aggregated cohort variants.
        gene_burdens : list[GeneBurden]
            Per-gene burden statistics.
        summary : CohortSummary
            Top-level cohort metrics.

        Returns
        -------
        list[Path]
            Paths to all generated report files.
        """
        self._output_dir.mkdir(parents=True, exist_ok=True)

        fmt = self._config.output_format
        if fmt == "csv":
            return self._write_csv(variants, gene_burdens, summary)
        return self._write_json(variants, gene_burdens, summary)

    def _write_json(
        self,
        variants: list[CohortVariant],
        gene_burdens: list[GeneBurden],
        summary: CohortSummary,
    ) -> list[Path]:
        """Write JSON format reports."""
        paths: list[Path] = []

        # Variants report
        variants_path = self._output_dir / f"{self._config.cohort_name}_variants.json"
        variant_records = [self._serialize_variant(v) for v in variants]
        self._write_json_file(variants_path, variant_records)
        paths.append(variants_path)

        # Gene burden report
        burden_path = self._output_dir / f"{self._config.cohort_name}_gene_burden.json"
        burden_records = [self._serialize_burden(b) for b in gene_burdens]
        self._write_json_file(burden_path, burden_records)
        paths.append(burden_path)

        # Summary report
        summary_path = self._output_dir / f"{self._config.cohort_name}_summary.json"
        summary_record = self._serialize_summary(summary)
        self._write_json_file(summary_path, summary_record)
        paths.append(summary_path)

        logger.info("JSON reports written to %s", self._output_dir)
        return paths

    def _write_csv(
        self,
        variants: list[CohortVariant],
        gene_burdens: list[GeneBurden],
        summary: CohortSummary,
    ) -> list[Path]:
        """Write CSV format reports."""
        paths: list[Path] = []

        # Variants CSV
        variants_path = self._output_dir / f"{self._config.cohort_name}_variants.csv"
        self._write_variants_csv(variants_path, variants)
        paths.append(variants_path)

        # Gene burden CSV
        burden_path = self._output_dir / f"{self._config.cohort_name}_gene_burden.csv"
        self._write_burden_csv(burden_path, gene_burdens)
        paths.append(burden_path)

        # Summary JSON (always JSON for structured metadata)
        summary_path = self._output_dir / f"{self._config.cohort_name}_summary.json"
        summary_record = self._serialize_summary(summary)
        self._write_json_file(summary_path, summary_record)
        paths.append(summary_path)

        logger.info("CSV reports written to %s", self._output_dir)
        return paths

    def _write_variants_csv(self, path: Path, variants: list[CohortVariant]) -> None:
        """Write cohort variants to CSV."""
        fieldnames = [
            "chrom",
            "pos",
            "ref",
            "alt",
            "gene_name",
            "consequence",
            "sample_count",
            "total_samples",
            "cohort_frequency",
            "allele_frequency",
            "max_classification",
            "evidence_tags",
            "samples",
        ]

        resolved = resolve_path(path)
        resolved.parent.mkdir(parents=True, exist_ok=True)
        with open(resolved, "w", newline="", encoding="utf-8") as f:
            writer = csv.DictWriter(f, fieldnames=fieldnames)
            writer.writeheader()
            for v in variants:
                writer.writerow(
                    {
                        "chrom": v.chrom,
                        "pos": v.pos,
                        "ref": v.ref,
                        "alt": v.alt,
                        "gene_name": v.gene_name or "",
                        "consequence": v.consequence.value,
                        "sample_count": v.sample_count,
                        "total_samples": v.total_samples,
                        "cohort_frequency": f"{v.cohort_frequency:.4f}",
                        "allele_frequency": (
                            f"{v.allele_frequency:.6f}"
                            if v.allele_frequency is not None
                            else ""
                        ),
                        "max_classification": v.max_classification.value,
                        "evidence_tags": ";".join(
                            sorted(t.value for t in v.all_evidence_tags)
                        ),
                        "samples": ";".join(v.sample_ids),
                    }
                )

    def _write_burden_csv(self, path: Path, burdens: list[GeneBurden]) -> None:
        """Write gene burden table to CSV."""
        fieldnames = [
            "gene_name",
            "total_variants",
            "pathogenic_count",
            "samples_affected",
            "total_samples",
            "penetrance",
            "most_severe",
        ]

        resolved = resolve_path(path)
        resolved.parent.mkdir(parents=True, exist_ok=True)
        with open(resolved, "w", newline="", encoding="utf-8") as f:
            writer = csv.DictWriter(f, fieldnames=fieldnames)
            writer.writeheader()
            for b in burdens:
                writer.writerow(
                    {
                        "gene_name": b.gene_name,
                        "total_variants": b.total_variants,
                        "pathogenic_count": b.pathogenic_count,
                        "samples_affected": b.samples_affected,
                        "total_samples": b.total_samples,
                        "penetrance": f"{b.penetrance:.4f}",
                        "most_severe": b.most_severe.value,
                    }
                )

    @staticmethod
    def _write_json_file(path: Path, data: Any) -> None:
        """Write data to a JSON file with consistent formatting."""
        resolved = resolve_path(path)
        resolved.parent.mkdir(parents=True, exist_ok=True)
        with open(resolved, "w", encoding="utf-8") as f:
            json.dump(data, f, indent=2, ensure_ascii=False)
            f.write("\n")

    @staticmethod
    def _serialize_variant(v: CohortVariant) -> dict[str, Any]:
        """Convert a CohortVariant to a JSON-serializable dict."""
        return {
            "chrom": v.chrom,
            "pos": v.pos,
            "ref": v.ref,
            "alt": v.alt,
            "gene_name": v.gene_name,
            "consequence": v.consequence.value,
            "sample_count": v.sample_count,
            "total_samples": v.total_samples,
            "cohort_frequency": round(v.cohort_frequency, 4),
            "is_singleton": v.is_singleton,
            "is_universal": v.is_universal,
            "allele_frequency": v.allele_frequency,
            "max_classification": v.max_classification.value,
            "evidence_tags": sorted(t.value for t in v.all_evidence_tags),
            "samples": [
                {
                    "sample_id": occ.sample_id,
                    "classification": occ.classified.classification.value,
                    "evidence_tags": sorted(
                        t.value for t in occ.classified.evidence_tags
                    ),
                }
                for occ in v.occurrences
            ],
        }

    @staticmethod
    def _serialize_burden(b: GeneBurden) -> dict[str, Any]:
        """Convert a GeneBurden to a JSON-serializable dict."""
        return {
            "gene_name": b.gene_name,
            "total_variants": b.total_variants,
            "pathogenic_count": b.pathogenic_count,
            "samples_affected": b.samples_affected,
            "total_samples": b.total_samples,
            "penetrance": round(b.penetrance, 4),
            "most_severe": b.most_severe.value,
        }

    @staticmethod
    def _serialize_summary(s: CohortSummary) -> dict[str, Any]:
        """Convert a CohortSummary to a JSON-serializable dict."""
        return {
            "cohort_name": s.cohort_name,
            "total_samples": s.total_samples,
            "total_variants": s.total_variants,
            "shared_variants": s.shared_variants,
            "singleton_variants": s.singleton_variants,
            "universal_variants": s.universal_variants,
            "pathogenic_variants": s.pathogenic_variants,
            "likely_pathogenic_variants": s.likely_pathogenic_variants,
            "genes_affected": s.genes_affected,
            "top_recurrent_genes": list(s.top_recurrent_genes),
            "samples_processed": list(s.samples_processed),
        }

generate(variants, gene_burdens, summary)

Write all cohort report files.

Parameters

variants : list[CohortVariant] Aggregated cohort variants. gene_burdens : list[GeneBurden] Per-gene burden statistics. summary : CohortSummary Top-level cohort metrics.

Returns

list[Path] Paths to all generated report files.

Source code in vartriage/cohort/report.py
def generate(
    self,
    variants: list[CohortVariant],
    gene_burdens: list[GeneBurden],
    summary: CohortSummary,
) -> list[Path]:
    """Write all cohort report files.

    Parameters
    ----------
    variants : list[CohortVariant]
        Aggregated cohort variants.
    gene_burdens : list[GeneBurden]
        Per-gene burden statistics.
    summary : CohortSummary
        Top-level cohort metrics.

    Returns
    -------
    list[Path]
        Paths to all generated report files.
    """
    self._output_dir.mkdir(parents=True, exist_ok=True)

    fmt = self._config.output_format
    if fmt == "csv":
        return self._write_csv(variants, gene_burdens, summary)
    return self._write_json(variants, gene_burdens, summary)

Writes cohort analysis results to disk in JSON or CSV format.

Methods:

  • generate(variants, gene_burdens, summary) -> list[Path] - Write all report files, returns paths.

Data Models

CohortConfig

Frozen dataclass. Configuration for multi-sample cohort analysis. See Cohort Analysis Guide for the full field reference.

CohortVariant

Frozen dataclass. A variant aggregated across multiple samples.

Key fields: chrom, pos, ref, alt, gene_name, consequence, sample_count, total_samples, occurrences, max_classification, all_evidence_tags, allele_frequency.

Properties: key, cohort_frequency, is_singleton, is_universal, sample_ids.

SampleOccurrence

Frozen dataclass. Record of a variant's appearance in one sample.

Fields: sample_id, vcf_path, classified (ClassifiedVariant).

GeneBurden

Frozen dataclass. Per-gene variant burden across the cohort.

Fields: gene_name, total_variants, pathogenic_count, samples_affected, total_samples, most_severe.

Properties: penetrance (float, fraction of cohort affected).

CohortSummary

Frozen dataclass. Aggregate statistics for a completed cohort run.

Fields: cohort_name, total_samples, total_variants, shared_variants, singleton_variants, universal_variants, pathogenic_variants, likely_pathogenic_variants, genes_affected, top_recurrent_genes, samples_processed.