初始化项目,由ModelHub XC社区提供模型
Model: EmpathicRobotics/vla-1.7b-qwen3-v2 Source: Original Platform
This commit is contained in:
36
.gitattributes
vendored
Normal file
36
.gitattributes
vendored
Normal file
@@ -0,0 +1,36 @@
|
||||
*.7z filter=lfs diff=lfs merge=lfs -text
|
||||
*.arrow filter=lfs diff=lfs merge=lfs -text
|
||||
*.bin filter=lfs diff=lfs merge=lfs -text
|
||||
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
||||
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
||||
*.ftz filter=lfs diff=lfs merge=lfs -text
|
||||
*.gz filter=lfs diff=lfs merge=lfs -text
|
||||
*.h5 filter=lfs diff=lfs merge=lfs -text
|
||||
*.joblib filter=lfs diff=lfs merge=lfs -text
|
||||
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
||||
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
||||
*.model filter=lfs diff=lfs merge=lfs -text
|
||||
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
||||
*.npy filter=lfs diff=lfs merge=lfs -text
|
||||
*.npz filter=lfs diff=lfs merge=lfs -text
|
||||
*.onnx filter=lfs diff=lfs merge=lfs -text
|
||||
*.ot filter=lfs diff=lfs merge=lfs -text
|
||||
*.parquet filter=lfs diff=lfs merge=lfs -text
|
||||
*.pb filter=lfs diff=lfs merge=lfs -text
|
||||
*.pickle filter=lfs diff=lfs merge=lfs -text
|
||||
*.pkl filter=lfs diff=lfs merge=lfs -text
|
||||
*.pt filter=lfs diff=lfs merge=lfs -text
|
||||
*.pth filter=lfs diff=lfs merge=lfs -text
|
||||
*.rar filter=lfs diff=lfs merge=lfs -text
|
||||
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
||||
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
||||
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
||||
*.tar filter=lfs diff=lfs merge=lfs -text
|
||||
*.tflite filter=lfs diff=lfs merge=lfs -text
|
||||
*.tgz filter=lfs diff=lfs merge=lfs -text
|
||||
*.wasm filter=lfs diff=lfs merge=lfs -text
|
||||
*.xz filter=lfs diff=lfs merge=lfs -text
|
||||
*.zip filter=lfs diff=lfs merge=lfs -text
|
||||
*.zst filter=lfs diff=lfs merge=lfs -text
|
||||
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
||||
tokenizer.json filter=lfs diff=lfs merge=lfs -text
|
||||
262
README.md
Normal file
262
README.md
Normal file
@@ -0,0 +1,262 @@
|
||||
---
|
||||
license: apache-2.0
|
||||
language:
|
||||
- en
|
||||
tags:
|
||||
- robotics
|
||||
- vla
|
||||
- vision-language-action
|
||||
- 3d-pose
|
||||
- qwen3
|
||||
- megatron
|
||||
- multimodal
|
||||
pipeline_tag: text-generation
|
||||
library_name: transformers
|
||||
---
|
||||
|
||||
# VLA 1.7B — Qwen3 v2
|
||||
|
||||
A 1.7B parameter Vision-Language-Action model, migrated to a **Qwen3** backbone and
|
||||
trained on a **5-source, ~32B-token multimodal mix** (video, 3D pose, audio, image+caption).
|
||||
This is the project's first Qwen3-based VLA model, and the first trained after fixing
|
||||
the "stuck in one modality" failure mode found in the previous model.
|
||||
|
||||
## Key facts
|
||||
|
||||
| | |
|
||||
|---|---|
|
||||
| **Architecture** | Qwen3 (28 layers, hidden 2048, intermediate 6144, 16 attn heads / 8 KV heads (GQA), qk-layernorm, RoPE θ=1e6, tied embeddings) |
|
||||
| **Parameters** | 1.94B (including embeddings for 257,920 vocab) |
|
||||
| **Vocab size** | 257,920 (Qwen3 base ~151,669 + 106,232 VLA tokens, padded) |
|
||||
| **Tokenizer** | [EmpathicRobotics/tokenizer-vla-qwen3](https://huggingface.co/EmpathicRobotics/tokenizer-vla-qwen3) |
|
||||
| **Training data** | ~32.01B tokens across 5 sources: FineVideo-VLA v6, MixtureVitae-Omni, OmniVideo-100K, synth-llava, emotional-roleplay |
|
||||
| **Training** | 7,632 iters (1 epoch), 64 nodes × 4 GH200 GPUs, global batch 1024, seq len 4096 |
|
||||
| **Final loss** | Train: 1.694, Val: 1.7526 (PPL 5.77), Test: 1.7722 (PPL 5.88) |
|
||||
| **Precision** | bf16 |
|
||||
| **Context length** | 4,096 tokens |
|
||||
|
||||
## What this model does
|
||||
|
||||
Given a text prompt (activity description, image seed2 block, or partial modality
|
||||
sequence), the model generates an interleaved multimodal token sequence spanning
|
||||
6 categories it was trained on:
|
||||
|
||||
```
|
||||
<seed2_N> ... # 1 FPS semantic image/video keyframes (vocab 8192)
|
||||
<cosmos_N> ... </cosmos> # 8-frame spatial video tokens (vocab 64000)
|
||||
<snac_N> ... </snac> # SNAC audio codec tokens (12,288)
|
||||
<speech> ... </speech> # inline spoken-dialogue text
|
||||
<caption> ... </caption> # inline visual caption text
|
||||
<agent> <fps_30> <pelvis> ... </agent> # 3D human pose, 17 H36M joints
|
||||
```
|
||||
|
||||
## Progress vs. the previous model
|
||||
|
||||
The first model ([vla-1.7b-pab-spline-adaptive](https://huggingface.co/EmpathicRobotics/vla-1.7b-pab-spline-adaptive))
|
||||
passed agent-completion but **failed modality transitions**: it stayed in `seed2` mode and
|
||||
never transitioned to `cosmos`/`avclm`/`agent` from text alone. This model no longer has
|
||||
that failure mode — it transitions freely across all 6 trained categories, in both
|
||||
greedy and sampled decoding.
|
||||
|
||||
Strongest evidence: given **only** 32 real `<seed2_N>` tokens from a held-out image record
|
||||
(no other text hint), the model generated a topically-correct caption closely matching the
|
||||
real ground truth, then closed `</caption><|im_end|>` cleanly — genuine image↔text
|
||||
cross-modal binding, not template noise.
|
||||
|
||||
It can also produce full agent (3D pose) blocks that decode to valid, non-degenerate
|
||||
coordinates, and — verified for the first time on this model's own generation, not just
|
||||
training data — `cosmos` video tokens that decode to a real, playable video via
|
||||
[`Cosmos-Tokenizer-DV8x16x16`](https://github.com/NVIDIA/Cosmos-Tokenizer), and `snac`
|
||||
audio tokens that decode to a real, non-silent waveform via
|
||||
[SNAC](https://github.com/hubertsiuzdak/snac) (`hubertsiuzdak/snac_24khz`).
|
||||
|
||||
All 3 non-text modalities the model actually produces in volume (`cosmos`, `snac`, `seed2`)
|
||||
now have a working decoder in the project repo (`tools/decode/`) and have each been
|
||||
round-tripped on real ground-truth tokens. `seed2` is generative rather than a
|
||||
deterministic codec round-trip — see Known limitations.
|
||||
|
||||
## Known limitations
|
||||
|
||||
- **Greedy decoding can degenerate into repeated-token loops** inside long `cosmos`
|
||||
runs (e.g. the same token repeating 6-8 times), which can burn the generation budget
|
||||
before reaching `<fps_N>`/`agent`. Sampling with `repetition_penalty>1` mitigates this.
|
||||
- **Sampling trades accuracy for diversity**: in the image-captioning test, sampled
|
||||
generation occasionally hallucinated details (e.g. an invented name) not present in
|
||||
the source image; greedy decoding did not.
|
||||
- **`cosmos` tokens dominate generation**: aggregated across all test prompts, `cosmos`
|
||||
is 61-77% of all non-text VLA tokens produced (vs. a minority share for
|
||||
agent/seed2/snac combined). This is largely structural (one cosmos chunk costs a fixed
|
||||
200 tokens vs. ~1-4 tokens for the others), but it does mean cosmos runs can consume
|
||||
most of a generation's token budget before reaching `<fps_N>`/`agent`.
|
||||
- **`avc_lm` tokens are essentially unused** — discarded at the data-flatten stage before
|
||||
training (to control token count), so the model rarely if ever produces them.
|
||||
- **`seed2`→image reconstruction is generative, not a deterministic round-trip.**
|
||||
Seed2Tokenizer has no pixel decoder of its own; reconstruction conditions a diffusion
|
||||
img2img pipeline (`StableUnCLIPImg2ImgPipeline`) on the token embeddings to *generate*
|
||||
a plausible image, unlike `cosmos`/`snac`'s lossy-but-deterministic codec decoders — two
|
||||
runs of the same tokens can come out visually different. Verified end-to-end on 32 real
|
||||
ground-truth `<seed2_N>` tokens (`tools/decode/decode_seed2.py`) — the diffusion weights
|
||||
now come from a community mirror (`sd2-community/stable-diffusion-2-1-unclip`), since
|
||||
the original `stabilityai/stable-diffusion-2-1-unclip` was removed from HuggingFace.
|
||||
- **Evaluation so far is qualitative** (manual inspection of generated tokens/decoded
|
||||
media) — no MPJPE, BLEU/CIDEr, or closed-loop task-success metric has been run yet.
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
import torch
|
||||
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
"EmpathicRobotics/vla-1.7b-qwen3-v2",
|
||||
torch_dtype=torch.bfloat16,
|
||||
device_map="auto",
|
||||
trust_remote_code=True,
|
||||
)
|
||||
tokenizer = AutoTokenizer.from_pretrained("EmpathicRobotics/vla-1.7b-qwen3-v2")
|
||||
|
||||
prompt = (
|
||||
"### Context: Person raises both arms above head.\n"
|
||||
"<seed2_3758> <seed2_2157> <cosmos_58567> "
|
||||
"<fps_30> <pelvis> <pelvis_t_0> <pelvis_x_128> <pelvis_y_128> <pelvis_z_128>"
|
||||
)
|
||||
input_ids = tokenizer.encode(prompt, return_tensors="pt").to(model.device)
|
||||
output = model.generate(
|
||||
input_ids, max_new_tokens=500,
|
||||
do_sample=True, temperature=0.8, top_p=0.9, repetition_penalty=1.3,
|
||||
)
|
||||
print(tokenizer.decode(output[0]))
|
||||
```
|
||||
|
||||
### Encoding real media into tokens (so you can actually prompt the model)
|
||||
|
||||
The `## Usage` prompt above uses pre-picked token ids as a demo. To send the
|
||||
model *real* media -- e.g. "here's a photo, continue the scene" or "here's
|
||||
a real motion clip, keep going" -- encode it first with the 4 encoders below
|
||||
(**verified working 2026-07-23**, each tested end-to-end: real media ->
|
||||
tokens -> decoded/compared back against the original). Bundled in this repo
|
||||
the same way as the decoders (`tools/encode/`), no separate `git clone`
|
||||
needed.
|
||||
|
||||
```bash
|
||||
# Image -> <seed2_N> tokens (32 ids, auto-downloads the Q-Former checkpoint
|
||||
# from ontocord/seed2 if not cached locally)
|
||||
python tools/encode/encode_seed2.py --image photo.jpg
|
||||
|
||||
# 8 video frames -> <cosmos_N> tokens (200 ids -- this model's OLD
|
||||
# window=8/square-crop convention, NOT the newer 2026-07-23 aspect-preserving
|
||||
# one; auto-downloads encoder.jit from nvidia/Cosmos-Tokenizer-DV8x16x16)
|
||||
python tools/encode/encode_cosmos.py --frames f0.png f1.png f2.png f3.png f4.png f5.png f6.png f7.png
|
||||
|
||||
# Audio/video file -> <snac_N> tokens (listen-format, <snac> wrapper --
|
||||
# this model never saw the newer <listen>/<speak> convention or speak-format L2)
|
||||
python tools/encode/encode_snac.py --input clip.wav
|
||||
|
||||
# Real 3D pose (8 frames x 17 joints x xyz, metres, root-centred) -> <agent>
|
||||
# tokens -- for "give the model a real motion capture / pose-pipeline output,
|
||||
# have it continue" (same behavior already verified: agent completion PASS)
|
||||
python tools/encode/encode_agent.py --input pose.npy # shape (8, 17, 3)
|
||||
```
|
||||
|
||||
Splice the printed token block into your prompt (e.g. after `### Context:
|
||||
...`) the same way the `## Usage` example does, then call `model.generate()`
|
||||
as shown there.
|
||||
|
||||
### Decoding generated tokens back to media
|
||||
|
||||
The decoder scripts + their vendored dependencies are bundled directly in
|
||||
**this repo** (`tools/`) -- one `snapshot_download` gets everything, no
|
||||
separate `git clone` needed. (Also mirrored at
|
||||
[github.com/TieuDaoChanNhan/finevideo-vla](https://github.com/TieuDaoChanNhan/finevideo-vla)
|
||||
if you'd rather browse/clone the code on its own.) **Verified working
|
||||
2026-07-23** with no cluster/internal access required, each tested end-to-end
|
||||
on real tokens this model actually generated.
|
||||
|
||||
```bash
|
||||
python -c "
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download('EmpathicRobotics/vla-1.7b-qwen3-v2', allow_patterns=['tools/*', 'tools/**/*'])
|
||||
"
|
||||
pip install scipy numpy torch torchvision imageio-ffmpeg soundfile snac huggingface_hub
|
||||
cd <snapshot-download-cache-dir-printed-above>
|
||||
```
|
||||
|
||||
**Agent tokens -> 3D pose** (pure Python, no extra downloads):
|
||||
```bash
|
||||
python tools/eval/decode_agent_tokens.py --input generated_tokens.txt --output poses.json
|
||||
```
|
||||
|
||||
**Cosmos tokens -> video** (auto-downloads the ~350MB decoder checkpoint from
|
||||
[nvidia/Cosmos-Tokenizer-DV8x16x16](https://huggingface.co/nvidia/Cosmos-Tokenizer-DV8x16x16)
|
||||
on first run):
|
||||
```bash
|
||||
python tools/decode/decode_cosmos.py --tokens 58345,57843,... --output out.mp4
|
||||
# this model's cosmos chunks are exactly 200 raw ids each (8 frames, 160x160,
|
||||
# square-cropped -- the DV8x16x16 checkpoint's own token grid for that input
|
||||
# size). A later dataset pivot (2026-07-23, aspect-preserving/896 tokens)
|
||||
# does NOT apply to this model -- it was trained entirely on the 200-token/
|
||||
# square-crop convention.
|
||||
```
|
||||
|
||||
**SNAC tokens -> audio** (auto-downloads `hubertsiuzdak/snac_24khz` from HF):
|
||||
```bash
|
||||
python tools/decode/decode_snac.py --tokens 130911,134940,... --format listen --output out.wav
|
||||
# this model only ever saw "listen" format (3 tokens/base-frame, <snac>
|
||||
# wrapper) -- do NOT use --format speak, that's a newer (2026-07-23)
|
||||
# convention this model was never trained on.
|
||||
```
|
||||
|
||||
**Seed2 tokens -> image** (auto-downloads the ~2.6GB Q-Former checkpoint from
|
||||
the tokenizer's own public repo,
|
||||
[ontocord/seed2](https://huggingface.co/ontocord/seed2), plus a ~5GB
|
||||
diffusion img2img pipeline on first run -- this one is a generative
|
||||
*reconstruction*, not a deterministic decode, so expect run-to-run and
|
||||
prompt-to-prompt variation in the exact pixels even for the same tokens):
|
||||
```bash
|
||||
python tools/decode/decode_seed2.py --tokens 6750,680,2472,... --output out.png
|
||||
# exactly 32 raw ids per image (Seed2Tokenizer's fixed Q-former query length)
|
||||
```
|
||||
|
||||
## Training details
|
||||
|
||||
### Loss curve
|
||||
|
||||
| Iter | Loss |
|
||||
|---|---|
|
||||
| 50 | 6.472 |
|
||||
| 500 | 2.840 |
|
||||
| 1000 | 2.154 |
|
||||
| 2000 | 1.953 |
|
||||
| 4000 | 1.826 |
|
||||
| 6000 | 1.767 |
|
||||
| 7600 | 1.694 |
|
||||
| 7632 (val) | 1.7526 (PPL 5.77) |
|
||||
| 7632 (test) | 1.7722 (PPL 5.88) |
|
||||
|
||||
### Config
|
||||
|
||||
- **Batch**: GBS 1024, seq_len 4096 → 32.01B tokens trained (exactly 1 epoch)
|
||||
- **Infrastructure**: 64 nodes × 4 GH200 GPUs (256 total), ~284 TFLOP/s/GPU, ~21,800 tok/s/GPU
|
||||
- **Framework**: Megatron-LM via oellm-autoexp
|
||||
|
||||
### Data mix
|
||||
|
||||
| Source | Tokens |
|
||||
|---|---|
|
||||
| MixtureVitae-Omni | 20.39B |
|
||||
| FineVideo-VLA v6 | 10.93B |
|
||||
| OmniVideo-100K (video) | 0.54B |
|
||||
| synth-llava | 0.10B |
|
||||
| emotional-roleplay (SNAC TTS) | 0.05B |
|
||||
| **Total** | **~32.01B** |
|
||||
|
||||
## Citation
|
||||
|
||||
```bibtex
|
||||
@misc{empathicrobotics2026vlaqwen3,
|
||||
title={VLA 1.7B Qwen3 v2: Multi-Source Multimodal Vision-Language-Action Pretraining},
|
||||
author={EmpathicRobotics},
|
||||
year={2026},
|
||||
url={https://huggingface.co/EmpathicRobotics/vla-1.7b-qwen3-v2}
|
||||
}
|
||||
```
|
||||
86
chat_template.jinja
Normal file
86
chat_template.jinja
Normal file
@@ -0,0 +1,86 @@
|
||||
{%- if tools %}
|
||||
{{- '<|im_start|>system\n' }}
|
||||
{%- if messages[0].role == 'system' %}
|
||||
{{- messages[0].content + '\n\n' }}
|
||||
{%- endif %}
|
||||
{{- "# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>" }}
|
||||
{%- for tool in tools %}
|
||||
{{- "\n" }}
|
||||
{{- tool | tojson }}
|
||||
{%- endfor %}
|
||||
{{- "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call><|im_end|>\n" }}
|
||||
{%- else %}
|
||||
{%- if messages[0].role == 'system' %}
|
||||
{{- '<|im_start|>system\n' + messages[0].content + '<|im_end|>\n' }}
|
||||
{%- endif %}
|
||||
{%- endif %}
|
||||
{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}
|
||||
{%- for message in messages[::-1] %}
|
||||
{%- set index = (messages|length - 1) - loop.index0 %}
|
||||
{%- if ns.multi_step_tool and message.role == "user" and message.content is string and not(message.content.startswith('<tool_response>') and message.content.endswith('</tool_response>')) %}
|
||||
{%- set ns.multi_step_tool = false %}
|
||||
{%- set ns.last_query_index = index %}
|
||||
{%- endif %}
|
||||
{%- endfor %}
|
||||
{%- for message in messages %}
|
||||
{%- if message.content is string %}
|
||||
{%- set content = message.content %}
|
||||
{%- else %}
|
||||
{%- set content = '' %}
|
||||
{%- endif %}
|
||||
{%- if (message.role == "user") or (message.role == "system" and not loop.first) %}
|
||||
{{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
|
||||
{%- elif message.role == "assistant" %}
|
||||
{%- set reasoning_content = '' %}
|
||||
{%- if message.reasoning_content is string %}
|
||||
{%- set reasoning_content = message.reasoning_content %}
|
||||
{%- else %}
|
||||
{%- if '</think>' in content %}
|
||||
{%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
|
||||
{%- set content = content.split('</think>')[-1].lstrip('\n') %}
|
||||
{%- endif %}
|
||||
{%- endif %}
|
||||
{%- if loop.index0 > ns.last_query_index %}
|
||||
{%- if loop.last or (not loop.last and reasoning_content) %}
|
||||
{{- '<|im_start|>' + message.role + '\n<think>\n' + reasoning_content.strip('\n') + '\n</think>\n\n' + content.lstrip('\n') }}
|
||||
{%- else %}
|
||||
{{- '<|im_start|>' + message.role + '\n' + content }}
|
||||
{%- endif %}
|
||||
{%- else %}
|
||||
{{- '<|im_start|>' + message.role + '\n' + content }}
|
||||
{%- endif %}
|
||||
{%- if message.tool_calls %}
|
||||
{%- for tool_call in message.tool_calls %}
|
||||
{%- if (loop.first and content) or (not loop.first) %}
|
||||
{{- '\n' }}
|
||||
{%- endif %}
|
||||
{%- if tool_call.function %}
|
||||
{%- set tool_call = tool_call.function %}
|
||||
{%- endif %}
|
||||
{{- '<tool_call>\n{"name": "' }}
|
||||
{{- tool_call.name }}
|
||||
{{- '", "arguments": ' }}
|
||||
{%- if tool_call.arguments is string %}
|
||||
{{- tool_call.arguments }}
|
||||
{%- else %}
|
||||
{{- tool_call.arguments | tojson }}
|
||||
{%- endif %}
|
||||
{{- '}\n</tool_call>' }}
|
||||
{%- endfor %}
|
||||
{%- endif %}
|
||||
{{- '<|im_end|>\n' }}
|
||||
{%- elif message.role == "tool" %}
|
||||
{%- if loop.first or (messages[loop.index0 - 1].role != "tool") %}
|
||||
{{- '<|im_start|>user' }}
|
||||
{%- endif %}
|
||||
{{- '\n<tool_response>\n' }}
|
||||
{{- content }}
|
||||
{{- '\n</tool_response>' }}
|
||||
{%- if loop.last or (messages[loop.index0 + 1].role != "tool") %}
|
||||
{{- '<|im_end|>\n' }}
|
||||
{%- endif %}
|
||||
{%- endif %}
|
||||
{%- endfor %}
|
||||
{%- if add_generation_prompt %}
|
||||
{{- '<|im_start|>assistant\n<think>\n' }}
|
||||
{%- endif %}
|
||||
64
config.json
Normal file
64
config.json
Normal file
@@ -0,0 +1,64 @@
|
||||
{
|
||||
"architectures": [
|
||||
"Qwen3ForCausalLM"
|
||||
],
|
||||
"attention_bias": false,
|
||||
"attention_dropout": 0.0,
|
||||
"bos_token_id": null,
|
||||
"dtype": "bfloat16",
|
||||
"eos_token_id": 151645,
|
||||
"head_dim": 128,
|
||||
"hidden_act": "silu",
|
||||
"hidden_size": 2048,
|
||||
"initializer_range": 0.02,
|
||||
"intermediate_size": 6144,
|
||||
"layer_norm_eps": 1e-06,
|
||||
"layer_types": [
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention",
|
||||
"full_attention"
|
||||
],
|
||||
"max_position_embeddings": 4096,
|
||||
"max_window_layers": 28,
|
||||
"mlp_bias": false,
|
||||
"model_type": "qwen3",
|
||||
"num_attention_heads": 16,
|
||||
"num_hidden_layers": 28,
|
||||
"num_key_value_heads": 8,
|
||||
"pad_token_id": null,
|
||||
"qk_layernorm": true,
|
||||
"rms_norm_eps": 1e-06,
|
||||
"rope_parameters": null,
|
||||
"rope_theta": 1000000,
|
||||
"sliding_window": null,
|
||||
"tie_word_embeddings": true,
|
||||
"transformers_version": "5.3.0",
|
||||
"use_cache": true,
|
||||
"use_sliding_window": false,
|
||||
"vocab_size": 257920
|
||||
}
|
||||
3
model.safetensors
Normal file
3
model.safetensors
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:e69e84c52a5299f5b8fd6b100306f204055e08fedd510de57941f80b485657e9
|
||||
size 3875296152
|
||||
309
pipeline_pose/phase5_adaptive_pchip.py
Normal file
309
pipeline_pose/phase5_adaptive_pchip.py
Normal file
@@ -0,0 +1,309 @@
|
||||
"""
|
||||
Phase 5 — Adaptive PCHIP per-joint tokenizer.
|
||||
|
||||
For each window from Phase 4 (WINDOW_FRAMES frames -- 8 originally, 24 as of
|
||||
2026-07-22, see REPORT.md #38), each of the 17 joints gets an independent
|
||||
PCHIP compression with adaptive control-point count based on per-joint
|
||||
curvature over the WHOLE window (all WINDOW_FRAMES frames considered, not a
|
||||
fixed sub-grid):
|
||||
|
||||
Tier 2 (2 CPs: start + end) — max curvature < tau_low
|
||||
Tier 4 (4 CPs: start + end + top-2 interior) — tau_low <= curvature < tau_high
|
||||
Tier MAX_CPS (MAX_CPS CPs: start+end+top-(MAX_CPS-2) interior, chosen by
|
||||
curvature out of ALL WINDOW_FRAMES candidates) — curvature >= tau_high
|
||||
|
||||
MAX_CPS is fixed at 8 regardless of WINDOW_FRAMES -- widening the window
|
||||
gives the top tier more candidate positions to pick its best 8 from, not more
|
||||
tokens/joint. See MAX_CPS's own docstring below for why this matters.
|
||||
|
||||
Token stream per window (WINDOW_FRAMES=24 example; t values now range over
|
||||
however many of the window's real frame indices got chosen as CPs, e.g.
|
||||
t_0 and t_23 for tier 2, or t_0/t_5/t_14/t_23 for tier 4 if frames 5 and 14
|
||||
had the highest curvature):
|
||||
<fps_30>
|
||||
<pelvis> <pelvis_t_0> <pelvis_x_N> <pelvis_y_N> <pelvis_z_N>
|
||||
<pelvis_t_23> <pelvis_x_N> <pelvis_y_N> <pelvis_z_N> </pelvis>
|
||||
<r_hip> <r_hip_t_0> <r_hip_x_N> ... </r_hip>
|
||||
...
|
||||
|
||||
Quantization: [-2.0 m, +2.0 m] -> [0, 255] (precision ~15.7 mm)
|
||||
Time tokens: real frame index within the window, 0 to WINDOW_FRAMES-1 --
|
||||
requires <{joint}_t_N> tokens up to WINDOW_FRAMES-1 in the
|
||||
tokenizer vocab (only 0-7 existed before 2026-07-22).
|
||||
|
||||
Input: outputs/yolo_cleaned_30fps/{video_id}_cleaned.jsonl
|
||||
Output: outputs/agent_tokens_adaptive/{video_id}_tokens.jsonl
|
||||
Each line: {"video_id", "window_id", "fps", "token_str", "cp_counts"}
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
|
||||
# ── Constants ─────────────────────────────────────────────────────────────────
|
||||
|
||||
TARGET_FPS = 30
|
||||
WINDOW_FRAMES = 8
|
||||
N_JOINTS = 17
|
||||
COORD_RANGE = 2.0
|
||||
STRIDE = 8
|
||||
# 2026-07-22 (REPORT.md #38): cap on control points per joint, independent of
|
||||
# WINDOW_FRAMES. Before this change WINDOW_FRAMES==MAX_CPS==8 always (the top
|
||||
# tier was literally "use every frame"), so widening the window to 24 frames
|
||||
# would have silently tripled worst-case tokens/joint (24 CPs instead of 8)
|
||||
# with no code change needed to trigger it. Keeping MAX_CPS fixed at 8 means
|
||||
# the top tier now means "pick the best 8 of WINDOW_FRAMES candidates by
|
||||
# curvature" instead of "use all of them" -- same worst-case token cost as
|
||||
# before, but the 8 chosen points can be anywhere in the (now wider) window
|
||||
# instead of forced onto a fixed 8-slot grid. This is the whole point of
|
||||
# "Option 2" (dense pose, no subsampling before curve-fitting) agreed with
|
||||
# the user: Phase 3 keeps every real frame, Phase 5 decides freely which
|
||||
# frames matter most.
|
||||
MAX_CPS = 8
|
||||
|
||||
TAU_LOW = 0.005
|
||||
TAU_HIGH = 0.05
|
||||
|
||||
JOINT_NAMES = [
|
||||
"pelvis", "r_hip", "r_knee", "r_ankle",
|
||||
"l_hip", "l_knee", "l_ankle",
|
||||
"spine", "thorax", "nose", "head_top",
|
||||
"l_shoulder", "l_elbow", "l_wrist",
|
||||
"r_shoulder", "r_elbow", "r_wrist",
|
||||
]
|
||||
|
||||
# ── Quantization ──────────────────────────────────────────────────────────────
|
||||
|
||||
def quantize(v: float) -> int:
|
||||
return int(np.clip(round((v + COORD_RANGE) / (2.0 * COORD_RANGE) * 255), 0, 255))
|
||||
|
||||
|
||||
def dequantize(n: int) -> float:
|
||||
return n / 255.0 * (2.0 * COORD_RANGE) - COORD_RANGE
|
||||
|
||||
|
||||
# ── Per-joint adaptive CP selection ───────────────────────────────────────────
|
||||
|
||||
def joint_curvature(trajectory: np.ndarray) -> float:
|
||||
"""Max curvature (acceleration norm) for a single joint trajectory (8, 3)."""
|
||||
if trajectory.shape[0] < 3:
|
||||
return 0.0
|
||||
vel = np.diff(trajectory, axis=0)
|
||||
acc = np.diff(vel, axis=0)
|
||||
return float(np.max(np.linalg.norm(acc, axis=1)))
|
||||
|
||||
|
||||
def select_cp_indices(trajectory: np.ndarray, tau_low: float, tau_high: float) -> np.ndarray:
|
||||
"""Choose which frame indices become control points for one joint.
|
||||
|
||||
Tier sizes are fixed at 2 / 4 / MAX_CPS regardless of how many frames are
|
||||
in the window (see MAX_CPS's docstring) -- the top tier picks the
|
||||
MAX_CPS-2 highest-curvature *interior* frames out of every candidate in
|
||||
the window, not a fixed grid position. With WINDOW_FRAMES==MAX_CPS==8
|
||||
(the original config) this is exactly equivalent to the old
|
||||
`np.arange(WINDOW_FRAMES)` behavior, since "top 6 of 6 interior
|
||||
candidates" is all of them.
|
||||
"""
|
||||
curv = joint_curvature(trajectory)
|
||||
n_frames = trajectory.shape[0]
|
||||
|
||||
if curv < tau_low:
|
||||
return np.array([0, n_frames - 1])
|
||||
|
||||
n_interior = 2 if curv < tau_high else (MAX_CPS - 2)
|
||||
n_interior = min(n_interior, max(n_frames - 2, 0))
|
||||
|
||||
vel = np.diff(trajectory, axis=0)
|
||||
acc = np.diff(vel, axis=0)
|
||||
acc_norms = np.linalg.norm(acc, axis=1) # (n_frames-2,)
|
||||
|
||||
# acc[i] corresponds to frame i+1 (second derivative offset)
|
||||
interior_curv = np.zeros(n_frames)
|
||||
for i in range(len(acc_norms)):
|
||||
interior_curv[i + 1] = acc_norms[i]
|
||||
|
||||
# Exclude endpoints (already included), pick top-n_interior interior frames
|
||||
interior_curv[0] = -1.0
|
||||
interior_curv[-1] = -1.0
|
||||
top_n = np.argsort(interior_curv)[-n_interior:] if n_interior > 0 else np.array([], dtype=int)
|
||||
|
||||
indices = np.unique(np.sort(np.concatenate(([0], top_n, [n_frames - 1]))))
|
||||
return indices.astype(int)
|
||||
|
||||
|
||||
# ── Token builder ─────────────────────────────────────────────────────────────
|
||||
|
||||
def build_token_str(
|
||||
states: np.ndarray,
|
||||
fps: int = TARGET_FPS,
|
||||
tau_low: float = TAU_LOW,
|
||||
tau_high: float = TAU_HIGH,
|
||||
) -> tuple:
|
||||
"""
|
||||
states : (8, 17, 3) float32, root-centred metric coordinates
|
||||
Returns (token_str, cp_counts_dict)
|
||||
"""
|
||||
parts = [f"<fps_{fps}>"]
|
||||
cp_counts = {}
|
||||
|
||||
for j in range(N_JOINTS):
|
||||
name = JOINT_NAMES[j]
|
||||
trajectory = states[:, j, :] # (8, 3)
|
||||
cp_idx = select_cp_indices(trajectory, tau_low, tau_high)
|
||||
cp_counts[name] = len(cp_idx)
|
||||
|
||||
parts.append(f"<{name}>")
|
||||
for fi in cp_idx:
|
||||
x, y, z = trajectory[fi]
|
||||
parts.append(f"<{name}_t_{fi}>")
|
||||
parts.append(f"<{name}_x_{quantize(x)}>")
|
||||
parts.append(f"<{name}_y_{quantize(y)}>")
|
||||
parts.append(f"<{name}_z_{quantize(z)}>")
|
||||
parts.append(f"</{name}>")
|
||||
|
||||
return " ".join(parts), cp_counts
|
||||
|
||||
|
||||
# ── Per-file processing ──────────────────────────────────────────────────────
|
||||
|
||||
def process_file(
|
||||
input_path: str,
|
||||
output_jsonl: str,
|
||||
video_id: str,
|
||||
stride: int = STRIDE,
|
||||
tau_low: float = TAU_LOW,
|
||||
tau_high: float = TAU_HIGH,
|
||||
) -> int:
|
||||
records = []
|
||||
|
||||
with open(input_path, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
|
||||
data = json.loads(line)
|
||||
window_id = int(data["window_id"])
|
||||
|
||||
if window_id % stride != 0:
|
||||
continue
|
||||
|
||||
states = np.array(data["states"], dtype=np.float32)
|
||||
|
||||
if states.shape != (WINDOW_FRAMES, N_JOINTS, 3):
|
||||
continue
|
||||
if np.isnan(states).any():
|
||||
continue
|
||||
|
||||
token_str, cp_counts = build_token_str(states, TARGET_FPS, tau_low, tau_high)
|
||||
|
||||
records.append({
|
||||
"video_id": video_id,
|
||||
"window_id": window_id,
|
||||
"fps": TARGET_FPS,
|
||||
"token_str": token_str,
|
||||
"cp_counts": cp_counts,
|
||||
})
|
||||
|
||||
if not records:
|
||||
return 0
|
||||
|
||||
tmp = output_jsonl + ".tmp"
|
||||
os.makedirs(os.path.dirname(output_jsonl), exist_ok=True)
|
||||
with open(tmp, "w", encoding="utf-8") as f:
|
||||
for rec in records:
|
||||
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
|
||||
os.replace(tmp, output_jsonl)
|
||||
|
||||
return len(records)
|
||||
|
||||
|
||||
# ── Entry point ───────────────────────────────────────────────────────────────
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(
|
||||
description="Phase 5 — Adaptive PCHIP per-joint tokenizer."
|
||||
)
|
||||
p.add_argument("--input-dir", required=True,
|
||||
help="Directory with *_cleaned.jsonl from Phase 4.")
|
||||
p.add_argument("--output-dir", required=True,
|
||||
help="Directory to write *_tokens.jsonl files.")
|
||||
p.add_argument("--stride", type=int, default=STRIDE,
|
||||
help=f"Keep windows where window_id %% stride == 0. Default: {STRIDE}")
|
||||
p.add_argument("--tau-low", type=float, default=TAU_LOW,
|
||||
help=f"Curvature threshold for 2-CP tier. Default: {TAU_LOW}")
|
||||
p.add_argument("--tau-high", type=float, default=TAU_HIGH,
|
||||
help=f"Curvature threshold for 8-CP tier. Default: {TAU_HIGH}")
|
||||
p.add_argument("--file-list", default=None,
|
||||
help="Optional text file listing specific *_cleaned.jsonl paths.")
|
||||
p.add_argument("--window-frames", type=int, default=WINDOW_FRAMES,
|
||||
help=f"Must match Phase 3/4's --window-size. Default: {WINDOW_FRAMES}. "
|
||||
f"2026-07-22: use 24 to match the wider cosmos chunk window -- "
|
||||
f"see REPORT.md #38. MAX_CPS (top-tier control-point count) is "
|
||||
f"NOT tied to this and stays 8 either way.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
global WINDOW_FRAMES
|
||||
args = parse_args()
|
||||
WINDOW_FRAMES = args.window_frames
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
if args.file_list:
|
||||
with open(args.file_list) as f:
|
||||
all_files = [l.strip() for l in f if l.strip()]
|
||||
else:
|
||||
all_files = sorted(glob.glob(os.path.join(args.input_dir, "*_cleaned.jsonl")))
|
||||
|
||||
task_id = int(os.environ.get("SLURM_ARRAY_TASK_ID", 0))
|
||||
num_tasks = int(os.environ.get("SLURM_ARRAY_TASK_COUNT", 1))
|
||||
|
||||
my_files = [f for i, f in enumerate(all_files) if i % num_tasks == task_id]
|
||||
total = len(my_files)
|
||||
|
||||
print(f"\n[Worker {task_id}/{num_tasks}] {total} files to process.")
|
||||
print("=" * 60)
|
||||
|
||||
processed = skipped = empty = 0
|
||||
tier_counts = {2: 0, 4: 0, 8: 0}
|
||||
|
||||
for idx, input_path in enumerate(my_files, start=1):
|
||||
base = os.path.basename(input_path)
|
||||
video_id = base[: -len("_cleaned.jsonl")]
|
||||
|
||||
out_jsonl = os.path.join(args.output_dir, f"{video_id}_tokens.jsonl")
|
||||
|
||||
if os.path.exists(out_jsonl):
|
||||
skipped += 1
|
||||
print(f"[{idx}/{total}] {video_id} — already done", end="\r")
|
||||
continue
|
||||
|
||||
try:
|
||||
n = process_file(
|
||||
input_path, out_jsonl, video_id,
|
||||
stride=args.stride,
|
||||
tau_low=args.tau_low,
|
||||
tau_high=args.tau_high,
|
||||
)
|
||||
if n > 0:
|
||||
processed += 1
|
||||
pct = (processed + skipped + empty) / total * 100
|
||||
print(f"[{idx}/{total}] {pct:.1f}% | {video_id} — {n} windows")
|
||||
else:
|
||||
empty += 1
|
||||
except Exception as e:
|
||||
print(f"[{idx}/{total}] ERROR {video_id} — {e}")
|
||||
for p in (out_jsonl + ".tmp",):
|
||||
if os.path.exists(p):
|
||||
os.remove(p)
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print(f"[Worker {task_id}] done — processed: {processed}, skipped: {skipped}, empty: {empty}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
582
pipeline_pose/snac_finevideo.py
Normal file
582
pipeline_pose/snac_finevideo.py
Normal file
@@ -0,0 +1,582 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
SNAC tokenization for FineVideo-VLA activities.
|
||||
|
||||
Reads final_dataset_adaptive JSONL files, extracts audio from .mp4 videos,
|
||||
and tokenizes each activity segment with SNAC_24kHz in "listen" format.
|
||||
|
||||
Listen format — 3 tokens per SNAC base frame (base rate = 12.5 Hz → 37.5 tokens/sec):
|
||||
token_1 = codes[0][i] + 128266 → <snac_128266> .. <snac_132361>
|
||||
token_2 = codes[1][2*i] + 128266 + 4096 → <snac_132362> .. <snac_136457>
|
||||
token_3 = codes[1][2*i+1] + 128266 +16384 → <snac_144650> .. <snac_148745>
|
||||
|
||||
Same Orpheus offset scheme as MixtureVitae-Omni → tokens are directly compatible.
|
||||
Total unique SNAC token strings: 3 × 4096 = 12,288.
|
||||
|
||||
Output: {OUTPUT_DIR}/{video_id}_snac.jsonl
|
||||
One line per activity:
|
||||
{
|
||||
"video_id":"...", "activity_id":"...", "start_sec":1.0, "end_sec":8.9,
|
||||
"has_agent": true,
|
||||
"snac_by_chunk": {
|
||||
"0": ["<snac_130055>", "<snac_133001>", "<snac_145000>", ...], // ~9-10 tokens
|
||||
"1": [...],
|
||||
...
|
||||
}
|
||||
}
|
||||
|
||||
snac_by_chunk keys are chunk_idx (integer as string), aligned to the same 8-frame
|
||||
grid as cosmos/avclm/agent. Phase7 reads chunk_idx → snac tokens directly.
|
||||
|
||||
Two modes:
|
||||
--build-tasks Scan all final_dataset_adaptive files, write snac_task_list.json.
|
||||
Run once on login node (or task 0) before the array job.
|
||||
(default) Load task list, process this SLURM task's slice of videos.
|
||||
|
||||
SLURM usage:
|
||||
SLURM_ARRAY_TASK_ID = task index (0-based)
|
||||
SLURM_ARRAY_TASK_COUNT = total number of tasks in the array
|
||||
|
||||
Local test:
|
||||
python pipeline_pose/snac_finevideo.py --build-tasks
|
||||
python pipeline_pose/snac_finevideo.py # task_id=0, num_tasks=1
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import multiprocessing
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
# ── Paths ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
VIDEO_DIR = "/e/data1/datasets/playground/mmlaion/shared/nguyen38/videos_staging"
|
||||
# Still scanned for activity time_range_sec + has_agent -- those don't depend
|
||||
# on window size. chunk_timing (per-window breakdown) is IGNORED (see
|
||||
# _scan_one_rank_file docstring) because this file is window=8-based (the
|
||||
# pre-pivot merge) and Phase 6 hasn't rerun at window=24 yet; n_chunks is
|
||||
# instead recomputed independently from CHUNK_SIZE below, same pattern as
|
||||
# data_prep/omnivideo_100k/snac_omnivideo.py.
|
||||
INPUT_GLOB = ("/e/data1/datasets/playground/mmlaion/shared/nguyen38/FineVideo-VLA/"
|
||||
"final_dataset_adaptive/final_vla_adaptive_rank_*.jsonl")
|
||||
OUTPUT_DIR = ("/e/data1/datasets/playground/mmlaion/shared/nguyen38/FineVideo-VLA/snac_tokens_w24")
|
||||
TASK_CACHE = ("/e/data1/datasets/playground/mmlaion/shared/nguyen38/FineVideo-VLA/"
|
||||
"snac_task_list_w24.json")
|
||||
HF_CACHE = "/e/project1/reformo/nguyen38/jupiter_cache/huggingface"
|
||||
SNAC_MODEL = "hubertsiuzdak/snac_24khz"
|
||||
SAMPLE_RATE = 24000
|
||||
TARGET_FPS = 30
|
||||
CHUNK_SIZE = 24 # 2026-07-23 window=24 pivot -- must match step_a_tokenize_video.py's CHUNK_SIZE
|
||||
|
||||
# ── SNAC listen-format offsets (matches MixtureVitae-Omni) ───────────────────
|
||||
|
||||
OFFSET_L0 = 128266 # codes[0] base
|
||||
OFFSET_L1A = 128266 + 4096 # codes[1] even frames → 132362
|
||||
OFFSET_L1B = 128266 + 4 * 4096 # codes[1] odd frames → 144650
|
||||
|
||||
# 2026-07-23: speak-format offsets for codes[2] (fine, 50Hz, 4 sub-positions
|
||||
# per base frame) -- corrected to match the REAL scheme in Huu/Chien's
|
||||
# production snac_gpu.py on Leonardo (pipeline_video/snac_gpu.py), the
|
||||
# Orpheus-standard SNAC packing layout. Sub-codes 0/1 sit between L1A and
|
||||
# L1B (in what was wrongly assumed to be an unused gap); sub-codes 2/3 sit
|
||||
# after L1B. See data_prep/laion_emotional_roleplay/tokenize_snac.py's
|
||||
# OFFSET_L2 docstring for the full correction history. Must match exactly.
|
||||
OFFSET_L2 = [136458, 140554, 148746, 152842]
|
||||
|
||||
# ── Logging ───────────────────────────────────────────────────────────────────
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s %(levelname)s %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
stream=sys.stdout,
|
||||
)
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Audio extraction
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
def _find_ffmpeg() -> str:
|
||||
"""Return path to ffmpeg binary, trying imageio_ffmpeg as fallback."""
|
||||
try:
|
||||
subprocess.run(["ffmpeg", "-version"], capture_output=True, check=True)
|
||||
return "ffmpeg"
|
||||
except (FileNotFoundError, subprocess.CalledProcessError):
|
||||
pass
|
||||
try:
|
||||
import imageio_ffmpeg
|
||||
return imageio_ffmpeg.get_ffmpeg_exe()
|
||||
except ImportError:
|
||||
pass
|
||||
raise RuntimeError("ffmpeg not found. Install ffmpeg or imageio_ffmpeg.")
|
||||
|
||||
_FFMPEG = None
|
||||
|
||||
def get_ffmpeg() -> str:
|
||||
global _FFMPEG
|
||||
if _FFMPEG is None:
|
||||
_FFMPEG = _find_ffmpeg()
|
||||
return _FFMPEG
|
||||
|
||||
|
||||
def extract_full_audio(video_path: str) -> np.ndarray | None:
|
||||
"""
|
||||
Extract full mono 24 kHz PCM audio from a video file.
|
||||
Pipes raw float32 PCM from ffmpeg directly to a numpy array — no temp files.
|
||||
Returns float32 array or None on any failure.
|
||||
"""
|
||||
cmd = [
|
||||
get_ffmpeg(), "-y",
|
||||
"-i", video_path,
|
||||
"-vn", # strip video
|
||||
"-ac", "1", # mono
|
||||
"-ar", str(SAMPLE_RATE), # 24 kHz
|
||||
"-f", "f32le", # raw float32 PCM
|
||||
"-",
|
||||
]
|
||||
try:
|
||||
result = subprocess.run(cmd, capture_output=True, timeout=300)
|
||||
if result.returncode != 0 or not result.stdout:
|
||||
return None
|
||||
audio = np.frombuffer(result.stdout, dtype=np.float32).copy()
|
||||
return audio if len(audio) > 0 else None
|
||||
except Exception as e:
|
||||
log.debug(f"ffmpeg failed for {video_path}: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def slice_audio(
|
||||
audio: np.ndarray,
|
||||
start_sec: float,
|
||||
end_sec: float,
|
||||
sr: int = SAMPLE_RATE,
|
||||
) -> np.ndarray:
|
||||
"""Slice a float32 audio array to [start_sec, end_sec]."""
|
||||
s = max(0, int(start_sec * sr))
|
||||
e = min(len(audio), int(end_sec * sr))
|
||||
return audio[s:e]
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# SNAC tokenization
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
def encode_listen(audio: np.ndarray, model, device: str) -> list[str]:
|
||||
"""
|
||||
Encode a float32 audio array with SNAC_24kHz, return listen-format tokens.
|
||||
|
||||
Listen format (3 tokens per base frame):
|
||||
<snac_{codes[0][i] + 128266}>
|
||||
<snac_{codes[1][2i] + 132362}>
|
||||
<snac_{codes[1][2i+1] + 144650}>
|
||||
|
||||
SNAC_24kHz hierarchy:
|
||||
codes[0] — base codebook, 12.5 Hz
|
||||
codes[1] — mid codebook, 25.0 Hz (2× codes[0])
|
||||
codes[2] — fine codebook, 50.0 Hz (4× codes[0], not used in listen)
|
||||
|
||||
Listen format ignores codes[2] (fine detail) to keep token count low
|
||||
(~37.5 tokens/sec vs 87.5 for full speak format). Matches MV-Omni.
|
||||
"""
|
||||
tensor = torch.from_numpy(audio).unsqueeze(0).unsqueeze(0).to(device) # (1,1,T)
|
||||
with torch.inference_mode():
|
||||
codes = model.encode(tensor) # list: [codes[0], codes[1], codes[2]]
|
||||
|
||||
c0 = codes[0] # (1, N0)
|
||||
c1 = codes[1] # (1, N1), N1 == 2*N0
|
||||
|
||||
n0 = c0.shape[1]
|
||||
tokens: list[str] = []
|
||||
for i in range(n0):
|
||||
i1a = 2 * i
|
||||
i1b = 2 * i + 1
|
||||
if i1b >= c1.shape[1]:
|
||||
break # boundary guard: shouldn't happen for valid audio
|
||||
tokens.append(f"<snac_{c0[0, i].item() + OFFSET_L0}>")
|
||||
tokens.append(f"<snac_{c1[0, i1a].item() + OFFSET_L1A}>")
|
||||
tokens.append(f"<snac_{c1[0, i1b].item() + OFFSET_L1B}>")
|
||||
return tokens
|
||||
|
||||
|
||||
def encode_speak(audio: np.ndarray, model, device: str) -> list[str]:
|
||||
"""
|
||||
Encode a float32 audio array with SNAC_24kHz, return full speak-format
|
||||
tokens (7 tokens per base frame, +133% vs encode_listen()) -- 2026-07-22
|
||||
(REPORT.md #37), decided after a real audio A/B
|
||||
(tools/snac_l2_experiment.py) showed audibly better reconstruction.
|
||||
|
||||
Speak format (2026-07-23, Leo-matched order -- L2 interleaved between
|
||||
L1a and L1b, not appended after):
|
||||
<snac_{codes[0][i] + OFFSET_L0}>
|
||||
<snac_{codes[1][2i] + OFFSET_L1A}>
|
||||
<snac_{codes[2][4i] + OFFSET_L2[0]}>
|
||||
<snac_{codes[2][4i+1] + OFFSET_L2[1]}>
|
||||
<snac_{codes[1][2i+1] + OFFSET_L1B}>
|
||||
<snac_{codes[2][4i+2] + OFFSET_L2[2]}>
|
||||
<snac_{codes[2][4i+3] + OFFSET_L2[3]}>
|
||||
"""
|
||||
tensor = torch.from_numpy(audio).unsqueeze(0).unsqueeze(0).to(device)
|
||||
with torch.inference_mode():
|
||||
codes = model.encode(tensor)
|
||||
|
||||
c0, c1, c2 = codes[0], codes[1], codes[2]
|
||||
n0 = c0.shape[1]
|
||||
tokens: list[str] = []
|
||||
for i in range(n0):
|
||||
i1a, i1b = 2 * i, 2 * i + 1
|
||||
i2 = [4 * i + k for k in range(4)]
|
||||
if i1b >= c1.shape[1] or i2[-1] >= c2.shape[1]:
|
||||
break
|
||||
tokens.append(f"<snac_{c0[0, i].item() + OFFSET_L0}>")
|
||||
tokens.append(f"<snac_{c1[0, i1a].item() + OFFSET_L1A}>")
|
||||
tokens.append(f"<snac_{c2[0, i2[0]].item() + OFFSET_L2[0]}>")
|
||||
tokens.append(f"<snac_{c2[0, i2[1]].item() + OFFSET_L2[1]}>")
|
||||
tokens.append(f"<snac_{c1[0, i1b].item() + OFFSET_L1B}>")
|
||||
tokens.append(f"<snac_{c2[0, i2[2]].item() + OFFSET_L2[2]}>")
|
||||
tokens.append(f"<snac_{c2[0, i2[3]].item() + OFFSET_L2[3]}>")
|
||||
return tokens
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Task list building (pre-processing step)
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
def _scan_one_rank_file(fpath: str) -> dict:
|
||||
"""
|
||||
Scan one final_dataset_adaptive rank file for activity time boundaries.
|
||||
Returns {video_id: [activity_dict, ...]} for ALL activities with a valid
|
||||
time_range_sec.
|
||||
|
||||
Each activity dict: activity_id, start_sec, end_sec, has_agent.
|
||||
|
||||
2026-07-23: no longer reads chunk_timing (this file is the pre-window=24-
|
||||
pivot merge; its chunk_timing reflects the OLD 8-frame grid). n_chunks is
|
||||
instead recomputed in process_video() from (end_sec-start_sec) and the
|
||||
current CHUNK_SIZE=24 -- same independent-recompute pattern
|
||||
data_prep/omnivideo_100k/snac_omnivideo.py uses, so this script no longer
|
||||
needs to wait on Phase 6 merge to re-run at window=24 before it can align
|
||||
correctly.
|
||||
|
||||
We tokenize ALL activities (not just agent) because:
|
||||
- Non-agent activities have seed2+cosmos → seed2+cosmos+snac trains modality transitions
|
||||
- Agent-only activities = only 14% of total; skipping the rest wastes 86% of this GPU run
|
||||
"""
|
||||
tasks: dict = {}
|
||||
try:
|
||||
with open(fpath, "r", errors="replace") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
rec = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
vid = rec.get("video_id", "")
|
||||
if not vid:
|
||||
continue
|
||||
for scene in rec.get("scenes", []):
|
||||
for act in scene.get("activities", []):
|
||||
tr = act.get("time_range_sec")
|
||||
if not tr or len(tr) < 2:
|
||||
continue
|
||||
has_agent = "<agent>" in act.get("video_tokens", "")
|
||||
tasks.setdefault(vid, []).append({
|
||||
"activity_id": act.get("activity_id", ""),
|
||||
"start_sec": float(tr[0]),
|
||||
"end_sec": float(tr[1]),
|
||||
"has_agent": has_agent,
|
||||
})
|
||||
except Exception as e:
|
||||
log.warning(f"Error scanning {fpath}: {e}")
|
||||
return tasks
|
||||
|
||||
|
||||
def build_task_list(input_glob: str, cache_path: str, workers: int = 8) -> dict:
|
||||
"""
|
||||
Scan all final_dataset_adaptive rank files in parallel to build a task list.
|
||||
Saves result to cache_path as JSON.
|
||||
Returns {video_id: [activity_dicts]}.
|
||||
|
||||
This is I/O-heavy (~657 GB total) — use multiprocessing to parallelize.
|
||||
Estimated wall time: 5–15 min with 8 workers on shared filesystem.
|
||||
"""
|
||||
rank_files = sorted(glob.glob(input_glob))
|
||||
if not rank_files:
|
||||
raise FileNotFoundError(f"No files matched: {input_glob}")
|
||||
log.info(f"Scanning {len(rank_files)} rank files with {workers} workers...")
|
||||
t0 = time.time()
|
||||
|
||||
with multiprocessing.Pool(workers) as pool:
|
||||
partial_results = pool.map(_scan_one_rank_file, rank_files)
|
||||
|
||||
# merge
|
||||
all_tasks: dict = {}
|
||||
for partial in partial_results:
|
||||
for vid, acts in partial.items():
|
||||
all_tasks.setdefault(vid, []).extend(acts)
|
||||
|
||||
# deduplicate activities by activity_id (in case of overlap across ranks)
|
||||
for vid in all_tasks:
|
||||
seen = set()
|
||||
deduped = []
|
||||
for act in all_tasks[vid]:
|
||||
key = act["activity_id"]
|
||||
if key not in seen:
|
||||
seen.add(key)
|
||||
deduped.append(act)
|
||||
all_tasks[vid] = deduped
|
||||
|
||||
log.info(
|
||||
f"Task list built: {len(all_tasks)} videos, "
|
||||
f"{sum(len(v) for v in all_tasks.values())} activities "
|
||||
f"({time.time()-t0:.0f}s)"
|
||||
)
|
||||
os.makedirs(os.path.dirname(cache_path), exist_ok=True)
|
||||
with open(cache_path, "w") as f:
|
||||
json.dump(all_tasks, f)
|
||||
log.info(f"Saved task list → {cache_path}")
|
||||
return all_tasks
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Chunk alignment
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
def split_snac_by_chunks(tokens: list[str], n_chunks: int) -> dict[int, list[str]]:
|
||||
"""
|
||||
Split a flat SNAC listen token list evenly across n_chunks video chunks.
|
||||
|
||||
Why: SNAC rate (12.5 Hz base × 3 tokens = 37.5 tok/s) does not divide evenly
|
||||
by the video chunk rate (30fps / 8 = 3.75 Hz). Encoding the full activity once
|
||||
preserves audio context; then we split by chunk count snapping to 3-token
|
||||
boundaries (one SNAC base frame = 3 listen tokens).
|
||||
|
||||
Per 8-frame chunk at 30fps (0.267s): ~3.33 SNAC base frames → 9–10 listen tokens.
|
||||
"""
|
||||
n_tokens = len(tokens)
|
||||
n_base = n_tokens // 3 # truncate to complete base frames
|
||||
tokens = tokens[:n_base * 3]
|
||||
|
||||
result: dict[int, list[str]] = {}
|
||||
for k in range(n_chunks):
|
||||
start_base = round(k * n_base / n_chunks)
|
||||
end_base = round((k + 1) * n_base / n_chunks)
|
||||
result[k] = tokens[start_base * 3 : end_base * 3]
|
||||
return result
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Per-video processing
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
def process_video(
|
||||
video_id: str,
|
||||
activities: list[dict],
|
||||
model,
|
||||
device: str,
|
||||
video_dir: str,
|
||||
output_dir: str,
|
||||
skip_existing: bool,
|
||||
encode_fn=encode_listen,
|
||||
) -> dict:
|
||||
"""
|
||||
Tokenize all activities for one video.
|
||||
|
||||
Steps:
|
||||
1. Check skip — if output file exists and skip_existing, return immediately.
|
||||
2. Extract full audio from .mp4 once (1 ffmpeg call per video).
|
||||
3. For each activity: slice audio by time_range_sec, run SNAC encode once,
|
||||
then split the flat token list across chunks (preserving audio context).
|
||||
4. Write all results to {output_dir}/{video_id}_snac.jsonl.
|
||||
|
||||
Output per activity: snac_by_chunk {chunk_idx → [tokens]}
|
||||
Phase6 merge uses this directly to inject SNAC tokens per 24-frame chunk,
|
||||
aligned with the cosmos/avclm/agent tokens that fire at the same chunk
|
||||
boundaries (n_chunks recomputed from CHUNK_SIZE, not read from a file).
|
||||
|
||||
Returns stats: {ok, skipped_vid, failed_audio, failed_snac, tokens}
|
||||
"""
|
||||
out_path = os.path.join(output_dir, f"{video_id}_snac.jsonl")
|
||||
if skip_existing and os.path.exists(out_path):
|
||||
return {"ok": 0, "skipped_vid": len(activities), "failed_audio": 0,
|
||||
"failed_snac": 0, "tokens": 0}
|
||||
|
||||
video_path = os.path.join(video_dir, f"{video_id}.mp4")
|
||||
if not os.path.exists(video_path):
|
||||
return {"ok": 0, "skipped_vid": 0, "failed_audio": len(activities),
|
||||
"failed_snac": 0, "tokens": 0}
|
||||
|
||||
# Extract full audio once
|
||||
full_audio = extract_full_audio(video_path)
|
||||
if full_audio is None:
|
||||
log.warning(f"No audio: {video_path}")
|
||||
return {"ok": 0, "skipped_vid": 0, "failed_audio": len(activities),
|
||||
"failed_snac": 0, "tokens": 0}
|
||||
|
||||
stats = {"ok": 0, "skipped_vid": 0, "failed_audio": 0, "failed_snac": 0, "tokens": 0}
|
||||
rows = []
|
||||
|
||||
for act in activities:
|
||||
segment = slice_audio(full_audio, act["start_sec"], act["end_sec"])
|
||||
if len(segment) < int(SAMPLE_RATE * 0.1): # skip segments < 100 ms
|
||||
stats["failed_audio"] += 1
|
||||
continue
|
||||
try:
|
||||
flat_tokens = encode_fn(segment, model, device)
|
||||
except Exception as e:
|
||||
log.warning(f"SNAC failed {video_id}/{act['activity_id']}: {e}")
|
||||
stats["failed_snac"] += 1
|
||||
continue
|
||||
if not flat_tokens:
|
||||
stats["failed_snac"] += 1
|
||||
continue
|
||||
|
||||
# Split flat token list into per-chunk dicts, aligned to the same
|
||||
# 24-frame grid as cosmos/avclm/agent. n_chunks recomputed
|
||||
# independently from CHUNK_SIZE (not from a merged file's
|
||||
# chunk_timing -- see _scan_one_rank_file docstring).
|
||||
total_frames = max(1, round((act["end_sec"] - act["start_sec"]) * TARGET_FPS))
|
||||
n_chunks = math.ceil(total_frames / CHUNK_SIZE)
|
||||
by_chunk = split_snac_by_chunks(flat_tokens, n_chunks)
|
||||
snac_by_chunk = {str(k): v for k, v in by_chunk.items()}
|
||||
|
||||
rows.append({
|
||||
"video_id": video_id,
|
||||
"activity_id": act["activity_id"],
|
||||
"start_sec": round(act["start_sec"], 4),
|
||||
"end_sec": round(act["end_sec"], 4),
|
||||
"has_agent": act.get("has_agent", False),
|
||||
"snac_by_chunk": snac_by_chunk,
|
||||
})
|
||||
stats["ok"] += 1
|
||||
stats["tokens"] += len(flat_tokens)
|
||||
|
||||
if rows:
|
||||
with open(out_path, "w") as f:
|
||||
for row in rows:
|
||||
f.write(json.dumps(row) + "\n")
|
||||
|
||||
return stats
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Main
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
def parse_args():
|
||||
p = argparse.ArgumentParser(description="SNAC tokenization for FineVideo-VLA")
|
||||
p.add_argument("--build-tasks", action="store_true",
|
||||
help="Scan final_dataset_adaptive and write snac_task_list.json, then exit.")
|
||||
p.add_argument("--input-glob", default=INPUT_GLOB)
|
||||
p.add_argument("--output-dir", default=OUTPUT_DIR)
|
||||
p.add_argument("--video-dir", default=VIDEO_DIR)
|
||||
p.add_argument("--task-cache", default=TASK_CACHE)
|
||||
p.add_argument("--hf-cache", default=HF_CACHE)
|
||||
p.add_argument("--scan-workers", type=int, default=8,
|
||||
help="CPU workers for --build-tasks scan (default 8)")
|
||||
p.add_argument("--no-skip", action="store_true",
|
||||
help="Re-process videos even if output file exists")
|
||||
p.add_argument("--format", choices=["listen", "speak"], default="listen",
|
||||
help="listen = L0+L1 only (current production, 3 tok/base-frame); "
|
||||
"speak = full L0+L1+L2 (2026-07-22, 7 tok/base-frame, +133%% tokens). "
|
||||
"speak requires L2 tokens added to the tokenizer vocab first.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
skip_existing = not args.no_skip
|
||||
|
||||
# ── Set HF cache ─────────────────────────────────────────────────────────
|
||||
os.environ.setdefault("HF_HOME", args.hf_cache)
|
||||
os.makedirs(args.hf_cache, exist_ok=True)
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# ── Mode: build task list ─────────────────────────────────────────────────
|
||||
if args.build_tasks:
|
||||
build_task_list(args.input_glob, args.task_cache, workers=args.scan_workers)
|
||||
return
|
||||
|
||||
# ── Mode: tokenize ────────────────────────────────────────────────────────
|
||||
|
||||
# SLURM array vars
|
||||
task_id = int(os.environ.get("SLURM_ARRAY_TASK_ID", "0"))
|
||||
num_tasks = int(os.environ.get("SLURM_ARRAY_TASK_COUNT", "1"))
|
||||
|
||||
# Load task list (must exist — run --build-tasks first)
|
||||
if not os.path.exists(args.task_cache):
|
||||
log.error(
|
||||
f"Task list not found: {args.task_cache}\n"
|
||||
f"Run first: python pipeline_pose/snac_finevideo.py --build-tasks"
|
||||
)
|
||||
sys.exit(1)
|
||||
with open(args.task_cache) as f:
|
||||
all_tasks = json.load(f)
|
||||
|
||||
all_vids = sorted(all_tasks.keys())
|
||||
my_vids = all_vids[task_id::num_tasks]
|
||||
log.info(
|
||||
f"Task {task_id}/{num_tasks}: {len(my_vids)}/{len(all_vids)} videos "
|
||||
f"skip_existing={skip_existing}"
|
||||
)
|
||||
|
||||
# ── Load SNAC model ───────────────────────────────────────────────────────
|
||||
from snac import SNAC # imported here to avoid load cost during --build-tasks
|
||||
device = "cuda:0" if torch.cuda.is_available() else "cpu"
|
||||
log.info(f"Loading SNAC model ({SNAC_MODEL}) on {device}...")
|
||||
t_load = time.time()
|
||||
os.environ["HF_HOME"] = args.hf_cache # ensure hub/ subdir is found
|
||||
model = SNAC.from_pretrained(SNAC_MODEL,
|
||||
local_files_only=True).eval().to(device)
|
||||
log.info(f"SNAC loaded ({time.time()-t_load:.1f}s)")
|
||||
|
||||
# ── Process videos ────────────────────────────────────────────────────────
|
||||
cumul = {"ok": 0, "skipped_vid": 0, "failed_audio": 0, "failed_snac": 0, "tokens": 0}
|
||||
t_start = time.time()
|
||||
|
||||
encode_fn = encode_speak if args.format == "speak" else encode_listen
|
||||
|
||||
for idx, vid in enumerate(my_vids, 1):
|
||||
s = process_video(
|
||||
vid, all_tasks[vid], model, device,
|
||||
args.video_dir, args.output_dir, skip_existing,
|
||||
encode_fn=encode_fn,
|
||||
)
|
||||
for k in cumul:
|
||||
cumul[k] += s[k]
|
||||
|
||||
if idx % 100 == 0 or idx == len(my_vids):
|
||||
elapsed = time.time() - t_start
|
||||
rate = idx / elapsed
|
||||
eta = (len(my_vids) - idx) / rate if rate > 0 else 0
|
||||
log.info(
|
||||
f"[{idx:5d}/{len(my_vids)}] vid={vid} "
|
||||
f"ok={s['ok']} skip={s['skipped_vid']} "
|
||||
f"fail_audio={s['failed_audio']} fail_snac={s['failed_snac']} "
|
||||
f"rate={rate:.1f}vid/s ETA={eta/60:.0f}m "
|
||||
f"total_tokens={cumul['tokens']:,}"
|
||||
)
|
||||
|
||||
elapsed = time.time() - t_start
|
||||
log.info(
|
||||
f"DONE task {task_id}: "
|
||||
f"ok={cumul['ok']} skipped={cumul['skipped_vid']} "
|
||||
f"fail_audio={cumul['failed_audio']} fail_snac={cumul['failed_snac']} "
|
||||
f"tokens={cumul['tokens']:,} wall={elapsed:.0f}s"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
3
tokenizer.json
Normal file
3
tokenizer.json
Normal file
@@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:d97a3c88b39e2dffeee8d470d767d48e2777c0d1395aaeee32b3709cb35e68a5
|
||||
size 31655733
|
||||
30
tokenizer_config.json
Normal file
30
tokenizer_config.json
Normal file
@@ -0,0 +1,30 @@
|
||||
{
|
||||
"add_prefix_space": false,
|
||||
"backend": "tokenizers",
|
||||
"bos_token": null,
|
||||
"clean_up_tokenization_spaces": false,
|
||||
"eos_token": "<|im_end|>",
|
||||
"errors": "replace",
|
||||
"extra_special_tokens": [
|
||||
"<|im_start|>",
|
||||
"<|im_end|>",
|
||||
"<|object_ref_start|>",
|
||||
"<|object_ref_end|>",
|
||||
"<|box_start|>",
|
||||
"<|box_end|>",
|
||||
"<|quad_start|>",
|
||||
"<|quad_end|>",
|
||||
"<|vision_start|>",
|
||||
"<|vision_end|>",
|
||||
"<|vision_pad|>",
|
||||
"<|image_pad|>",
|
||||
"<|video_pad|>"
|
||||
],
|
||||
"is_local": true,
|
||||
"local_files_only": false,
|
||||
"model_max_length": 1010000,
|
||||
"pad_token": "<|endoftext|>",
|
||||
"split_special_tokens": false,
|
||||
"tokenizer_class": "Qwen2Tokenizer",
|
||||
"unk_token": null
|
||||
}
|
||||
202
tools/decode/decode_cosmos.py
Normal file
202
tools/decode/decode_cosmos.py
Normal file
@@ -0,0 +1,202 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
General-purpose Cosmos detokenizer -- turns `<cosmos_N>` tokens (from any of
|
||||
our flattened datasets: FineVideo-VLA, OmniVideo-100K, ...) back into an
|
||||
actual video.
|
||||
|
||||
Verified 20/07/2026 (see PROGRESS_VI.md entry) against a real 200-token
|
||||
chunk from omnivideo_100k_final: for the DV8x16x16 checkpoint, encode() of
|
||||
an 8-frame/160x160 chunk produces indices of shape (1, 2, 10, 10) == 200
|
||||
flat tokens -- this is CHUNK_TOKENS/CHUNK_GRID below, hardcoded for this
|
||||
checkpoint (would need updating if a different Cosmos-Tokenizer-* variant
|
||||
is ever used).
|
||||
|
||||
This is a lossy, generative-but-deterministic *reconstruction* (the neural
|
||||
decoder's output, not the original pixels) -- expect visible blur/artifacts,
|
||||
this is the accepted tradeoff for a much lower token count than storing raw
|
||||
pixels. Contrast with decode_avclm.py in this same directory, which is
|
||||
byte-exact (BPE over raw H.264 bytes, not a neural codec).
|
||||
|
||||
Token values in `<cosmos_N>` are the RAW encoder codebook indices with no
|
||||
vocab offset added (verified: flatten scripts across the project just wrap
|
||||
the same digit string in `<cosmos_{n}>`, never add an offset) -- so N can be
|
||||
fed straight into decode() after reshaping.
|
||||
|
||||
Usage:
|
||||
# From a raw list of ints:
|
||||
python tools/decode/decode_cosmos.py --tokens 18697,55801,44451,... --output out.mp4
|
||||
|
||||
# From a flattened JSONL record (extracts the Nth 200-token chunk):
|
||||
python tools/decode/decode_cosmos.py --input-jsonl path.jsonl --record-id VIDEO_ID \
|
||||
--chunk-index 0 --output out.mp4
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
# 2026-07-23: PROTOTYPE_DIR only exists on the internal cluster (not part of
|
||||
# the public github.com/TieuDaoChanNhan/finevideo-vla repo -- 0 files under
|
||||
# prototype/ are git-tracked). External users have no way to reach a
|
||||
# hardcoded local path, so the checkpoint is now fetched from NVIDIA's public
|
||||
# HF repo on first use instead (cached under the standard HF_HOME, same as
|
||||
# any other from_pretrained() call) -- falls back to the local cluster copy
|
||||
# first if it happens to be present, to avoid a redundant download here.
|
||||
PROTOTYPE_DIR = "/e/project1/reformo/nguyen38/prototype"
|
||||
_LOCAL_CHECKPOINT_DEC = os.path.join(
|
||||
PROTOTYPE_DIR, "pretrained_ckpts/Cosmos-Tokenizer-DV8x16x16/decoder.jit"
|
||||
)
|
||||
COSMOS_HF_REPO = "nvidia/Cosmos-Tokenizer-DV8x16x16"
|
||||
CHUNK_GRID = (2, 10, 10) # (T', H', W') per 8-frame/160x160 input chunk, this checkpoint
|
||||
CHUNK_TOKENS = CHUNK_GRID[0] * CHUNK_GRID[1] * CHUNK_GRID[2] # 200
|
||||
|
||||
|
||||
def _resolve_checkpoint_dec() -> str:
|
||||
if os.path.exists(_LOCAL_CHECKPOINT_DEC):
|
||||
return _LOCAL_CHECKPOINT_DEC
|
||||
from huggingface_hub import hf_hub_download
|
||||
print(f"Local checkpoint not found -- downloading decoder.jit from {COSMOS_HF_REPO} "
|
||||
f"(~350MB, cached for future runs)...")
|
||||
return hf_hub_download(repo_id=COSMOS_HF_REPO, filename="decoder.jit")
|
||||
|
||||
_COSMOS_ATOMIC_RE = re.compile(r"<cosmos_(\d+)>")
|
||||
_COSMOS_RAW_BLOCK_RE = re.compile(r"<cosmos>(.*?)</cosmos>", re.DOTALL)
|
||||
|
||||
|
||||
def extract_chunk_tokens(text: str, chunk_index: int) -> list:
|
||||
"""Pull out the Nth CHUNK_TOKENS-sized slice of cosmos ids from a record.
|
||||
Supports both formats found in this project:
|
||||
- flattened/atomic: `<cosmos_N> <cosmos_N> ...` (post-flatten output,
|
||||
e.g. omnivideo_100k_final, FineVideo's megatron_dataset_*)
|
||||
- raw pre-flatten block: `<cosmos>N N N...</cosmos>` (Step A's own
|
||||
output before any flatten script runs, e.g. FineVideo's
|
||||
training_ready_rank_*.jsonl activity.video_tokens, or OmniVideo's
|
||||
omnivideo_100k_video_flat) -- chunk_index selects which <cosmos>
|
||||
block (each block is already exactly one chunk, verified 200 tokens
|
||||
both for FineVideo and OmniVideo real data).
|
||||
|
||||
In the flattened/atomic format, chunks with <50% cosmos dropout keep-rate
|
||||
have gaps (some chunks entirely missing cosmos), so chunk_index there
|
||||
means "Nth cosmos chunk present in the stream", not "Nth temporal chunk
|
||||
of the video" -- the two only coincide if dropout happened to keep every
|
||||
chunk up to that point.
|
||||
"""
|
||||
raw_blocks = _COSMOS_RAW_BLOCK_RE.findall(text)
|
||||
if raw_blocks:
|
||||
if chunk_index >= len(raw_blocks):
|
||||
raise ValueError(f"Requested chunk {chunk_index} but only {len(raw_blocks)} <cosmos> blocks in this record")
|
||||
chunk = [int(x) for x in raw_blocks[chunk_index].split() if x.isdigit()]
|
||||
if len(chunk) != CHUNK_TOKENS:
|
||||
raise ValueError(f"<cosmos> block {chunk_index} has {len(chunk)} tokens, expected {CHUNK_TOKENS}")
|
||||
return chunk
|
||||
|
||||
all_ids = [int(x) for x in _COSMOS_ATOMIC_RE.findall(text)]
|
||||
start = chunk_index * CHUNK_TOKENS
|
||||
end = start + CHUNK_TOKENS
|
||||
chunk = all_ids[start:end]
|
||||
if len(chunk) != CHUNK_TOKENS:
|
||||
raise ValueError(
|
||||
f"Requested chunk {chunk_index} needs tokens [{start}:{end}) but only "
|
||||
f"{len(all_ids)} cosmos tokens total in this record."
|
||||
)
|
||||
return chunk
|
||||
|
||||
|
||||
def _record_text(rec: dict) -> str:
|
||||
"""Flat records (OmniVideo-100K, Megatron-flattened FineVideo): {"text": ...}.
|
||||
Raw FineVideo Step A records (training_ready_rank_*.jsonl) instead nest
|
||||
per-activity `video_tokens` under scenes[].activities[] -- concatenate
|
||||
all of them for the record so chunk_index can walk the whole video."""
|
||||
if "text" in rec:
|
||||
return rec["text"]
|
||||
if "scenes" in rec:
|
||||
return "".join(
|
||||
act.get("video_tokens", "")
|
||||
for scene in rec["scenes"]
|
||||
for act in scene.get("activities", [])
|
||||
)
|
||||
raise KeyError("Record has neither 'text' nor 'scenes' -- unrecognized schema")
|
||||
|
||||
|
||||
def load_tokens_from_jsonl(path: str, record_id: str, chunk_index: int) -> list:
|
||||
with open(path, encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
rec = json.loads(line)
|
||||
rid = rec.get("video_id", rec.get("id"))
|
||||
if rid == record_id:
|
||||
return extract_chunk_tokens(_record_text(rec), chunk_index)
|
||||
raise KeyError(f"record_id={record_id!r} not found in {path}")
|
||||
|
||||
|
||||
def decode_cosmos_chunk(token_ids: list, output_path: str, fps: int = 6) -> None:
|
||||
"""token_ids: exactly CHUNK_TOKENS (200) raw cosmos codebook indices."""
|
||||
if len(token_ids) != CHUNK_TOKENS:
|
||||
raise ValueError(f"Expected exactly {CHUNK_TOKENS} tokens, got {len(token_ids)}")
|
||||
output_path = os.path.abspath(output_path) # resolve before any cwd assumptions below
|
||||
|
||||
# 2026-07-23: cosmos_tokenizer itself is also vendored (tools/decode/vendor/),
|
||||
# not pip-installable -- see vendor/cosmos_tokenizer/NOTICE.md. No os.chdir()
|
||||
# needed anymore (that was only for PROTOTYPE_DIR's relative lookups).
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "vendor"))
|
||||
import imageio_ffmpeg
|
||||
import torch
|
||||
import torchvision.transforms as T
|
||||
from cosmos_tokenizer.video_lib import CausalVideoTokenizer
|
||||
|
||||
checkpoint_dec = _resolve_checkpoint_dec()
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
dec = CausalVideoTokenizer(checkpoint_dec=checkpoint_dec).to(device)
|
||||
|
||||
indices = torch.tensor(token_ids, dtype=torch.int64, device=device).view(1, *CHUNK_GRID)
|
||||
with torch.no_grad():
|
||||
out = dec.decode(indices) # (1, 3, T, H, W), range ~[-1, 1]
|
||||
|
||||
out = ((out.float() + 1.0) / 2.0).clamp(0, 1).squeeze(0) # (3, T, H, W)
|
||||
n_frames = out.shape[1]
|
||||
|
||||
frame_dir = f"/tmp/cosmos_decode_{os.getpid()}"
|
||||
os.makedirs(frame_dir, exist_ok=True)
|
||||
to_pil = T.ToPILImage()
|
||||
for i in range(n_frames):
|
||||
to_pil(out[:, i, :, :].cpu()).save(f"{frame_dir}/frame_{i:02d}.png")
|
||||
|
||||
ffmpeg_bin = imageio_ffmpeg.get_ffmpeg_exe()
|
||||
subprocess.run(
|
||||
[ffmpeg_bin, "-y", "-framerate", str(fps), "-i", f"{frame_dir}/frame_%02d.png",
|
||||
"-vf", "scale=320:320:flags=neighbor", "-pix_fmt", "yuv420p", output_path],
|
||||
check=True, capture_output=True,
|
||||
)
|
||||
for i in range(n_frames):
|
||||
os.remove(f"{frame_dir}/frame_{i:02d}.png")
|
||||
os.rmdir(frame_dir)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--tokens", help="Comma-separated list of exactly 200 raw cosmos ids")
|
||||
ap.add_argument("--input-jsonl", help="Flattened JSONL file to pull tokens from")
|
||||
ap.add_argument("--record-id", help="video_id/id field to select within --input-jsonl")
|
||||
ap.add_argument("--chunk-index", type=int, default=0,
|
||||
help="Which 200-token chunk (in order of appearance) to decode, 0-indexed")
|
||||
ap.add_argument("--fps", type=int, default=6, help="Output mp4 framerate (decoded frames are few, slow fps for visibility)")
|
||||
ap.add_argument("--output", required=True)
|
||||
args = ap.parse_args()
|
||||
|
||||
if args.tokens:
|
||||
token_ids = [int(x) for x in args.tokens.split(",")]
|
||||
elif args.input_jsonl and args.record_id:
|
||||
token_ids = load_tokens_from_jsonl(args.input_jsonl, args.record_id, args.chunk_index)
|
||||
else:
|
||||
ap.error("Provide either --tokens or (--input-jsonl and --record-id)")
|
||||
|
||||
decode_cosmos_chunk(token_ids, args.output, fps=args.fps)
|
||||
print(f"Saved: {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
257
tools/decode/decode_seed2.py
Normal file
257
tools/decode/decode_seed2.py
Normal file
@@ -0,0 +1,257 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Seed2 detokenizer -- turns `<seed2_N>` tokens back into an actual image.
|
||||
|
||||
Unlike decode_cosmos.py/decode_avclm.py (both have a dedicated neural video
|
||||
decoder), Seed2Tokenizer (prototype/seed2/, vendored `ontocord/seed2`) is a
|
||||
Q-Former/BLIP-2-style *understanding* tokenizer -- it has no pixel decoder of
|
||||
its own. Its own README documents the only reconstruction path: look up each
|
||||
token's codebook embedding, then condition a diffusion img2img pipeline
|
||||
(`stabilityai/stable-diffusion-2-1-unclip`) on those embeddings to *generate*
|
||||
a plausible image. This is fundamentally different from cosmos/avc_lm's
|
||||
reconstruction (same lossy-but-deterministic neural codec used to encode) --
|
||||
here the "decode" is itself a generative sample, not a round-trip of the
|
||||
original pixels. Two different tokens can plausibly decode to visually
|
||||
different images even if they'd caption similarly, and the same tokens can
|
||||
decode to slightly different images across runs unless a fixed seed/latent
|
||||
is used (this script fixes the latent, matching the tokenizer's own
|
||||
`self.latents` buffer, so decoding is reproducible run-to-run).
|
||||
|
||||
Token ids are raw codebook indices, 0-8191, no vocab offset (verified:
|
||||
tools/tokenizer/build_tokenizers.py's `<seed2_{i}> for i in range(8192)`
|
||||
and Seed2Tokenizer.num_image_tokens == 8192 agree; unlike snac's +128266
|
||||
offset or cosmos's chunking, seed2 tokens map directly).
|
||||
|
||||
First run downloads `stabilityai/stable-diffusion-2-1-unclip` (~5GB) from
|
||||
HF -- needs internet access (compute nodes reportedly do not have it, see
|
||||
REPORT.md's checkpoint-conversion note; run this on a node that does, or
|
||||
pre-populate $HF_HOME first).
|
||||
|
||||
Usage:
|
||||
# From a raw list of ints:
|
||||
python tools/decode/decode_seed2.py --tokens 6750,2157,4657,... --output out.png
|
||||
|
||||
# From a flattened JSONL / text file containing <seed2_N> tokens anywhere:
|
||||
python tools/decode/decode_seed2.py --text-file path.txt --output out.png
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
|
||||
PROTOTYPE_DIR = "/e/project1/reformo/nguyen38/prototype"
|
||||
# 2026-07-23: prototype/ is NOT part of the public github.com/TieuDaoChanNhan/
|
||||
# finevideo-vla repo (0 files git-tracked), so external users can't reach
|
||||
# _LOCAL_SEED2_DIR below. Found that Seed2Tokenizer's own vendored README
|
||||
# (prototype/seed2/README.md) already documents its true public home --
|
||||
# `git clone https://huggingface.co/ontocord/seed2` -- confirmed for real via
|
||||
# HfApi().model_info("ontocord/seed2"): same 2 checkpoint files
|
||||
# (ae.safetensors, model.safetensors) + seed2_tokenizer.py are already public
|
||||
# there. No redistribution/licensing decision needed on our end (unlike a
|
||||
# vendor-our-own-checkpoint approach) -- just point at the existing repo.
|
||||
SEED2_HF_REPO = "ontocord/seed2"
|
||||
_LOCAL_SEED2_DIR = os.path.join(PROTOTYPE_DIR, "seed2")
|
||||
|
||||
|
||||
def _resolve_seed2_dir() -> str:
|
||||
if os.path.isdir(_LOCAL_SEED2_DIR):
|
||||
return _LOCAL_SEED2_DIR
|
||||
from huggingface_hub import snapshot_download
|
||||
print(f"Local seed2 checkpoint not found -- downloading from {SEED2_HF_REPO} "
|
||||
f"(~2.6GB, cached for future runs)...")
|
||||
return snapshot_download(repo_id=SEED2_HF_REPO)
|
||||
|
||||
|
||||
# stabilityai/stable-diffusion-2-1-unclip returns a genuine 404 (page title
|
||||
# literally "404 - Hugging Face", not a gated-access page) as of 2026-07-22 --
|
||||
# confirmed via HF search API that it no longer appears under the stabilityai
|
||||
# org at all (Stability AI removed it, not access-gated). Using the community
|
||||
# re-upload instead: same weights/pipeline class (StableUnCLIPImg2ImgPipeline,
|
||||
# safetensors, openrail++ license), created 2025-11-14 specifically as a mirror.
|
||||
DIFFUSION_NAME = "sd2-community/stable-diffusion-2-1-unclip"
|
||||
NUM_IMAGE_TOKENS = 8192
|
||||
SEED2_QUERY_LEN = 32 # fixed Q-former query length trained into Seed2Tokenizer, 1 image's worth
|
||||
|
||||
_SEED2_ATOMIC_RE = re.compile(r"<seed2_(\d+)>")
|
||||
_SEED2_BLOCK_RE = re.compile(r"<seed2>(.*?)</seed2>", re.DOTALL)
|
||||
|
||||
|
||||
def extract_seed2_tokens(text: str) -> list:
|
||||
"""Pull every <seed2_N> id inside every <seed2>...</seed2> block, in order,
|
||||
concatenated into one flat list. Falls back to scanning the whole text if
|
||||
no <seed2>...</seed2> wrapper is present. Only meaningful as one image's
|
||||
worth of tokens if the text has exactly one block (or zero, unwrapped) --
|
||||
for multi-block text (e.g. a multi-activity generation), use
|
||||
extract_seed2_blocks() instead so each image decodes separately."""
|
||||
blocks = _SEED2_BLOCK_RE.findall(text)
|
||||
source = " ".join(blocks) if blocks else text
|
||||
return [int(x) for x in _SEED2_ATOMIC_RE.findall(source)]
|
||||
|
||||
|
||||
def extract_seed2_blocks(text: str) -> list:
|
||||
"""Like extract_seed2_tokens(), but keeps each <seed2>...</seed2> block's
|
||||
ids separate (one list per image) instead of concatenating them -- needed
|
||||
because Seed2Tokenizer only ever decodes exactly SEED2_QUERY_LEN=32 tokens
|
||||
as one image (see decode_seed2_tokens()). Falls back to treating the whole
|
||||
text as one block if no wrapper is present."""
|
||||
blocks = _SEED2_BLOCK_RE.findall(text)
|
||||
if not blocks:
|
||||
return [[int(x) for x in _SEED2_ATOMIC_RE.findall(text)]] if _SEED2_ATOMIC_RE.search(text) else []
|
||||
return [[int(x) for x in _SEED2_ATOMIC_RE.findall(b)] for b in blocks]
|
||||
|
||||
|
||||
def _load_seed2_tokenizer():
|
||||
"""Same import/shim sequence as data_prep/synth_llava/tokenize_seed2.py --
|
||||
reused verbatim rather than re-derived (see that file's docstring for why
|
||||
each shim exists: a transformers-version helper-function move, and a
|
||||
Qformer.cls=None crash).
|
||||
|
||||
NOTE: tokenize_seed2.py only ever calls .encode_image() (only needs
|
||||
prototype/pipeline.py's thin Seed2Tokenizer wrapper), but decode needs
|
||||
.decode()/.from_pretrained(), which only exist on the real HF PreTrainedModel
|
||||
class in seed2_tokenizer.py itself -- pipeline.py's wrapper doesn't have
|
||||
them. Copying the import verbatim silently returned the wrong class here
|
||||
(caught 2026-07-22: AttributeError, no from_pretrained) -- must return
|
||||
_seed2_tokenizer.Seed2Tokenizer directly, not prototype/pipeline.py's.
|
||||
|
||||
2026-07-23: no longer hardcodes PROTOTYPE_DIR -- uses whatever
|
||||
_resolve_seed2_dir() finds (local cluster copy, or a fresh download from
|
||||
the public ontocord/seed2 HF repo). Returns that dir alongside the class
|
||||
so callers know where to point Seed2Tokenizer.from_pretrained().
|
||||
|
||||
chdir target differs by branch: init_tokenizer() inside seed2_tokenizer.py
|
||||
does BertTokenizer.from_pretrained("./seed2/bert-base-uncased") -- a
|
||||
relative lookup that expects cwd to be seed2_dir's PARENT (matching how
|
||||
the local cluster copy is laid out: PROTOTYPE_DIR/seed2/bert-base-uncased).
|
||||
Chdir'ing into seed2_dir itself (one level too deep) breaks that lookup --
|
||||
caught 2026-07-23 testing the new encode_seed2.py against the local
|
||||
branch specifically (OSError: can't load './seed2/bert-base-uncased').
|
||||
The downloaded ontocord/seed2 snapshot has no bert-base-uncased subfolder
|
||||
at all (verified via its real file listing) yet works anyway -- empirically
|
||||
that branch's relative lookup resolves some other way (not fully
|
||||
root-caused), so only the local branch needs the parent-dir chdir fix."""
|
||||
seed2_dir = _resolve_seed2_dir()
|
||||
os.chdir(PROTOTYPE_DIR if seed2_dir == _LOCAL_SEED2_DIR else seed2_dir)
|
||||
|
||||
import transformers.modeling_utils as _modeling_utils
|
||||
import transformers.pytorch_utils as _pytorch_utils
|
||||
for _name in ("apply_chunking_to_forward", "find_pruneable_heads_and_indices", "prune_linear_layer"):
|
||||
if not hasattr(_modeling_utils, _name):
|
||||
setattr(_modeling_utils, _name, getattr(_pytorch_utils, _name))
|
||||
|
||||
sys.path.insert(0, seed2_dir)
|
||||
import seed2_tokenizer as _seed2_tokenizer
|
||||
|
||||
def _safe_get_output_embeddings(self):
|
||||
return None if self.cls is None else self.cls.predictions.decoder
|
||||
|
||||
def _safe_set_output_embeddings(self, new_embeddings):
|
||||
if self.cls is not None:
|
||||
self.cls.predictions.decoder = new_embeddings
|
||||
|
||||
for _cls in (_seed2_tokenizer.BertLMHeadModel, _seed2_tokenizer.BertForMaskedLM):
|
||||
_cls.get_output_embeddings = _safe_get_output_embeddings
|
||||
_cls.set_output_embeddings = _safe_set_output_embeddings
|
||||
|
||||
return _seed2_tokenizer.Seed2Tokenizer, seed2_dir
|
||||
|
||||
|
||||
def decode_seed2_tokens(token_ids: list, output_path: str, guidance_scale: float = 10.0,
|
||||
num_inference_steps: int = 20) -> None:
|
||||
if not token_ids:
|
||||
raise ValueError("No seed2 tokens to decode")
|
||||
bad = [t for t in token_ids if not (0 <= t < NUM_IMAGE_TOKENS)]
|
||||
if bad:
|
||||
raise ValueError(f"Token ids out of range [0, {NUM_IMAGE_TOKENS}): {bad[:5]}...")
|
||||
# Seed2Tokenizer's Q-former was trained with a fixed 32 query tokens/image
|
||||
# (see pos_embed_image.repeat(query_output_up.shape[0], ...) in
|
||||
# seed2_tokenizer.py -- shape[0] must be 1 image's worth of queries, i.e.
|
||||
# len(token_ids)==32). Any other count desyncs from the diffusion
|
||||
# pipeline's own batch-size assumption deep in the UNet (hit for real
|
||||
# 2026-07-22: 3-token and 96-token spans both crashed with a "tensor a
|
||||
# must match tensor b" RuntimeError -- see
|
||||
# samples/qwen3_1.7b_vla_v2_eval/2026-07-22_full_eval/SUMMARY.md, tests
|
||||
# 02_agent_continuation and 07_full_chain_from_scratch). Fail clearly
|
||||
# instead of letting that opaque error surface from inside the UNet.
|
||||
if len(token_ids) != SEED2_QUERY_LEN:
|
||||
raise ValueError(
|
||||
f"Got {len(token_ids)} seed2 tokens, but Seed2Tokenizer only decodes exactly "
|
||||
f"{SEED2_QUERY_LEN} tokens at a time (1 image's fixed Q-former query length). "
|
||||
f"If this span was extracted from text containing multiple <seed2>...</seed2> blocks "
|
||||
f"(e.g. a multi-activity generation), decode each block separately."
|
||||
)
|
||||
|
||||
import torch
|
||||
from diffusers import StableUnCLIPImg2ImgPipeline
|
||||
|
||||
Seed2Tokenizer, seed2_dir = _load_seed2_tokenizer()
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
dtype = torch.float16 if device == "cuda" else torch.float32
|
||||
|
||||
print(f"Loading {DIFFUSION_NAME} (first run downloads ~5GB)...")
|
||||
pipe = StableUnCLIPImg2ImgPipeline.from_pretrained(DIFFUSION_NAME, torch_dtype=dtype).to(device)
|
||||
|
||||
print("Loading Seed2Tokenizer...")
|
||||
tokenizer = Seed2Tokenizer.from_pretrained(seed2_dir, torch_dtype=dtype).to(device)
|
||||
|
||||
# get_codebook_entry() does self.embedding(indices) with no batch handling of
|
||||
# its own -- dim 0 of `indices` becomes the batch dim downstream (repeat(),
|
||||
# etc.). A flat (N,) tensor is silently read as N separate 1-token images
|
||||
# instead of 1 image made of N query tokens, which desyncs from the diffusion
|
||||
# pipeline's own (batch=1, x2 for CFG) conditioning shape and crashes deep in
|
||||
# the UNet (caught 2026-07-22: "tensor a (2) must match tensor b (64)" for a
|
||||
# 32-token input, i.e. 32 x 2 leaking through). Needs an explicit batch dim.
|
||||
indices = torch.tensor(token_ids, dtype=torch.long, device=device).unsqueeze(0)
|
||||
image = tokenizer.decode(pipe, indices, guidance_scale=guidance_scale,
|
||||
num_inference_steps=num_inference_steps)[0]
|
||||
image.save(output_path)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--tokens", help="Comma-separated list of raw <seed2_N> ids (0-8191)")
|
||||
ap.add_argument("--text-file", help="Plain text / JSONL file containing <seed2_N> tokens anywhere in it")
|
||||
ap.add_argument("--guidance-scale", type=float, default=10.0)
|
||||
ap.add_argument("--num-inference-steps", type=int, default=20)
|
||||
ap.add_argument("--output", required=True)
|
||||
args = ap.parse_args()
|
||||
|
||||
# _load_seed2_tokenizer() os.chdir()s into PROTOTYPE_DIR (needs to be cwd for
|
||||
# its own relative "./seed2" lookups) -- resolve a relative --output against
|
||||
# the original cwd *before* that happens, or it silently lands in prototype/
|
||||
# instead (caught 2026-07-22: full generation succeeded, only image.save() failed).
|
||||
output_path = os.path.abspath(args.output)
|
||||
|
||||
if args.tokens:
|
||||
blocks = [[int(x) for x in args.tokens.split(",")]]
|
||||
elif args.text_file:
|
||||
blocks = extract_seed2_blocks(open(args.text_file, encoding="utf-8").read())
|
||||
else:
|
||||
ap.error("Provide --tokens or --text-file")
|
||||
|
||||
if not blocks:
|
||||
ap.error("No <seed2_N> tokens found")
|
||||
|
||||
if len(blocks) == 1:
|
||||
print(f"Decoding {len(blocks[0])} seed2 tokens...")
|
||||
decode_seed2_tokens(blocks[0], output_path, args.guidance_scale, args.num_inference_steps)
|
||||
print(f"Saved: {output_path}")
|
||||
return
|
||||
|
||||
# Multiple <seed2>...</seed2> blocks (e.g. a multi-activity generation) --
|
||||
# decode each as its own image rather than concatenating and erroring
|
||||
# (added 2026-07-22 after a 96-token/3-block span crashed decode_seed2_tokens).
|
||||
print(f"Found {len(blocks)} separate <seed2>...</seed2> blocks -- decoding each as its own image.")
|
||||
stem, ext = os.path.splitext(output_path)
|
||||
for i, block in enumerate(blocks):
|
||||
block_output = f"{stem}_{i}{ext}"
|
||||
print(f"\nBlock {i}: {len(block)} tokens -> {block_output}")
|
||||
try:
|
||||
decode_seed2_tokens(block, block_output, args.guidance_scale, args.num_inference_steps)
|
||||
print(f"Saved: {block_output}")
|
||||
except ValueError as e:
|
||||
print(f"Skipped block {i}: {e}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
215
tools/decode/decode_snac.py
Normal file
215
tools/decode/decode_snac.py
Normal file
@@ -0,0 +1,215 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
SNAC detokenizer -- turns `<snac_N>` "listen format" tokens back into a real
|
||||
audio waveform (24kHz mono), using the same `hubertsiuzdak/snac_24khz` model
|
||||
used to encode them (data_prep/laion_emotional_roleplay/tokenize_snac.py,
|
||||
pipeline_pose/snac_finevideo.py).
|
||||
|
||||
Listen format encodes only 2 of the model's 3 hierarchical codebook levels
|
||||
(base 12.5Hz level 0, and one 25Hz level 1 -- the finest 50Hz level 2 is
|
||||
dropped entirely to save tokens). 3 tokens per base frame, in fixed order
|
||||
(L0, L1_even, L1_odd), with these offsets added to the raw codebook index
|
||||
(0-4095 each):
|
||||
|
||||
OFFSET_L0 = 128266 (raw + 128266)
|
||||
OFFSET_L1A = 132362 (raw + 128266 + 4096)
|
||||
OFFSET_L1B = 144650 (raw + 128266 + 4*4096)
|
||||
|
||||
Since level 2 was never encoded, this decoder reconstructs it as all-zero
|
||||
codes (index 0 in each of the 4x-oversampled slots) -- `SNAC.decode()` needs
|
||||
all 3 levels present, so this is the only reconstruction available without
|
||||
retraining/re-deriving the missing fine detail. Expect this to sound coarser
|
||||
than the original clip (level 2 carries the finest timbral detail), same
|
||||
lossy tradeoff already documented for `decode_cosmos.py`'s video reconstruction.
|
||||
|
||||
Usage:
|
||||
# From a raw list of ints (must be a multiple of 3):
|
||||
python tools/decode/decode_snac.py --tokens 128266,132850,145181,... --output out.wav
|
||||
|
||||
# From a flattened JSONL record, pulling every <snac>...</snac> block:
|
||||
python tools/decode/decode_snac.py --input-jsonl path.jsonl --record-id ID --output out.wav
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
import sys
|
||||
|
||||
OFFSET_L0 = 128266
|
||||
OFFSET_L1A = 128266 + 4096
|
||||
OFFSET_L1B = 128266 + 4 * 4096
|
||||
# Speak-format (full 3-level) offsets, 2026-07-23 -- matches the Leo/Orpheus
|
||||
# scheme in pipeline_video/snac_gpu.py and the corrected
|
||||
# data_prep/laion_emotional_roleplay/tokenize_snac.py::encode_speak(). Group
|
||||
# order per base frame: L0, L1a, L2_0, L2_1, L1b, L2_2, L2_3 (7 tokens).
|
||||
OFFSET_L2 = [136458, 140554, 148746, 152842]
|
||||
SAMPLE_RATE = 24000
|
||||
SNAC_MODEL = "hubertsiuzdak/snac_24khz"
|
||||
|
||||
_SNAC_ATOMIC_RE = re.compile(r"<snac_(\d+)>")
|
||||
# 2026-07-23: production wrapper is <listen>/<speak>, not <snac> -- <snac>
|
||||
# kept for backward compat with any older data/samples still using it.
|
||||
_SNAC_BLOCK_RE = re.compile(r"<(?:snac|listen|speak)>(.*?)</(?:snac|listen|speak)>", re.DOTALL)
|
||||
|
||||
|
||||
def extract_snac_tokens(text: str) -> list:
|
||||
"""Pull every <snac_N> id inside every <listen>...</listen> or
|
||||
<speak>...</speak> (or legacy <snac>...</snac>) block, in order. Falls
|
||||
back to scanning the whole text if no such wrapper is present."""
|
||||
blocks = _SNAC_BLOCK_RE.findall(text)
|
||||
source = " ".join(blocks) if blocks else text
|
||||
return [int(x) for x in _SNAC_ATOMIC_RE.findall(source)]
|
||||
|
||||
|
||||
def load_tokens_from_jsonl(path: str, record_id: str) -> list:
|
||||
with open(path, encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
rec = json.loads(line)
|
||||
rid = rec.get("video_id", rec.get("id"))
|
||||
if rid == record_id:
|
||||
return extract_snac_tokens(rec.get("text", ""))
|
||||
raise KeyError(f"record_id={record_id!r} not found in {path}")
|
||||
|
||||
|
||||
def decode_snac_tokens(token_ids: list, output_path: str) -> None:
|
||||
"""token_ids: raw <snac_N> ids (with offsets), length must be a multiple of 3."""
|
||||
if len(token_ids) % 3 != 0:
|
||||
raise ValueError(f"Expected a multiple of 3 tokens (L0,L1a,L1b triplets), got {len(token_ids)}")
|
||||
if not token_ids:
|
||||
raise ValueError("No snac tokens to decode")
|
||||
|
||||
import torch
|
||||
import soundfile as sf
|
||||
from snac import SNAC
|
||||
|
||||
n0 = len(token_ids) // 3
|
||||
c0 = torch.zeros(1, n0, dtype=torch.long)
|
||||
c1 = torch.zeros(1, 2 * n0, dtype=torch.long)
|
||||
c2 = torch.zeros(1, 4 * n0, dtype=torch.long) # level 2 was never encoded -- zero-fill
|
||||
|
||||
# This function assumes strict positional cycling (index%3==0 -> L0,
|
||||
# ==1 -> L1a, ==2 -> L1b) matches the actual band each id belongs to --
|
||||
# true for well-formed <snac>...</snac> blocks, but NOT guaranteed for a
|
||||
# free-floating/unwrapped token span assembled by extract_snac_tokens()'s
|
||||
# whole-text fallback (e.g. a generation that never closes </snac>,
|
||||
# spliced with earlier/later snac fragments elsewhere in the text). A
|
||||
# mismatched id produces an out-of-[0,4095) codebook index, which SNAC's
|
||||
# embedding lookup on GPU turns into an opaque
|
||||
# "CUDA error: device-side assert triggered" instead of a real error
|
||||
# message (hit for real 2026-07-22, roleplay_speech/sample eval run --
|
||||
# see samples/qwen3_1.7b_vla_v2_eval/2026-07-22_full_eval/SUMMARY.md).
|
||||
# Validate up front so the failure is legible.
|
||||
for i in range(n0):
|
||||
raw_l0, raw_l1a, raw_l1b = token_ids[3 * i], token_ids[3 * i + 1], token_ids[3 * i + 2]
|
||||
r0, r1a, r1b = raw_l0 - OFFSET_L0, raw_l1a - OFFSET_L1A, raw_l1b - OFFSET_L1B
|
||||
for pos, (name, tok, raw) in enumerate([("L0", raw_l0, r0), ("L1a", raw_l1a, r1a), ("L1b", raw_l1b, r1b)]):
|
||||
if not (0 <= raw < 4096):
|
||||
raise ValueError(
|
||||
f"Triplet {i} position {pos} ({name}): token <snac_{tok}> decodes to raw "
|
||||
f"codebook index {raw}, outside valid [0, 4096). This id likely belongs to a "
|
||||
f"different band than its position implies (offsets: L0={OFFSET_L0}, "
|
||||
f"L1a={OFFSET_L1A}, L1b={OFFSET_L1B}) -- often means the input tokens aren't "
|
||||
f"one clean <snac>...</snac> block (e.g. spliced fragments from an unwrapped "
|
||||
f"or unclosed generation)."
|
||||
)
|
||||
c0[0, i] = r0
|
||||
c1[0, 2 * i] = r1a
|
||||
c1[0, 2 * i + 1] = r1b
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model = SNAC.from_pretrained(SNAC_MODEL).eval().to(device)
|
||||
|
||||
with torch.inference_mode():
|
||||
audio = model.decode([c0.to(device), c1.to(device), c2.to(device)]) # (1, 1, samples)
|
||||
|
||||
waveform = audio.squeeze().float().cpu().numpy()
|
||||
sf.write(output_path, waveform, SAMPLE_RATE)
|
||||
|
||||
|
||||
def decode_speak_tokens(token_ids: list, output_path: str) -> None:
|
||||
"""token_ids: raw <snac_N> ids (with offsets), length must be a multiple of 7,
|
||||
group order per base frame: L0, L1a, L2_0, L2_1, L1b, L2_2, L2_3 -- matches
|
||||
encode_speak() in data_prep/laion_emotional_roleplay/tokenize_snac.py and
|
||||
pipeline_pose/snac_finevideo.py. Unlike decode_snac_tokens() (listen-only,
|
||||
zero-fills level 2), this reconstructs the REAL level-2 codes."""
|
||||
if len(token_ids) % 7 != 0:
|
||||
raise ValueError(f"Expected a multiple of 7 tokens (speak-format groups), got {len(token_ids)}")
|
||||
if not token_ids:
|
||||
raise ValueError("No snac tokens to decode")
|
||||
|
||||
import torch
|
||||
import soundfile as sf
|
||||
from snac import SNAC
|
||||
|
||||
n0 = len(token_ids) // 7
|
||||
c0 = torch.zeros(1, n0, dtype=torch.long)
|
||||
c1 = torch.zeros(1, 2 * n0, dtype=torch.long)
|
||||
c2 = torch.zeros(1, 4 * n0, dtype=torch.long)
|
||||
|
||||
offsets = [OFFSET_L0, OFFSET_L1A, OFFSET_L2[0], OFFSET_L2[1], OFFSET_L1B, OFFSET_L2[2], OFFSET_L2[3]]
|
||||
names = ["L0", "L1a", "L2_0", "L2_1", "L1b", "L2_2", "L2_3"]
|
||||
for i in range(n0):
|
||||
raws = []
|
||||
for pos in range(7):
|
||||
tok = token_ids[7 * i + pos]
|
||||
raw = tok - offsets[pos]
|
||||
if not (0 <= raw < 4096):
|
||||
raise ValueError(
|
||||
f"Group {i} position {pos} ({names[pos]}): token <snac_{tok}> decodes to raw "
|
||||
f"codebook index {raw}, outside valid [0, 4096). Likely a band/position "
|
||||
f"mismatch (offsets: {dict(zip(names, offsets))})."
|
||||
)
|
||||
raws.append(raw)
|
||||
c0[0, i] = raws[0]
|
||||
c1[0, 2 * i] = raws[1]
|
||||
c2[0, 4 * i] = raws[2]
|
||||
c2[0, 4 * i + 1] = raws[3]
|
||||
c1[0, 2 * i + 1] = raws[4]
|
||||
c2[0, 4 * i + 2] = raws[5]
|
||||
c2[0, 4 * i + 3] = raws[6]
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model = SNAC.from_pretrained(SNAC_MODEL).eval().to(device)
|
||||
|
||||
with torch.inference_mode():
|
||||
audio = model.decode([c0.to(device), c1.to(device), c2.to(device)])
|
||||
|
||||
waveform = audio.squeeze().float().cpu().numpy()
|
||||
sf.write(output_path, waveform, SAMPLE_RATE)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--tokens", help="Comma-separated list of raw <snac_N> ids")
|
||||
ap.add_argument("--input-jsonl", help="Flattened JSONL file to pull tokens from")
|
||||
ap.add_argument("--record-id", help="video_id/id field to select within --input-jsonl")
|
||||
ap.add_argument("--text-file", help="Plain text file containing <snac_N> tokens anywhere in it")
|
||||
ap.add_argument("--format", choices=["listen", "speak"], default="listen",
|
||||
help="listen = 3 tok/frame (L0+L1, zero-filled L2); speak = 7 tok/frame (real L2)")
|
||||
ap.add_argument("--output", required=True)
|
||||
args = ap.parse_args()
|
||||
|
||||
if args.tokens:
|
||||
token_ids = [int(x) for x in args.tokens.split(",")]
|
||||
elif args.input_jsonl and args.record_id:
|
||||
token_ids = load_tokens_from_jsonl(args.input_jsonl, args.record_id)
|
||||
elif args.text_file:
|
||||
token_ids = extract_snac_tokens(open(args.text_file, encoding="utf-8").read())
|
||||
else:
|
||||
ap.error("Provide --tokens, --text-file, or (--input-jsonl and --record-id)")
|
||||
|
||||
group_size = 7 if args.format == "speak" else 3
|
||||
print(f"Decoding {len(token_ids)} snac tokens ({args.format} format, "
|
||||
f"{len(token_ids) // group_size} base frames, "
|
||||
f"~{len(token_ids) / group_size / 12.5:.2f}s @ 12.5Hz base rate)...")
|
||||
if args.format == "speak":
|
||||
decode_speak_tokens(token_ids, args.output)
|
||||
else:
|
||||
decode_snac_tokens(token_ids, args.output)
|
||||
print(f"Saved: {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
9
tools/decode/vendor/cosmos_tokenizer/NOTICE.md
vendored
Normal file
9
tools/decode/vendor/cosmos_tokenizer/NOTICE.md
vendored
Normal file
@@ -0,0 +1,9 @@
|
||||
Vendored from NVIDIA's [Cosmos-Tokenizer](https://github.com/NVIDIA/Cosmos-Tokenizer),
|
||||
Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES, licensed under Apache-2.0
|
||||
(see individual file headers). Vendored here (2026-07-23) so
|
||||
`tools/decode/decode_cosmos.py` doesn't require the internal cluster's
|
||||
`prototype/` directory (not part of this repo) -- only inference code
|
||||
(`video_lib.py` + its direct dependencies) is included, not training code.
|
||||
Model checkpoints are downloaded separately from
|
||||
[nvidia/Cosmos-Tokenizer-DV8x16x16](https://huggingface.co/nvidia/Cosmos-Tokenizer-DV8x16x16)
|
||||
on first use, not vendored here.
|
||||
14
tools/decode/vendor/cosmos_tokenizer/__init__.py
vendored
Normal file
14
tools/decode/vendor/cosmos_tokenizer/__init__.py
vendored
Normal file
@@ -0,0 +1,14 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
197
tools/decode/vendor/cosmos_tokenizer/image_cli.py
vendored
Normal file
197
tools/decode/vendor/cosmos_tokenizer/image_cli.py
vendored
Normal file
@@ -0,0 +1,197 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""A CLI to run ImageTokenizer on plain images based on torch.jit.
|
||||
|
||||
Usage:
|
||||
python3 -m cosmos_tokenizer.image_cli \
|
||||
--image_pattern 'path/to/input/folder/*.jpg' \
|
||||
--output_dir ./reconstructions \
|
||||
--checkpoint_enc ./pretrained_ckpts/CosmosCI_f8x8/encoder.jit \
|
||||
--checkpoint_dec ./pretrained_ckpts/CosmosCI_f8x8/decoder.jit
|
||||
|
||||
Optionally, you can run the model in pure PyTorch mode:
|
||||
python3 -m cosmos_tokenizer.image_cli \
|
||||
--image_pattern 'path/to/input/folder/*.jpg' \
|
||||
--mode torch \
|
||||
--tokenizer_type CI \
|
||||
--spatial_compression 8 \
|
||||
--checkpoint_enc ./pretrained_ckpts/CosmosCI_f8x8/encoder.jit \
|
||||
--checkpoint_dec ./pretrained_ckpts/CosmosCI_f8x8/decoder.jit
|
||||
"""
|
||||
|
||||
import os
|
||||
from argparse import ArgumentParser, Namespace
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
from loguru import logger as logging
|
||||
from cosmos_tokenizer.networks import TokenizerConfigs
|
||||
|
||||
from cosmos_tokenizer.image_lib import ImageTokenizer
|
||||
from cosmos_tokenizer.utils import (
|
||||
get_filepaths,
|
||||
get_output_filepath,
|
||||
read_image,
|
||||
resize_image,
|
||||
write_image,
|
||||
)
|
||||
|
||||
|
||||
def _parse_args() -> tuple[Namespace, dict[str, Any]]:
|
||||
parser = ArgumentParser(
|
||||
description="A CLI for running ImageTokenizer on plain images."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image_pattern",
|
||||
type=str,
|
||||
default="path/to/images/*.jpg",
|
||||
help="Glob pattern.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoint",
|
||||
type=str,
|
||||
default=None,
|
||||
help="JIT full Autoencoder model filepath.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoint_enc",
|
||||
type=str,
|
||||
default=None,
|
||||
help="JIT Encoder model filepath.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoint_dec",
|
||||
type=str,
|
||||
default=None,
|
||||
help="JIT Decoder model filepath.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tokenizer_type",
|
||||
type=str,
|
||||
choices=["CI", "DI"],
|
||||
help="Specifies the tokenizer type.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--spatial_compression",
|
||||
type=int,
|
||||
choices=[8, 16],
|
||||
default=8,
|
||||
help="The spatial compression factor.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mode",
|
||||
type=str,
|
||||
choices=["torch", "jit"],
|
||||
default="jit",
|
||||
help="Specify the backend: native 'torch' or 'jit' (default: 'jit')",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--short_size",
|
||||
type=int,
|
||||
default=None,
|
||||
help="The size to resample inputs. None, by default.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
default="bfloat16",
|
||||
help="Sets the precision. Default bfloat16.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device",
|
||||
type=str,
|
||||
default="cuda",
|
||||
help="Device for invoking the model.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_dir", type=str, default=None, help="Output directory."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save_input",
|
||||
action="store_true",
|
||||
help="If on, the input image will be be outputed too.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
logging.info("Initializes args ...")
|
||||
args = _parse_args()
|
||||
if args.mode == "torch" and args.tokenizer_type not in ["CI", "DI"]:
|
||||
logging.error("'torch' backend requires the tokenizer_type of 'CI' or 'DI'.")
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def _run_eval() -> None:
|
||||
"""Invokes the evaluation pipeline."""
|
||||
|
||||
if (
|
||||
args.checkpoint_enc is None
|
||||
and args.checkpoint_dec is None
|
||||
and args.checkpoint is None
|
||||
):
|
||||
logging.warning(
|
||||
"Aborting. Both encoder or decoder JIT required. Or provide the full autoencoder JIT model."
|
||||
)
|
||||
return
|
||||
|
||||
if args.mode == "torch":
|
||||
tokenizer_config = TokenizerConfigs[args.tokenizer_type].value
|
||||
tokenizer_config.update(dict(spatial_compression=args.spatial_compression))
|
||||
else:
|
||||
tokenizer_config = None
|
||||
|
||||
logging.info(
|
||||
f"Loading a torch.jit model `{os.path.dirname(args.checkpoint or args.checkpoint_enc or args.checkpoint_dec)}` ..."
|
||||
)
|
||||
autoencoder = ImageTokenizer(
|
||||
checkpoint=args.checkpoint,
|
||||
checkpoint_enc=args.checkpoint_enc,
|
||||
checkpoint_dec=args.checkpoint_dec,
|
||||
tokenizer_config=tokenizer_config,
|
||||
device=args.device,
|
||||
dtype=args.dtype,
|
||||
)
|
||||
|
||||
filepaths = get_filepaths(args.image_pattern)
|
||||
logging.info(f"Found {len(filepaths)} images from {args.image_pattern}.")
|
||||
|
||||
for filepath in filepaths:
|
||||
logging.info(f"Reading image {filepath} ...")
|
||||
image = read_image(filepath)
|
||||
image = resize_image(image, short_size=args.short_size)
|
||||
batch_image = np.expand_dims(image, axis=0)
|
||||
|
||||
logging.info("Invoking the autoencoder model in ... ")
|
||||
output_image = autoencoder(batch_image)[0]
|
||||
|
||||
output_filepath = get_output_filepath(filepath, output_dir=args.output_dir)
|
||||
logging.info(f"Outputing {output_filepath} ...")
|
||||
write_image(output_filepath, output_image)
|
||||
|
||||
if args.save_input:
|
||||
ext = os.path.splitext(output_filepath)[-1]
|
||||
input_filepath = output_filepath.replace(ext, "_input" + ext)
|
||||
write_image(input_filepath, image)
|
||||
|
||||
|
||||
@logging.catch(reraise=True)
|
||||
def main() -> None:
|
||||
_run_eval()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
128
tools/decode/vendor/cosmos_tokenizer/image_lib.py
vendored
Normal file
128
tools/decode/vendor/cosmos_tokenizer/image_lib.py
vendored
Normal file
@@ -0,0 +1,128 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""A library for image tokenizers inference."""
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from typing import Any
|
||||
|
||||
from cosmos_tokenizer.utils import (
|
||||
load_model,
|
||||
load_encoder_model,
|
||||
load_decoder_model,
|
||||
numpy2tensor,
|
||||
pad_image_batch,
|
||||
tensor2numpy,
|
||||
unpad_image_batch,
|
||||
)
|
||||
|
||||
|
||||
class ImageTokenizer(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
checkpoint: str = None,
|
||||
checkpoint_enc: str = None,
|
||||
checkpoint_dec: str = None,
|
||||
tokenizer_config: dict[str, Any] = None,
|
||||
device: str = "cuda",
|
||||
dtype: str = "bfloat16",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self._device = device
|
||||
self._dtype = getattr(torch, dtype)
|
||||
self._full_model = (
|
||||
load_model(checkpoint, tokenizer_config, device).to(self._dtype)
|
||||
if checkpoint is not None
|
||||
else None
|
||||
)
|
||||
self._enc_model = (
|
||||
load_encoder_model(checkpoint_enc, tokenizer_config, device).to(self._dtype)
|
||||
if checkpoint_enc is not None
|
||||
else None
|
||||
)
|
||||
self._dec_model = (
|
||||
load_decoder_model(checkpoint_dec, tokenizer_config, device).to(self._dtype)
|
||||
if checkpoint_dec is not None
|
||||
else None
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
def autoencode(self, input_tensor: torch.Tensor) -> torch.Tensor:
|
||||
"""Reconstrcuts a batch of image tensors after embedding into a latent.
|
||||
|
||||
Args:
|
||||
input_tensor: The input image Bx3xHxW layout, range [-1..1].
|
||||
Returns:
|
||||
The reconstructed tensor, layout Bx3xHxW, range [-1..1].
|
||||
"""
|
||||
if self._full_model is not None:
|
||||
output_tensor = self._full_model(input_tensor)
|
||||
output_tensor = (
|
||||
output_tensor[0] if isinstance(output_tensor, tuple) else output_tensor
|
||||
)
|
||||
else:
|
||||
output_latent = self.encode(input_tensor)[0]
|
||||
output_tensor = self.decode(output_latent)
|
||||
return output_tensor
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, input_latent: torch.Tensor) -> torch.Tensor:
|
||||
"""Decodes an image from a provided latent embedding.
|
||||
|
||||
Args:
|
||||
input_latent: The continuous latent Bx16xhxw for CI,
|
||||
or the discrete indices Bxhxw for DI.
|
||||
Returns:
|
||||
The output tensor in Bx3xHxW, range [-1..1].
|
||||
"""
|
||||
return self._dec_model(input_latent)
|
||||
|
||||
@torch.no_grad()
|
||||
def encode(self, input_tensor: torch.Tensor) -> tuple[torch.Tensor]:
|
||||
"""Encodes an image into a latent embedding or code.
|
||||
|
||||
Args:
|
||||
input_tensor: The input tensor Bx3xHxW layout, range [-1..1].
|
||||
Returns:
|
||||
For continuous image (CI) tokenizer, the tuple contains:
|
||||
- The latent embedding, Bx16x(h)x(w), where the compression
|
||||
rate is (H/h x W/w), and channel dimension of 16.
|
||||
For discrete image (DI) tokenizer, the tuple contains:
|
||||
- The indices, Bx(h)x(w), from a codebook of size 64K, which
|
||||
corresponds to FSQ levels of (8,8,8,5,5,5).
|
||||
- The discrete code, Bx6x(h)x(w), where the compression rate is
|
||||
again (H/h x W/w), and channel dimension of 6.
|
||||
"""
|
||||
output_latent = self._enc_model(input_tensor)
|
||||
if isinstance(output_latent, torch.Tensor):
|
||||
return output_latent
|
||||
return output_latent[:-1]
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, image: np.ndarray) -> np.ndarray:
|
||||
"""Reconstructs an image using a pre-trained tokenizer.
|
||||
|
||||
Args:
|
||||
image: The input image BxHxWxC layout, range [0..255].
|
||||
Returns:
|
||||
The reconstructed image in range [0..255], layout BxHxWxC.
|
||||
"""
|
||||
padded_input_image, crop_region = pad_image_batch(image)
|
||||
input_tensor = numpy2tensor(
|
||||
padded_input_image, dtype=self._dtype, device=self._device
|
||||
)
|
||||
output_tensor = self.autoencode(input_tensor)
|
||||
padded_output_image = tensor2numpy(output_tensor)
|
||||
return unpad_image_batch(padded_output_image, crop_region)
|
||||
63
tools/decode/vendor/cosmos_tokenizer/modules/__init__.py
vendored
Normal file
63
tools/decode/vendor/cosmos_tokenizer/modules/__init__.py
vendored
Normal file
@@ -0,0 +1,63 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
from enum import Enum
|
||||
|
||||
from cosmos_tokenizer.modules.distributions import (
|
||||
GaussianDistribution,
|
||||
IdentityDistribution,
|
||||
)
|
||||
from cosmos_tokenizer.modules.layers2d import Decoder, Encoder
|
||||
from cosmos_tokenizer.modules.layers3d import (
|
||||
DecoderBase,
|
||||
DecoderFactorized,
|
||||
EncoderBase,
|
||||
EncoderFactorized,
|
||||
)
|
||||
from cosmos_tokenizer.modules.quantizers import (
|
||||
FSQuantizer,
|
||||
LFQuantizer,
|
||||
ResidualFSQuantizer,
|
||||
VectorQuantizer,
|
||||
)
|
||||
|
||||
|
||||
class EncoderType(Enum):
|
||||
Default = Encoder
|
||||
|
||||
|
||||
class DecoderType(Enum):
|
||||
Default = Decoder
|
||||
|
||||
|
||||
class Encoder3DType(Enum):
|
||||
BASE = EncoderBase
|
||||
FACTORIZED = EncoderFactorized
|
||||
|
||||
|
||||
class Decoder3DType(Enum):
|
||||
BASE = DecoderBase
|
||||
FACTORIZED = DecoderFactorized
|
||||
|
||||
|
||||
class ContinuousFormulation(Enum):
|
||||
VAE = GaussianDistribution
|
||||
AE = IdentityDistribution
|
||||
|
||||
|
||||
class DiscreteQuantizer(Enum):
|
||||
VQ = VectorQuantizer
|
||||
LFQ = LFQuantizer
|
||||
FSQ = FSQuantizer
|
||||
RESFSQ = ResidualFSQuantizer
|
||||
41
tools/decode/vendor/cosmos_tokenizer/modules/distributions.py
vendored
Normal file
41
tools/decode/vendor/cosmos_tokenizer/modules/distributions.py
vendored
Normal file
@@ -0,0 +1,41 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""The distribution modes to use for continuous image tokenizers."""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class IdentityDistribution(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, parameters):
|
||||
return parameters, (torch.tensor([0.0]), torch.tensor([0.0]))
|
||||
|
||||
|
||||
class GaussianDistribution(torch.nn.Module):
|
||||
def __init__(self, min_logvar: float = -30.0, max_logvar: float = 20.0):
|
||||
super().__init__()
|
||||
self.min_logvar = min_logvar
|
||||
self.max_logvar = max_logvar
|
||||
|
||||
def sample(self, mean, logvar):
|
||||
std = torch.exp(0.5 * logvar)
|
||||
return mean + std * torch.randn_like(mean)
|
||||
|
||||
def forward(self, parameters):
|
||||
mean, logvar = torch.chunk(parameters, 2, dim=1)
|
||||
logvar = torch.clamp(logvar, self.min_logvar, self.max_logvar)
|
||||
return self.sample(mean, logvar), (mean, logvar)
|
||||
368
tools/decode/vendor/cosmos_tokenizer/modules/layers2d.py
vendored
Normal file
368
tools/decode/vendor/cosmos_tokenizer/modules/layers2d.py
vendored
Normal file
@@ -0,0 +1,368 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""The model definition for Continuous 2D layers
|
||||
|
||||
Adapted from: https://github.com/CompVis/stable-diffusion/blob/
|
||||
21f890f9da3cfbeaba8e2ac3c425ee9e998d5229/ldm/modules/diffusionmodules/model.py
|
||||
|
||||
[Copyright (c) 2022 Robin Rombach and Patrick Esser and contributors]
|
||||
https://github.com/CompVis/stable-diffusion/blob/
|
||||
21f890f9da3cfbeaba8e2ac3c425ee9e998d5229/LICENSE
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
|
||||
# pytorch_diffusion + derived encoder decoder
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from loguru import logger as logging
|
||||
|
||||
from cosmos_tokenizer.modules.patching import Patcher, UnPatcher
|
||||
from cosmos_tokenizer.modules.utils import Normalize, nonlinearity
|
||||
|
||||
|
||||
class Upsample(nn.Module):
|
||||
def __init__(self, in_channels: int):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = x.repeat_interleave(2, dim=2).repeat_interleave(2, dim=3)
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class Downsample(nn.Module):
|
||||
def __init__(self, in_channels: int):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=3, stride=2, padding=0
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
pad = (0, 1, 0, 1)
|
||||
x = F.pad(x, pad, mode="constant", value=0)
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class ResnetBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
in_channels: int,
|
||||
out_channels: int = None,
|
||||
dropout: float,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
|
||||
self.norm1 = Normalize(in_channels)
|
||||
self.conv1 = nn.Conv2d(
|
||||
in_channels, out_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
self.norm2 = Normalize(out_channels)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.conv2 = nn.Conv2d(
|
||||
out_channels, out_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
self.nin_shortcut = (
|
||||
nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
|
||||
if in_channels != out_channels
|
||||
else nn.Identity()
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
h = x
|
||||
h = self.norm1(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv1(h)
|
||||
|
||||
h = self.norm2(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.dropout(h)
|
||||
h = self.conv2(h)
|
||||
|
||||
x = self.nin_shortcut(x)
|
||||
|
||||
return x + h
|
||||
|
||||
|
||||
class AttnBlock(nn.Module):
|
||||
def __init__(self, in_channels: int):
|
||||
super().__init__()
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
self.k = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
self.v = nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
self.proj_out = nn.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
# TODO (freda): Consider reusing implementations in Attn `imaginaire`,
|
||||
# since than one is gonna be based on TransformerEngine's attn op,
|
||||
# w/c could ease CP implementations.
|
||||
h_ = x
|
||||
h_ = self.norm(h_)
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
b, c, h, w = q.shape
|
||||
q = q.reshape(b, c, h * w)
|
||||
q = q.permute(0, 2, 1)
|
||||
k = k.reshape(b, c, h * w)
|
||||
w_ = torch.bmm(q, k)
|
||||
w_ = w_ * (int(c) ** (-0.5))
|
||||
w_ = F.softmax(w_, dim=2)
|
||||
|
||||
# attend to values
|
||||
v = v.reshape(b, c, h * w)
|
||||
w_ = w_.permute(0, 2, 1)
|
||||
h_ = torch.bmm(v, w_)
|
||||
h_ = h_.reshape(b, c, h, w)
|
||||
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return x + h_
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
channels: int,
|
||||
channels_mult: list[int],
|
||||
num_res_blocks: int,
|
||||
attn_resolutions: list[int],
|
||||
dropout: float,
|
||||
resolution: int,
|
||||
z_channels: int,
|
||||
spatial_compression: int,
|
||||
**ignore_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_resolutions = len(channels_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
|
||||
# Patcher.
|
||||
patch_size = ignore_kwargs.get("patch_size", 1)
|
||||
self.patcher = Patcher(
|
||||
patch_size, ignore_kwargs.get("patch_method", "rearrange")
|
||||
)
|
||||
in_channels = in_channels * patch_size * patch_size
|
||||
|
||||
# calculate the number of downsample operations
|
||||
self.num_downsamples = int(math.log2(spatial_compression)) - int(
|
||||
math.log2(patch_size)
|
||||
)
|
||||
assert (
|
||||
self.num_downsamples <= self.num_resolutions
|
||||
), f"we can only downsample {self.num_resolutions} times at most"
|
||||
|
||||
# downsampling
|
||||
self.conv_in = torch.nn.Conv2d(
|
||||
in_channels, channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
curr_res = resolution // patch_size
|
||||
in_ch_mult = (1,) + tuple(channels_mult)
|
||||
self.in_ch_mult = in_ch_mult
|
||||
self.down = nn.ModuleList()
|
||||
for i_level in range(self.num_resolutions):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_in = channels * in_ch_mult[i_level]
|
||||
block_out = channels * channels_mult[i_level]
|
||||
for _ in range(self.num_res_blocks):
|
||||
block.append(
|
||||
ResnetBlock(
|
||||
in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
dropout=dropout,
|
||||
)
|
||||
)
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(AttnBlock(block_in))
|
||||
down = nn.Module()
|
||||
down.block = block
|
||||
down.attn = attn
|
||||
if i_level < self.num_downsamples:
|
||||
down.downsample = Downsample(block_in)
|
||||
curr_res = curr_res // 2
|
||||
self.down.append(down)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock(
|
||||
in_channels=block_in, out_channels=block_in, dropout=dropout
|
||||
)
|
||||
self.mid.attn_1 = AttnBlock(block_in)
|
||||
self.mid.block_2 = ResnetBlock(
|
||||
in_channels=block_in, out_channels=block_in, dropout=dropout
|
||||
)
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = torch.nn.Conv2d(
|
||||
block_in, z_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.patcher(x)
|
||||
|
||||
# downsampling
|
||||
hs = [self.conv_in(x)]
|
||||
for i_level in range(self.num_resolutions):
|
||||
for i_block in range(self.num_res_blocks):
|
||||
h = self.down[i_level].block[i_block](hs[-1])
|
||||
if len(self.down[i_level].attn) > 0:
|
||||
h = self.down[i_level].attn[i_block](h)
|
||||
hs.append(h)
|
||||
if i_level < self.num_downsamples:
|
||||
hs.append(self.down[i_level].downsample(hs[-1]))
|
||||
|
||||
# middle
|
||||
h = hs[-1]
|
||||
h = self.mid.block_1(h)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h)
|
||||
|
||||
# end
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
return h
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
out_channels: int,
|
||||
channels: int,
|
||||
channels_mult: list[int],
|
||||
num_res_blocks: int,
|
||||
attn_resolutions: int,
|
||||
dropout: float,
|
||||
resolution: int,
|
||||
z_channels: int,
|
||||
spatial_compression: int,
|
||||
**ignore_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_resolutions = len(channels_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
|
||||
# UnPatcher.
|
||||
patch_size = ignore_kwargs.get("patch_size", 1)
|
||||
self.unpatcher = UnPatcher(
|
||||
patch_size, ignore_kwargs.get("patch_method", "rearrange")
|
||||
)
|
||||
out_ch = out_channels * patch_size * patch_size
|
||||
|
||||
# calculate the number of upsample operations
|
||||
self.num_upsamples = int(math.log2(spatial_compression)) - int(
|
||||
math.log2(patch_size)
|
||||
)
|
||||
assert (
|
||||
self.num_upsamples <= self.num_resolutions
|
||||
), f"we can only upsample {self.num_resolutions} times at most"
|
||||
|
||||
block_in = channels * channels_mult[self.num_resolutions - 1]
|
||||
curr_res = (resolution // patch_size) // 2 ** (self.num_resolutions - 1)
|
||||
self.z_shape = (1, z_channels, curr_res, curr_res)
|
||||
logging.info(
|
||||
"Working with z of shape {} = {} dimensions.".format(
|
||||
self.z_shape, np.prod(self.z_shape)
|
||||
)
|
||||
)
|
||||
|
||||
# z to block_in
|
||||
self.conv_in = torch.nn.Conv2d(
|
||||
z_channels, block_in, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
# middle
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock(
|
||||
in_channels=block_in, out_channels=block_in, dropout=dropout
|
||||
)
|
||||
self.mid.attn_1 = AttnBlock(block_in)
|
||||
self.mid.block_2 = ResnetBlock(
|
||||
in_channels=block_in, out_channels=block_in, dropout=dropout
|
||||
)
|
||||
|
||||
# upsampling
|
||||
self.up = nn.ModuleList()
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_out = channels * channels_mult[i_level]
|
||||
for _ in range(self.num_res_blocks + 1):
|
||||
block.append(
|
||||
ResnetBlock(
|
||||
in_channels=block_in,
|
||||
out_channels=block_out,
|
||||
dropout=dropout,
|
||||
)
|
||||
)
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(AttnBlock(block_in))
|
||||
up = nn.Module()
|
||||
up.block = block
|
||||
up.attn = attn
|
||||
if i_level >= (self.num_resolutions - self.num_upsamples):
|
||||
up.upsample = Upsample(block_in)
|
||||
curr_res = curr_res * 2
|
||||
self.up.insert(0, up)
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = torch.nn.Conv2d(
|
||||
block_in, out_ch, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
def forward(self, z: torch.Tensor) -> torch.Tensor:
|
||||
h = self.conv_in(z)
|
||||
|
||||
# middle
|
||||
h = self.mid.block_1(h)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h)
|
||||
|
||||
# upsampling
|
||||
for i_level in reversed(range(self.num_resolutions)):
|
||||
for i_block in range(self.num_res_blocks + 1):
|
||||
h = self.up[i_level].block[i_block](h)
|
||||
if len(self.up[i_level].attn) > 0:
|
||||
h = self.up[i_level].attn[i_block](h)
|
||||
if i_level >= (self.num_resolutions - self.num_upsamples):
|
||||
h = self.up[i_level].upsample(h)
|
||||
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
h = self.unpatcher(h)
|
||||
return h
|
||||
1040
tools/decode/vendor/cosmos_tokenizer/modules/layers3d.py
vendored
Normal file
1040
tools/decode/vendor/cosmos_tokenizer/modules/layers3d.py
vendored
Normal file
File diff suppressed because it is too large
Load Diff
356
tools/decode/vendor/cosmos_tokenizer/modules/patching.py
vendored
Normal file
356
tools/decode/vendor/cosmos_tokenizer/modules/patching.py
vendored
Normal file
@@ -0,0 +1,356 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""The patcher and unpatcher implementation for 2D and 3D data.
|
||||
|
||||
The idea of Haar wavelet is to compute LL, LH, HL, HH component as two 1D convolutions.
|
||||
One on the rows and one on the columns.
|
||||
For example, in 1D signal, we have [a, b], then the low-freq compoenent is [a + b] / 2 and high-freq is [a - b] / 2.
|
||||
We can use a 1D convolution with kernel [1, 1] and stride 2 to represent the L component.
|
||||
For H component, we can use a 1D convolution with kernel [1, -1] and stride 2.
|
||||
Although in principle, we typically only do additional Haar wavelet over the LL component. But here we do it for all
|
||||
as we need to support downsampling for more than 2x.
|
||||
For example, 4x downsampling can be done by 2x Haar and additional 2x Haar, and the shape would be.
|
||||
[3, 256, 256] -> [12, 128, 128] -> [48, 64, 64]
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
_WAVELETS = {
|
||||
"haar": torch.tensor([0.7071067811865476, 0.7071067811865476]),
|
||||
"rearrange": torch.tensor([1.0, 1.0]),
|
||||
}
|
||||
_PERSISTENT = False
|
||||
|
||||
|
||||
class Patcher(torch.nn.Module):
|
||||
"""A module to convert image tensors into patches using torch operations.
|
||||
|
||||
The main difference from `class Patching` is that this module implements
|
||||
all operations using torch, rather than python or numpy, for efficiency purpose.
|
||||
|
||||
It's bit-wise identical to the Patching module outputs, with the added
|
||||
benefit of being torch.jit scriptable.
|
||||
"""
|
||||
|
||||
def __init__(self, patch_size=1, patch_method="haar"):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.patch_method = patch_method
|
||||
self.register_buffer(
|
||||
"wavelets", _WAVELETS[patch_method], persistent=_PERSISTENT
|
||||
)
|
||||
self.range = range(int(torch.log2(torch.tensor(self.patch_size)).item()))
|
||||
self.register_buffer(
|
||||
"_arange",
|
||||
torch.arange(_WAVELETS[patch_method].shape[0]),
|
||||
persistent=_PERSISTENT,
|
||||
)
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, x):
|
||||
if self.patch_method == "haar":
|
||||
return self._haar(x)
|
||||
elif self.patch_method == "rearrange":
|
||||
return self._arrange(x)
|
||||
else:
|
||||
raise ValueError("Unknown patch method: " + self.patch_method)
|
||||
|
||||
def _dwt(self, x, mode="reflect", rescale=False):
|
||||
dtype = x.dtype
|
||||
h = self.wavelets
|
||||
|
||||
n = h.shape[0]
|
||||
g = x.shape[1]
|
||||
hl = h.flip(0).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hh = (h * ((-1) ** self._arange)).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hh = hh.to(dtype=dtype)
|
||||
hl = hl.to(dtype=dtype)
|
||||
|
||||
x = F.pad(x, pad=(n - 2, n - 1, n - 2, n - 1), mode=mode).to(dtype)
|
||||
xl = F.conv2d(x, hl.unsqueeze(2), groups=g, stride=(1, 2))
|
||||
xh = F.conv2d(x, hh.unsqueeze(2), groups=g, stride=(1, 2))
|
||||
xll = F.conv2d(xl, hl.unsqueeze(3), groups=g, stride=(2, 1))
|
||||
xlh = F.conv2d(xl, hh.unsqueeze(3), groups=g, stride=(2, 1))
|
||||
xhl = F.conv2d(xh, hl.unsqueeze(3), groups=g, stride=(2, 1))
|
||||
xhh = F.conv2d(xh, hh.unsqueeze(3), groups=g, stride=(2, 1))
|
||||
|
||||
out = torch.cat([xll, xlh, xhl, xhh], dim=1)
|
||||
if rescale:
|
||||
out = out / 2
|
||||
return out
|
||||
|
||||
def _haar(self, x):
|
||||
for _ in self.range:
|
||||
x = self._dwt(x, rescale=True)
|
||||
return x
|
||||
|
||||
def _arrange(self, x):
|
||||
x = rearrange(
|
||||
x,
|
||||
"b c (h p1) (w p2) -> b (c p1 p2) h w",
|
||||
p1=self.patch_size,
|
||||
p2=self.patch_size,
|
||||
).contiguous()
|
||||
return x
|
||||
|
||||
|
||||
class Patcher3D(Patcher):
|
||||
"""A 3D discrete wavelet transform for video data, expects 5D tensor, i.e. a batch of videos."""
|
||||
|
||||
def __init__(self, patch_size=1, patch_method="haar"):
|
||||
super().__init__(patch_method=patch_method, patch_size=patch_size)
|
||||
self.register_buffer(
|
||||
"patch_size_buffer",
|
||||
patch_size * torch.ones([1], dtype=torch.int32),
|
||||
persistent=_PERSISTENT,
|
||||
)
|
||||
|
||||
def _dwt(self, x, wavelet, mode="reflect", rescale=False):
|
||||
dtype = x.dtype
|
||||
h = self.wavelets
|
||||
|
||||
n = h.shape[0]
|
||||
g = x.shape[1]
|
||||
hl = h.flip(0).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hh = (h * ((-1) ** self._arange)).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hh = hh.to(dtype=dtype)
|
||||
hl = hl.to(dtype=dtype)
|
||||
|
||||
# Handles temporal axis.
|
||||
x = F.pad(
|
||||
x, pad=(max(0, n - 2), n - 1, n - 2, n - 1, n - 2, n - 1), mode=mode
|
||||
).to(dtype)
|
||||
xl = F.conv3d(x, hl.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1))
|
||||
xh = F.conv3d(x, hh.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1))
|
||||
|
||||
# Handles spatial axes.
|
||||
xll = F.conv3d(xl, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1))
|
||||
xlh = F.conv3d(xl, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1))
|
||||
xhl = F.conv3d(xh, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1))
|
||||
xhh = F.conv3d(xh, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1))
|
||||
|
||||
xlll = F.conv3d(xll, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xllh = F.conv3d(xll, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xlhl = F.conv3d(xlh, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xlhh = F.conv3d(xlh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xhll = F.conv3d(xhl, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xhlh = F.conv3d(xhl, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xhhl = F.conv3d(xhh, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
xhhh = F.conv3d(xhh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2))
|
||||
|
||||
out = torch.cat([xlll, xllh, xlhl, xlhh, xhll, xhlh, xhhl, xhhh], dim=1)
|
||||
if rescale:
|
||||
out = out / (2 * torch.sqrt(torch.tensor(2.0)))
|
||||
return out
|
||||
|
||||
def _haar(self, x):
|
||||
xi, xv = torch.split(x, [1, x.shape[2] - 1], dim=2)
|
||||
x = torch.cat([xi.repeat_interleave(self.patch_size, dim=2), xv], dim=2)
|
||||
for _ in self.range:
|
||||
x = self._dwt(x, "haar", rescale=True)
|
||||
return x
|
||||
|
||||
def _arrange(self, x):
|
||||
xi, xv = torch.split(x, [1, x.shape[2] - 1], dim=2)
|
||||
x = torch.cat([xi.repeat_interleave(self.patch_size, dim=2), xv], dim=2)
|
||||
x = rearrange(
|
||||
x,
|
||||
"b c (t p1) (h p2) (w p3) -> b (c p1 p2 p3) t h w",
|
||||
p1=self.patch_size,
|
||||
p2=self.patch_size,
|
||||
p3=self.patch_size,
|
||||
).contiguous()
|
||||
return x
|
||||
|
||||
|
||||
class UnPatcher(torch.nn.Module):
|
||||
"""A module to convert patches into image tensorsusing torch operations.
|
||||
|
||||
The main difference from `class Unpatching` is that this module implements
|
||||
all operations using torch, rather than python or numpy, for efficiency purpose.
|
||||
|
||||
It's bit-wise identical to the Unpatching module outputs, with the added
|
||||
benefit of being torch.jit scriptable.
|
||||
"""
|
||||
|
||||
def __init__(self, patch_size=1, patch_method="haar"):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.patch_method = patch_method
|
||||
self.register_buffer(
|
||||
"wavelets", _WAVELETS[patch_method], persistent=_PERSISTENT
|
||||
)
|
||||
self.range = range(int(torch.log2(torch.tensor(self.patch_size)).item()))
|
||||
self.register_buffer(
|
||||
"_arange",
|
||||
torch.arange(_WAVELETS[patch_method].shape[0]),
|
||||
persistent=_PERSISTENT,
|
||||
)
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, x):
|
||||
if self.patch_method == "haar":
|
||||
return self._ihaar(x)
|
||||
elif self.patch_method == "rearrange":
|
||||
return self._iarrange(x)
|
||||
else:
|
||||
raise ValueError("Unknown patch method: " + self.patch_method)
|
||||
|
||||
def _idwt(self, x, wavelet="haar", mode="reflect", rescale=False):
|
||||
dtype = x.dtype
|
||||
h = self.wavelets
|
||||
n = h.shape[0]
|
||||
|
||||
g = x.shape[1] // 4
|
||||
hl = h.flip([0]).reshape(1, 1, -1).repeat([g, 1, 1])
|
||||
hh = (h * ((-1) ** self._arange)).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hh = hh.to(dtype=dtype)
|
||||
hl = hl.to(dtype=dtype)
|
||||
|
||||
xll, xlh, xhl, xhh = torch.chunk(x.to(dtype), 4, dim=1)
|
||||
|
||||
# Inverse transform.
|
||||
yl = torch.nn.functional.conv_transpose2d(
|
||||
xll, hl.unsqueeze(3), groups=g, stride=(2, 1), padding=(n - 2, 0)
|
||||
)
|
||||
yl += torch.nn.functional.conv_transpose2d(
|
||||
xlh, hh.unsqueeze(3), groups=g, stride=(2, 1), padding=(n - 2, 0)
|
||||
)
|
||||
yh = torch.nn.functional.conv_transpose2d(
|
||||
xhl, hl.unsqueeze(3), groups=g, stride=(2, 1), padding=(n - 2, 0)
|
||||
)
|
||||
yh += torch.nn.functional.conv_transpose2d(
|
||||
xhh, hh.unsqueeze(3), groups=g, stride=(2, 1), padding=(n - 2, 0)
|
||||
)
|
||||
y = torch.nn.functional.conv_transpose2d(
|
||||
yl, hl.unsqueeze(2), groups=g, stride=(1, 2), padding=(0, n - 2)
|
||||
)
|
||||
y += torch.nn.functional.conv_transpose2d(
|
||||
yh, hh.unsqueeze(2), groups=g, stride=(1, 2), padding=(0, n - 2)
|
||||
)
|
||||
|
||||
if rescale:
|
||||
y = y * 2
|
||||
return y
|
||||
|
||||
def _ihaar(self, x):
|
||||
for _ in self.range:
|
||||
x = self._idwt(x, "haar", rescale=True)
|
||||
return x
|
||||
|
||||
def _iarrange(self, x):
|
||||
x = rearrange(
|
||||
x,
|
||||
"b (c p1 p2) h w -> b c (h p1) (w p2)",
|
||||
p1=self.patch_size,
|
||||
p2=self.patch_size,
|
||||
)
|
||||
return x
|
||||
|
||||
|
||||
class UnPatcher3D(UnPatcher):
|
||||
"""A 3D inverse discrete wavelet transform for video wavelet decompositions."""
|
||||
|
||||
def __init__(self, patch_size=1, patch_method="haar"):
|
||||
super().__init__(patch_method=patch_method, patch_size=patch_size)
|
||||
|
||||
def _idwt(self, x, wavelet="haar", mode="reflect", rescale=False):
|
||||
dtype = x.dtype
|
||||
h = self.wavelets
|
||||
n = h.shape[0]
|
||||
|
||||
g = x.shape[1] // 8 # split into 8 spatio-temporal filtered tesnors.
|
||||
hl = h.flip([0]).reshape(1, 1, -1).repeat([g, 1, 1])
|
||||
hh = (h * ((-1) ** self._arange)).reshape(1, 1, -1).repeat(g, 1, 1)
|
||||
hl = hl.to(dtype=dtype)
|
||||
hh = hh.to(dtype=dtype)
|
||||
|
||||
xlll, xllh, xlhl, xlhh, xhll, xhlh, xhhl, xhhh = torch.chunk(x, 8, dim=1)
|
||||
|
||||
# Height height transposed convolutions.
|
||||
xll = F.conv_transpose3d(
|
||||
xlll, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
xll += F.conv_transpose3d(
|
||||
xllh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
|
||||
xlh = F.conv_transpose3d(
|
||||
xlhl, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
xlh += F.conv_transpose3d(
|
||||
xlhh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
|
||||
xhl = F.conv_transpose3d(
|
||||
xhll, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
xhl += F.conv_transpose3d(
|
||||
xhlh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
|
||||
xhh = F.conv_transpose3d(
|
||||
xhhl, hl.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
xhh += F.conv_transpose3d(
|
||||
xhhh, hh.unsqueeze(2).unsqueeze(3), groups=g, stride=(1, 1, 2)
|
||||
)
|
||||
|
||||
# Handles width transposed convolutions.
|
||||
xl = F.conv_transpose3d(
|
||||
xll, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)
|
||||
)
|
||||
xl += F.conv_transpose3d(
|
||||
xlh, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)
|
||||
)
|
||||
xh = F.conv_transpose3d(
|
||||
xhl, hl.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)
|
||||
)
|
||||
xh += F.conv_transpose3d(
|
||||
xhh, hh.unsqueeze(2).unsqueeze(4), groups=g, stride=(1, 2, 1)
|
||||
)
|
||||
|
||||
# Handles time axis transposed convolutions.
|
||||
x = F.conv_transpose3d(
|
||||
xl, hl.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1)
|
||||
)
|
||||
x += F.conv_transpose3d(
|
||||
xh, hh.unsqueeze(3).unsqueeze(4), groups=g, stride=(2, 1, 1)
|
||||
)
|
||||
|
||||
if rescale:
|
||||
x = x * (2 * torch.sqrt(torch.tensor(2.0)))
|
||||
return x
|
||||
|
||||
def _ihaar(self, x):
|
||||
for _ in self.range:
|
||||
x = self._idwt(x, "haar", rescale=True)
|
||||
x = x[:, :, self.patch_size - 1 :, ...]
|
||||
return x
|
||||
|
||||
def _iarrange(self, x):
|
||||
x = rearrange(
|
||||
x,
|
||||
"b (c p1 p2 p3) t h w -> b c (t p1) (h p2) (w p3)",
|
||||
p1=self.patch_size,
|
||||
p2=self.patch_size,
|
||||
p3=self.patch_size,
|
||||
)
|
||||
x = x[:, :, self.patch_size - 1 :, ...]
|
||||
return x
|
||||
546
tools/decode/vendor/cosmos_tokenizer/modules/quantizers.py
vendored
Normal file
546
tools/decode/vendor/cosmos_tokenizer/modules/quantizers.py
vendored
Normal file
@@ -0,0 +1,546 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Quantizers for discrete image and video tokenization."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import reduce
|
||||
from loguru import logger as logging
|
||||
|
||||
from cosmos_tokenizer.modules.utils import (
|
||||
default,
|
||||
entropy,
|
||||
pack_one,
|
||||
rearrange,
|
||||
round_ste,
|
||||
unpack_one,
|
||||
)
|
||||
|
||||
|
||||
class ResidualFSQuantizer(nn.Module):
|
||||
"""Residual Finite Scalar Quantization
|
||||
|
||||
Follows Algorithm 1. in https://arxiv.org/pdf/2107.03312.pdf
|
||||
"""
|
||||
|
||||
def __init__(self, levels: list[int], num_quantizers: int, **ignore_kwargs):
|
||||
super().__init__()
|
||||
self.dtype = ignore_kwargs.get("dtype", torch.float32)
|
||||
self.layers = nn.ModuleList(
|
||||
[FSQuantizer(levels=levels) for _ in range(num_quantizers)]
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
indices_stack = []
|
||||
residual = x
|
||||
quantized_out = 0
|
||||
loss_out = 0
|
||||
for i, layer in enumerate(self.layers):
|
||||
quant_indices, z, loss = layer(residual)
|
||||
indices_stack.append(quant_indices)
|
||||
residual = residual - z.detach()
|
||||
quantized_out = quantized_out + z
|
||||
loss_out = loss_out + loss
|
||||
self.residual = residual
|
||||
indices = torch.stack(indices_stack, dim=1)
|
||||
return indices, quantized_out.to(self.dtype), loss_out.to(self.dtype)
|
||||
|
||||
def indices_to_codes(self, indices_stack: torch.Tensor) -> torch.Tensor:
|
||||
quantized_out = 0
|
||||
for layer, indices in zip(self.layers, indices_stack.transpose(0, 1)):
|
||||
quantized_out += layer.indices_to_codes(indices)
|
||||
return quantized_out
|
||||
|
||||
|
||||
class FSQuantizer(nn.Module):
|
||||
"""Finite Scalar Quantization: VQ-VAE Made Simple - https://arxiv.org/abs/2309.15505
|
||||
|
||||
Code adapted from Jax version in Appendix A.1.
|
||||
|
||||
Adapted from: https://github.com/lucidrains/vector-quantize-pytorch/blob/9502a1f447876d53fd37685b226bf28f250dc4a3/
|
||||
vector_quantize_pytorch/finite_scalar_quantization.py
|
||||
[Copyright (c) 2020 Phil Wang]
|
||||
https://github.com/lucidrains/vector-quantize-pytorch/blob/9502a1f447876d53fd37685b226bf28f250dc4a3/LICENSE
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
levels: list[int],
|
||||
dim: Optional[int] = None,
|
||||
num_codebooks=1,
|
||||
keep_num_codebooks_dim: Optional[bool] = None,
|
||||
scale: Optional[float] = None,
|
||||
**ignore_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.dtype = ignore_kwargs.get("dtype", torch.bfloat16)
|
||||
_levels = torch.tensor(levels, dtype=torch.int32)
|
||||
self.register_buffer("_levels", _levels, persistent=False)
|
||||
|
||||
_basis = torch.cumprod(
|
||||
torch.tensor([1] + levels[:-1]), dim=0, dtype=torch.int32
|
||||
)
|
||||
self.register_buffer("_basis", _basis, persistent=False)
|
||||
|
||||
self.scale = scale
|
||||
|
||||
codebook_dim = len(levels)
|
||||
self.codebook_dim = codebook_dim
|
||||
|
||||
effective_codebook_dim = codebook_dim * num_codebooks
|
||||
self.num_codebooks = num_codebooks
|
||||
self.effective_codebook_dim = effective_codebook_dim
|
||||
|
||||
keep_num_codebooks_dim = default(keep_num_codebooks_dim, num_codebooks > 1)
|
||||
assert not (num_codebooks > 1 and not keep_num_codebooks_dim)
|
||||
self.keep_num_codebooks_dim = keep_num_codebooks_dim
|
||||
|
||||
self.dim = default(dim, len(_levels) * num_codebooks)
|
||||
|
||||
has_projections = self.dim != effective_codebook_dim
|
||||
self.project_in = (
|
||||
nn.Linear(self.dim, effective_codebook_dim)
|
||||
if has_projections
|
||||
else nn.Identity()
|
||||
)
|
||||
self.project_out = (
|
||||
nn.Linear(effective_codebook_dim, self.dim)
|
||||
if has_projections
|
||||
else nn.Identity()
|
||||
)
|
||||
self.has_projections = has_projections
|
||||
|
||||
self.codebook_size = self._levels.prod().item()
|
||||
|
||||
implicit_codebook = self.indices_to_codes(
|
||||
torch.arange(self.codebook_size), project_out=False
|
||||
)
|
||||
self.register_buffer("implicit_codebook", implicit_codebook, persistent=False)
|
||||
|
||||
def bound(self, z: torch.Tensor, eps: float = 1e-3) -> torch.Tensor:
|
||||
"""Bound `z`, an array of shape (..., d)."""
|
||||
half_l = (self._levels - 1) * (1 + eps) / 2
|
||||
offset = torch.where(self._levels % 2 == 0, 0.5, 0.0)
|
||||
shift = (offset / half_l).atanh()
|
||||
return (z + shift).tanh() * half_l - offset
|
||||
|
||||
def quantize(self, z: torch.Tensor) -> torch.Tensor:
|
||||
"""Quantizes z, returns quantized zhat, same shape as z."""
|
||||
quantized = round_ste(self.bound(z))
|
||||
half_width = self._levels // 2 # Renormalize to [-1, 1].
|
||||
return quantized / half_width
|
||||
|
||||
def _scale_and_shift(self, zhat_normalized: torch.Tensor) -> torch.Tensor:
|
||||
half_width = self._levels // 2
|
||||
return (zhat_normalized * half_width) + half_width
|
||||
|
||||
def _scale_and_shift_inverse(self, zhat: torch.Tensor) -> torch.Tensor:
|
||||
half_width = self._levels // 2
|
||||
return (zhat - half_width) / half_width
|
||||
|
||||
def codes_to_indices(self, zhat: torch.Tensor) -> torch.Tensor:
|
||||
"""Converts a `code` to an index in the codebook."""
|
||||
assert zhat.shape[-1] == self.codebook_dim
|
||||
zhat = self._scale_and_shift(zhat).float()
|
||||
return (zhat * self._basis).sum(dim=-1).to(torch.int32)
|
||||
|
||||
def indices_to_codes(self, indices: torch.Tensor, project_out=True) -> torch.Tensor:
|
||||
"""Inverse of `codes_to_indices`."""
|
||||
is_img_or_video = indices.ndim >= (3 + int(self.keep_num_codebooks_dim))
|
||||
indices = rearrange(indices, "... -> ... 1")
|
||||
codes_non_centered = (indices // self._basis) % self._levels
|
||||
codes = self._scale_and_shift_inverse(codes_non_centered)
|
||||
|
||||
if self.keep_num_codebooks_dim:
|
||||
codes = rearrange(codes, "... c d -> ... (c d)")
|
||||
|
||||
if project_out:
|
||||
codes = self.project_out(codes)
|
||||
|
||||
if is_img_or_video:
|
||||
codes = rearrange(codes, "b ... d -> b d ...")
|
||||
|
||||
return codes.to(self.dtype)
|
||||
|
||||
def forward(self, z: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
einstein notation
|
||||
b - batch
|
||||
n - sequence (or flattened spatial dimensions)
|
||||
d - feature dimension, which is also log2(codebook size)
|
||||
c - number of codebook dim
|
||||
"""
|
||||
is_img_or_video = z.ndim >= 4
|
||||
|
||||
# standardize image or video into (batch, seq, dimension)
|
||||
|
||||
if is_img_or_video:
|
||||
z = rearrange(z, "b d ... -> b ... d")
|
||||
z, ps = pack_one(z, "b * d")
|
||||
|
||||
assert (
|
||||
z.shape[-1] == self.dim
|
||||
), f"expected dimension of {self.dim} but found dimension of {z.shape[-1]}"
|
||||
|
||||
z = self.project_in(z)
|
||||
|
||||
z = rearrange(z, "b n (c d) -> b n c d", c=self.num_codebooks)
|
||||
|
||||
codes = self.quantize(z)
|
||||
indices = self.codes_to_indices(codes)
|
||||
|
||||
codes = rearrange(codes, "b n c d -> b n (c d)")
|
||||
|
||||
out = self.project_out(codes)
|
||||
|
||||
# reconstitute image or video dimensions
|
||||
|
||||
if is_img_or_video:
|
||||
out = unpack_one(out, ps, "b * d")
|
||||
out = rearrange(out, "b ... d -> b d ...")
|
||||
indices = unpack_one(indices, ps, "b * c")
|
||||
dummy_loss = torch.zeros_like(out.mean(dim=[1, 2, 3], keepdim=True))
|
||||
else:
|
||||
dummy_loss = torch.zeros_like(out.mean(dim=[1, 2], keepdim=True)).unsqueeze(
|
||||
1
|
||||
)
|
||||
|
||||
if not self.keep_num_codebooks_dim:
|
||||
indices = rearrange(indices, "... 1 -> ...")
|
||||
|
||||
return (indices, out.to(self.dtype), dummy_loss)
|
||||
|
||||
|
||||
class VectorQuantizer(nn.Module):
|
||||
"""Improved version over VectorQuantizer. Mostly
|
||||
avoids costly matrix multiplications and allows for post-hoc remapping of indices.
|
||||
|
||||
Adapted from: https://github.com/CompVis/taming-transformers/blob/3ba01b241669f5ade541ce990f7650a3b8f65318/
|
||||
taming/modules/vqvae/quantize.py
|
||||
|
||||
[Copyright (c) 2020 Patrick Esser and Robin Rombach and Björn Ommer]
|
||||
https://github.com/CompVis/taming-transformers/blob/3ba01b241669f5ade541ce990f7650a3b8f65318/License.txt
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_embeddings: int,
|
||||
embedding_dim: int,
|
||||
beta: float = 0.25,
|
||||
remap: str = None,
|
||||
unknown_index: str = "random",
|
||||
sane_index_shape: bool = False,
|
||||
legacy: bool = True,
|
||||
use_norm=False,
|
||||
**ignore_kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.n_e = num_embeddings
|
||||
self.e_dim = embedding_dim
|
||||
self.beta = beta
|
||||
self.legacy = legacy
|
||||
self.norm = lambda x: F.normalize(x, dim=-1) if use_norm else x
|
||||
|
||||
self.embedding = nn.Embedding(self.n_e, self.e_dim)
|
||||
self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e)
|
||||
|
||||
self.remap = remap
|
||||
if self.remap is not None:
|
||||
self.register_buffer("used", torch.tensor(np.load(self.remap)))
|
||||
self.re_embed = self.used.shape[0]
|
||||
self.unknown_index = unknown_index
|
||||
if self.unknown_index == "extra":
|
||||
self.unknown_index = self.re_embed
|
||||
self.re_embed = self.re_embed + 1
|
||||
print(
|
||||
f"Remapping {self.n_e} indices to {self.re_embed} indices. "
|
||||
f"Using {self.unknown_index} for unknown indices."
|
||||
)
|
||||
else:
|
||||
self.re_embed = num_embeddings
|
||||
|
||||
self.sane_index_shape = sane_index_shape
|
||||
self.dtype = ignore_kwargs.get("dtype", torch.float32)
|
||||
|
||||
def remap_to_used(self, inds):
|
||||
ishape = inds.shape
|
||||
assert len(ishape) > 1
|
||||
inds = inds.reshape(ishape[0], -1)
|
||||
used = self.used.to(inds)
|
||||
match = (inds[:, :, None] == used[None, None, ...]).long()
|
||||
new = match.argmax(-1)
|
||||
unknown = match.sum(2) < 1
|
||||
if self.unknown_index == "random":
|
||||
new[unknown] = torch.randint(0, self.re_embed, size=new[unknown].shape).to(
|
||||
device=new.device
|
||||
)
|
||||
else:
|
||||
new[unknown] = self.unknown_index
|
||||
return new.reshape(ishape)
|
||||
|
||||
def unmap_to_all(self, inds):
|
||||
ishape = inds.shape
|
||||
assert len(ishape) > 1
|
||||
inds = inds.reshape(ishape[0], -1)
|
||||
used = self.used.to(inds)
|
||||
if self.re_embed > self.used.shape[0]: # extra token
|
||||
inds[inds >= self.used.shape[0]] = 0 # simply set to zero
|
||||
back = torch.gather(used[None, :][inds.shape[0] * [0], :], 1, inds)
|
||||
return back.reshape(ishape)
|
||||
|
||||
def forward(self, z, temp=None, rescale_logits=False, return_logits=False):
|
||||
assert temp is None or temp == 1.0, "Only for interface compatible with Gumbel"
|
||||
assert rescale_logits is False, "Only for interface compatible with Gumbel"
|
||||
assert return_logits is False, "Only for interface compatible with Gumbel"
|
||||
z = rearrange(z, "b c h w -> b h w c").contiguous()
|
||||
z_flattened = z.view(-1, self.e_dim)
|
||||
|
||||
d = (
|
||||
torch.sum(z_flattened**2, dim=1, keepdim=True)
|
||||
+ torch.sum(self.embedding.weight**2, dim=1)
|
||||
- 2
|
||||
* torch.einsum(
|
||||
"bd,dn->bn",
|
||||
z_flattened,
|
||||
rearrange(self.embedding.weight, "n d -> d n"),
|
||||
)
|
||||
)
|
||||
|
||||
encoding_indices = torch.argmin(d, dim=1).unsqueeze(1)
|
||||
encodings = torch.zeros(encoding_indices.shape[0], self.n_e, device=z.device)
|
||||
encodings.scatter_(1, encoding_indices, 1)
|
||||
z_q = torch.matmul(encodings, self.embedding.weight).view(z.shape)
|
||||
min_encodings = None
|
||||
|
||||
z_q, z = self.norm(z_q), self.norm(z)
|
||||
|
||||
# compute loss for embedding
|
||||
commit_loss = torch.mean((z_q - z.detach()) ** 2, dim=[1, 2, 3], keepdim=True)
|
||||
emb_loss = torch.mean((z_q.detach() - z) ** 2, dim=[1, 2, 3], keepdim=True)
|
||||
if not self.legacy:
|
||||
loss = self.beta * emb_loss + commit_loss
|
||||
else:
|
||||
loss = emb_loss + self.beta * commit_loss
|
||||
|
||||
# preserve gradients
|
||||
z_q = z + (z_q - z).detach()
|
||||
avg_probs = torch.mean(encodings, dim=0)
|
||||
perplexity = torch.exp(-torch.sum(avg_probs * torch.log(avg_probs + 1e-10)))
|
||||
|
||||
# reshape back to match original input shape
|
||||
z_q = rearrange(z_q, "b h w c -> b c h w").contiguous()
|
||||
|
||||
if self.remap is not None:
|
||||
min_encoding_indices = encoding_indices.squeeze(1).reshape(
|
||||
z.shape[0], -1
|
||||
) # add batch axis
|
||||
min_encoding_indices = self.remap_to_used(encoding_indices.squeeze(1))
|
||||
min_encoding_indices = min_encoding_indices.reshape(-1, 1) # flatten
|
||||
|
||||
if self.sane_index_shape:
|
||||
min_encoding_indices = min_encoding_indices.reshape(
|
||||
z_q.shape[0], z_q.shape[2], z_q.shape[3]
|
||||
)
|
||||
|
||||
# TODO: return (indices, z_q, loss)
|
||||
return (
|
||||
z_q,
|
||||
loss,
|
||||
(
|
||||
encoding_indices.squeeze(1),
|
||||
min_encodings,
|
||||
commit_loss.mean().detach(),
|
||||
self.beta * emb_loss.mean().detach(),
|
||||
perplexity.mean().detach(),
|
||||
),
|
||||
)
|
||||
|
||||
def get_codebook_entry(self, indices, shape):
|
||||
# shape specifying (batch, height, width, channel)
|
||||
if self.remap is not None:
|
||||
indices = indices.reshape(shape[0], -1) # add batch axis
|
||||
indices = self.unmap_to_all(indices)
|
||||
indices = indices.reshape(-1) # flatten again
|
||||
|
||||
# get quantized latent vectors
|
||||
z_q = self.embedding(indices)
|
||||
|
||||
if shape is not None:
|
||||
z_q = z_q.view(shape)
|
||||
# reshape back to match original input shape
|
||||
z_q = z_q.permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
return z_q
|
||||
|
||||
|
||||
class LFQuantizer(nn.Module):
|
||||
"""Lookup-Free Quantization
|
||||
|
||||
Adapted from: https://github.com/lucidrains/vector-quantize-pytorch/blob/9502a1f447876d53fd37685b226bf28f250dc4a3/
|
||||
vector_quantize_pytorch/lookup_free_quantization.py
|
||||
[Copyright (c) 2020 Phil Wang]
|
||||
https://github.com/lucidrains/vector-quantize-pytorch/blob/9502a1f447876d53fd37685b226bf28f250dc4a3/LICENSE
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
codebook_size: int,
|
||||
codebook_dim: int,
|
||||
embed_dim: Optional[int] = None, # if None, use codebook_dim
|
||||
entropy_loss_weight=0.1,
|
||||
commitment_loss_weight=0.25,
|
||||
default_temp: float = 0.01,
|
||||
entropy_loss: bool = False,
|
||||
**ignore_kwargs,
|
||||
):
|
||||
"""Lookup-Free Quantization
|
||||
|
||||
Args:
|
||||
codebook_size (int): The number of entries in the codebook.
|
||||
codebook_dim (int): The number of bits in each code.
|
||||
embed_dim (Optional[int], optional): The dimension of the input embedding. Defaults to None.
|
||||
entropy_loss_weight (float, optional): Whether to use entropy loss. Defaults to 0.1.
|
||||
commitment_loss_weight (float, optional): Weight for commitment loss. Defaults to 0.25.
|
||||
default_temp (float, optional): The temprature to use. Defaults to 0.01.
|
||||
entropy_loss (bool, optional): Flag for entropy loss. Defaults to False.
|
||||
"""
|
||||
super().__init__()
|
||||
self.entropy_loss = entropy_loss
|
||||
self.codebook_dim = codebook_dim
|
||||
self.default_temp = default_temp
|
||||
self.entrop_loss_weight = entropy_loss_weight
|
||||
self.commitment_loss_weight = commitment_loss_weight
|
||||
embed_dim = embed_dim or codebook_dim
|
||||
|
||||
has_projections = embed_dim != codebook_dim
|
||||
self.project_in = (
|
||||
nn.Linear(embed_dim, codebook_dim) if has_projections else nn.Identity()
|
||||
)
|
||||
self.project_out = (
|
||||
nn.Linear(codebook_dim, embed_dim) if has_projections else nn.Identity()
|
||||
)
|
||||
logging.info(
|
||||
f"LFQ: has_projections={has_projections}, dim_in={embed_dim}, codebook_dim={codebook_dim}"
|
||||
)
|
||||
|
||||
self.dtype = ignore_kwargs.get("dtype", torch.float32)
|
||||
|
||||
if entropy_loss:
|
||||
assert (
|
||||
2**codebook_dim == codebook_size
|
||||
), "codebook size must be 2 ** codebook_dim"
|
||||
self.codebook_size = codebook_size
|
||||
|
||||
self.register_buffer(
|
||||
"mask",
|
||||
2 ** torch.arange(codebook_dim - 1, -1, -1),
|
||||
persistent=False,
|
||||
)
|
||||
self.register_buffer("zero", torch.tensor(0.0), persistent=False)
|
||||
|
||||
all_codes = torch.arange(codebook_size)
|
||||
bits = ((all_codes[..., None].int() & self.mask) != 0).float()
|
||||
codebook = 2 * bits - 1.0
|
||||
|
||||
self.register_buffer(
|
||||
"codebook", codebook, persistent=False
|
||||
) # [codebook_size, codebook_dim]
|
||||
|
||||
def forward(self, z: torch.Tensor, temp: float = None) -> torch.Tensor:
|
||||
temp = temp or self.default_temp
|
||||
|
||||
z = rearrange(z, "b d ... -> b ... d")
|
||||
z, ps = pack_one(z, "b * d")
|
||||
z = self.project_in(z)
|
||||
|
||||
# split out number of codebooks
|
||||
z = rearrange(z, "b n (c d) -> b n c d", c=self.num_codebooks)
|
||||
|
||||
# quantization
|
||||
original_input = z
|
||||
|
||||
codebook_value = torch.ones_like(z)
|
||||
z_q = torch.where(z > 0, codebook_value, -codebook_value)
|
||||
|
||||
# preserve gradients
|
||||
z_q = z + (z_q - z).detach()
|
||||
|
||||
# commit loss
|
||||
commit_loss = ((original_input - z_q.detach()) ** 2).mean(dim=[1, 2, 3])
|
||||
|
||||
z_q = rearrange(z_q, "b n c d -> b n (c d)")
|
||||
z_q = self.project_out(z_q)
|
||||
|
||||
# reshape
|
||||
z_q = unpack_one(z_q, ps, "b * d")
|
||||
z_q = rearrange(z_q, "b ... d -> b d ...")
|
||||
|
||||
loss = self.commitment_loss_weight * commit_loss
|
||||
|
||||
# entropy loss (eq-5)
|
||||
if self.entropy_loss:
|
||||
# indices
|
||||
indices = reduce((z > 0).int() * self.mask.int(), "b n c d -> b n c", "sum")
|
||||
indices = unpack_one(indices, ps, "b * c")
|
||||
indices = rearrange(indices, "... 1 -> ...")
|
||||
|
||||
distance = -2 * torch.einsum(
|
||||
"... i d, j d -> ... i j",
|
||||
original_input,
|
||||
self.codebook.to(original_input.dtype),
|
||||
)
|
||||
prob = (-distance / temp).softmax(dim=-1)
|
||||
per_sample_entropy = entropy(prob).mean(dim=[1, 2])
|
||||
avg_prob = reduce(prob, "... c d -> c d", "mean")
|
||||
codebook_entropy = entropy(avg_prob).mean()
|
||||
entropy_aux_loss = per_sample_entropy - codebook_entropy
|
||||
|
||||
loss += self.entrop_loss_weight * entropy_aux_loss
|
||||
|
||||
# TODO: return (indices, z_q, loss)
|
||||
return (
|
||||
z_q,
|
||||
loss.unsqueeze(1).unsqueeze(1).unsqueeze(1),
|
||||
(
|
||||
indices,
|
||||
self.commitment_loss_weight * commit_loss.mean().detach(),
|
||||
self.entrop_loss_weight * entropy_aux_loss.mean().detach(),
|
||||
self.entrop_loss_weight * per_sample_entropy.mean().detach(),
|
||||
self.entrop_loss_weight * codebook_entropy.mean().detach(),
|
||||
),
|
||||
)
|
||||
else:
|
||||
return (
|
||||
z_q,
|
||||
loss.unsqueeze(1).unsqueeze(1).unsqueeze(1),
|
||||
self.commitment_loss_weight * commit_loss.mean().detach(),
|
||||
)
|
||||
|
||||
|
||||
class InvQuantizerJit(nn.Module):
|
||||
"""Use for decoder_jit to trace quantizer in discrete tokenizer"""
|
||||
|
||||
def __init__(self, quantizer):
|
||||
super().__init__()
|
||||
self.quantizer = quantizer
|
||||
|
||||
def forward(self, indices: torch.Tensor):
|
||||
codes = self.quantizer.indices_to_codes(indices)
|
||||
return codes.to(self.quantizer.dtype)
|
||||
117
tools/decode/vendor/cosmos_tokenizer/modules/utils.py
vendored
Normal file
117
tools/decode/vendor/cosmos_tokenizer/modules/utils.py
vendored
Normal file
@@ -0,0 +1,117 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Shared utilities for the networks module."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from einops import pack, rearrange, unpack
|
||||
|
||||
|
||||
def time2batch(x: torch.Tensor) -> tuple[torch.Tensor, int]:
|
||||
batch_size = x.shape[0]
|
||||
return rearrange(x, "b c t h w -> (b t) c h w"), batch_size
|
||||
|
||||
|
||||
def batch2time(x: torch.Tensor, batch_size: int) -> torch.Tensor:
|
||||
return rearrange(x, "(b t) c h w -> b c t h w", b=batch_size)
|
||||
|
||||
|
||||
def space2batch(x: torch.Tensor) -> tuple[torch.Tensor, int]:
|
||||
batch_size, height = x.shape[0], x.shape[-2]
|
||||
return rearrange(x, "b c t h w -> (b h w) c t"), batch_size, height
|
||||
|
||||
|
||||
def batch2space(x: torch.Tensor, batch_size: int, height: int) -> torch.Tensor:
|
||||
return rearrange(x, "(b h w) c t -> b c t h w", b=batch_size, h=height)
|
||||
|
||||
|
||||
def cast_tuple(t: Any, length: int = 1) -> Any:
|
||||
return t if isinstance(t, tuple) else ((t,) * length)
|
||||
|
||||
|
||||
def replication_pad(x):
|
||||
return torch.cat([x[:, :, :1, ...], x], dim=2)
|
||||
|
||||
|
||||
def divisible_by(num: int, den: int) -> bool:
|
||||
return (num % den) == 0
|
||||
|
||||
|
||||
def is_odd(n: int) -> bool:
|
||||
return not divisible_by(n, 2)
|
||||
|
||||
|
||||
def nonlinearity(x):
|
||||
return x * torch.sigmoid(x)
|
||||
|
||||
|
||||
def Normalize(in_channels, num_groups=32):
|
||||
return torch.nn.GroupNorm(
|
||||
num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True
|
||||
)
|
||||
|
||||
|
||||
class CausalNormalize(torch.nn.Module):
|
||||
def __init__(self, in_channels, num_groups=1):
|
||||
super().__init__()
|
||||
self.norm = torch.nn.GroupNorm(
|
||||
num_groups=num_groups,
|
||||
num_channels=in_channels,
|
||||
eps=1e-6,
|
||||
affine=True,
|
||||
)
|
||||
self.num_groups = num_groups
|
||||
|
||||
def forward(self, x):
|
||||
# if num_groups !=1, we apply a spatio-temporal groupnorm for backward compatibility purpose.
|
||||
# All new models should use num_groups=1, otherwise causality is not guaranteed.
|
||||
if self.num_groups == 1:
|
||||
x, batch_size = time2batch(x)
|
||||
return batch2time(self.norm(x), batch_size)
|
||||
return self.norm(x)
|
||||
|
||||
|
||||
def exists(v):
|
||||
return v is not None
|
||||
|
||||
|
||||
def default(*args):
|
||||
for arg in args:
|
||||
if exists(arg):
|
||||
return arg
|
||||
return None
|
||||
|
||||
|
||||
def pack_one(t, pattern):
|
||||
return pack([t], pattern)
|
||||
|
||||
|
||||
def unpack_one(t, ps, pattern):
|
||||
return unpack(t, ps, pattern)[0]
|
||||
|
||||
|
||||
def round_ste(z: torch.Tensor) -> torch.Tensor:
|
||||
"""Round with straight through gradients."""
|
||||
zhat = z.round()
|
||||
return z + (zhat - z).detach()
|
||||
|
||||
|
||||
def log(t, eps=1e-5):
|
||||
return t.clamp(min=eps).log()
|
||||
|
||||
|
||||
def entropy(prob):
|
||||
return (-prob * log(prob)).sum(dim=-1)
|
||||
52
tools/decode/vendor/cosmos_tokenizer/networks/__init__.py
vendored
Normal file
52
tools/decode/vendor/cosmos_tokenizer/networks/__init__.py
vendored
Normal file
@@ -0,0 +1,52 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from enum import Enum
|
||||
|
||||
from cosmos_tokenizer.networks.configs import (
|
||||
continuous_image as continuous_image_dict,
|
||||
)
|
||||
from cosmos_tokenizer.networks.configs import (
|
||||
discrete_image as discrete_image_dict,
|
||||
)
|
||||
from cosmos_tokenizer.networks.configs import (
|
||||
continuous_video as continuous_video_dict,
|
||||
)
|
||||
from cosmos_tokenizer.networks.configs import (
|
||||
discrete_video as discrete_video_dict,
|
||||
)
|
||||
|
||||
from cosmos_tokenizer.networks.continuous_image import ContinuousImageTokenizer
|
||||
from cosmos_tokenizer.networks.discrete_image import DiscreteImageTokenizer
|
||||
from cosmos_tokenizer.networks.continuous_video import (
|
||||
CausalContinuousVideoTokenizer,
|
||||
)
|
||||
from cosmos_tokenizer.networks.discrete_video import (
|
||||
CausalDiscreteVideoTokenizer,
|
||||
)
|
||||
|
||||
|
||||
class TokenizerConfigs(Enum):
|
||||
CI = continuous_image_dict
|
||||
DI = discrete_image_dict
|
||||
CV = continuous_video_dict
|
||||
DV = discrete_video_dict
|
||||
|
||||
|
||||
class TokenizerModels(Enum):
|
||||
CI = ContinuousImageTokenizer
|
||||
DI = DiscreteImageTokenizer
|
||||
CV = CausalContinuousVideoTokenizer
|
||||
DV = CausalDiscreteVideoTokenizer
|
||||
146
tools/decode/vendor/cosmos_tokenizer/networks/configs.py
vendored
Normal file
146
tools/decode/vendor/cosmos_tokenizer/networks/configs.py
vendored
Normal file
@@ -0,0 +1,146 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""The default image and video tokenizer configs."""
|
||||
|
||||
from cosmos_tokenizer.modules import (
|
||||
ContinuousFormulation,
|
||||
DiscreteQuantizer,
|
||||
EncoderType,
|
||||
DecoderType,
|
||||
Encoder3DType,
|
||||
Decoder3DType,
|
||||
)
|
||||
|
||||
continuous_image = dict(
|
||||
# The attention resolution for res blocks.
|
||||
attn_resolutions=[32],
|
||||
# The base number of channels.
|
||||
channels=128,
|
||||
# The channel multipler for each resolution.
|
||||
channels_mult=[2, 4, 4],
|
||||
dropout=0.0,
|
||||
in_channels=3,
|
||||
# The spatial compression ratio.
|
||||
spatial_compression=16,
|
||||
# The number of layers in each res block.
|
||||
num_res_blocks=2,
|
||||
out_channels=3,
|
||||
resolution=1024,
|
||||
patch_size=4,
|
||||
patch_method="haar",
|
||||
# The output latent dimension (channels).
|
||||
latent_channels=16,
|
||||
# The encoder output channels just before sampling.
|
||||
# Which is also the decoder's input channels.
|
||||
z_channels=16,
|
||||
# A factor over the z_channels, to get the total channels the encoder should output.
|
||||
# For a VAE for instance, we want to output the mean and variance, so we need 2 * z_channels.
|
||||
z_factor=1,
|
||||
name="CI",
|
||||
# What formulation to use, either "AE" or "VAE".
|
||||
# Chose VAE here, since the pre-trained ckpt were of a VAE formulation.
|
||||
formulation=ContinuousFormulation.AE.name,
|
||||
# Specify type of encoder ["Default", "LiteVAE"]
|
||||
encoder=EncoderType.Default.name,
|
||||
# Specify type of decoder ["Default"]
|
||||
decoder=DecoderType.Default.name,
|
||||
)
|
||||
|
||||
discrete_image = dict(
|
||||
# The attention resolution for res blocks.
|
||||
attn_resolutions=[32],
|
||||
# The base number of channels.
|
||||
channels=128,
|
||||
# The channel multipler for each resolution.
|
||||
channels_mult=[2, 4, 4],
|
||||
dropout=0.0,
|
||||
in_channels=3,
|
||||
# The spatial compression ratio.
|
||||
spatial_compression=16,
|
||||
# The number of layers in each res block.
|
||||
num_res_blocks=2,
|
||||
out_channels=3,
|
||||
resolution=1024,
|
||||
patch_size=4,
|
||||
patch_method="haar",
|
||||
# The encoder output channels just before sampling.
|
||||
z_channels=256,
|
||||
# A factor over the z_channels, to get the total channels the encoder should output.
|
||||
# for discrete tokenization, often we directly use the vector, so z_factor=1.
|
||||
z_factor=1,
|
||||
# The quantizer of choice, VQ, LFQ, FSQ, or ResFSQ.
|
||||
quantizer=DiscreteQuantizer.FSQ.name,
|
||||
# The embedding dimension post-quantization, which is also the input channels of the decoder.
|
||||
# Which is also the output
|
||||
embedding_dim=6,
|
||||
# The number of levels to use for fine-scalar quantization.
|
||||
levels=[8, 8, 8, 5, 5, 5],
|
||||
# The number of quantizers to use for residual fine-scalar quantization.
|
||||
num_quantizers=4,
|
||||
name="DI",
|
||||
# Specify type of encoder ["Default", "LiteVAE"]
|
||||
encoder=EncoderType.Default.name,
|
||||
# Specify type of decoder ["Default"]
|
||||
decoder=DecoderType.Default.name,
|
||||
)
|
||||
|
||||
continuous_video = dict(
|
||||
attn_resolutions=[32],
|
||||
channels=128,
|
||||
channels_mult=[2, 4, 4],
|
||||
dropout=0.0,
|
||||
in_channels=3,
|
||||
num_res_blocks=2,
|
||||
out_channels=3,
|
||||
resolution=1024,
|
||||
patch_size=4,
|
||||
patch_method="haar",
|
||||
latent_channels=16,
|
||||
z_channels=16,
|
||||
z_factor=1,
|
||||
num_groups=1,
|
||||
legacy_mode=False,
|
||||
spatial_compression=8,
|
||||
temporal_compression=8,
|
||||
formulation=ContinuousFormulation.AE.name,
|
||||
encoder=Encoder3DType.FACTORIZED.name,
|
||||
decoder=Decoder3DType.FACTORIZED.name,
|
||||
name="CV",
|
||||
)
|
||||
|
||||
discrete_video = dict(
|
||||
attn_resolutions=[32],
|
||||
channels=128,
|
||||
channels_mult=[2, 4, 4],
|
||||
dropout=0.0,
|
||||
in_channels=3,
|
||||
num_res_blocks=2,
|
||||
out_channels=3,
|
||||
resolution=1024,
|
||||
patch_size=4,
|
||||
patch_method="haar",
|
||||
z_channels=16,
|
||||
z_factor=1,
|
||||
num_groups=1,
|
||||
legacy_mode=False,
|
||||
spatial_compression=16,
|
||||
temporal_compression=8,
|
||||
quantizer=DiscreteQuantizer.FSQ.name,
|
||||
embedding_dim=6,
|
||||
levels=[8, 8, 8, 5, 5, 5],
|
||||
encoder=Encoder3DType.FACTORIZED.name,
|
||||
decoder=Decoder3DType.FACTORIZED.name,
|
||||
name="DV",
|
||||
)
|
||||
104
tools/decode/vendor/cosmos_tokenizer/networks/continuous_image.py
vendored
Normal file
104
tools/decode/vendor/cosmos_tokenizer/networks/continuous_image.py
vendored
Normal file
@@ -0,0 +1,104 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""The continuous image tokenizer with VAE or AE formulation for 2D data."""
|
||||
|
||||
from collections import OrderedDict, namedtuple
|
||||
|
||||
import torch
|
||||
from loguru import logger as logging
|
||||
from torch import nn
|
||||
|
||||
from cosmos_tokenizer.modules import (
|
||||
ContinuousFormulation,
|
||||
DecoderType,
|
||||
EncoderType,
|
||||
)
|
||||
|
||||
NetworkEval = namedtuple("NetworkEval", ["reconstructions", "posteriors", "latent"])
|
||||
|
||||
|
||||
class ContinuousImageTokenizer(nn.Module):
|
||||
def __init__(
|
||||
self, z_channels: int, z_factor: int, latent_channels: int, **kwargs
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.name = kwargs.get("name", "ContinuousImageTokenizer")
|
||||
self.latent_channels = latent_channels
|
||||
|
||||
encoder_name = kwargs.get("encoder", EncoderType.Default.name)
|
||||
self.encoder = EncoderType[encoder_name].value(
|
||||
z_channels=z_factor * z_channels, **kwargs
|
||||
)
|
||||
|
||||
decoder_name = kwargs.get("decoder", DecoderType.Default.name)
|
||||
self.decoder = DecoderType[decoder_name].value(z_channels=z_channels, **kwargs)
|
||||
|
||||
self.quant_conv = torch.nn.Conv2d(
|
||||
z_factor * z_channels, z_factor * latent_channels, 1
|
||||
)
|
||||
self.post_quant_conv = torch.nn.Conv2d(latent_channels, z_channels, 1)
|
||||
|
||||
formulation_name = kwargs.get("formulation", ContinuousFormulation.AE.name)
|
||||
self.distribution = ContinuousFormulation[formulation_name].value()
|
||||
logging.info(
|
||||
f"{self.name} based on {formulation_name} formulation, with {kwargs}."
|
||||
)
|
||||
|
||||
num_parameters = sum(param.numel() for param in self.parameters())
|
||||
logging.info(f"model={self.name}, num_parameters={num_parameters:,}")
|
||||
logging.info(
|
||||
f"z_channels={z_channels}, latent_channels={self.latent_channels}."
|
||||
)
|
||||
|
||||
def encoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("encoder", self.encoder),
|
||||
("quant_conv", self.quant_conv),
|
||||
("distribution", self.distribution),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def decoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("post_quant_conv", self.post_quant_conv),
|
||||
("decoder", self.decoder),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def last_decoder_layer(self):
|
||||
return self.decoder.conv_out
|
||||
|
||||
def encode(self, x):
|
||||
h = self.encoder(x)
|
||||
moments = self.quant_conv(h)
|
||||
return self.distribution(moments)
|
||||
|
||||
def decode(self, z):
|
||||
z = self.post_quant_conv(z)
|
||||
dec = self.decoder(z)
|
||||
return dec
|
||||
|
||||
def forward(self, input) -> dict[str, torch.Tensor] | NetworkEval:
|
||||
latent, posteriors = self.encode(input)
|
||||
dec = self.decode(latent)
|
||||
if self.training:
|
||||
return dict(reconstructions=dec, posteriors=posteriors, latent=latent)
|
||||
return NetworkEval(reconstructions=dec, posteriors=posteriors, latent=latent)
|
||||
118
tools/decode/vendor/cosmos_tokenizer/networks/continuous_video.py
vendored
Normal file
118
tools/decode/vendor/cosmos_tokenizer/networks/continuous_video.py
vendored
Normal file
@@ -0,0 +1,118 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""The causal continuous video tokenizer with VAE or AE formulation for 3D data.."""
|
||||
from collections import OrderedDict, namedtuple
|
||||
|
||||
from loguru import logger as logging
|
||||
from torch import nn
|
||||
|
||||
from cosmos_tokenizer.modules import (
|
||||
ContinuousFormulation,
|
||||
Decoder3DType,
|
||||
Encoder3DType,
|
||||
)
|
||||
from cosmos_tokenizer.modules.layers3d import CausalConv3d
|
||||
|
||||
NetworkEval = namedtuple("NetworkEval", ["reconstructions", "posteriors", "latent"])
|
||||
|
||||
|
||||
class CausalContinuousVideoTokenizer(nn.Module):
|
||||
def __init__(
|
||||
self, z_channels: int, z_factor: int, latent_channels: int, **kwargs
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.name = kwargs.get("name", "CausalContinuousVideoTokenizer")
|
||||
self.latent_channels = latent_channels
|
||||
|
||||
encoder_name = kwargs.get("encoder", Encoder3DType.BASE.name)
|
||||
self.encoder = Encoder3DType[encoder_name].value(
|
||||
z_channels=z_factor * z_channels, **kwargs
|
||||
)
|
||||
if kwargs.get("temporal_compression", 4) == 4:
|
||||
kwargs["channels_mult"] = [2, 4]
|
||||
decoder_name = kwargs.get("decoder", Decoder3DType.BASE.name)
|
||||
self.decoder = Decoder3DType[decoder_name].value(
|
||||
z_channels=z_channels, **kwargs
|
||||
)
|
||||
|
||||
self.quant_conv = CausalConv3d(
|
||||
z_factor * z_channels,
|
||||
z_factor * latent_channels,
|
||||
kernel_size=1,
|
||||
padding=0,
|
||||
)
|
||||
self.post_quant_conv = CausalConv3d(
|
||||
latent_channels, z_channels, kernel_size=1, padding=0
|
||||
)
|
||||
|
||||
formulation_name = kwargs.get("formulation", ContinuousFormulation.AE.name)
|
||||
self.distribution = ContinuousFormulation[formulation_name].value()
|
||||
logging.info(
|
||||
f"{self.name} based on {formulation_name} formulation, with {kwargs}."
|
||||
)
|
||||
|
||||
num_parameters = sum(param.numel() for param in self.parameters())
|
||||
logging.info(f"model={self.name}, num_parameters={num_parameters:,}")
|
||||
logging.info(
|
||||
f"z_channels={z_channels}, latent_channels={self.latent_channels}."
|
||||
)
|
||||
|
||||
def encoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("encoder", self.encoder),
|
||||
("quant_conv", self.quant_conv),
|
||||
("distribution", self.distribution),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def decoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("post_quant_conv", self.post_quant_conv),
|
||||
("decoder", self.decoder),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def last_decoder_layer(self):
|
||||
return self.decoder.conv_out
|
||||
|
||||
def encode(self, x):
|
||||
h = self.encoder(x)
|
||||
moments = self.quant_conv(h)
|
||||
return self.distribution(moments)
|
||||
|
||||
def decode(self, z):
|
||||
z = self.post_quant_conv(z)
|
||||
return self.decoder(z)
|
||||
|
||||
def forward(self, input):
|
||||
latent, posteriors = self.encode(input)
|
||||
reconstructions = self.decode(latent)
|
||||
if self.training:
|
||||
return dict(
|
||||
reconstructions=reconstructions,
|
||||
posteriors=posteriors,
|
||||
latent=latent,
|
||||
)
|
||||
return NetworkEval(
|
||||
reconstructions=reconstructions,
|
||||
posteriors=posteriors,
|
||||
latent=latent,
|
||||
)
|
||||
129
tools/decode/vendor/cosmos_tokenizer/networks/discrete_image.py
vendored
Normal file
129
tools/decode/vendor/cosmos_tokenizer/networks/discrete_image.py
vendored
Normal file
@@ -0,0 +1,129 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""The network definition for discrete image tokenization with VQ, LFQ, FSQ or ResidualFSQ."""
|
||||
from collections import OrderedDict, namedtuple
|
||||
|
||||
import torch
|
||||
from loguru import logger as logging
|
||||
from torch import nn
|
||||
|
||||
from cosmos_tokenizer.modules import DecoderType, DiscreteQuantizer, EncoderType
|
||||
from cosmos_tokenizer.modules.quantizers import InvQuantizerJit
|
||||
|
||||
NetworkEval = namedtuple("NetworkEval", ["reconstructions", "quant_loss", "quant_info"])
|
||||
|
||||
|
||||
class DiscreteImageTokenizer(nn.Module):
|
||||
def __init__(self, z_channels: int, embedding_dim: int, **kwargs) -> None:
|
||||
super().__init__()
|
||||
self.name = kwargs.get("name", "DiscreteImageTokenizer")
|
||||
self.embedding_dim = embedding_dim
|
||||
|
||||
encoder_name = kwargs.get("encoder", EncoderType.Default.name)
|
||||
self.encoder = EncoderType[encoder_name].value(z_channels=z_channels, **kwargs)
|
||||
|
||||
decoder_name = kwargs.get("decoder", DecoderType.Default.name)
|
||||
self.decoder = DecoderType[decoder_name].value(z_channels=z_channels, **kwargs)
|
||||
self.quant_conv = nn.Conv2d(z_channels, embedding_dim, 1)
|
||||
self.post_quant_conv = nn.Conv2d(embedding_dim, z_channels, 1)
|
||||
|
||||
quantizer_name = kwargs.get("quantizer", DiscreteQuantizer.RESFSQ.name)
|
||||
if quantizer_name == DiscreteQuantizer.VQ.name:
|
||||
assert (
|
||||
"num_embeddings" in kwargs
|
||||
), f"`num_embeddings` must be provided for {quantizer_name}."
|
||||
kwargs.update(dict(embedding_dim=embedding_dim))
|
||||
elif quantizer_name == DiscreteQuantizer.LFQ.name:
|
||||
assert (
|
||||
"codebook_size" in kwargs
|
||||
), f"`codebook_size` must be provided for {quantizer_name}."
|
||||
assert (
|
||||
"codebook_dim" in kwargs
|
||||
), f"`codebook_dim` must be provided for {quantizer_name}."
|
||||
elif quantizer_name == DiscreteQuantizer.FSQ.name:
|
||||
assert (
|
||||
"levels" in kwargs
|
||||
), f"`levels` must be provided for {quantizer_name}."
|
||||
elif quantizer_name == DiscreteQuantizer.RESFSQ.name:
|
||||
assert (
|
||||
"levels" in kwargs
|
||||
), f"`levels` must be provided for {quantizer_name}.name."
|
||||
assert (
|
||||
"num_quantizers" in kwargs
|
||||
), f"`num_quantizers` must be provided for {quantizer_name}."
|
||||
self.quantizer = DiscreteQuantizer[quantizer_name].value(**kwargs)
|
||||
logging.info(f"{self.name} based on {quantizer_name}-VAE, with {kwargs}.")
|
||||
|
||||
num_parameters = sum(param.numel() for param in self.parameters())
|
||||
logging.info(f"model={self.name}, num_parameters={num_parameters:,}")
|
||||
logging.info(f"z_channels={z_channels}, embedding_dim={self.embedding_dim}.")
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
setattr(self.quantizer, "dtype", kwargs.get("dtype", torch.bfloat16))
|
||||
return super(DiscreteImageTokenizer, self).to(*args, **kwargs)
|
||||
|
||||
def encoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("encoder", self.encoder),
|
||||
("quant_conv", self.quant_conv),
|
||||
("quantizer", self.quantizer),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def decoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("inv_quant", InvQuantizerJit(self.quantizer)),
|
||||
("post_quant_conv", self.post_quant_conv),
|
||||
("decoder", self.decoder),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def last_decoder_layer(self):
|
||||
return self.decoder.conv_out
|
||||
|
||||
def encode(self, x):
|
||||
h = self.encoder(x)
|
||||
h = self.quant_conv(h)
|
||||
return self.quantizer(h)
|
||||
|
||||
def decode(self, quant):
|
||||
quant = self.post_quant_conv(quant)
|
||||
return self.decoder(quant)
|
||||
|
||||
def decode_code(self, code_b):
|
||||
quant_b = self.quantizer.indices_to_codes(code_b)
|
||||
quant_b = self.post_quant_conv(quant_b)
|
||||
return self.decoder(quant_b)
|
||||
|
||||
def forward(self, input):
|
||||
quant_info, quant_codes, quant_loss = self.encode(input)
|
||||
reconstructions = self.decode(quant_codes)
|
||||
if self.training:
|
||||
return dict(
|
||||
reconstructions=reconstructions,
|
||||
quant_loss=quant_loss,
|
||||
quant_info=quant_info,
|
||||
)
|
||||
return NetworkEval(
|
||||
reconstructions=reconstructions,
|
||||
quant_loss=quant_loss,
|
||||
quant_info=quant_info,
|
||||
)
|
||||
145
tools/decode/vendor/cosmos_tokenizer/networks/discrete_video.py
vendored
Normal file
145
tools/decode/vendor/cosmos_tokenizer/networks/discrete_video.py
vendored
Normal file
@@ -0,0 +1,145 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""The network definition for discrete video tokenizer with VQ, LFQ, FSQ or ResidualFSQ. """
|
||||
from collections import OrderedDict, namedtuple
|
||||
|
||||
import torch
|
||||
from loguru import logger as logging
|
||||
from torch import nn
|
||||
|
||||
from cosmos_tokenizer.modules import (
|
||||
Decoder3DType,
|
||||
DiscreteQuantizer,
|
||||
Encoder3DType,
|
||||
)
|
||||
from cosmos_tokenizer.modules.layers3d import CausalConv3d
|
||||
from cosmos_tokenizer.modules.quantizers import InvQuantizerJit
|
||||
|
||||
NetworkEval = namedtuple("NetworkEval", ["reconstructions", "quant_loss", "quant_info"])
|
||||
|
||||
|
||||
class CausalDiscreteVideoTokenizer(nn.Module):
|
||||
def __init__(
|
||||
self, z_channels: int, z_factor: int, embedding_dim: int, **kwargs
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.name = kwargs.get("name", "CausalDiscreteVideoTokenizer")
|
||||
self.embedding_dim = embedding_dim
|
||||
|
||||
encoder_name = kwargs.get("encoder", Encoder3DType.BASE.name)
|
||||
self.encoder = Encoder3DType[encoder_name].value(
|
||||
z_channels=z_factor * z_channels, **kwargs
|
||||
)
|
||||
|
||||
decoder_name = kwargs.get("decoder", Decoder3DType.BASE.name)
|
||||
self.decoder = Decoder3DType[decoder_name].value(
|
||||
z_channels=z_channels, **kwargs
|
||||
)
|
||||
|
||||
self.quant_conv = CausalConv3d(
|
||||
z_factor * z_channels, embedding_dim, kernel_size=1, padding=0
|
||||
)
|
||||
self.post_quant_conv = CausalConv3d(
|
||||
embedding_dim, z_channels, kernel_size=1, padding=0
|
||||
)
|
||||
|
||||
quantizer_name = kwargs.get("quantizer", DiscreteQuantizer.RESFSQ.name)
|
||||
if quantizer_name == DiscreteQuantizer.VQ.name:
|
||||
assert (
|
||||
"num_embeddings" in kwargs
|
||||
), f"`num_embeddings` must be provided for {quantizer_name}."
|
||||
kwargs.update(dict(embedding_dim=embedding_dim))
|
||||
elif quantizer_name == DiscreteQuantizer.LFQ.name:
|
||||
assert (
|
||||
"codebook_size" in kwargs
|
||||
), f"`codebook_size` must be provided for {quantizer_name}."
|
||||
assert (
|
||||
"codebook_dim" in kwargs
|
||||
), f"`codebook_dim` must be provided for {quantizer_name}."
|
||||
elif quantizer_name == DiscreteQuantizer.FSQ.name:
|
||||
assert (
|
||||
"levels" in kwargs
|
||||
), f"`levels` must be provided for {quantizer_name}."
|
||||
elif quantizer_name == DiscreteQuantizer.RESFSQ.name:
|
||||
assert (
|
||||
"levels" in kwargs
|
||||
), f"`levels` must be provided for {quantizer_name}."
|
||||
assert (
|
||||
"num_quantizers" in kwargs
|
||||
), f"`num_quantizers` must be provided for {quantizer_name}."
|
||||
self.quantizer = DiscreteQuantizer[quantizer_name].value(**kwargs)
|
||||
logging.info(f"{self.name} based on {quantizer_name}-VAE, with {kwargs}.")
|
||||
|
||||
num_parameters = sum(param.numel() for param in self.parameters())
|
||||
logging.info(f"model={self.name}, num_parameters={num_parameters:,}")
|
||||
logging.info(f"z_channels={z_channels}, embedding_dim={self.embedding_dim}.")
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
setattr(self.quantizer, "dtype", kwargs.get("dtype", torch.bfloat16))
|
||||
return super(CausalDiscreteVideoTokenizer, self).to(*args, **kwargs)
|
||||
|
||||
def encoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("encoder", self.encoder),
|
||||
("quant_conv", self.quant_conv),
|
||||
("quantizer", self.quantizer),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def decoder_jit(self):
|
||||
return nn.Sequential(
|
||||
OrderedDict(
|
||||
[
|
||||
("inv_quant", InvQuantizerJit(self.quantizer)),
|
||||
("post_quant_conv", self.post_quant_conv),
|
||||
("decoder", self.decoder),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def last_decoder_layer(self):
|
||||
return self.decoder.conv_out
|
||||
|
||||
def encode(self, x):
|
||||
h = self.encoder(x)
|
||||
h = self.quant_conv(h)
|
||||
return self.quantizer(h)
|
||||
|
||||
def decode(self, quant):
|
||||
quant = self.post_quant_conv(quant)
|
||||
return self.decoder(quant)
|
||||
|
||||
def decode_code(self, code_b):
|
||||
quant_b = self.quantizer.indices_to_codes(code_b)
|
||||
quant_b = self.post_quant_conv(quant_b)
|
||||
return self.decoder(quant_b)
|
||||
|
||||
def forward(self, input):
|
||||
quant_info, quant_codes, quant_loss = self.encode(input)
|
||||
reconstructions = self.decode(quant_codes)
|
||||
if self.training:
|
||||
return dict(
|
||||
reconstructions=reconstructions,
|
||||
quant_loss=quant_loss,
|
||||
quant_info=quant_info,
|
||||
)
|
||||
return NetworkEval(
|
||||
reconstructions=reconstructions,
|
||||
quant_loss=quant_loss,
|
||||
quant_info=quant_info,
|
||||
)
|
||||
408
tools/decode/vendor/cosmos_tokenizer/utils.py
vendored
Normal file
408
tools/decode/vendor/cosmos_tokenizer/utils.py
vendored
Normal file
@@ -0,0 +1,408 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Utility functions for the inference libraries."""
|
||||
|
||||
import os
|
||||
from glob import glob
|
||||
from typing import Any
|
||||
|
||||
import mediapy as media
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from cosmos_tokenizer.networks import TokenizerModels
|
||||
|
||||
_DTYPE, _DEVICE = torch.bfloat16, "cuda"
|
||||
_UINT8_MAX_F = float(torch.iinfo(torch.uint8).max)
|
||||
_SPATIAL_ALIGN = 16
|
||||
_TEMPORAL_ALIGN = 8
|
||||
|
||||
|
||||
def load_model(
|
||||
jit_filepath: str = None,
|
||||
tokenizer_config: dict[str, Any] = None,
|
||||
device: str = "cuda",
|
||||
) -> torch.nn.Module | torch.jit.ScriptModule:
|
||||
"""Loads a torch.nn.Module from a filepath.
|
||||
|
||||
Args:
|
||||
jit_filepath: The filepath to the JIT-compiled model.
|
||||
device: The device to load the model onto, default=cuda.
|
||||
Returns:
|
||||
The JIT compiled model loaded to device and on eval mode.
|
||||
"""
|
||||
if tokenizer_config is None:
|
||||
return load_jit_model(jit_filepath, device)
|
||||
full_model, ckpts = _load_pytorch_model(jit_filepath, tokenizer_config, device)
|
||||
full_model.load_state_dict(ckpts.state_dict(), strict=False)
|
||||
return full_model.eval().to(device)
|
||||
|
||||
|
||||
def load_encoder_model(
|
||||
jit_filepath: str = None,
|
||||
tokenizer_config: dict[str, Any] = None,
|
||||
device: str = "cuda",
|
||||
) -> torch.nn.Module | torch.jit.ScriptModule:
|
||||
"""Loads a torch.nn.Module from a filepath.
|
||||
|
||||
Args:
|
||||
jit_filepath: The filepath to the JIT-compiled model.
|
||||
device: The device to load the model onto, default=cuda.
|
||||
Returns:
|
||||
The JIT compiled model loaded to device and on eval mode.
|
||||
"""
|
||||
if tokenizer_config is None:
|
||||
return load_jit_model(jit_filepath, device)
|
||||
full_model, ckpts = _load_pytorch_model(jit_filepath, tokenizer_config, device)
|
||||
encoder_model = full_model.encoder_jit()
|
||||
encoder_model.load_state_dict(ckpts.state_dict(), strict=False)
|
||||
return encoder_model.eval().to(device)
|
||||
|
||||
|
||||
def load_decoder_model(
|
||||
jit_filepath: str = None,
|
||||
tokenizer_config: dict[str, Any] = None,
|
||||
device: str = "cuda",
|
||||
) -> torch.nn.Module | torch.jit.ScriptModule:
|
||||
"""Loads a torch.nn.Module from a filepath.
|
||||
|
||||
Args:
|
||||
jit_filepath: The filepath to the JIT-compiled model.
|
||||
device: The device to load the model onto, default=cuda.
|
||||
Returns:
|
||||
The JIT compiled model loaded to device and on eval mode.
|
||||
"""
|
||||
if tokenizer_config is None:
|
||||
return load_jit_model(jit_filepath, device)
|
||||
full_model, ckpts = _load_pytorch_model(jit_filepath, tokenizer_config, device)
|
||||
decoder_model = full_model.decoder_jit()
|
||||
decoder_model.load_state_dict(ckpts.state_dict(), strict=False)
|
||||
return decoder_model.eval().to(device)
|
||||
|
||||
|
||||
def _load_pytorch_model(
|
||||
jit_filepath: str = None, tokenizer_config: str = None, device: str = "cuda"
|
||||
) -> torch.nn.Module:
|
||||
"""Loads a torch.nn.Module from a filepath.
|
||||
|
||||
Args:
|
||||
jit_filepath: The filepath to the JIT-compiled model.
|
||||
device: The device to load the model onto, default=cuda.
|
||||
Returns:
|
||||
The JIT compiled model loaded to device and on eval mode.
|
||||
"""
|
||||
tokenizer_name = tokenizer_config["name"]
|
||||
model = TokenizerModels[tokenizer_name].value(**tokenizer_config)
|
||||
ckpts = torch.jit.load(jit_filepath, map_location=device)
|
||||
return model, ckpts
|
||||
|
||||
|
||||
def load_jit_model(
|
||||
jit_filepath: str = None, device: str = "cuda"
|
||||
) -> torch.jit.ScriptModule:
|
||||
"""Loads a torch.jit.ScriptModule from a filepath.
|
||||
|
||||
Args:
|
||||
jit_filepath: The filepath to the JIT-compiled model.
|
||||
device: The device to load the model onto, default=cuda.
|
||||
Returns:
|
||||
The JIT compiled model loaded to device and on eval mode.
|
||||
"""
|
||||
model = torch.jit.load(jit_filepath, map_location=device)
|
||||
return model.eval().to(device)
|
||||
|
||||
|
||||
def save_jit_model(
|
||||
model: torch.jit.ScriptModule | torch.jit.RecursiveScriptModule = None,
|
||||
jit_filepath: str = None,
|
||||
) -> None:
|
||||
"""Saves a torch.jit.ScriptModule or torch.jit.RecursiveScriptModule to file.
|
||||
|
||||
Args:
|
||||
model: JIT compiled model loaded onto `config.checkpoint.jit.device`.
|
||||
jit_filepath: The filepath to the JIT-compiled model.
|
||||
"""
|
||||
torch.jit.save(model, jit_filepath)
|
||||
|
||||
|
||||
def get_filepaths(input_pattern) -> list[str]:
|
||||
"""Returns a list of filepaths from a pattern."""
|
||||
filepaths = sorted(glob(str(input_pattern)))
|
||||
return list(set(filepaths))
|
||||
|
||||
|
||||
def get_output_filepath(filepath: str, output_dir: str = None) -> str:
|
||||
"""Returns the output filepath for the given input filepath."""
|
||||
output_dir = output_dir or f"{os.path.dirname(filepath)}/reconstructions"
|
||||
output_filepath = f"{output_dir}/{os.path.basename(filepath)}"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
return output_filepath
|
||||
|
||||
|
||||
def read_image(filepath: str) -> np.ndarray:
|
||||
"""Reads an image from a filepath.
|
||||
|
||||
Args:
|
||||
filepath: The filepath to the image.
|
||||
|
||||
Returns:
|
||||
The image as a numpy array, layout HxWxC, range [0..255], uint8 dtype.
|
||||
"""
|
||||
image = media.read_image(filepath)
|
||||
# convert the grey scale image to RGB
|
||||
# since our tokenizers always assume 3-channel RGB image
|
||||
if image.ndim == 2:
|
||||
image = np.stack([image] * 3, axis=-1)
|
||||
# convert RGBA to RGB
|
||||
if image.shape[-1] == 4:
|
||||
image = image[..., :3]
|
||||
return image
|
||||
|
||||
|
||||
def read_video(filepath: str) -> np.ndarray:
|
||||
"""Reads a video from a filepath.
|
||||
|
||||
Args:
|
||||
filepath: The filepath to the video.
|
||||
Returns:
|
||||
The video as a numpy array, layout TxHxWxC, range [0..255], uint8 dtype.
|
||||
"""
|
||||
video = media.read_video(filepath)
|
||||
# convert the grey scale frame to RGB
|
||||
# since our tokenizers always assume 3-channel video
|
||||
if video.ndim == 3:
|
||||
video = np.stack([video] * 3, axis=-1)
|
||||
# convert RGBA to RGB
|
||||
if video.shape[-1] == 4:
|
||||
video = video[..., :3]
|
||||
return video
|
||||
|
||||
|
||||
def resize_image(image: np.ndarray, short_size: int = None) -> np.ndarray:
|
||||
"""Resizes an image to have the short side of `short_size`.
|
||||
|
||||
Args:
|
||||
image: The image to resize, layout HxWxC, of any range.
|
||||
short_size: The size of the short side.
|
||||
Returns:
|
||||
The resized image.
|
||||
"""
|
||||
if short_size is None:
|
||||
return image
|
||||
height, width = image.shape[-3:-1]
|
||||
if height <= width:
|
||||
height_new, width_new = short_size, int(width * short_size / height + 0.5)
|
||||
width_new = width_new if width_new % 2 == 0 else width_new + 1
|
||||
else:
|
||||
height_new, width_new = (
|
||||
int(height * short_size / width + 0.5),
|
||||
short_size,
|
||||
)
|
||||
height_new = height_new if height_new % 2 == 0 else height_new + 1
|
||||
return media.resize_image(image, shape=(height_new, width_new))
|
||||
|
||||
|
||||
def resize_video(video: np.ndarray, short_size: int = None) -> np.ndarray:
|
||||
"""Resizes a video to have the short side of `short_size`.
|
||||
|
||||
Args:
|
||||
video: The video to resize, layout TxHxWxC, of any range.
|
||||
short_size: The size of the short side.
|
||||
Returns:
|
||||
The resized video.
|
||||
"""
|
||||
if short_size is None:
|
||||
return video
|
||||
height, width = video.shape[-3:-1]
|
||||
if height <= width:
|
||||
height_new, width_new = short_size, int(width * short_size / height + 0.5)
|
||||
width_new = width_new if width_new % 2 == 0 else width_new + 1
|
||||
else:
|
||||
height_new, width_new = (
|
||||
int(height * short_size / width + 0.5),
|
||||
short_size,
|
||||
)
|
||||
height_new = height_new if height_new % 2 == 0 else height_new + 1
|
||||
return media.resize_video(video, shape=(height_new, width_new))
|
||||
|
||||
|
||||
def write_image(filepath: str, image: np.ndarray):
|
||||
"""Writes an image to a filepath."""
|
||||
return media.write_image(filepath, image)
|
||||
|
||||
|
||||
def write_video(filepath: str, video: np.ndarray, fps: int = 24) -> None:
|
||||
"""Writes a video to a filepath."""
|
||||
return media.write_video(filepath, video, fps=fps)
|
||||
|
||||
|
||||
def numpy2tensor(
|
||||
input_image: np.ndarray,
|
||||
dtype: torch.dtype = _DTYPE,
|
||||
device: str = _DEVICE,
|
||||
range_min: int = -1,
|
||||
) -> torch.Tensor:
|
||||
"""Converts image(dtype=np.uint8) to `dtype` in range [0..255].
|
||||
|
||||
Args:
|
||||
input_image: A batch of images in range [0..255], BxHxWx3 layout.
|
||||
Returns:
|
||||
A torch.Tensor of layout Bx3xHxW in range [-1..1], dtype.
|
||||
"""
|
||||
ndim = input_image.ndim
|
||||
indices = list(range(1, ndim))[-1:] + list(range(1, ndim))[:-1]
|
||||
image = input_image.transpose((0,) + tuple(indices)) / _UINT8_MAX_F
|
||||
if range_min == -1:
|
||||
image = 2.0 * image - 1.0
|
||||
return torch.from_numpy(image).to(dtype).to(device)
|
||||
|
||||
|
||||
def tensor2numpy(input_tensor: torch.Tensor, range_min: int = -1) -> np.ndarray:
|
||||
"""Converts tensor in [-1,1] to image(dtype=np.uint8) in range [0..255].
|
||||
|
||||
Args:
|
||||
input_tensor: Input image tensor of Bx3xHxW layout, range [-1..1].
|
||||
Returns:
|
||||
A numpy image of layout BxHxWx3, range [0..255], uint8 dtype.
|
||||
"""
|
||||
if range_min == -1:
|
||||
input_tensor = (input_tensor.float() + 1.0) / 2.0
|
||||
ndim = input_tensor.ndim
|
||||
output_image = input_tensor.clamp(0, 1).cpu().numpy()
|
||||
output_image = output_image.transpose((0,) + tuple(range(2, ndim)) + (1,))
|
||||
return (output_image * _UINT8_MAX_F + 0.5).astype(np.uint8)
|
||||
|
||||
|
||||
def pad_image_batch(
|
||||
batch: np.ndarray, spatial_align: int = _SPATIAL_ALIGN
|
||||
) -> tuple[np.ndarray, list[int]]:
|
||||
"""Pads a batch of images to be divisible by `spatial_align`.
|
||||
|
||||
Args:
|
||||
batch: The batch of images to pad, layout BxHxWx3, in any range.
|
||||
align: The alignment to pad to.
|
||||
Returns:
|
||||
The padded batch and the crop region.
|
||||
"""
|
||||
height, width = batch.shape[1:3]
|
||||
align = spatial_align
|
||||
height_to_pad = (align - height % align) if height % align != 0 else 0
|
||||
width_to_pad = (align - width % align) if width % align != 0 else 0
|
||||
|
||||
crop_region = [
|
||||
height_to_pad >> 1,
|
||||
width_to_pad >> 1,
|
||||
height + (height_to_pad >> 1),
|
||||
width + (width_to_pad >> 1),
|
||||
]
|
||||
batch = np.pad(
|
||||
batch,
|
||||
(
|
||||
(0, 0),
|
||||
(height_to_pad >> 1, height_to_pad - (height_to_pad >> 1)),
|
||||
(width_to_pad >> 1, width_to_pad - (width_to_pad >> 1)),
|
||||
(0, 0),
|
||||
),
|
||||
mode="constant",
|
||||
)
|
||||
return batch, crop_region
|
||||
|
||||
|
||||
def pad_video_batch(
|
||||
batch: np.ndarray,
|
||||
temporal_align: int = _TEMPORAL_ALIGN,
|
||||
spatial_align: int = _SPATIAL_ALIGN,
|
||||
) -> tuple[np.ndarray, list[int]]:
|
||||
"""Pads a batch of videos to be divisible by `temporal_align` or `spatial_align`.
|
||||
|
||||
Zero pad spatially. Reflection pad temporally to handle causality better.
|
||||
Args:
|
||||
batch: The batch of videos to pad., layout BxFxHxWx3, in any range.
|
||||
align: The alignment to pad to.
|
||||
Returns:
|
||||
The padded batch and the crop region.
|
||||
"""
|
||||
num_frames, height, width = batch.shape[-4:-1]
|
||||
align = spatial_align
|
||||
height_to_pad = (align - height % align) if height % align != 0 else 0
|
||||
width_to_pad = (align - width % align) if width % align != 0 else 0
|
||||
|
||||
align = temporal_align
|
||||
frames_to_pad = (
|
||||
(align - (num_frames - 1) % align) if (num_frames - 1) % align != 0 else 0
|
||||
)
|
||||
|
||||
crop_region = [
|
||||
frames_to_pad >> 1,
|
||||
height_to_pad >> 1,
|
||||
width_to_pad >> 1,
|
||||
num_frames + (frames_to_pad >> 1),
|
||||
height + (height_to_pad >> 1),
|
||||
width + (width_to_pad >> 1),
|
||||
]
|
||||
batch = np.pad(
|
||||
batch,
|
||||
(
|
||||
(0, 0),
|
||||
(0, 0),
|
||||
(height_to_pad >> 1, height_to_pad - (height_to_pad >> 1)),
|
||||
(width_to_pad >> 1, width_to_pad - (width_to_pad >> 1)),
|
||||
(0, 0),
|
||||
),
|
||||
mode="constant",
|
||||
)
|
||||
batch = np.pad(
|
||||
batch,
|
||||
(
|
||||
(0, 0),
|
||||
(frames_to_pad >> 1, frames_to_pad - (frames_to_pad >> 1)),
|
||||
(0, 0),
|
||||
(0, 0),
|
||||
(0, 0),
|
||||
),
|
||||
mode="edge",
|
||||
)
|
||||
return batch, crop_region
|
||||
|
||||
|
||||
def unpad_video_batch(batch: np.ndarray, crop_region: list[int]) -> np.ndarray:
|
||||
"""Unpads video with `crop_region`.
|
||||
|
||||
Args:
|
||||
batch: A batch of numpy videos, layout BxFxHxWxC.
|
||||
crop_region: [f1,y1,x1,f2,y2,x2] first, top, left, last, bot, right crop indices.
|
||||
|
||||
Returns:
|
||||
np.ndarray: Cropped numpy video, layout BxFxHxWxC.
|
||||
"""
|
||||
assert len(crop_region) == 6, "crop_region should be len of 6."
|
||||
f1, y1, x1, f2, y2, x2 = crop_region
|
||||
return batch[..., f1:f2, y1:y2, x1:x2, :]
|
||||
|
||||
|
||||
def unpad_image_batch(batch: np.ndarray, crop_region: list[int]) -> np.ndarray:
|
||||
"""Unpads image with `crop_region`.
|
||||
|
||||
Args:
|
||||
batch: A batch of numpy images, layout BxHxWxC.
|
||||
crop_region: [y1,x1,y2,x2] top, left, bot, right crop indices.
|
||||
|
||||
Returns:
|
||||
np.ndarray: Cropped numpy image, layout BxHxWxC.
|
||||
"""
|
||||
assert len(crop_region) == 4, "crop_region should be len of 4."
|
||||
y1, x1, y2, x2 = crop_region
|
||||
return batch[..., y1:y2, x1:x2, :]
|
||||
217
tools/decode/vendor/cosmos_tokenizer/video_cli.py
vendored
Normal file
217
tools/decode/vendor/cosmos_tokenizer/video_cli.py
vendored
Normal file
@@ -0,0 +1,217 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""A CLI to run CausalVideoTokenizer on plain videos based on torch.jit.
|
||||
|
||||
Usage:
|
||||
python3 -m cosmos_tokenizer.video_cli \
|
||||
--video_pattern 'path/to/video/samples/*.mp4' \
|
||||
--output_dir ./reconstructions \
|
||||
--checkpoint_enc ./pretrained_ckpts/CosmosCV_f4x8x8/encoder.jit \
|
||||
--checkpoint_dec ./pretrained_ckpts/CosmosCV_f4x8x8/decoder.jit
|
||||
|
||||
Optionally, you can run the model in pure PyTorch mode:
|
||||
python3 -m cosmos_tokenizer.video_cli \
|
||||
--video_pattern 'path/to/video/samples/*.mp4' \
|
||||
--mode=torch \
|
||||
--tokenizer_type=CV \
|
||||
--temporal_compression=4 \
|
||||
--spatial_compression=8 \
|
||||
--checkpoint_enc ./pretrained_ckpts/CosmosCV_f4x8x8/encoder.jit \
|
||||
--checkpoint_dec ./pretrained_ckpts/CosmosCV_f4x8x8/decoder.jit
|
||||
"""
|
||||
|
||||
import os
|
||||
from argparse import ArgumentParser, Namespace
|
||||
from typing import Any
|
||||
import sys
|
||||
|
||||
import numpy as np
|
||||
from loguru import logger as logging
|
||||
|
||||
from cosmos_tokenizer.networks import TokenizerConfigs
|
||||
from cosmos_tokenizer.utils import (
|
||||
get_filepaths,
|
||||
get_output_filepath,
|
||||
read_video,
|
||||
resize_video,
|
||||
write_video,
|
||||
)
|
||||
from cosmos_tokenizer.video_lib import CausalVideoTokenizer
|
||||
|
||||
|
||||
def _parse_args() -> tuple[Namespace, dict[str, Any]]:
|
||||
parser = ArgumentParser(description="A CLI for CausalVideoTokenizer.")
|
||||
parser.add_argument(
|
||||
"--video_pattern",
|
||||
type=str,
|
||||
default="path/to/videos/*.mp4",
|
||||
help="Glob pattern.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoint",
|
||||
type=str,
|
||||
default=None,
|
||||
help="JIT full Autoencoder model filepath.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoint_enc",
|
||||
type=str,
|
||||
default=None,
|
||||
help="JIT Encoder model filepath.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoint_dec",
|
||||
type=str,
|
||||
default=None,
|
||||
help="JIT Decoder model filepath.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tokenizer_type",
|
||||
type=str,
|
||||
choices=["CV", "DV"],
|
||||
help="Specifies the tokenizer type.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--spatial_compression",
|
||||
type=int,
|
||||
choices=[8, 16],
|
||||
default=8,
|
||||
help="The spatial compression factor.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--temporal_compression",
|
||||
type=int,
|
||||
choices=[4, 8],
|
||||
default=4,
|
||||
help="The temporal compression factor.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mode",
|
||||
type=str,
|
||||
choices=["torch", "jit"],
|
||||
default="jit",
|
||||
help="Specify the backend: native 'torch' or 'jit' (default: 'jit')",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--short_size",
|
||||
type=int,
|
||||
default=None,
|
||||
help="The size to resample inputs. None, by default.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--temporal_window",
|
||||
type=int,
|
||||
default=17,
|
||||
help="The temporal window to operate at a time.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
default="bfloat16",
|
||||
help="Sets the precision, default bfloat16.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device",
|
||||
type=str,
|
||||
default="cuda",
|
||||
help="Device for invoking the model.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_dir", type=str, default=None, help="Output directory."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_fps",
|
||||
type=float,
|
||||
default=24.0,
|
||||
help="Output frames-per-second (FPS).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save_input",
|
||||
action="store_true",
|
||||
help="If on, the input video will be be outputted too.",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
logging.info("Initializes args ...")
|
||||
args = _parse_args()
|
||||
if args.mode == "torch" and args.tokenizer_type not in ["CV", "DV"]:
|
||||
logging.error("'torch' backend requires the tokenizer_type of 'CV' or 'DV'.")
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def _run_eval() -> None:
|
||||
"""Invokes JIT-compiled CausalVideoTokenizer on an input video."""
|
||||
|
||||
if (
|
||||
args.checkpoint_enc is None
|
||||
and args.checkpoint_dec is None
|
||||
and args.checkpoint is None
|
||||
):
|
||||
logging.warning(
|
||||
"Aborting. Both encoder or decoder JIT required. Or provide the full autoencoder JIT model."
|
||||
)
|
||||
return
|
||||
|
||||
if args.mode == "torch":
|
||||
tokenizer_config = TokenizerConfigs[args.tokenizer_type].value
|
||||
tokenizer_config.update(dict(spatial_compression=args.spatial_compression))
|
||||
tokenizer_config.update(dict(temporal_compression=args.temporal_compression))
|
||||
else:
|
||||
tokenizer_config = None
|
||||
|
||||
logging.info(
|
||||
f"Loading a torch.jit model `{os.path.dirname(args.checkpoint or args.checkpoint_enc or args.checkpoint_dec)}` ..."
|
||||
)
|
||||
autoencoder = CausalVideoTokenizer(
|
||||
checkpoint=args.checkpoint,
|
||||
checkpoint_enc=args.checkpoint_enc,
|
||||
checkpoint_dec=args.checkpoint_dec,
|
||||
tokenizer_config=tokenizer_config,
|
||||
device=args.device,
|
||||
dtype=args.dtype,
|
||||
)
|
||||
|
||||
logging.info(f"Looking for files matching video_pattern={args.video_pattern} ...")
|
||||
filepaths = get_filepaths(args.video_pattern)
|
||||
logging.info(f"Found {len(filepaths)} videos from {args.video_pattern}.")
|
||||
|
||||
for filepath in filepaths:
|
||||
logging.info(f"Reading video {filepath} ...")
|
||||
video = read_video(filepath)
|
||||
video = resize_video(video, short_size=args.short_size)
|
||||
|
||||
logging.info("Invoking the autoencoder model in ... ")
|
||||
batch_video = video[np.newaxis, ...]
|
||||
output_video = autoencoder(batch_video, temporal_window=args.temporal_window)[0]
|
||||
logging.info("Constructing output filepath ...")
|
||||
output_filepath = get_output_filepath(filepath, output_dir=args.output_dir)
|
||||
logging.info(f"Outputing {output_filepath} ...")
|
||||
write_video(output_filepath, output_video, fps=args.output_fps)
|
||||
if args.save_input:
|
||||
ext = os.path.splitext(output_filepath)[-1]
|
||||
input_filepath = output_filepath.replace(ext, "_input" + ext)
|
||||
write_video(input_filepath, video, fps=args.output_fps)
|
||||
|
||||
|
||||
@logging.catch(reraise=True)
|
||||
def main() -> None:
|
||||
_run_eval()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
153
tools/decode/vendor/cosmos_tokenizer/video_lib.py
vendored
Normal file
153
tools/decode/vendor/cosmos_tokenizer/video_lib.py
vendored
Normal file
@@ -0,0 +1,153 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""A library for Causal Video Tokenizer inference."""
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from typing import Any
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from cosmos_tokenizer.utils import (
|
||||
load_model,
|
||||
load_encoder_model,
|
||||
load_decoder_model,
|
||||
numpy2tensor,
|
||||
pad_video_batch,
|
||||
tensor2numpy,
|
||||
unpad_video_batch,
|
||||
)
|
||||
|
||||
|
||||
class CausalVideoTokenizer(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
checkpoint: str = None,
|
||||
checkpoint_enc: str = None,
|
||||
checkpoint_dec: str = None,
|
||||
tokenizer_config: dict[str, Any] = None,
|
||||
device: str = "cuda",
|
||||
dtype: str = "bfloat16",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self._device = device
|
||||
self._dtype = getattr(torch, dtype)
|
||||
self._full_model = (
|
||||
load_model(checkpoint, tokenizer_config, device).to(self._dtype)
|
||||
if checkpoint is not None
|
||||
else None
|
||||
)
|
||||
self._enc_model = (
|
||||
load_encoder_model(checkpoint_enc, tokenizer_config, device).to(self._dtype)
|
||||
if checkpoint_enc is not None
|
||||
else None
|
||||
)
|
||||
self._dec_model = (
|
||||
load_decoder_model(checkpoint_dec, tokenizer_config, device).to(self._dtype)
|
||||
if checkpoint_dec is not None
|
||||
else None
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
def autoencode(self, input_tensor: torch.Tensor) -> torch.Tensor:
|
||||
"""Reconstrcuts a batch of video tensors after embedding into a latent.
|
||||
|
||||
Args:
|
||||
video: The input video Bx3xTxHxW layout, range [-1..1].
|
||||
Returns:
|
||||
The reconstructed video, layout Bx3xTxHxW, range [-1..1].
|
||||
"""
|
||||
if self._full_model is not None:
|
||||
output_tensor = self._full_model(input_tensor)
|
||||
output_tensor = (
|
||||
output_tensor[0] if isinstance(output_tensor, tuple) else output_tensor
|
||||
)
|
||||
else:
|
||||
output_latent = self.encode(input_tensor)[0]
|
||||
output_tensor = self.decode(output_latent)
|
||||
return output_tensor
|
||||
|
||||
@torch.no_grad()
|
||||
def encode(self, input_tensor: torch.Tensor) -> tuple[torch.Tensor]:
|
||||
"""Encodes a numpy video into a CausalVideo latent or code.
|
||||
|
||||
Args:
|
||||
input_tensor: The input tensor Bx3xTxHxW layout, range [-1..1].
|
||||
Returns:
|
||||
For causal continuous video (CV) tokenizer, the tuple contains:
|
||||
- The latent embedding, Bx16x(t)x(h)x(w), where the compression
|
||||
rate is (T/t x H/h x W/w), and channel dimension of 16.
|
||||
For causal discrete video (DV) tokenizer, the tuple contains:
|
||||
1) The indices, Bx(t)x(h)x(w), from a codebook of size 64K, which
|
||||
is formed by FSQ levels of (8,8,8,5,5,5).
|
||||
2) The discrete code, Bx6x(t)x(h)x(w), where the compression rate
|
||||
is again (T/t x H/h x W/w), and channel dimension of 6.
|
||||
"""
|
||||
assert input_tensor.ndim == 5, "input video should be of 5D."
|
||||
|
||||
output_latent = self._enc_model(input_tensor)
|
||||
if isinstance(output_latent, torch.Tensor):
|
||||
return output_latent
|
||||
return output_latent[:-1]
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, input_latent: torch.Tensor) -> torch.Tensor:
|
||||
"""Encodes a numpy video into a CausalVideo latent.
|
||||
|
||||
Args:
|
||||
input_latent: The continuous latent Bx16xtxhxw for CV,
|
||||
or the discrete indices Bxtxhxw for DV.
|
||||
Returns:
|
||||
The reconstructed tensor, layout [B,3,1+(T-1)*8,H*16,W*16] in range [-1..1].
|
||||
"""
|
||||
assert (
|
||||
input_latent.ndim >= 4
|
||||
), "input latent should be of 5D for continuous and 4D for discrete."
|
||||
return self._dec_model(input_latent)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
video: np.ndarray,
|
||||
temporal_window: int = 17,
|
||||
) -> np.ndarray:
|
||||
"""Reconstructs video using a pre-trained CausalTokenizer autoencoder.
|
||||
Given a video of arbitrary length, the forward invokes the CausalVideoTokenizer
|
||||
in a sliding manner with a `temporal_window` size.
|
||||
|
||||
Args:
|
||||
video: The input video BxTxHxWx3 layout, range [0..255].
|
||||
temporal_window: The length of the temporal window to process, default=25.
|
||||
Returns:
|
||||
The reconstructed video in range [0..255], layout BxTxHxWx3.
|
||||
"""
|
||||
assert video.ndim == 5, "input video should be of 5D."
|
||||
num_frames = video.shape[1] # can be of any length.
|
||||
output_video_list = []
|
||||
for idx in tqdm(range(0, (num_frames - 1) // temporal_window + 1)):
|
||||
# Input video for the current window.
|
||||
start, end = idx * temporal_window, (idx + 1) * temporal_window
|
||||
input_video = video[:, start:end, ...]
|
||||
|
||||
# Spatio-temporally pad input_video so it's evenly divisible.
|
||||
padded_input_video, crop_region = pad_video_batch(input_video)
|
||||
input_tensor = numpy2tensor(
|
||||
padded_input_video, dtype=self._dtype, device=self._device
|
||||
)
|
||||
output_tensor = self.autoencode(input_tensor)
|
||||
padded_output_video = tensor2numpy(output_tensor)
|
||||
output_video = unpad_video_batch(padded_output_video, crop_region)
|
||||
|
||||
output_video_list.append(output_video)
|
||||
return np.concatenate(output_video_list, axis=1)
|
||||
104
tools/encode/encode_agent.py
Normal file
104
tools/encode/encode_agent.py
Normal file
@@ -0,0 +1,104 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Agent (3D pose) encoder -- turns a real 17-joint pose window into `<agent>`
|
||||
tokens, the reverse of tools/eval/decode_agent_tokens.py. Reuses
|
||||
pipeline_pose/phase5_adaptive_pchip.py's build_token_str() verbatim (pure
|
||||
numpy, no cluster dependency at all) rather than re-deriving the adaptive
|
||||
PCHIP control-point selection.
|
||||
|
||||
Use case: you have a REAL 3D pose sequence (motion capture, or your own
|
||||
video run through an HRNet+MotionBERT-style pipeline) and want the model to
|
||||
continue/predict from it -- this is exactly the "agent completion" behavior
|
||||
already verified for this model (give a partial <agent> block, it completes
|
||||
all 17 joints). Unlike seed2/cosmos/snac, this is the only encoder where
|
||||
"raw input" isn't a stock media file -- it's already-estimated 3D joint
|
||||
positions, which is a fair thing to require (you can't derive metric 3D pose
|
||||
from nothing; some upstream pose-estimation step is unavoidable no matter
|
||||
who's doing the encoding).
|
||||
|
||||
Input contract (IMPORTANT, easy to get wrong):
|
||||
- shape (8, 17, 3) float -- exactly 8 frames (this model's WINDOW_FRAMES),
|
||||
NOT 24 (that's the newer 2026-07-23 pipeline convention this model never
|
||||
saw), 17 joints in the exact order below, xyz in METRES.
|
||||
- ROOT-CENTERED: pelvis (joint 0) must be at [0,0,0] in every frame --
|
||||
subtract the pelvis position from all 17 joints per-frame yourself first
|
||||
if your source data isn't already root-relative (see
|
||||
pipeline_pose/phase3_kinematics_processor.py's split_root_motion() for
|
||||
the exact convention this project uses).
|
||||
- Values should stay within [-2.0, +2.0]m per axis (COORD_RANGE) --
|
||||
quantize() clips silently outside that range, so a badly-scaled pose
|
||||
(e.g. millimetres instead of metres) will silently flatten to the
|
||||
boundary rather than erroring. No automatic unit detection is attempted.
|
||||
|
||||
Joint order:
|
||||
pelvis, r_hip, r_knee, r_ankle, l_hip, l_knee, l_ankle, spine, thorax,
|
||||
nose, head_top, l_shoulder, l_elbow, l_wrist, r_shoulder, r_elbow, r_wrist
|
||||
|
||||
Usage:
|
||||
python tools/encode/encode_agent.py --input pose.json
|
||||
# pose.json: {"states": [[[x,y,z], ...17 joints...], ...8 frames...]}
|
||||
python tools/encode/encode_agent.py --input pose.npy
|
||||
# pose.npy: numpy array, shape (8, 17, 3)
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "..", "pipeline_pose"))
|
||||
from phase5_adaptive_pchip import build_token_str, JOINT_NAMES, N_JOINTS, TARGET_FPS, COORD_RANGE # noqa: E402
|
||||
|
||||
WINDOW_FRAMES = 8 # this model's convention -- NOT the newer 24-frame pipeline
|
||||
|
||||
|
||||
def load_states(path: str):
|
||||
import numpy as np
|
||||
|
||||
if path.endswith(".npy"):
|
||||
states = np.load(path)
|
||||
else:
|
||||
with open(path) as f:
|
||||
data = json.load(f)
|
||||
states = np.array(data["states"], dtype=np.float32)
|
||||
|
||||
if states.shape != (WINDOW_FRAMES, N_JOINTS, 3):
|
||||
raise ValueError(
|
||||
f"Expected shape ({WINDOW_FRAMES}, {N_JOINTS}, 3), got {states.shape}. "
|
||||
f"This model was trained on 8-frame windows -- 24-frame input (the newer "
|
||||
f"pipeline convention) will NOT tokenize correctly here."
|
||||
)
|
||||
|
||||
pelvis = states[:, 0, :]
|
||||
if not (abs(pelvis).max() < 1e-4):
|
||||
print(f"WARNING: pelvis (joint 0) is not at origin (max |pelvis|={abs(pelvis).max():.4f}m) -- "
|
||||
f"auto-centering now. If this wasn't intended, check your source data's convention.",
|
||||
file=sys.stderr)
|
||||
states = states - pelvis[:, None, :]
|
||||
|
||||
bad = states[abs(states) > COORD_RANGE]
|
||||
if bad.size > 0:
|
||||
print(f"WARNING: {bad.size} coordinate value(s) outside [-{COORD_RANGE}, {COORD_RANGE}]m -- "
|
||||
f"will be silently clipped by quantize(). Check units (expected metres).",
|
||||
file=sys.stderr)
|
||||
|
||||
return states
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--input", required=True, help=".json (with a 'states' key) or .npy file, shape (8,17,3)")
|
||||
ap.add_argument("--fps", type=int, default=TARGET_FPS)
|
||||
args = ap.parse_args()
|
||||
|
||||
states = load_states(args.input)
|
||||
token_str, cp_counts = build_token_str(states, fps=args.fps)
|
||||
|
||||
print(f"Encoded {N_JOINTS} joints, {sum(cp_counts.values())} total control points:")
|
||||
for name in JOINT_NAMES:
|
||||
print(f" {name}: {cp_counts[name]} CPs")
|
||||
print()
|
||||
print("<agent> " + token_str + " </agent>")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
95
tools/encode/encode_cosmos.py
Normal file
95
tools/encode/encode_cosmos.py
Normal file
@@ -0,0 +1,95 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Cosmos video-chunk encoder -- turns 8 real video frames into `<cosmos_N>`
|
||||
tokens (200 raw ids, no offset), the reverse of tools/decode/decode_cosmos.py.
|
||||
Reproduces the OLD (vla-1.7b-qwen3-v2 training-era) preprocessing convention
|
||||
EXACTLY, not the current/newer aspect-preserving pipeline convention
|
||||
(2026-07-23 pivot, 896 tokens/chunk) -- recovered from git history
|
||||
(commit edf25393, before 38d8e5f2 switched to aspect-preserving):
|
||||
|
||||
Resize((160,160)) # direct squash/stretch to 160x160, NOT an
|
||||
# aspect-preserving crop -- distorts aspect
|
||||
# ratio on purpose (this model's own training
|
||||
# convention, don't "fix" it here)
|
||||
ToTensor()
|
||||
Normalize(mean=[0.5,0.5,0.5], std=[0.5,0.5,0.5]) # -> range [-1, 1]
|
||||
|
||||
8 frames -> stack -> permute(1,0,2,3) -> (3,8,160,160) -> unsqueeze(0) ->
|
||||
(1,3,8,160,160) -> CausalVideoTokenizer.encode() -> (1,2,10,10) == 200 ids,
|
||||
checkpoint nvidia/Cosmos-Tokenizer-DV8x16x16 (encoder.jit -- same repo
|
||||
decode_cosmos.py already downloads decoder.jit from).
|
||||
|
||||
Usage:
|
||||
python tools/encode/encode_cosmos.py --frames f0.png f1.png ... f7.png --output tokens.txt
|
||||
# exactly 8 frame image paths, in temporal order
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "decode", "vendor"))
|
||||
|
||||
TARGET_SIZE = 160
|
||||
N_FRAMES = 8
|
||||
COSMOS_HF_REPO = "nvidia/Cosmos-Tokenizer-DV8x16x16"
|
||||
_LOCAL_CHECKPOINT_ENC = "/e/project1/reformo/nguyen38/prototype/pretrained_ckpts/Cosmos-Tokenizer-DV8x16x16/encoder.jit"
|
||||
|
||||
|
||||
def _resolve_checkpoint_enc() -> str:
|
||||
if os.path.exists(_LOCAL_CHECKPOINT_ENC):
|
||||
return _LOCAL_CHECKPOINT_ENC
|
||||
from huggingface_hub import hf_hub_download
|
||||
print(f"Local checkpoint not found -- downloading encoder.jit from {COSMOS_HF_REPO} "
|
||||
f"(~350MB, cached for future runs)...")
|
||||
return hf_hub_download(repo_id=COSMOS_HF_REPO, filename="encoder.jit")
|
||||
|
||||
|
||||
def encode_frames(frame_paths: list) -> list:
|
||||
if len(frame_paths) != N_FRAMES:
|
||||
raise ValueError(f"Expected exactly {N_FRAMES} frame paths, got {len(frame_paths)}")
|
||||
|
||||
import torch
|
||||
import torchvision.transforms as T
|
||||
from PIL import Image
|
||||
from cosmos_tokenizer.video_lib import CausalVideoTokenizer
|
||||
|
||||
transform = T.Compose([
|
||||
T.Resize((TARGET_SIZE, TARGET_SIZE)),
|
||||
T.ToTensor(),
|
||||
T.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
|
||||
])
|
||||
frames = [transform(Image.open(p).convert("RGB")) for p in frame_paths]
|
||||
video = torch.stack(frames, dim=0) # (T, 3, H, W)
|
||||
video = video.permute(1, 0, 2, 3) # (3, T, H, W)
|
||||
video = video.unsqueeze(0) # (1, 3, T, H, W)
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
enc = CausalVideoTokenizer(checkpoint_enc=_resolve_checkpoint_enc()).to(device)
|
||||
with torch.no_grad():
|
||||
indices = enc.encode(video.to(device))[0] # (1, 2, 10, 10)
|
||||
|
||||
ids = indices.reshape(-1).tolist()
|
||||
if len(ids) != 200:
|
||||
raise ValueError(f"Expected 200 raw ids, got {len(ids)} -- checkpoint/shape mismatch")
|
||||
return ids
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--frames", nargs=8, required=True, metavar="FRAME",
|
||||
help="Exactly 8 frame image paths, in temporal order")
|
||||
ap.add_argument("--output", help="Optional: write comma-separated ids to this file")
|
||||
args = ap.parse_args()
|
||||
|
||||
ids = encode_frames(args.frames)
|
||||
out = ",".join(str(i) for i in ids)
|
||||
print(f"{len(ids)} cosmos tokens:")
|
||||
print(out)
|
||||
if args.output:
|
||||
with open(args.output, "w") as f:
|
||||
f.write(out)
|
||||
print(f"Saved: {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
66
tools/encode/encode_seed2.py
Normal file
66
tools/encode/encode_seed2.py
Normal file
@@ -0,0 +1,66 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Seed2 encoder -- turns a real image into `<seed2_N>` tokens (0-8191, no
|
||||
offset), the reverse of tools/decode/decode_seed2.py. Reuses that module's
|
||||
_load_seed2_tokenizer() (same runtime patches: the transformers import move
|
||||
+ BertLMHeadModel.cls=None guard) rather than re-deriving them -- the public
|
||||
ontocord/seed2 repo's own seed2_tokenizer.py still has both bugs unpatched
|
||||
(verified 2026-07-23), so any fresh download needs these regardless of
|
||||
whether you're encoding or decoding.
|
||||
|
||||
Preprocessing: Seed2Tokenizer.encode_image() does its own internal resize to
|
||||
224x224 (CLIP-style Resize+Normalize, see seed2_tokenizer.py's `self.processor`)
|
||||
-- pass a PIL image straight through, no manual resize needed first.
|
||||
|
||||
Usage:
|
||||
python tools/encode/encode_seed2.py --image photo.jpg
|
||||
# prints 32 raw ids; wrap as <seed2> <seed2_N> ... </seed2> to splice
|
||||
# into a prompt for this model (v2's convention -- no offset needed)
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "decode"))
|
||||
from decode_seed2 import _load_seed2_tokenizer, NUM_IMAGE_TOKENS # noqa: E402
|
||||
|
||||
|
||||
def encode_image(image_path: str) -> list:
|
||||
from PIL import Image
|
||||
|
||||
# _load_seed2_tokenizer() os.chdir()s -- resolve a relative image_path
|
||||
# against the original cwd *before* that happens, same class of bug
|
||||
# already fixed once in decode_seed2.py's own --output handling.
|
||||
image_path = os.path.abspath(image_path)
|
||||
|
||||
Seed2Tokenizer, seed2_dir = _load_seed2_tokenizer()
|
||||
tokenizer = Seed2Tokenizer.from_pretrained(seed2_dir).eval()
|
||||
if hasattr(tokenizer, "cuda") and __import__("torch").cuda.is_available():
|
||||
tokenizer = tokenizer.cuda()
|
||||
|
||||
image = Image.open(image_path).convert("RGB")
|
||||
ids = tokenizer.encode_image(image_pil=image)
|
||||
ids = ids.view(-1).tolist()
|
||||
|
||||
bad = [t for t in ids if not (0 <= t < NUM_IMAGE_TOKENS)]
|
||||
if bad:
|
||||
raise ValueError(f"encode_image produced out-of-range ids: {bad[:5]}... (expected [0, {NUM_IMAGE_TOKENS}))")
|
||||
return ids
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--image", required=True, help="Path to a real image file")
|
||||
args = ap.parse_args()
|
||||
|
||||
ids = encode_image(args.image)
|
||||
print(f"{len(ids)} seed2 tokens:")
|
||||
print(",".join(str(i) for i in ids))
|
||||
print()
|
||||
print("As a prompt fragment:")
|
||||
print("<seed2> " + " ".join(f"<seed2_{i}>" for i in ids) + " </seed2>")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
78
tools/encode/encode_snac.py
Normal file
78
tools/encode/encode_snac.py
Normal file
@@ -0,0 +1,78 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
SNAC audio encoder -- turns a real audio/video file into `<snac_N>` tokens
|
||||
(listen format, 3 tokens/base-frame), the reverse of tools/decode/decode_snac.py.
|
||||
Reuses the exact encode_listen() logic already in pipeline_pose/snac_finevideo.py
|
||||
(unchanged since inception) rather than re-deriving it.
|
||||
|
||||
This model (vla-1.7b-qwen3-v2) only ever saw listen-format audio wrapped in
|
||||
the generic <snac> tag -- NOT the newer (2026-07-23) <listen>/<speak>
|
||||
convention or the speak-format L2 tokens. Output here matches that: always
|
||||
listen-format, always <snac> wrapper.
|
||||
|
||||
Usage:
|
||||
python tools/encode/encode_snac.py --input clip.wav --output tokens.txt
|
||||
python tools/encode/encode_snac.py --input video.mp4 --output tokens.txt
|
||||
# any format ffmpeg can read (audio extracted automatically, works on
|
||||
# video files too -- just uses the audio track)
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "..", "pipeline_pose"))
|
||||
|
||||
SAMPLE_RATE = 24000
|
||||
SNAC_MODEL = "hubertsiuzdak/snac_24khz"
|
||||
|
||||
|
||||
def _get_ffmpeg() -> str:
|
||||
import imageio_ffmpeg
|
||||
return imageio_ffmpeg.get_ffmpeg_exe()
|
||||
|
||||
|
||||
def extract_audio(input_path: str):
|
||||
import numpy as np
|
||||
|
||||
cmd = [_get_ffmpeg(), "-y", "-i", input_path, "-vn", "-ac", "1", "-ar", str(SAMPLE_RATE), "-f", "f32le", "-"]
|
||||
result = subprocess.run(cmd, capture_output=True, timeout=300)
|
||||
if result.returncode != 0 or not result.stdout:
|
||||
raise RuntimeError(f"ffmpeg failed to extract audio from {input_path}: {result.stderr.decode(errors='replace')[:500]}")
|
||||
audio = np.frombuffer(result.stdout, dtype=np.float32).copy()
|
||||
if len(audio) == 0:
|
||||
raise RuntimeError(f"No audio extracted from {input_path} -- does it have an audio track?")
|
||||
return audio
|
||||
|
||||
|
||||
def encode_file(input_path: str) -> list:
|
||||
import torch
|
||||
from snac import SNAC
|
||||
from snac_finevideo import encode_listen # reused verbatim, not re-derived
|
||||
|
||||
audio = extract_audio(input_path)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
model = SNAC.from_pretrained(SNAC_MODEL).eval().to(device)
|
||||
tokens = encode_listen(audio, model, device) # already "<snac_N>" strings, listen-format
|
||||
return tokens
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--input", required=True, help="Audio or video file (any ffmpeg-readable format)")
|
||||
ap.add_argument("--output", help="Optional: write the <snac> block to this file")
|
||||
args = ap.parse_args()
|
||||
|
||||
tokens = encode_file(args.input)
|
||||
duration_s = len(tokens) / 3 / 12.5
|
||||
print(f"{len(tokens)} snac tokens ({len(tokens) // 3} base frames, ~{duration_s:.2f}s @ 12.5Hz base rate)")
|
||||
block = "<snac> " + " ".join(tokens) + " </snac>"
|
||||
print(block)
|
||||
if args.output:
|
||||
with open(args.output, "w") as f:
|
||||
f.write(block)
|
||||
print(f"Saved: {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
268
tools/eval/decode_agent_tokens.py
Normal file
268
tools/eval/decode_agent_tokens.py
Normal file
@@ -0,0 +1,268 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Decode adaptive PCHIP agent tokens into 3D pose trajectories.
|
||||
|
||||
Takes the raw token string that the VLA model produces and reconstructs
|
||||
the full (n_frames, 17, 3) skeleton trajectory via PCHIP interpolation.
|
||||
|
||||
Usage:
|
||||
# Decode tokens from stdin
|
||||
echo "<fps_30> <pelvis> <pelvis_t_0> ..." | python decode_agent_tokens.py
|
||||
|
||||
# Decode from a text file
|
||||
python decode_agent_tokens.py --input generated_tokens.txt
|
||||
|
||||
# Decode and save JSON
|
||||
python decode_agent_tokens.py --input tokens.txt --output poses.json
|
||||
|
||||
Token format (one 8-frame window):
|
||||
<fps_30>
|
||||
<pelvis> <pelvis_t_0> <pelvis_x_N> <pelvis_y_N> <pelvis_z_N>
|
||||
<pelvis_t_7> <pelvis_x_N> <pelvis_y_N> <pelvis_z_N> </pelvis>
|
||||
<r_hip> <r_hip_t_0> <r_hip_x_N> <r_hip_y_N> <r_hip_z_N> ... </r_hip>
|
||||
...17 joints...
|
||||
|
||||
Dequantization: coord = N / 255.0 * 4.0 - 2.0 (metres, range [-2, 2])
|
||||
Time tokens: frame index 0-7 within the 8-frame window
|
||||
Reconstruction: PCHIP interpolation over control points -> 8 frames
|
||||
|
||||
H36M joint order (17 joints):
|
||||
0 pelvis 4 l_hip 8 thorax 12 l_elbow 16 r_wrist
|
||||
1 r_hip 5 l_knee 9 nose 13 l_wrist
|
||||
2 r_knee 6 l_ankle 10 head_top 14 r_shoulder
|
||||
3 r_ankle 7 spine 11 l_shoulder 15 r_elbow
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
import sys
|
||||
|
||||
import numpy as np
|
||||
from scipy.interpolate import PchipInterpolator
|
||||
|
||||
WINDOW_FRAMES = 8
|
||||
COORD_RANGE = 2.0
|
||||
|
||||
JOINT_NAMES = [
|
||||
"pelvis", "r_hip", "r_knee", "r_ankle",
|
||||
"l_hip", "l_knee", "l_ankle",
|
||||
"spine", "thorax", "nose", "head_top",
|
||||
"l_shoulder", "l_elbow", "l_wrist",
|
||||
"r_shoulder", "r_elbow", "r_wrist",
|
||||
]
|
||||
JOINT_INDEX = {name: i for i, name in enumerate(JOINT_NAMES)}
|
||||
N_JOINTS = len(JOINT_NAMES)
|
||||
|
||||
|
||||
def dequantize(n: int) -> float:
|
||||
return n / 255.0 * (2.0 * COORD_RANGE) - COORD_RANGE
|
||||
|
||||
|
||||
def parse_window(tokens: list[str]) -> dict:
|
||||
"""Parse tokens for a single 8-frame window into per-joint control points.
|
||||
|
||||
Args:
|
||||
tokens: list of token strings like ['<fps_30>', '<pelvis>', '<pelvis_t_0>', ...]
|
||||
|
||||
Returns:
|
||||
dict with fps (int) and joints (dict mapping joint name to
|
||||
t_indices ndarray and cp_coords ndarray of shape (n_cp, 3)).
|
||||
"""
|
||||
fps = 30
|
||||
if tokens and tokens[0].startswith("<fps_"):
|
||||
fps = int(re.match(r"<fps_(\d+)>", tokens[0]).group(1))
|
||||
tokens = tokens[1:]
|
||||
|
||||
joints = {}
|
||||
i = 0
|
||||
while i < len(tokens):
|
||||
m = re.match(r"^<([a-z_]+)>$", tokens[i])
|
||||
if not m or m.group(1) not in JOINT_INDEX:
|
||||
i += 1
|
||||
continue
|
||||
|
||||
name = m.group(1)
|
||||
close = f"</{name}>"
|
||||
i += 1
|
||||
t_indices = []
|
||||
coords = []
|
||||
|
||||
while i < len(tokens) and tokens[i] != close:
|
||||
tm = re.match(rf"<{name}_t_(\d+)>$", tokens[i])
|
||||
if tm and i + 3 < len(tokens):
|
||||
t_indices.append(int(tm.group(1)))
|
||||
xm = re.match(rf"<{name}_x_(\d+)>$", tokens[i + 1])
|
||||
ym = re.match(rf"<{name}_y_(\d+)>$", tokens[i + 2])
|
||||
zm = re.match(rf"<{name}_z_(\d+)>$", tokens[i + 3])
|
||||
if xm and ym and zm:
|
||||
coords.append([
|
||||
dequantize(int(xm.group(1))),
|
||||
dequantize(int(ym.group(1))),
|
||||
dequantize(int(zm.group(1))),
|
||||
])
|
||||
i += 4
|
||||
continue
|
||||
i += 1
|
||||
|
||||
if i < len(tokens) and tokens[i] == close:
|
||||
i += 1
|
||||
|
||||
if t_indices:
|
||||
joints[name] = {
|
||||
"t_indices": np.array(t_indices, dtype=int),
|
||||
"cp_coords": np.array(coords, dtype=np.float32),
|
||||
}
|
||||
|
||||
return {"fps": fps, "joints": joints}
|
||||
|
||||
|
||||
def reconstruct(parsed: dict) -> np.ndarray:
|
||||
"""PCHIP-interpolate sparse control points into a full per-frame trajectory.
|
||||
|
||||
Window length is inferred from the data itself (max t_index + 1 seen
|
||||
across all joints) rather than a fixed constant -- 2026-07-22 (REPORT.md
|
||||
#38): windows are now 24 frames (t up to 23), but this decoder is also
|
||||
used on older 8-frame data (t up to 7), so it must handle both without
|
||||
being told which convention a given token string uses.
|
||||
|
||||
Returns ndarray of shape (window_frames, 17, 3) in metres, root-centred.
|
||||
"""
|
||||
max_t = WINDOW_FRAMES - 1
|
||||
for jdata in parsed["joints"].values():
|
||||
if len(jdata["t_indices"]):
|
||||
max_t = max(max_t, int(jdata["t_indices"].max()))
|
||||
window_frames = max_t + 1
|
||||
|
||||
t_out = np.arange(window_frames, dtype=np.float64)
|
||||
traj = np.zeros((window_frames, N_JOINTS, 3), dtype=np.float32)
|
||||
|
||||
for name, jdata in parsed["joints"].items():
|
||||
j = JOINT_INDEX[name]
|
||||
t_cp = jdata["t_indices"].astype(np.float64)
|
||||
cp = jdata["cp_coords"]
|
||||
|
||||
if len(t_cp) < 2:
|
||||
traj[:, j, :] = cp[0]
|
||||
continue
|
||||
|
||||
for d in range(3):
|
||||
traj[:, j, d] = PchipInterpolator(t_cp, cp[:, d])(t_out)
|
||||
|
||||
return traj
|
||||
|
||||
|
||||
def decode(token_str: str) -> list[np.ndarray]:
|
||||
"""Decode a token string into a list of (8, 17, 3) trajectories.
|
||||
|
||||
Handles both single windows and multiple consecutive windows.
|
||||
"""
|
||||
all_tokens = re.findall(r"<[^>]+>", token_str)
|
||||
if not all_tokens:
|
||||
return []
|
||||
|
||||
# Split on <fps_N> boundaries — each is one window
|
||||
window_starts = [i for i, t in enumerate(all_tokens) if t.startswith("<fps_")]
|
||||
|
||||
if not window_starts:
|
||||
parsed = parse_window(all_tokens)
|
||||
return [reconstruct(parsed)]
|
||||
|
||||
trajectories = []
|
||||
for wi, start in enumerate(window_starts):
|
||||
end = window_starts[wi + 1] if wi + 1 < len(window_starts) else len(all_tokens)
|
||||
parsed = parse_window(all_tokens[start:end])
|
||||
trajectories.append(reconstruct(parsed))
|
||||
|
||||
return trajectories
|
||||
|
||||
|
||||
def to_json(trajectories: list[np.ndarray], fps: int = 30) -> dict:
|
||||
"""Convert decoded trajectories to a JSON-serialisable dict.
|
||||
|
||||
Per-window frame count is read from each trajectory's own shape (not a
|
||||
fixed constant) -- see reconstruct()'s docstring, 2026-07-22."""
|
||||
windows = []
|
||||
cum_frames = 0
|
||||
for i, traj in enumerate(trajectories):
|
||||
n_frames_this = traj.shape[0]
|
||||
motion = np.linalg.norm(traj[-1] - traj[0], axis=-1)
|
||||
top_movers = sorted(
|
||||
[(JOINT_NAMES[j], round(float(motion[j]), 4)) for j in range(N_JOINTS)],
|
||||
key=lambda x: x[1], reverse=True,
|
||||
)
|
||||
n_missing = sum(1 for name in JOINT_NAMES if name not in
|
||||
{JOINT_NAMES[j] for j in range(N_JOINTS) if np.any(traj[:, j, :] != 0)})
|
||||
|
||||
windows.append({
|
||||
"window": i,
|
||||
"time_sec": round(cum_frames / fps, 4),
|
||||
"trajectory": traj.tolist(),
|
||||
"value_range_m": [round(float(traj.min()), 4), round(float(traj.max()), 4)],
|
||||
"top_movers": top_movers[:5],
|
||||
"joints_all_zero": n_missing,
|
||||
})
|
||||
cum_frames += n_frames_this
|
||||
|
||||
total_frames = sum(traj.shape[0] for traj in trajectories)
|
||||
return {
|
||||
"n_windows": len(trajectories),
|
||||
"total_frames": total_frames,
|
||||
"duration_sec": round(total_frames / fps, 4),
|
||||
"shape": [len(trajectories)] + list(trajectories[0].shape) if trajectories else [],
|
||||
"value_range_m": [round(float(min(t.min() for t in trajectories)), 4),
|
||||
round(float(max(t.max() for t in trajectories)), 4)] if trajectories else [0, 0],
|
||||
"joint_names": JOINT_NAMES,
|
||||
"windows": windows,
|
||||
}
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser(description="Decode agent tokens to 3D poses.")
|
||||
p.add_argument("--input", "-i", default=None,
|
||||
help="File containing agent tokens (default: read stdin)")
|
||||
p.add_argument("--output", "-o", default=None,
|
||||
help="Save decoded poses to JSON file")
|
||||
args = p.parse_args()
|
||||
|
||||
if args.input:
|
||||
with open(args.input, "r") as f:
|
||||
token_str = f.read()
|
||||
else:
|
||||
token_str = sys.stdin.read()
|
||||
|
||||
token_str = token_str.strip()
|
||||
if not token_str:
|
||||
print("No tokens provided.", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
trajectories = decode(token_str)
|
||||
if not trajectories:
|
||||
print("Could not parse any agent windows from input.", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
result = to_json(trajectories)
|
||||
|
||||
print(f"Decoded {result['n_windows']} windows "
|
||||
f"({result['total_frames']} frames, {result['duration_sec']}s)")
|
||||
print(f"Shape: {result['shape']} (windows, frames, joints, xyz)")
|
||||
print(f"Value range: {result['value_range_m']} m")
|
||||
|
||||
for w in result["windows"][:3]:
|
||||
print(f"\n Window {w['window']} (t={w['time_sec']}s):")
|
||||
if w["joints_all_zero"] > 0:
|
||||
print(f" WARNING: {w['joints_all_zero']} joints are all-zero (missing)")
|
||||
print(f" Top movers: ", end="")
|
||||
print(", ".join(f"{name} {d:.3f}m" for name, d in w["top_movers"]))
|
||||
|
||||
if result["n_windows"] > 3:
|
||||
print(f"\n ... {result['n_windows'] - 3} more windows")
|
||||
|
||||
if args.output:
|
||||
with open(args.output, "w") as f:
|
||||
json.dump(result, f, indent=2)
|
||||
print(f"\nSaved to: {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user