AlonSaguy commited on
Commit
f6ed18d
·
verified ·
1 Parent(s): 57f506d

Publish 2026_W39_unit

Browse files
README.md CHANGED
@@ -24,12 +24,11 @@ The model has three stages:
24
  2. a **25-component Gaussian mixture** over the standardised latents, whose
25
  components read as putative cell types. Component shapes are global; the **mixture weights
26
  depend on the molecular context** (MERFISH + AGEA PCA volumes) at the unit's position;
27
- 3. a distance-weighted **k=20 nearest-neighbour projection** onto training units maps latents
28
- back to 10 interpretable waveform features: `depolarisation_slope`, `recovery_slope`, `repolarisation_slope`, `spatial_spread_um`, `tip_val`, `spike_width_secs`, `predepolarisation_width_secs`, `spike_amplitude`, `peak_to_trough_ratio_log`, `polarity`.
29
 
30
  > **The input is a position.** `predict` takes `x, y, z` (IBL/Allen frame, metres) and returns the
31
- > expected phenotype there: the local component weights times each component's expected features.
32
- > `mixture_weights` returns the local putative cell-type composition itself.
33
 
34
  ## Quickstart
35
 
@@ -41,6 +40,7 @@ model = load_pretrained("int-brain-lab/ea-encoder-unit", revision="2026_W39")
41
  df = pd.read_parquet("example/positions_sample.parquet") # positions (x, y, z)
42
  out = model.predict(df) # one pred_<feature> column each
43
  weights = model.mixture_weights(df) # putative cell-type composition
 
44
  model.selftest() # reproduces the shipped outputs
45
  ```
46
 
@@ -52,10 +52,9 @@ as in `config.json`), then assigned to a putative cell type with `model.assign(z
52
  - `autoencoder.pt`, `shared_latent_scaler.joblib` -- the phenotype encoder and its latent scaler.
53
  - `global_gmm.joblib` -- the mixture's component means and covariances.
54
  - `context_transform.joblib`, `context_weight_model_bundle.pt` -- context -> mixture weights.
55
- - `knn_bank.npz` -- standardised latents and phenotype features of the training units (the kNN
56
- exemplars; part of the model, like the stored points of any nearest-neighbour model).
57
- - `component_feature_expectations.npz` -- each component's expected features, which `predict`
58
- mixes with the local weights.
59
  - `agea_vol_pca.npy`, `merfish_vol_pca.npy` -- the molecular-context volumes, the same frozen
60
  volumes as the channel-level model `int-brain-lab/ea-encoder-channel` of this vintage.
61
  - `split.json`, `config.json`, `preprocessing/unit_stats.npz`, `results/` -- the probe split, the
 
24
  2. a **25-component Gaussian mixture** over the standardised latents, whose
25
  components read as putative cell types. Component shapes are global; the **mixture weights
26
  depend on the molecular context** (MERFISH + AGEA PCA volumes) at the unit's position;
27
+ 3. a **context-local exemplar readout** maps the model back to 10 interpretable waveform features: `depolarisation_slope`, `recovery_slope`, `repolarisation_slope`, `spatial_spread_um`, `tip_val`, `spike_width_secs`, `predepolarisation_width_secs`, `spike_amplitude`, `peak_to_trough_ratio_log`, `polarity`. At a position, each component is represented by its training members whose molecular context resembles the position's (the 4096 nearest in a 3-d context key, shrunk toward all of the component's members -- the more so the farther the context lies from the training data, and entirely where there is no context). So the context sets the mixture weights **and** selects which training exemplars represent each component; the latent density depends on it only through the weights.
 
28
 
29
  > **The input is a position.** `predict` takes `x, y, z` (IBL/Allen frame, metres) and returns the
30
+ > expected phenotype there: the local component weights times each component's expected features
31
+ > there. `mixture_weights` returns the local putative cell-type composition itself.
32
 
33
  ## Quickstart
34
 
 
40
  df = pd.read_parquet("example/positions_sample.parquet") # positions (x, y, z)
41
  out = model.predict(df) # one pred_<feature> column each
42
  weights = model.mixture_weights(df) # putative cell-type composition
43
+ draws = model.sample(df, n_samples=16) # training units drawn from the predictive distribution
44
  model.selftest() # reproduces the shipped outputs
45
  ```
46
 
 
52
  - `autoencoder.pt`, `shared_latent_scaler.joblib` -- the phenotype encoder and its latent scaler.
53
  - `global_gmm.joblib` -- the mixture's component means and covariances.
54
  - `context_transform.joblib`, `context_weight_model_bundle.pt` -- context -> mixture weights.
55
+ - `knn_bank.npz` -- standardised latents, phenotype features, GMM components and context keys of the training units (the exemplars; part of the model, like the stored points of any nearest-neighbour model).
56
+ - `component_feature_expectations.npz` -- each component's expected features over the whole
57
+ brain.
 
58
  - `agea_vol_pca.npy`, `merfish_vol_pca.npy` -- the molecular-context volumes, the same frozen
59
  volumes as the channel-level model `int-brain-lab/ea-encoder-channel` of this vintage.
60
  - `split.json`, `config.json`, `preprocessing/unit_stats.npz`, `results/` -- the probe split, the
autoencoder.pt CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:df36b2645806cfa5b87bbcdcf40375c597d13917096d4e8f1c68268de54e2f68
3
- size 22725861
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2346c6d07e78136054f78f08f075108d7a42c2e1beb42b91c6337464b38b686b
3
+ size 22726117
checksums.json CHANGED
@@ -8,13 +8,13 @@
8
  },
9
  {
10
  "path": "autoencoder.pt",
11
- "hash": "3b24b551896a9e3457f2688838b44dde0166cf10",
12
- "bytes": 22725861
13
  },
14
  {
15
  "path": "component_feature_expectations.npz",
16
- "hash": "fd157f4216b410d88c65e4dc488b5b68b6fb4da0",
17
- "bytes": 1983
18
  },
19
  {
20
  "path": "conditioning_summary.json",
@@ -23,8 +23,8 @@
23
  },
24
  {
25
  "path": "config.json",
26
- "hash": "5b6b1ca80f31d9a047ba365546f4c3d9e3190868",
27
- "bytes": 2910
28
  },
29
  {
30
  "path": "context_transform.joblib",
@@ -38,8 +38,8 @@
38
  },
39
  {
40
  "path": "ephysatlas_model.json",
41
- "hash": "ec11272198df580311b04f780d881cbfdbae690d",
42
- "bytes": 2779
43
  },
44
  {
45
  "path": "example/expected_latents.npy",
@@ -48,8 +48,8 @@
48
  },
49
  {
50
  "path": "example/expected_predictions.parquet",
51
- "hash": "dcb3aa478eba00d80f881c8c80697195ada4e46a",
52
- "bytes": 32885
53
  },
54
  {
55
  "path": "example/positions_sample.parquet",
@@ -73,8 +73,8 @@
73
  },
74
  {
75
  "path": "knn_bank.npz",
76
- "hash": "97f669f75d2e8c7e775933d798585abd2501dd5f",
77
- "bytes": 14794605
78
  },
79
  {
80
  "path": "merfish_vol_pca.npy",
@@ -86,10 +86,15 @@
86
  "hash": "b2202de2690f78ac77893569b44fcf96640390e3",
87
  "bytes": 4683
88
  },
 
 
 
 
 
89
  {
90
  "path": "results/summary.json",
91
- "hash": "ec12ba0eb657743eecea98569ba977c0e4fa4664",
92
- "bytes": 2210
93
  },
94
  {
95
  "path": "shared_latent_scaler.joblib",
 
8
  },
9
  {
10
  "path": "autoencoder.pt",
11
+ "hash": "184f905a06df05208ea6ef7ed15a0ce96e315189",
12
+ "bytes": 22726117
13
  },
14
  {
15
  "path": "component_feature_expectations.npz",
16
+ "hash": "e17db88bed240b4b7bf6907481cb43b9e21de24e",
17
+ "bytes": 1974
18
  },
19
  {
20
  "path": "conditioning_summary.json",
 
23
  },
24
  {
25
  "path": "config.json",
26
+ "hash": "6edf96900240ca5a309974290b526a98bc79fdbd",
27
+ "bytes": 3134
28
  },
29
  {
30
  "path": "context_transform.joblib",
 
38
  },
39
  {
40
  "path": "ephysatlas_model.json",
41
+ "hash": "3a4df0fb1d0ed75e50a4be2357de11220c9b0e30",
42
+ "bytes": 3411
43
  },
44
  {
45
  "path": "example/expected_latents.npy",
 
48
  },
49
  {
50
  "path": "example/expected_predictions.parquet",
51
+ "hash": "01809e9956fc0f26fec921398e696ab87a9690e4",
52
+ "bytes": 32897
53
  },
54
  {
55
  "path": "example/positions_sample.parquet",
 
73
  },
74
  {
75
  "path": "knn_bank.npz",
76
+ "hash": "31801285d6c244e8764e3a467253f4fb452efcef",
77
+ "bytes": 14966230
78
  },
79
  {
80
  "path": "merfish_vol_pca.npy",
 
86
  "hash": "b2202de2690f78ac77893569b44fcf96640390e3",
87
  "bytes": 4683
88
  },
89
+ {
90
+ "path": "readout_summary.json",
91
+ "hash": "c3a1c76392649d00fab09c8c7f4463569dcd271b",
92
+ "bytes": 1209
93
+ },
94
  {
95
  "path": "results/summary.json",
96
+ "hash": "f8a976e54314e3b152c86bce68815b10bc797707",
97
+ "bytes": 2793
98
  },
99
  {
100
  "path": "shared_latent_scaler.joblib",
component_feature_expectations.npz CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:3b38760939c6be13d8ce001d0df1665a39413602ec846b50a27158016db814e9
3
- size 1983
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a52e08355cd58ce06fe7a02a82f3966769253668006f5b5dda9a2f0729122dda
3
+ size 1974
config.json CHANGED
@@ -55,6 +55,17 @@
55
  "rare_component_power": 0.5,
56
  "rare_component_weight_cap": 10.0,
57
  "knn_decoder_k": 20,
 
 
 
 
 
 
 
 
 
 
 
58
  "mirror_x_to_single_hemisphere": true,
59
  "mirror_x_sign": -1.0,
60
  "region_gaussian_variance_floor": 0.001,
 
55
  "rare_component_power": 0.5,
56
  "rare_component_weight_cap": 10.0,
57
  "knn_decoder_k": 20,
58
+ "readout_key_dim": 3,
59
+ "readout_key_alphas": [
60
+ 100.0,
61
+ 1000.0,
62
+ 10000.0,
63
+ 30000.0,
64
+ 100000.0
65
+ ],
66
+ "readout_neighbours": 4096,
67
+ "readout_shrinkage": 10.0,
68
+ "readout_off_data_quantile": 0.99,
69
  "mirror_x_to_single_hemisphere": true,
70
  "mirror_x_sign": -1.0,
71
  "region_gaussian_variance_floor": 0.001,
ephysatlas_model.json CHANGED
@@ -92,7 +92,7 @@
92
  "context": {
93
  "n_cell_pcs": 50,
94
  "n_gene_pcs": 50,
95
- "conditions": "mixture weights only; component means/covariances are global",
96
  "source_model": "int-brain-lab/ea-encoder-channel"
97
  },
98
  "preprocessing": {
@@ -100,6 +100,16 @@
100
  "stats_file": "preprocessing/unit_stats.npz",
101
  "mirror_x_to_single_hemisphere": true,
102
  "mirror_x_sign": -1.0
 
 
 
 
 
 
 
 
 
 
103
  }
104
  },
105
  "data_source": {
 
92
  "context": {
93
  "n_cell_pcs": 50,
94
  "n_gene_pcs": 50,
95
+ "conditions": "mixture weights, and which training exemplars represent each component in the phenotype readout; component means/covariances are global",
96
  "source_model": "int-brain-lab/ea-encoder-channel"
97
  },
98
  "preprocessing": {
 
100
  "stats_file": "preprocessing/unit_stats.npz",
101
  "mirror_x_to_single_hemisphere": true,
102
  "mirror_x_sign": -1.0
103
+ },
104
+ "readout": {
105
+ "method": "context_local_members",
106
+ "description": "the molecular context sets the mixture weights gamma_k(x) and selects which training exemplars represent each component: its members nearest to x in a context key, shrunk toward all of its members -- further as the key leaves the training data, and entirely without context",
107
+ "key_dim": 3,
108
+ "key_ridge_alpha": 30000.0,
109
+ "neighbours": 4096,
110
+ "shrinkage": 10.0,
111
+ "off_data_quantile": 0.99,
112
+ "absent_context": "global member means"
113
  }
114
  },
115
  "data_source": {
example/expected_predictions.parquet CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:51da16f91e81be1d49c2a879c8ea96a1ee5e1c6ac48134546da78b0ec66e6258
3
- size 32885
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1956eff735fde1ad163940ffdce8b3134a17bee774ebf7b5d2b9ef5ec455f428
3
+ size 32897
knn_bank.npz CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:bfa52dac944ae2233a65c7e1d71ea79e0c61f412bb55cfcbc4ce2a3e6e1896ba
3
- size 14794605
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a8e49ee3ecf897fc1296d1d297fff06af8746a9fe3ef1d0ff9c0847a6b1986d2
3
+ size 14966230
readout_summary.json ADDED
@@ -0,0 +1,64 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "method": "context_local_members",
3
+ "key_features": [
4
+ "depolarisation_slope",
5
+ "recovery_slope",
6
+ "repolarisation_slope",
7
+ "spatial_spread_um",
8
+ "tip_val",
9
+ "spike_width_secs",
10
+ "predepolarisation_width_secs",
11
+ "spike_amplitude",
12
+ "peak_to_trough_ratio_log"
13
+ ],
14
+ "key_alpha": 30000.0,
15
+ "key_alpha_validation_mse": {
16
+ "100.0": 0.8998134946502477,
17
+ "1000.0": 0.8993491239299802,
18
+ "10000.0": 0.8983127253407119,
19
+ "30000.0": 0.8982345222443193,
20
+ "100000.0": 0.9005961719454074
21
+ },
22
+ "key_dim": 3,
23
+ "key_singular_values": [
24
+ 80.27060833307382,
25
+ 48.52855142725006,
26
+ 43.25424518501049,
27
+ 27.807332889109052,
28
+ 16.724424931676705,
29
+ 11.457588376291048,
30
+ 6.943266663920836,
31
+ 6.228926358749439,
32
+ 2.4549947322639083
33
+ ],
34
+ "neighbours": 4096,
35
+ "shrinkage": 10.0,
36
+ "off_data_quantile": 0.99,
37
+ "train_members_per_component": [
38
+ 1672,
39
+ 4436,
40
+ 3821,
41
+ 1876,
42
+ 1507,
43
+ 2522,
44
+ 2396,
45
+ 5296,
46
+ 2604,
47
+ 5307,
48
+ 463,
49
+ 2184,
50
+ 3571,
51
+ 2285,
52
+ 2279,
53
+ 1403,
54
+ 1539,
55
+ 1210,
56
+ 1981,
57
+ 3729,
58
+ 1566,
59
+ 1131,
60
+ 969,
61
+ 1045,
62
+ 1366
63
+ ]
64
+ }
results/summary.json CHANGED
@@ -12,10 +12,19 @@
12
  "gmm_components": 25,
13
  "gmm_covariance_type": "full",
14
  "gmm_geometry": "global means and covariance matrices",
15
- "context_conditioning": "mixture weights gamma only",
16
  "context_definition": "50 MERFISH PCs + 50 AGEA PCs",
17
  "knn_k": 20,
18
- "knn_policy": "distance-weighted empirical TRAIN-exemplar projection in standardized latent space"
 
 
 
 
 
 
 
 
 
19
  },
20
  "split_counts": {
21
  "train_units": 58158,
@@ -57,7 +66,7 @@
57
  "invariants": [
58
  "TEST PIDs are never used for AE/GMM/context training or as kNN retrieval exemplars.",
59
  "The GMM means and covariance matrices are global.",
60
- "Molecular context changes only the GMM mixture weights gamma_k(x).",
61
  "The kNN stage changes phenotype projection only; it does not change the conditional latent density.",
62
  "All spatial modeling uses the configured single-hemisphere mirroring convention."
63
  ]
 
12
  "gmm_components": 25,
13
  "gmm_covariance_type": "full",
14
  "gmm_geometry": "global means and covariance matrices",
15
+ "context_conditioning": "mixture weights gamma, and the TRAIN exemplars representing each component",
16
  "context_definition": "50 MERFISH PCs + 50 AGEA PCs",
17
  "knn_k": 20,
18
+ "knn_policy": "distance-weighted empirical TRAIN-exemplar projection in standardized latent space",
19
+ "phenotype_readout": {
20
+ "method": "context-local members: gamma_k(x) weights the mean of component k's TRAIN members nearest to x in a context key, shrunk toward all its members",
21
+ "key_dim": 3,
22
+ "key_ridge_alpha": 30000.0,
23
+ "neighbours": 4096,
24
+ "shrinkage": 10.0,
25
+ "off_data_quantile": 0.99,
26
+ "absent_context": "global member means"
27
+ }
28
  },
29
  "split_counts": {
30
  "train_units": 58158,
 
66
  "invariants": [
67
  "TEST PIDs are never used for AE/GMM/context training or as kNN retrieval exemplars.",
68
  "The GMM means and covariance matrices are global.",
69
+ "Molecular context sets the GMM mixture weights gamma_k(x) and selects which TRAIN exemplars represent each component in the phenotype readout; the conditional latent density depends on it only through gamma_k(x).",
70
  "The kNN stage changes phenotype projection only; it does not change the conditional latent density.",
71
  "All spatial modeling uses the configured single-hemisphere mirroring convention."
72
  ]