File size: 10,428 Bytes
1e553ed | 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 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 | # well_poincare_rl
Hierarchical, multi-step, hyperbolic (Poincaré ball) predictor for spatiotemporal
scientific fields, with PPO fine-tuning, continual learning (Replay + EWC), and
explicit data/checkpoint contracts for multi-contributor use.
**What this is for:** production-minded RL and representation work on *real*
scientific data — not demos that silently fall back to synthetic trajectories.
The same codebase also supports hierarchy recovery on real taxonomic trees
(PBDB), as a direct test of whether hyperbolic geometry helps where theory
predicts it should.
---
## Current status
| Capability | Status |
|---|---|
| Channel-agnostic encoder (variable C) | 🔒 Done — verified C=2, 11, 47 |
| Per-domain normalizer + same-C replay | 🔒 Done — end-to-end multi-C |
| Well adapter (channels-last → `(T,C,H,W)`) | 🔒 Done — matches real `the_well` interface |
| Live multi-stream Well (HF) | 🔒 setup / 📋 results — see below |
| Real PBDB hierarchy embedding | 🔒 Done — 2141 nodes, Poincaré > Euclidean at dim=8 |
| Multi-seed / multi-dim real PBDB sweep | Open |
| Attention-based channel fusion | Open |
| Soft hierarchical region priors in the ball | Open (gated on multi-domain latent data) |
---
## Key results
### Track A — Real multi-stream continual learning (Well)
Sequential training on live Hugging Face streams, no synthetic fallback.
🔒 **Checkpoint-verified:** provenance `REAL_STREAMED` for all three domains,
dataset names (`gray_scott_reaction_diffusion`, `active_matter`, `shear_flow`),
per-domain normalizer channel counts (2, 11, 4 — confirmed via each
normalizer's stored tensor shape), and model architecture (channel-agnostic
encoder, matching this session's redesign, not a stale checkpoint).
📋 **Reported (retention table below):** the loss values themselves.
`run_multistream.py`'s checkpoint format does not currently save loss/retention
history — only `model`, `params`, `normalizers`, `datasets` — so these numbers
were not independently re-derived from a saved artifact, only structurally
corroborated by everything above.
| After domain | gray_scott (C=2) | active_matter (C=11) | shear_flow (C=4) |
|---|---|---|---|
| Domain 1 | 0.3486 | — | — |
| Domain 2 | 0.3561 | 0.3299 | — |
| Domain 3 | 0.3248 | 0.3484 | 0.3358 |
Replay buffer by channel count after full run: `{2: 24, 11: 24, 4: 24}`
(🔒 structurally consistent with the checkpoint's saved normalizer shapes).
Artifact: `multistream_continual_gray_scott_active_matter_shear_flow.pt`
### Track B — Real PBDB hierarchy (Poincaré vs Euclidean)
Live taxonomy edges from Dinosauria + Mammalia (16,000 occurrence records).
🔒 **Checkpoint-verified, independently recomputed from raw data:** edge hash,
config hash, and identity hash were all recomputed from the raw occurrence
JSON using this repo's own `build_edge_list`/`hash_edge_list`/`hash_config`/
`combined_identity_hash` functions and matched the checkpoint filename and
its stored `results_table` — including `mean_delta_mrr` — to full float
precision. This is the strongest-verified result in the project.
| Setting | ΔMRR | Verification |
|---|---|---|
| dim=8, softmax, c=1.0 | 🔒 **+0.1151** | recomputed from checkpoint |
| dim=8, softmax, c=2.0 | 🔒 **+0.1576** | recomputed from checkpoint |
Individual Poincaré/Euclidean MRR values (e.g. 0.42 vs 0.305) appear in the
original run log and are arithmetically consistent with the deltas above, but
are 📋 reported, not 🔒 checkpoint-verified — the checkpoint's `results_table`
stores the delta and win count, not the two absolute MRR values separately.
- Provenance: `REAL_PBDB_TAXONOMY`
- Scale: 2141 nodes, 2172 edges
- Protocol: reconstruction (train = test edges) — standard for measuring
embedding capacity, not link-prediction generalization
- Artifact: content-addressed checkpoint under `checkpoints/` + `.meta.json`
**Scope (honest):** one seed, one dimension, metrics-only checkpoint (full
embedding weights not yet exported). Multi-seed / multi-dim confirmation on
this graph is the natural next step, not a re-verification of what's already
solid.
---
## Core contracts
### Data loading
- Real paths **hard-fail** on missing data, schema mismatch, or stream failure.
They never silently return synthetic data.
- Synthetic data is available only via explicit APIs (`get_synthetic_dataset`,
`--synthetic`). Provenance is always reported.
- Tensor layout is validated against declared contracts, never guessed from
shape heuristics.
### Trajectory validation
Checked **before** training. Trajectories shorter than `window + pred_steps`
raise `TrajectoryTooShortError` / `EmptyDatasetError` instead of being skipped
inside the loop (which previously could train an epoch on zero batches).
### Dataset-reuse registry (`DatasetRegistry`)
Prevents retraining on an already-consumed dataset under multi-contributor use.
```
(absent) --claim()--> IN_PROGRESS --mark_consumed()--> CONSUMED
\--mark_failed()--> FAILED --allow_retry()--> (absent)
```
- Claims are atomic (`O_CREAT|O_EXCL`).
- `CONSUMED` blocks future claims by default (`DatasetAlreadyUsedError`).
- Explicit override: `--allow-dataset-reuse`.
### Checkpoints (`CheckpointStore`)
Content-addressed: filename = `sha256(config_hash + code_hash + dataset_hash)`.
- Atomic writes (temp → `os.replace()` only after success).
- Partial writes (`.pt` without `.meta.json`) raise `CheckpointIntegrityError`.
- Identical `(config, code, dataset)` → `DUPLICATE_EXISTS`, no redundant write.
### Structured outcomes
Contract failures raise typed exceptions with stable `outcome_code` values
(e.g. `NO_DATA_DIRECTORY`, `CHANNEL_COUNT_MISMATCH`, `ALREADY_CONSUMED`).
---
## Architecture (as shipped)
| Component | Role |
|---|---|
| `MultiScaleEncoder` | Channel-agnostic: shared 1×1 stem per channel → mean fusion → spatial path |
| `HierarchicalHyperbolicPredictor` | Multi-step prediction in the Poincaré ball |
| `FieldNormalizer` | Per-domain z-score; hard-rejects channel-count mismatch |
| `ReplayBuffer` | Same-C (and optional spatial) filtering; clear error if mixed shapes stacked |
| `WellStreamAdapter` | Converts real `the_well` samples (channels-last input/output) → `{"fields": (T,C,H,W)}` |
| `DiagonalEWC` + teacher distillation | Continual learning across domains |
| `HierarchyEmbedding` | Poincaré / Euclidean node embeddings for taxonomy trees |
| Physics losses | Spatial/temporal consistency + generalized channel-coupling (all pairs) |
**Design note on variable channels:** ingestion (encoder + normalizer + replay)
is solved separately from latent organization. Soft hierarchical region priors
in the ball remain open and are gated on multi-domain data actually reaching
a shared latent space.
---
## Running
```bash
pip install torch geoopt optuna gymnasium h5py tqdm the_well
# Contract tests
python -m pytest tests/ -v
# Synthetic full pipeline (no network)
python -m src.run_full --synthetic
# Real local Well-format HDF5 (hard-fails if missing/invalid)
python -m src.run_full --data-root ./data/real
# Live multi-stream continual run (requires HF access, e.g. Kaggle)
python -m src.run_multistream \
--datasets gray_scott_reaction_diffusion active_matter shear_flow \
--max-samples 96 --epochs-per-domain 3
# Hierarchy embedding — synthetic (no network)
python -m src.run_hierarchy_embed --synthetic --tree-type balanced \
--branching-factor 3 --tree-depth 6 --dims 2 3 5 8 \
--loss-types softmax --burn-in-epochs 0 --c-values 1.0 2.0 \
--epochs 80 --lr 0.02 --seeds 0 1 2
# Hierarchy embedding — real PBDB (requires network)
python -m src.data_pbdb --discover --base-name Dinosauria --limit 5 # first
python -m src.run_hierarchy_embed \
--pbdb-taxa Dinosauria Mammalia \
--dims 8 --loss-types softmax --burn-in-epochs 0 \
--c-values 1.0 2.0 --epochs 80 --lr 0.02 --seeds 0 --optimizer radam
```
Defaults that matter for Track B with current stability guards:
`--burn-in-epochs 0` (ablation showed burn-in reduces ΔMRR under guarded
`clip_to_ball` / `clamp_curvature` setup).
---
## Artifacts
| File | What it is |
|---|---|
| `multistream_continual_gray_scott_active_matter_shear_flow.pt` | Track A model + params + per-domain normalizers + dataset list (no loss history) |
| `checkpoints/<identity_hash>.pt` + `.meta.json` | Track B metrics table + provenance (content-addressed) |
| Large `*_8000.json` under data cache | Raw PBDB fetch cache only — optional, regenerable |
Hierarchy checkpoints currently store `results_table` (metrics), not full
embedding weights. Weight export is a small follow-on if reloadable node
vectors are needed.
---
## Known limitations
- Multi-stream retention numbers are 📋 reported, not 🔒 checkpoint-verified —
see Track A above. The run's provenance, channel counts, and architecture
are independently confirmed; the specific loss trajectory is not yet, since
the current checkpoint format doesn't persist it.
- Real PBDB hierarchy result is one seed / one dimension.
- Channel fusion is mean-pooling (lossy). Attention fusion is the natural upgrade.
- Soft region structure in the Poincaré ball is intentionally deferred until
multi-domain data is routinely in a shared latent.
- PPO is implemented and smoke-tested; it is not the primary claim of the
current verified results.
- Contributors are assumed to exchange portable checkpoint + meta files;
a shared remote registry is a future decision if contributor count grows.
---
## Design principles (non-negotiable)
1. **Real-data paths hard-fail.** Synthetic only via explicit opt-in.
2. **No silent shape / channel mistakes.** Prefer loud, typed errors.
3. **Provenance is first-class.** Every load reports where data came from.
4. **Checkpoints are content-addressed and atomic.**
5. **Report what the numbers say**, including when Euclidean wins or a
previously recommended default (e.g. burn-in) is reversed by ablation.
6. **Verification tier is part of the claim.** "Verified" always specifies
*verified how* — checkpoint-recomputed and self-reported are not the same
thing, and this document says which is which rather than picking one word
for both.
---
## License / packaging
Apache 2.0 intended for model and code release. Package layout is
`src/`-based and pip-installable for local development. |