gfn-gssm-s_mnih-k2 / README.md
joaquinsturtz's picture
Update README.md
028569e verified
|
Raw History Blame Contribute Delete
6.32 kB
metadata
language: en
license: cc-by-nc-nd-4.0
library_name: gfn
pipeline_tag: other
tags:
  - gfn
  - physics-informed
  - geometric-deep-learning
  - g-ssm
  - synthetic-needle-in-haystack
  - long-context
model-index:
  - name: gfn-gssm-mnih-k2
    results:
      - task:
          type: other
          name: Synthetic Multi-Needle Retrieval
        dataset:
          name: synthetic-binary-needle-haystack
          type: synthetic
        metrics:
          - name: Accuracy
            type: accuracy
            value: 100

G-SSM mNIH Solver (Official Checkpoint)

DOI: 10.5281/zenodo.19141133 Models: Hugging Face GitHub: GFN Framework

This repository contains a Geodesic State Space Model (G-SSM) checkpoint optimized for the Synthetic Binary Multi-Needle-in-a-Haystack (s-mNIH) task.

Note on Benchmark Scope: This checkpoint is trained on the synthetic binary benchmark (vocab_size = 2: 0 represents background noise, 1 represents needle impulses). It serves as a proof-of-concept demonstrating continuous geodesic integration and phase-space thresholding.

Highlights

  • Architecture: Geodesic State Space Model (G-SSM).
  • Parameters: 8,109 (PyTorch verified).
  • Inductive Bias: Toroidal manifold ($S^1$) with phase-space thresholding.
  • Context Scalability: 100% accuracy verified up to 32,000 tokens context.
  • Inference Efficiency: Constant O(1) VRAM.

Inductive Bias and Physical Mechanism

In s-mNIH, long-range dependencies are tracked via continuous Hamiltonian flow on a compact Riemannian manifold rather than softmax-based quadratic self-attention:

  1. Phase Integration: The latent state evolves on a toroidal manifold ($S^1$). Background noise tokens (0) impart negligible force, keeping the state in its resting basin near $-\pi/2$.
  2. Impulse Accumulation: Each needle token (1) injects a discrete momentum kick along the geodesic.
  3. Threshold Transition (K=2): This checkpoint was trained for exactly $K=2$ needles. The accumulation of two needle impulses overcomes the potential barrier, transitioning the trajectory into an attractor basin near $+\pi/2$.
  4. Geometric Decoding: State predictions are evaluated by measuring the circular geodesic distance modulo $2\pi$ to $+\pi/2$ (active/retrieved) versus $-\pi/2$ (resting/unretrieved).

Technical Usage (Inference)

To run inference locally, install the GFN Framework.

1. Install GFN Framework

pip install gfn==2.7.2

2. Clone this repository

git lfs install
git clone https://huggingface.co/DepthMuun/gfn-gssm-mnih-k2
cd gfn-gssm-mnih-k2

3. Run Inference Script

Use the included inference.py script to test sequence retrieval:

python inference.py

Python API Example

import math
import torch
from gfn import gssm

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# 1. Load model (resolves config.json in the same directory)
model = gssm.load("mnih_model_final.pt", device=device)
model.eval()

# 2. Build synthetic needle-in-haystack sequence (e.g., L=10,000, K=2 needles)
L = 10000
K = 2

seq = torch.zeros(1, L, dtype=torch.long)
needle_pos = sorted(torch.randperm(max(1, L - 2))[:K].tolist())
for p in needle_pos:
    seq[0, p] = 1  # 1 represents needle impulse, 0 represents noise

# Ground truth: target activates (1) only after ALL K needles are observed
y_class = torch.zeros(1, L, dtype=torch.long)
y_class[0, needle_pos[-1]:] = 1

seq = seq.to(device)
y_class = y_class.to(device)

# 3. Forward Pass
with torch.no_grad():
    # G-SSM returns (logits, (pos, vel), info)
    logits, (pos, vel), info = model(seq)

    # 4. Toroidal Decoding Logic on S^1
    PI = math.pi
    TWO_PI = 2.0 * PI
    half_pi = PI * 0.5

    # Average over projection/heads dimension if 4D
    logits_avg = logits.mean(dim=2) if logits.ndim == 4 else logits

    # Circular distance modulo 2*pi to +PI/2 (Active state)
    dist_pos = torch.min(
        torch.abs(logits_avg - half_pi) % TWO_PI,
        TWO_PI - (torch.abs(logits_avg - half_pi) % TWO_PI)
    )

    # Circular distance modulo 2*pi to -PI/2 (Inactive state)
    dist_neg = torch.min(
        torch.abs(logits_avg + half_pi) % TWO_PI,
        TWO_PI - (torch.abs(logits_avg + half_pi) % TWO_PI)
    )

    # Binary prediction for every step in the sequence [1, L]
    preds = (dist_pos.mean(dim=-1) < dist_neg.mean(dim=-1)).long()

# 5. Evaluate Metrics
accuracy = (preds == y_class).float().mean().item() * 100
final_pred = preds[0, -1].item()
target_final = y_class[0, -1].item()

print(f"Needle Positions      : {needle_pos}")
print(f"Final Step Prediction : {final_pred} (Target: {target_final})")
print(f"Sequence Accuracy     : {accuracy:.2f}%")
print("Status                :", "PASSED" if accuracy == 100.0 else "FAILED")

Manual Assembly Fallback (Optional)

If loading raw PyTorch checkpoints without gssm.load:

import json
import torch
from gfn import gssm

with open("config.json", "r") as f:
    config = json.load(f)

model = gssm.create(config=config).to(device)
checkpoint = torch.load("mnih_model_final.pt", map_location=device, weights_only=False)
state_dict = checkpoint.get("state_dict") or checkpoint.get("model") or checkpoint

model_state = model.state_dict()
filtered_state = {k: v for k, v in state_dict.items() if k in model_state}
model.load_state_dict(filtered_state, strict=False)
model.eval()

Citation

If you use this work, please cite:

@article{sturtz2026gfn,
  title={Geometric Flow Networks: A Physics-Informed Paradigm for Sequential Intelligence},
  author={Stürtz, Joaquín},
  journal={Zenodo Preprints},
  year={2026},
  doi={10.5281/zenodo.19141133},
  url={https://doi.org/10.5281/zenodo.19141132}
}

Resources