--- license: mit library_name: pytorch tags: - reinforcement-learning - offline-rl - zero-shot-rl - ogbench - forward-backward - successor-measure - gciql - manipulation --- # FB / GC-IQL Bridge — Checkpoints Model checkpoints from a study of **future-reachability representations** (Forward–Backward / successor-measure methods, GC-IQL, contrastive RL, PSM/RLDP, TD-MPC2) on **OGBench** cube manipulation, antmaze, and scene. The collection backs the "bridge" analysis of *why* successor-measure objectives underperform scalar Bellman supervision on manipulation, and the proposed **Bellman-supervised FB** and **centered-measure** variants. - **2,529 checkpoints** (`.pt` torch, `.pkl` jax flax params), all training save-intervals. - **~253 GB**, spanning 2 frameworks × many methods × tasks × seeds. - Every run's original training config travels with it, and a top-level [`checkpoint_manifest.json`](./checkpoint_manifest.json) indexes everything. > **Datasets are not included.** These are trained agents, not data. The offline datasets are the > public **OGBench** releases — obtain them from OGBench, not here. ## Repository layout ``` checkpoints/ / # torch | jax / # e.g. fb-flowbc, RGFB-Cube-Single-State-d50-v0, tdmpc2 __s__/ step_200000.pt (torch) or params_500000.pkl (jax) ... ... config.yaml (torch: the run's .hydra/config.yaml) flags.json (jax: the run's flags.json) checkpoint_manifest.json # full index (see below) ``` Labels (`method`, `task`, `seed`) are read from each run's **authoritative config**, not parsed from paths. ## What's inside | Framework | Files | Notes | |---|---|---| | `torch` (`.pt`) | 1,530 | FB / FlowBC, PSM-FlowBC, RLDP-FlowBC, CRL-FlowBC, TD-MPC2 | | `jax` (`.pkl`) | 999 | GC-IQL flow variants: bilinear, centered-measure, RG-FB, HIQL/GCIVL | Representative methods: `fb-flowbc`, `psm-flowbc`, `rldp-flowbc`, `tdmpc2`, `RGFB-Cube-Single-State-d50-v0` (Bellman-supervised FB), `GCIQL-State-Flow-CenteredMeasure-{Aux,Only}-d50` (centered successor measure), `GCIQL-State-Flow-Bilinear-*`, `HIQL-AntMaze-Medium-*`. Tasks: `cube-single-play-v0` (state & pixels), `antmaze-medium-navigate-v0`, `scene-play-v0`. ## The manifest `checkpoint_manifest.json` is a JSON array with one record per checkpoint: ```json { "repo_path": "checkpoints/jax/RGFB-Cube-Single-State-d50-v0/cube-single-play-v0__s0__rgfb_v0_s0/params_500000.pkl", "framework": "jax", "method": "RGFB-Cube-Single-State-d50-v0", "task": "cube-single-play-v0", "seed": 0, "step": "500000", "size": 89070661 } ``` ## Downloading Grab a single checkpoint (only that file is fetched): ```python from huggingface_hub import hf_hub_download path = hf_hub_download( repo_id="amsks/bilinear-checkpoints", filename="checkpoints/jax/RGFB-Cube-Single-State-d50-v0/" "cube-single-play-v0__s0__rgfb_v0_s0/params_500000.pkl", ) ``` Find checkpoints by filtering the manifest: ```python import json from huggingface_hub import hf_hub_download manifest = json.load(open(hf_hub_download("amsks/bilinear-checkpoints", "checkpoint_manifest.json"))) # e.g. every RG-FB seed rgfb = [r for r in manifest if r["method"] == "RGFB-Cube-Single-State-d50-v0"] for r in rgfb: ckpt = hf_hub_download("amsks/bilinear-checkpoints", r["repo_path"]) ``` Download an entire method/run folder: ```python from huggingface_hub import snapshot_download snapshot_download( repo_id="amsks/bilinear-checkpoints", allow_patterns="checkpoints/jax/RGFB-Cube-Single-State-d50-v0/*", ) ``` ## Loading a checkpoint Checkpoints are training-state artifacts, not standalone modules — construct the matching agent from the training code, then load the weights. Each run folder ships the exact config used to train it. - **jax (`params_*.pkl`)** — pickled flax parameter pytree for the GC-IQL flow agent. Build the agent from the co-located `flags.json`, then restore the params. - **torch (`.pt`)** — a `torch.load`-able state dict for the FB / FlowBC / TD-MPC2 agent. Build the agent from the co-located `config.yaml` (`.hydra`), then `load_state_dict`. The agent definitions and loaders live in the **Factored-FB** codebase (`tools/wandb_mode_shim/gciql_flow.py` for jax GC-IQL; `agents/fb/…` for torch FB). See that repository for exact construction. ## Provenance & reproducibility - `step` is the training step of the save; multiple steps per run are included so training trajectories can be reconstructed. - Configs are byte-for-byte the ones used at train time (`config.yaml` / `flags.json`). - Three identical duplicate backups were de-duplicated; no run is lost. ## License MIT.