File size: 2,430 Bytes
0fff343
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
"""Smoke test for the H1 verification fixture (synthetic; no real-data dependency)."""

from __future__ import annotations

import numpy as np
import pandas as pd

from dsl import Cohort
from validate.h1 import IMMUNE_GENES, MMR_GENES, run_h1


def _synthetic_named_cohort(n_per_group: int = 60, seed: int = 0) -> Cohort:
    """A synthetic cohort where MSI-H tumours have low MMR and high immune expression."""
    rng = np.random.default_rng(seed)
    other_genes = [f"GENE_{i}" for i in range(20)]
    feature_cols = MMR_GENES + IMMUNE_GENES + other_genes

    rows = []
    sample_ids = []
    msi = []
    stages = []
    ages = []
    for i in range(n_per_group):
        row = rng.normal(loc=5.0, scale=0.5, size=len(feature_cols))
        for g in MMR_GENES:
            row[feature_cols.index(g)] = rng.normal(2.0, 0.3)
        for g in IMMUNE_GENES:
            row[feature_cols.index(g)] = rng.normal(8.0, 0.3)
        rows.append(row)
        sample_ids.append(f"H{i}")
        msi.append("MSI-H")
        stages.append(rng.choice(["I", "II", "III", "IV"]))
        ages.append(float(rng.integers(40, 85)))
    for i in range(n_per_group):
        row = rng.normal(loc=5.0, scale=0.5, size=len(feature_cols))
        for g in MMR_GENES:
            row[feature_cols.index(g)] = rng.normal(8.0, 0.3)
        for g in IMMUNE_GENES:
            row[feature_cols.index(g)] = rng.normal(2.0, 0.3)
        rows.append(row)
        sample_ids.append(f"S{i}")
        msi.append("MSS")
        stages.append(rng.choice(["I", "II", "III", "IV"]))
        ages.append(float(rng.integers(40, 85)))

    idx = pd.Index(sample_ids, name="sample_id")
    expr = pd.DataFrame(rows, index=idx, columns=feature_cols)
    clinical = pd.DataFrame({"stage": stages, "age": ages}, index=idx)
    labels = pd.DataFrame({"msi_status": msi, "tmb": [np.nan] * len(idx)}, index=idx)
    return Cohort(expression=expr, clinical=clinical, labels=labels)


def test_h1_recovers_known_direction_on_synthetic_data():
    cohort = _synthetic_named_cohort()
    res = run_h1(cohort)
    assert res.mmr_separates_correct_direction      # MSI-H low on MMR
    assert res.immune_separates_correct_direction   # MSI-H high on immune
    # With a clean synthetic signal Fit should crush it.
    assert res.fit.auroc > 0.95
    assert res.fit.balanced_acc > 0.9
    # n_used is the full synthetic cohort (no missing stage/age).
    assert res.n_used == 120