File size: 2,585 Bytes
0834211
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
64
65
66
67
68
69
70
71
72
---
license: cc-by-4.0
tags:
  - single-cell
  - scRNA-seq
  - batch-correction
  - scvi
library_name: pytorch
---

# BA-scVI trained on scMARK v2

A Batch-Adversarial scVI model trained on the [scMARK v2](https://zenodo.org/records/7795653)
benchmark (11 cancer scRNA-seq studies, 109,543 cells, 23 standardised cell types).

## Result

Scored with the K-Neighbors Intersection (KNI) metric from
[Diaz-Mejia et al., LMRL @ ICLR 2025](https://arxiv.org/abs/2503.20730)
(k=50, tau=0.8k, batch = `study_name`, label = `standard_true_celltype`),
over all 109,543 cells:

| model | KNI | cross-study acc | batch diversity |
|---|---|---|---|
| **this model** | **0.7153** | 0.7182 | 0.995 |
| BA-scVI (published) | 0.7110 | 0.712 | 0.999 |
| PCA baseline | 0.470 | - | - |

The published 0.7110 was independently reproduced from the authors' released
embedding using both their `calc_kni_score` and our implementation (0.7110 /
0.7114), so the two numbers are directly comparable.

## Architecture

scVI-family VAE with ZINB likelihood and an adversarial discriminator on the
latent. Batch identity is injected into the **decoder only** (the encoder never
sees it), matching the configuration used for the published BA-scVI results.

- 10-d latent, 512 hidden units, 4 layers, dropout 0.1
- ZINB reconstruction, masked by each study's measured gene panel
- discriminator weight 100, confusion objective, normalised per batch level
- batch levels: 1 modality / 11 studies / 354 samples
- input: 35,804 genes (see `gene_list.txt`; order matters)

## Usage

```python
import torch
ck = torch.load("bascvi_scmark_epoch63.ckpt", map_location="cpu", weights_only=False)
gene_list = ck["hyper_parameters"]["gene_list"]   # 35,804 HGNC symbols, input order
```

Counts must be raw integers in `gene_list` order; the model applies
`log(1 + 1e4 * x / x.sum())` internally. Embed via
`forward(batch, encode=True, predict_mode=True)["qz_m"]`, which zeroes the batch
vector — the latent is what you score.

## Training data

Trained **only** on scMARK v2 (CC-BY 4.0). No other expression data was used.
The 35,804-gene input vocabulary is a superset reference gene list; genes a
study did not measure are masked out of the reconstruction loss as structural
zeros rather than treated as biological zeros.

## Caveats

- KNI is computed over all cells, including those the model trained on — this is
  the published benchmark's own methodology, and applies equally to the 0.7110
  reference.
- KNI rewards low-dimensional embeddings through its batch-diversity gate; all
  numbers above are at 10-d.