Instructions to use AlterraLaniakea/jepa-guided-diffusion with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Cosmos
How to use AlterraLaniakea/jepa-guided-diffusion with Cosmos:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
Add training metrics (CSV + TensorBoard); rename VL-JEPA -> JEPA-Guidance in the card
d17a21b verified | license: other | |
| tags: | |
| - video | |
| - world-model | |
| - jepa | |
| - vjepa2 | |
| - cosmos | |
| - wts | |
| - video-generation | |
| base_model: | |
| - facebook/vjepa2-vitl-fpc64-384 | |
| - meta-llama/Llama-3.2-1B | |
| - nvidia/Cosmos-Reason1-7B | |
| pipeline_tag: other | |
| # JEPA-Guidance β Cosmos context predictor (`jepa-guided-diffusion`) | |
| JEPA-Guidance is a lightweight **predictor bridge** that maps a WTS traffic video (plus a text | |
| query) into the frozen **Cosmos-Reason1 `(512, 1024)` context space**, so the predicted | |
| embedding can condition **Cosmos-Predict2.5** video generation *without running the 7B | |
| Reason1 text encoder at inference time*. | |
| - **Input:** video / frame directory + a WTS-format query. | |
| - **Output:** a `(512, 1024)` context tensor (the Cosmos cross-attention conditioning set). | |
| - **Objective:** permutation-invariant **set-match Hungarian MSE** β Cosmos consumes the | |
| context as an unordered set, so the predicted 512 tokens are matched to the target 512 | |
| tokens (Hungarian assignment) before MSE. | |
| ## Architecture | |
| | Part | Role | Source | | |
| |------|------|--------| | |
| | **X-Encoder** (frozen) | VJEPA 2.1 ViT-L/384 (`vjepa2_1_vit_large_384`), video mode β visual tokens | torch.hub (`facebookresearch/vjepa2`) | | |
| | **Compressor** (trained) | BLIP-2-style **QFormer**: 512 query tokens, 4 layers, cross-attn every layer | in `model.pt` | | |
| | **Backbone** (trained) | **last 4 layers of Llama-3.2-1B**, made bidirectional | arch from hub, weights in `model.pt` | | |
| | **Head** (trained) | `output_projection` β `(512, 1024)` | in `model.pt` | | |
| | **Y-Encoder** (frozen, target only) | Cosmos-Reason1-7B `full_concat` + `crossattn_proj` | not needed at inference | | |
| Only the compressor, the last 4 Llama layers, the vision projection, and the output | |
| projection are trained; VJEPA2 and Reason1 are frozen. | |
| ## Training metrics | |
| Full logs for the run that produced these weights are in `training_metrics/`. Training ran | |
| **28 epochs**; the published `best_model/` is **epoch 24**, selected on minimum validation | |
| set-match MSE. | |
| | Metric (epoch 24) | Value | | |
| |---|---| | |
| | Train set-match MSE | 5.8287 | | |
| | **Val set-match MSE** | **6.7279** | | |
| | Retrieval top-1 | 0.3606 | | |
| | Retrieval MRR | 0.4954 | | |
| | Set top-1 | 0.3613 | | |
| | Set MRR | 0.5049 | | |
| | Alignment diag | 0.1603 | | |
| | Alignment off-diag | 0.1561 | | |
| | LR at best epoch | 2.40e-05 | | |
| Validation set-match MSE fell from **21.50** (epoch 1) to **6.73** (epoch 24). | |
| ``` | |
| training_metrics/epoch_metrics.csv # 28 rows β per-epoch loss, MSE, LR, grad-norm, retrieval/alignment | |
| training_metrics/batch_metrics.csv # 4452 rows β per-step trace | |
| training_metrics/events.out.tfevents.* # TensorBoard event file | |
| ``` | |
| ```bash | |
| tensorboard --logdir training_metrics | |
| ``` | |
| Note that `align_diag` (0.160) and `align_offdiag` (0.156) sit close together: the cached | |
| Reason1 targets are strongly anisotropic, so a large shared mean dominates every vector. | |
| Training subtracts a precomputed all-token mean before normalization to recover the | |
| per-caption residual. | |
| ## Repo layout | |
| ``` | |
| config.yaml # full model/eval config the loader reads | |
| best_model/model.pt # trained weights (predictor only) | |
| best_model/model_config.yaml # model config sidecar | |
| best_model/query_tokenizer/ # Llama query tokenizer | |
| fixed/negative_<view>.pt # the two fixed negative embeddings (vehicle / overhead) | |
| training_metrics/ # per-epoch + per-batch CSVs and the TensorBoard events | |
| ``` | |
| ## Usage | |
| Install the [JEPA-Guidance](https://github.com/AlterraFa/VLJEPA-WTS) package (uv-managed), then | |
| pass this repo id as `--checkpoint` β it is downloaded and cached automatically: | |
| ```bash | |
| # whole test set: writes <scenario>/embedding.pt for the Cosmos handoff | |
| uv run jepa-guidance seq-infer --checkpoint AlterraLaniakea/jepa-guided-diffusion --test-root wts/test | |
| # single scenario -> print the (512,1024) embedding | |
| uv run jepa-guidance infer \ | |
| --checkpoint AlterraLaniakea/jepa-guided-diffusion \ | |
| --visual-path wts/test/<scenario>/input \ | |
| --output embedding \ | |
| --query "The current preset is from overhead view with 2 target subjects. Pedestrian: A man in his 30s stands on the road facing the oncoming vehicle. Vehicle: The vehicle goes straight at a constant speed." | |
| ``` | |
| **Query format.** For in-distribution embeddings, use the same query format as training | |
| (`seq-infer` builds this automatically from `caption.json`): | |
| ``` | |
| The current preset is from <overhead|vehicle> view with <N> target subject(s). Pedestrian: <ped caption> Vehicle: <veh caption> | |
| ``` | |
| - `<overhead|vehicle>` β `overhead` for CCTV, `vehicle` for dashcam (`video*`) clips. | |
| - `<N>` β subject types present: `1` (pedestrian or vehicle) or `2` (both). | |
| ## Requirements | |
| - **Gated Llama-3.2-1B**: run `hf auth login` (or export `HF_TOKEN`) before first use β the | |
| loader rebuilds the Llama backbone architecture from the hub and overlays the trained | |
| weights. (Set `hf_token: true` in `config.yaml`, or provide the token via env.) | |
| - **VJEPA2** is fetched from `facebookresearch/vjepa2` (torch.hub / GitHub source). | |
| - A GPU is recommended; VJEPA2 runs at crop size 384. | |
| ## Downstream | |
| Feed the predicted `embedding.pt` into Cosmos-Predict2.5 via `scripts/generate_cosmos.py` | |
| (see the repo README) to generate the WTS video for each scenario. | |