CoMind single-ego 2-view training code + README (what to change)
Browse files- README.md +133 -0
- cosmos_predict2/_src/predict2/models/video2world_model_rectified_flow.py +247 -0
- cosmos_predict2/_src/predict2_multiview/__init__.py +15 -0
- cosmos_predict2/_src/predict2_multiview/callbacks/every_n_draw_sample_multiviewvideo.py +485 -0
- cosmos_predict2/_src/predict2_multiview/callbacks/frame_loss_log.py +43 -0
- cosmos_predict2/_src/predict2_multiview/callbacks/log_weight.py +62 -0
- cosmos_predict2/_src/predict2_multiview/callbacks/nymeria_validation_viz.py +335 -0
- cosmos_predict2/_src/predict2_multiview/callbacks/sigma_loss_analysis_per_frame.py +338 -0
- cosmos_predict2/_src/predict2_multiview/conditioner.py +103 -0
- cosmos_predict2/_src/predict2_multiview/configs/__init__.py +15 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/__init__.py +15 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/config.py +47 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/__init__.py +15 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/callbacks.py +50 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/conditioner.py +720 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/dataloader.py +101 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/dataloader_local.py +111 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/model.py +57 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/net.py +190 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/optimizer.py +72 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/__init__.py +15 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/buttercup2p5_rectified_flow.py +229 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/buttercup2p5_rectified_flow_14b.py +243 -0
- cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/nymeria_pose_2actor.py +1005 -0
- cosmos_predict2/_src/predict2_multiview/datasets/__init__.py +15 -0
- cosmos_predict2/_src/predict2_multiview/datasets/comind_pairs.py +302 -0
- cosmos_predict2/_src/predict2_multiview/datasets/local.py +173 -0
- cosmos_predict2/_src/predict2_multiview/datasets/multiview.py +547 -0
- cosmos_predict2/_src/predict2_multiview/datasets/nymeria_pairs.py +1074 -0
- cosmos_predict2/_src/predict2_multiview/datasets/wdinfo_utils.py +79 -0
- cosmos_predict2/_src/predict2_multiview/models/multiview_pose_model_rectified_flow.py +261 -0
- cosmos_predict2/_src/predict2_multiview/models/multiview_vid2vid_model_rectified_flow.py +655 -0
- cosmos_predict2/_src/predict2_multiview/models/view_sampling.py +65 -0
- cosmos_predict2/_src/predict2_multiview/networks/multiview_cross_dit.py +1187 -0
- cosmos_predict2/_src/predict2_multiview/networks/multiview_dit.py +618 -0
- cosmos_predict2/_src/predict2_multiview/networks/multiview_pose_dit.py +526 -0
- cosmos_predict2/_src/predict2_multiview/scripts/inference.py +330 -0
- cosmos_predict2/_src/predict2_multiview/scripts/inference_cli.py +582 -0
- cosmos_predict2/_src/predict2_multiview/scripts/mv_visualize_helper.py +164 -0
- cosmos_predict2/_src/predict2_multiview/utils/optim_instantiate.py +234 -0
- sh/train_nymeria_longer.sh +38 -0
README.md
ADDED
|
@@ -0,0 +1,133 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Training code — CoMind 2-view generation (single-ego setup)
|
| 2 |
+
|
| 3 |
+
This bundle is our current 2-view egocentric video training stack (Cosmos-Predict2.5, `predict2_multiview`).
|
| 4 |
+
It is set up for a **2-view actor-actor** model with a lot of cross-view conditioning. You are going to
|
| 5 |
+
**strip it down to a single-ego generation setup that still outputs 2 views**, trained **from scratch on the
|
| 6 |
+
CoMind dataset**. This README lists exactly what to change and what to watch out for.
|
| 7 |
+
|
| 8 |
+
--------------------------------------------------------------------------------
|
| 9 |
+
## 0. What the current model conditions on (the parts you will REMOVE / keep)
|
| 10 |
+
|
| 11 |
+
Per view, each generated clip is conditioned on:
|
| 12 |
+
| Signal | Key | What it is | Single-ego plan |
|
| 13 |
+
|---|---|---|---|
|
| 14 |
+
| warped RGB | `control_input_warped` | RGB warped into the target view from a **combined pool (both views' past)** | **KEEP the channel, but change the DATA**: warp only from **that view's own past context** |
|
| 15 |
+
| pose | `control_input_pose` | skeletons. `person_pose=True` = identity-colored, **both people** visible | **KEEP the channel, change the DATA**: render **only the wearer's own** pose |
|
| 16 |
+
| in-context references | `reference_frames` (+ `reference_pose`, `reference_cam_w2c`) | clean appearance/pose anchor frames appended as extra tokens | **REMOVE entirely** |
|
| 17 |
+
| camera Plücker | `plucker_map` / `reference_plucker` | per-pixel rays in a **shared cross-view canonical frame** | **REMOVE** (its only purpose is cross-view alignment; single-ego doesn't need it) |
|
| 18 |
+
| view embedding | net `view_embeddings` | per-view identity added to tokens | **REMOVE** (drop `concat_view_embedding`) |
|
| 19 |
+
| cross-view self-attention | token layout `B (V·t) …` | both views attend each other | **Leave as-is** — you do NOT need to separate it; unified attention is fine, and with all cross-view conditioning removed the two views are effectively independent anyway |
|
| 20 |
+
|
| 21 |
+
So the single-ego model keeps only: **warped_cond (own past) + own pose**, generated per view, no refs / no
|
| 22 |
+
Plücker / no view-embedding, trained from the Cosmos base 2B.
|
| 23 |
+
|
| 24 |
+
--------------------------------------------------------------------------------
|
| 25 |
+
## 1. Code changes (all in `cosmos_predict2/_src/predict2_multiview/`)
|
| 26 |
+
|
| 27 |
+
Make a new experiment function in
|
| 28 |
+
`configs/vid2vid/experiment/nymeria_pose_2actor.py` (copy `comind_actoractor_personpose_shared` as a start),
|
| 29 |
+
and flip these toggles. The flags exist already — you're just turning conditioning OFF.
|
| 30 |
+
|
| 31 |
+
### 1a. model config (`models/multiview_pose_model_rectified_flow.py` fields, set in the experiment)
|
| 32 |
+
```python
|
| 33 |
+
cfg["model"]["config"]["num_reference_frames"] = 0 # was 4
|
| 34 |
+
cfg["model"]["config"]["enable_reference_plucker"] = False # was True
|
| 35 |
+
cfg["model"]["config"]["enable_reference_pose"] = False # was True
|
| 36 |
+
cfg["model"]["config"]["enable_plucker"] = False # drop cross-view camera Plücker
|
| 37 |
+
```
|
| 38 |
+
|
| 39 |
+
### 1b. net config (`networks/multiview_pose_dit.py`, set via `cfg["model"]["config"]["net"].update(...)`)
|
| 40 |
+
```python
|
| 41 |
+
cfg["model"]["config"]["net"].update(
|
| 42 |
+
enable_reference_frames = False,
|
| 43 |
+
num_reference_frames = 0,
|
| 44 |
+
shared_reference = False,
|
| 45 |
+
enable_reference_plucker= False,
|
| 46 |
+
enable_reference_pose = False,
|
| 47 |
+
enable_plucker = False,
|
| 48 |
+
concat_view_embedding = False, # <-- drops the view embedding
|
| 49 |
+
)
|
| 50 |
+
```
|
| 51 |
+
`pose_mode` stays `"vae_concat"` (pose is VAE-encoded and concatenated into `cond_embedder`; keep it). The
|
| 52 |
+
`cond_embedder` in-channels auto-compute as `warped_latent(16) + visibility(1) + pose_latent(16)` — leave that.
|
| 53 |
+
|
| 54 |
+
### 1c. checkpoint = from scratch
|
| 55 |
+
```python
|
| 56 |
+
cfg["checkpoint"]["load_path"] = _base_2b_multiview_ckpt() # Cosmos base 2B, NOT a warm-start
|
| 57 |
+
cfg["checkpoint"]["strict_resume"] = False
|
| 58 |
+
```
|
| 59 |
+
Because refs/Plücker/view-emb are gone and the pose/warp embedders are zero-init or freshly-shaped, **iter-1
|
| 60 |
+
loss will start HIGH** (~base video prior) — that's correct for a from-scratch run, unlike a warm-start.
|
| 61 |
+
|
| 62 |
+
### 1d. data loader
|
| 63 |
+
Point `override /data_train` and `override /data_val` at a CoMind data config that:
|
| 64 |
+
- reads **your regenerated** CoMind clips (own-past warp + own-only pose),
|
| 65 |
+
- sets `num_reference_frames=0`, `shared_reference=False`, `emit_reference_pose=False` (do NOT emit any
|
| 66 |
+
`reference_*`), `person_pose=` your choice (you're rendering own-only pose, so the identity-color flag is moot).
|
| 67 |
+
|
| 68 |
+
The CoMind loader lives in `datasets/comind_pairs.py` / the actor-actor path in `datasets/nymeria_pairs.py`
|
| 69 |
+
(`get_nymeria_actor_actor_loader`, class `NymeriaActorActorDataset`). Register a new `cs.store(...)` data
|
| 70 |
+
config with the flags above.
|
| 71 |
+
|
| 72 |
+
### 1e. validation viz
|
| 73 |
+
`callbacks/nymeria_validation_viz.py`: set `num_reference_frames=0, shared_reference=False,
|
| 74 |
+
emit_reference_pose=False` in the viz kwargs too, or the viz dataset build will look for refs that aren't there.
|
| 75 |
+
|
| 76 |
+
--------------------------------------------------------------------------------
|
| 77 |
+
## 2. DATA you must regenerate (this is the real work — NOT in this code bundle)
|
| 78 |
+
|
| 79 |
+
The training code just consumes `clip.npz` per view. For the single-ego setup you must rebuild the CoMind
|
| 80 |
+
`clip.npz` so that:
|
| 81 |
+
|
| 82 |
+
1. **`warped_cond`** = warp into the target view using **only that view's OWN past frames** as the source pool.
|
| 83 |
+
The current pipeline uses a **combined pool** (both views' past, max-coverage source selection). Change the
|
| 84 |
+
source pool to per-view own-past only. (Depth + campose per view are already there.)
|
| 85 |
+
2. **`pose` / `pose_person`** = render **only the wearer's own** skeleton in each view (not the partner's).
|
| 86 |
+
Currently `pose_person` renders both people identity-colored.
|
| 87 |
+
|
| 88 |
+
Everything else in `clip.npz` (`target_rgb`, `visibility_packed`, per-view `campose`) stays. You do **not**
|
| 89 |
+
need `refs_shared.npz` at all (no references).
|
| 90 |
+
|
| 91 |
+
--------------------------------------------------------------------------------
|
| 92 |
+
## 3. Things to know / gotchas
|
| 93 |
+
|
| 94 |
+
- **Env / launch**: `sh/train_nymeria_longer.sh`, run as
|
| 95 |
+
`CUDA_VISIBLE_DEVICES=0,1,2,3 EXP=<your_exp_name> NPROC=4 bash sh/train_nymeria_longer.sh`.
|
| 96 |
+
Conda env `ego_dh`. Needs `COSMOS_QWEN_TOKENIZER_DIR=/data/cosmos_reason1_7b` for the online text encoder,
|
| 97 |
+
`IMAGINAIRE_OUTPUT_ROOT=<out>`.
|
| 98 |
+
- **Online text encoding**: CoMind emits `ai_caption`; the model text-encodes it online each step (Qwen). Keep
|
| 99 |
+
the tokenizer dir env set or conditioning will crash on `t5_text_embeddings=None`.
|
| 100 |
+
- **DCP checkpoints** are reshardable across NPROC (4↔3↔2 resume works). `save_iter` in the experiment.
|
| 101 |
+
- **register the experiment**: append your function to the `experiments = [...]` list at the bottom of
|
| 102 |
+
`nymeria_pose_2actor.py`, and confirm it builds:
|
| 103 |
+
`python -c "..."` overriding `experiment=<name>` (see how the file builds configs) before launching 4-GPU.
|
| 104 |
+
- **`state_t`**: latent temporal length = `1 + (T-1)//4`. For 77 frames it's 20. It appears in the model config;
|
| 105 |
+
keep it consistent between data (`num_video_frames`) and model.
|
| 106 |
+
- **AR / noise recipe**: `predict2/models/video2world_model_rectified_flow.py` has an optional
|
| 107 |
+
`noisy_conditioning_*` + `conditional_frames_probs` recipe (teacher-forcing for autoregressive rollout). It is
|
| 108 |
+
OFF by default (`{1:1.0}`, prob 0). Ignore it unless you want AR — not needed for the single-ego baseline.
|
| 109 |
+
- **Cross-attention**: you asked whether to separate it — you don't need to. The DiT already runs unified
|
| 110 |
+
self-attention over `B (V·t) H W D`; with no cross-view conditioning the two views just don't share useful
|
| 111 |
+
signal, which is exactly the single-ego behavior. Leaving it unified is simplest and correct.
|
| 112 |
+
|
| 113 |
+
--------------------------------------------------------------------------------
|
| 114 |
+
## 4. File map (what's in this bundle)
|
| 115 |
+
|
| 116 |
+
```
|
| 117 |
+
cosmos_predict2/_src/predict2_multiview/
|
| 118 |
+
configs/vid2vid/experiment/nymeria_pose_2actor.py # all experiments (copy comind_* -> your single-ego exp)
|
| 119 |
+
configs/vid2vid/defaults/{conditioner,data,model}.py
|
| 120 |
+
datasets/nymeria_pairs.py # NymeriaActorActorDataset + loaders + cs.store data configs
|
| 121 |
+
datasets/comind_pairs.py # CoMind dataset/loader
|
| 122 |
+
networks/multiview_pose_dit.py # the DiT: all the enable_* / view-emb / reference embedders
|
| 123 |
+
models/multiview_pose_model_rectified_flow.py # conditioning preprocessing (refs/plucker/refpose/depth)
|
| 124 |
+
models/multiview_vid2vid_model_rectified_flow.py
|
| 125 |
+
callbacks/nymeria_validation_viz.py # validation grid
|
| 126 |
+
predict2/models/video2world_model_rectified_flow.py # base rectified-flow denoise (+optional AR noise)
|
| 127 |
+
sh/train_nymeria_longer.sh # launch
|
| 128 |
+
```
|
| 129 |
+
|
| 130 |
+
Summary of the diff you need: **turn OFF** `enable_reference_frames/plucker/pose`, `num_reference_frames=0`,
|
| 131 |
+
`shared_reference=False`, `enable_plucker=False`, `concat_view_embedding=False`; **from base 2B**; and feed
|
| 132 |
+
**CoMind clips rebuilt with own-past warp + own-only pose**. Keep warped_cond + pose channels and unified
|
| 133 |
+
cross-view attention.
|
cosmos_predict2/_src/predict2/models/video2world_model_rectified_flow.py
ADDED
|
@@ -0,0 +1,247 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
from enum import Enum
|
| 17 |
+
from typing import Callable, Dict, Optional, Tuple
|
| 18 |
+
|
| 19 |
+
import attrs
|
| 20 |
+
import torch
|
| 21 |
+
from megatron.core import parallel_state
|
| 22 |
+
from torch import Tensor
|
| 23 |
+
|
| 24 |
+
from cosmos_predict2._src.predict2.conditioner import DataType
|
| 25 |
+
from cosmos_predict2._src.predict2.configs.video2world.defaults.conditioner import Video2WorldCondition
|
| 26 |
+
from cosmos_predict2._src.predict2.models.denoise_prediction import DenoisePrediction
|
| 27 |
+
from cosmos_predict2._src.predict2.models.text2world_model_rectified_flow import (
|
| 28 |
+
Text2WorldCondition,
|
| 29 |
+
Text2WorldModelRectifiedFlow,
|
| 30 |
+
Text2WorldModelRectifiedFlowConfig,
|
| 31 |
+
)
|
| 32 |
+
|
| 33 |
+
NUM_CONDITIONAL_FRAMES_KEY: str = "num_conditional_frames"
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class ConditioningStrategy(str, Enum):
|
| 37 |
+
FRAME_REPLACE = "frame_replace" # First few frames of the video are replaced with the conditional frames
|
| 38 |
+
|
| 39 |
+
def __str__(self) -> str:
|
| 40 |
+
return self.value
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def apply_noisy_conditioning(cond_state_B_C_T_H_W, cond_mask_B_T, *, training, prob, scale, min_frames, n_views=1):
|
| 44 |
+
"""Add gaussian noise to the CONDITION latents only (network input), leaving the loss target clean.
|
| 45 |
+
Counts cond latents PER VIEW and uses the max (uneven view lengths would otherwise average below the
|
| 46 |
+
min_frames threshold and silently disable noise). See AR handoff recipe."""
|
| 47 |
+
if not training or prob <= 0.0:
|
| 48 |
+
return cond_state_B_C_T_H_W
|
| 49 |
+
B, _, T = cond_state_B_C_T_H_W.shape[:3]
|
| 50 |
+
if n_views > 1 and T % n_views == 0:
|
| 51 |
+
L = T // n_views
|
| 52 |
+
ncf_B = torch.stack([cond_mask_B_T[:, v * L:(v + 1) * L].sum(dim=1) for v in range(n_views)],
|
| 53 |
+
dim=0).max(dim=0).values
|
| 54 |
+
else:
|
| 55 |
+
ncf_B = cond_mask_B_T.sum(dim=1)
|
| 56 |
+
for b in range(B):
|
| 57 |
+
if ncf_B[b].item() >= min_frames and torch.rand(1).item() < prob:
|
| 58 |
+
cond_state_B_C_T_H_W[b] = cond_state_B_C_T_H_W[b] + scale * torch.randn_like(cond_state_B_C_T_H_W[b])
|
| 59 |
+
return cond_state_B_C_T_H_W
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
@attrs.define(slots=False)
|
| 63 |
+
class Video2WorldModelRectifiedFlowConfig(Text2WorldModelRectifiedFlowConfig):
|
| 64 |
+
min_num_conditional_frames: int = 1 # Minimum number of latent conditional frames
|
| 65 |
+
max_num_conditional_frames: int = 2 # Maximum number of latent conditional frames
|
| 66 |
+
conditional_frame_timestep: float = (
|
| 67 |
+
-1.0
|
| 68 |
+
) # Noise level used for conditional frames; default is -1 which will not take effective
|
| 69 |
+
conditioning_strategy: str = str(ConditioningStrategy.FRAME_REPLACE) # What strategy to use for conditioning
|
| 70 |
+
denoise_replace_gt_frames: bool = True # Whether to denoise the ground truth frames
|
| 71 |
+
conditional_frames_probs: Optional[Dict[int, float]] = None # Probability distribution for conditional frames
|
| 72 |
+
# AR: add noise to the CONDITION frames (network input only; loss target stays clean) to close the
|
| 73 |
+
# train-inference gap — at inference the AR cond frames are the model's own imperfect outputs.
|
| 74 |
+
noisy_conditioning_prob: float = 0.0 # per-sample prob to apply (0 = off)
|
| 75 |
+
noisy_conditioning_scale: float = 0.1 # gaussian noise std in latent units
|
| 76 |
+
noisy_conditioning_min_frames: int = 2 # only when >= this many cond latents (per view) — i.e. the AR case
|
| 77 |
+
|
| 78 |
+
def __attrs_post_init__(self):
|
| 79 |
+
super().__attrs_post_init__()
|
| 80 |
+
assert self.conditioning_strategy in [
|
| 81 |
+
str(ConditioningStrategy.FRAME_REPLACE),
|
| 82 |
+
]
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
class Video2WorldModelRectifiedFlow(Text2WorldModelRectifiedFlow):
|
| 86 |
+
def get_data_and_condition(
|
| 87 |
+
self, data_batch: dict[str, torch.Tensor]
|
| 88 |
+
) -> Tuple[Tensor, Tensor, Video2WorldCondition]:
|
| 89 |
+
# generate random number of conditional frames for training
|
| 90 |
+
raw_state, latent_state, condition = super().get_data_and_condition(data_batch)
|
| 91 |
+
condition = condition.set_video_condition(
|
| 92 |
+
gt_frames=latent_state.to(**self.tensor_kwargs),
|
| 93 |
+
random_min_num_conditional_frames=self.config.min_num_conditional_frames,
|
| 94 |
+
random_max_num_conditional_frames=self.config.max_num_conditional_frames,
|
| 95 |
+
num_conditional_frames=data_batch.get(NUM_CONDITIONAL_FRAMES_KEY, None),
|
| 96 |
+
conditional_frames_probs=self.config.conditional_frames_probs,
|
| 97 |
+
)
|
| 98 |
+
return raw_state, latent_state, condition
|
| 99 |
+
|
| 100 |
+
def denoise(
|
| 101 |
+
self,
|
| 102 |
+
noise: torch.Tensor,
|
| 103 |
+
xt_B_C_T_H_W: torch.Tensor,
|
| 104 |
+
timesteps_B_T: torch.Tensor,
|
| 105 |
+
condition: Text2WorldCondition,
|
| 106 |
+
) -> DenoisePrediction:
|
| 107 |
+
"""
|
| 108 |
+
Args:
|
| 109 |
+
xt (torch.Tensor): The input noise data.
|
| 110 |
+
sigma (torch.Tensor): The noise level.
|
| 111 |
+
condition (Text2WorldCondition): conditional information, generated from self.conditioner
|
| 112 |
+
|
| 113 |
+
Returns:
|
| 114 |
+
velocity prediction
|
| 115 |
+
"""
|
| 116 |
+
if condition.is_video:
|
| 117 |
+
condition_state_in_B_C_T_H_W = condition.gt_frames.type_as(xt_B_C_T_H_W)
|
| 118 |
+
if not condition.use_video_condition:
|
| 119 |
+
# When using random dropout, we zero out the ground truth frames
|
| 120 |
+
condition_state_in_B_C_T_H_W = condition_state_in_B_C_T_H_W * 0
|
| 121 |
+
|
| 122 |
+
_, C, _, _, _ = xt_B_C_T_H_W.shape
|
| 123 |
+
condition_video_mask = condition.condition_video_input_mask_B_C_T_H_W.repeat(1, C, 1, 1, 1).type_as(
|
| 124 |
+
xt_B_C_T_H_W
|
| 125 |
+
)
|
| 126 |
+
|
| 127 |
+
# AR: noise the condition latents (network input only). The loss-target replacement below uses the
|
| 128 |
+
# CLEAN gt_frames, so the target is unchanged. Counts cond latents per view (state_t) -> max.
|
| 129 |
+
if getattr(self.config, "noisy_conditioning_prob", 0.0) > 0.0:
|
| 130 |
+
cond_mask_B_T = condition.condition_video_input_mask_B_C_T_H_W[:, 0, :, 0, 0] # [B,T]
|
| 131 |
+
st = getattr(self.config, "state_t", 0)
|
| 132 |
+
nv = max(1, (cond_mask_B_T.shape[1] // st)) if st and st > 0 else 1
|
| 133 |
+
condition_state_in_B_C_T_H_W = apply_noisy_conditioning(
|
| 134 |
+
condition_state_in_B_C_T_H_W, cond_mask_B_T, training=self.training,
|
| 135 |
+
prob=self.config.noisy_conditioning_prob, scale=self.config.noisy_conditioning_scale,
|
| 136 |
+
min_frames=self.config.noisy_conditioning_min_frames, n_views=nv,
|
| 137 |
+
)
|
| 138 |
+
|
| 139 |
+
# Make the first few frames of x_t be the ground truth frames
|
| 140 |
+
xt_B_C_T_H_W = condition_state_in_B_C_T_H_W * condition_video_mask + xt_B_C_T_H_W * (
|
| 141 |
+
1 - condition_video_mask
|
| 142 |
+
)
|
| 143 |
+
|
| 144 |
+
if self.config.conditional_frame_timestep >= 0:
|
| 145 |
+
condition_video_mask_B_1_T_1_1 = condition_video_mask.mean(dim=[1, 3, 4], keepdim=True)
|
| 146 |
+
timestep_cond_B_1_T_1_1 = (
|
| 147 |
+
torch.ones_like(condition_video_mask_B_1_T_1_1) * self.config.conditional_frame_timestep
|
| 148 |
+
)
|
| 149 |
+
|
| 150 |
+
timesteps_B_1_T_1_1 = timestep_cond_B_1_T_1_1 * condition_video_mask_B_1_T_1_1 + timesteps_B_T * (
|
| 151 |
+
1 - condition_video_mask_B_1_T_1_1
|
| 152 |
+
)
|
| 153 |
+
|
| 154 |
+
timesteps_B_T = timesteps_B_1_T_1_1.squeeze()
|
| 155 |
+
timesteps_B_T = (
|
| 156 |
+
timesteps_B_T.unsqueeze(0) if timesteps_B_T.ndim == 1 else timesteps_B_T
|
| 157 |
+
) # add dimension for batch
|
| 158 |
+
|
| 159 |
+
# forward pass through the network
|
| 160 |
+
net_output_B_C_T_H_W = self.net(
|
| 161 |
+
x_B_C_T_H_W=xt_B_C_T_H_W.to(**self.tensor_kwargs), # Eq. 7 of https://arxiv.org/pdf/2206.00364.pdf
|
| 162 |
+
timesteps_B_T=timesteps_B_T, # Eq. 7 of https://arxiv.org/pdf/2206.00364.pdf
|
| 163 |
+
**condition.to_dict(),
|
| 164 |
+
).float()
|
| 165 |
+
|
| 166 |
+
if condition.is_video and self.config.denoise_replace_gt_frames:
|
| 167 |
+
gt_frames_x0 = condition.gt_frames.type_as(net_output_B_C_T_H_W)
|
| 168 |
+
gt_frames_velocity = noise - gt_frames_x0
|
| 169 |
+
net_output_B_C_T_H_W = gt_frames_velocity * condition_video_mask + net_output_B_C_T_H_W * (
|
| 170 |
+
1 - condition_video_mask
|
| 171 |
+
)
|
| 172 |
+
|
| 173 |
+
return net_output_B_C_T_H_W
|
| 174 |
+
|
| 175 |
+
def get_velocity_fn_from_batch(
|
| 176 |
+
self,
|
| 177 |
+
data_batch: Dict,
|
| 178 |
+
guidance: float = 1.5,
|
| 179 |
+
is_negative_prompt: bool = False,
|
| 180 |
+
) -> Callable:
|
| 181 |
+
"""
|
| 182 |
+
Generates a callable function `x0_fn` based on the provided data batch and guidance factor.
|
| 183 |
+
|
| 184 |
+
This function first processes the input data batch through a conditioning workflow (`conditioner`) to obtain conditioned and unconditioned states. It then defines a nested function `x0_fn` which applies a denoising operation on an input `noise_x` at a given noise level `sigma` using both the conditioned and unconditioned states.
|
| 185 |
+
|
| 186 |
+
Args:
|
| 187 |
+
- data_batch (Dict): A batch of data used for conditioning. The format and content of this dictionary should align with the expectations of the `self.conditioner`
|
| 188 |
+
- guidance (float, optional): A scalar value that modulates the influence of the conditioned state relative to the unconditioned state in the output. Defaults to 1.5.
|
| 189 |
+
- is_negative_prompt (bool): use negative prompt t5 in uncondition if true
|
| 190 |
+
|
| 191 |
+
Returns:
|
| 192 |
+
- Callable: A function `x0_fn(noise_x, sigma)` that takes two arguments, `noise_x` and `sigma`, and return velocity predictoin
|
| 193 |
+
|
| 194 |
+
The returned function is suitable for use in scenarios where a denoised state is required based on both conditioned and unconditioned inputs, with an adjustable level of guidance influence.
|
| 195 |
+
"""
|
| 196 |
+
|
| 197 |
+
if NUM_CONDITIONAL_FRAMES_KEY in data_batch:
|
| 198 |
+
num_conditional_frames = data_batch[NUM_CONDITIONAL_FRAMES_KEY]
|
| 199 |
+
else:
|
| 200 |
+
num_conditional_frames = 1
|
| 201 |
+
|
| 202 |
+
if is_negative_prompt:
|
| 203 |
+
condition, uncondition = self.conditioner.get_condition_with_negative_prompt(data_batch)
|
| 204 |
+
else:
|
| 205 |
+
condition, uncondition = self.conditioner.get_condition_uncondition(data_batch)
|
| 206 |
+
|
| 207 |
+
is_image_batch = self.is_image_batch(data_batch)
|
| 208 |
+
condition = condition.edit_data_type(DataType.IMAGE if is_image_batch else DataType.VIDEO)
|
| 209 |
+
uncondition = uncondition.edit_data_type(DataType.IMAGE if is_image_batch else DataType.VIDEO)
|
| 210 |
+
_, x0, _ = self.get_data_and_condition(data_batch)
|
| 211 |
+
# override condition with inference mode; num_conditional_frames used Here!
|
| 212 |
+
condition = condition.set_video_condition(
|
| 213 |
+
gt_frames=x0,
|
| 214 |
+
random_min_num_conditional_frames=self.config.min_num_conditional_frames,
|
| 215 |
+
random_max_num_conditional_frames=self.config.max_num_conditional_frames,
|
| 216 |
+
num_conditional_frames=num_conditional_frames,
|
| 217 |
+
conditional_frames_probs=self.config.conditional_frames_probs,
|
| 218 |
+
)
|
| 219 |
+
uncondition = uncondition.set_video_condition(
|
| 220 |
+
gt_frames=x0,
|
| 221 |
+
random_min_num_conditional_frames=self.config.min_num_conditional_frames,
|
| 222 |
+
random_max_num_conditional_frames=self.config.max_num_conditional_frames,
|
| 223 |
+
num_conditional_frames=num_conditional_frames,
|
| 224 |
+
conditional_frames_probs=self.config.conditional_frames_probs,
|
| 225 |
+
)
|
| 226 |
+
condition = condition.edit_for_inference(is_cfg_conditional=True, num_conditional_frames=num_conditional_frames)
|
| 227 |
+
uncondition = uncondition.edit_for_inference(
|
| 228 |
+
is_cfg_conditional=False, num_conditional_frames=num_conditional_frames
|
| 229 |
+
)
|
| 230 |
+
|
| 231 |
+
_, condition, _, _ = self.broadcast_split_for_model_parallelsim(x0, condition, None, None)
|
| 232 |
+
_, uncondition, _, _ = self.broadcast_split_for_model_parallelsim(x0, uncondition, None, None)
|
| 233 |
+
|
| 234 |
+
if parallel_state.is_initialized():
|
| 235 |
+
pass
|
| 236 |
+
else:
|
| 237 |
+
assert not self.net.is_context_parallel_enabled, (
|
| 238 |
+
"parallel_state is not initialized, context parallel should be turned off."
|
| 239 |
+
)
|
| 240 |
+
|
| 241 |
+
def velocity_fn(noise: torch.Tensor, noise_x: torch.Tensor, timestep: torch.Tensor) -> torch.Tensor:
|
| 242 |
+
cond_v = self.denoise(noise, noise_x, timestep, condition)
|
| 243 |
+
uncond_v = self.denoise(noise, noise_x, timestep, uncondition)
|
| 244 |
+
velocity_pred = cond_v + guidance * (cond_v - uncond_v)
|
| 245 |
+
return velocity_pred
|
| 246 |
+
|
| 247 |
+
return velocity_fn
|
cosmos_predict2/_src/predict2_multiview/__init__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
cosmos_predict2/_src/predict2_multiview/callbacks/every_n_draw_sample_multiviewvideo.py
ADDED
|
@@ -0,0 +1,485 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
from contextlib import nullcontext
|
| 17 |
+
from functools import partial
|
| 18 |
+
from typing import Optional
|
| 19 |
+
|
| 20 |
+
import torch
|
| 21 |
+
import torch.distributed as dist
|
| 22 |
+
import torch.nn.functional as F
|
| 23 |
+
import torchvision
|
| 24 |
+
import wandb
|
| 25 |
+
from einops import rearrange, repeat
|
| 26 |
+
|
| 27 |
+
from cosmos_predict2._src.imaginaire.utils import log, misc
|
| 28 |
+
from cosmos_predict2._src.imaginaire.utils.easy_io import easy_io
|
| 29 |
+
from cosmos_predict2._src.imaginaire.utils.parallel_state_helper import is_tp_cp_pp_rank0
|
| 30 |
+
from cosmos_predict2._src.imaginaire.visualize.video import save_img_or_video
|
| 31 |
+
from cosmos_predict2._src.predict2.callbacks.every_n_draw_sample import (
|
| 32 |
+
EveryNDrawSample,
|
| 33 |
+
convert_to_primitive,
|
| 34 |
+
is_primitive,
|
| 35 |
+
resize_image,
|
| 36 |
+
)
|
| 37 |
+
from cosmos_predict2._src.predict2.models.video2world_model import NUM_CONDITIONAL_FRAMES_KEY
|
| 38 |
+
from cosmos_predict2._src.predict2_multiview.models.multiview_vid2vid_model_rectified_flow import (
|
| 39 |
+
MultiviewVid2VidModelRectifiedFlow,
|
| 40 |
+
)
|
| 41 |
+
|
| 42 |
+
try:
|
| 43 |
+
import ffmpegcv
|
| 44 |
+
except Exception as e: # ImportError cannot catch all problems
|
| 45 |
+
log.info(e)
|
| 46 |
+
ffmpegcv = None
|
| 47 |
+
import cv2
|
| 48 |
+
import numpy as np
|
| 49 |
+
|
| 50 |
+
try:
|
| 51 |
+
import imageio
|
| 52 |
+
except Exception as e: # ImportError cannot catch all problems
|
| 53 |
+
log.info(e)
|
| 54 |
+
imageio = None
|
| 55 |
+
|
| 56 |
+
CONTROL_WEIGHT_KEY = "control_weight"
|
| 57 |
+
|
| 58 |
+
# view index order for visualization of 7-view autonomous driving dataset
|
| 59 |
+
camera_to_view_id = {
|
| 60 |
+
"camera_cross_left_120fov": 5,
|
| 61 |
+
"camera_cross_right_120fov": 1,
|
| 62 |
+
"camera_front_tele_30fov": 6,
|
| 63 |
+
"camera_front_wide_120fov": 0,
|
| 64 |
+
"camera_rear_left_70fov": 4,
|
| 65 |
+
"camera_rear_right_70fov": 2,
|
| 66 |
+
"camera_rear_tele_30fov": 3,
|
| 67 |
+
}
|
| 68 |
+
|
| 69 |
+
visualization_camera_order = [
|
| 70 |
+
"camera_rear_left_70fov",
|
| 71 |
+
"camera_cross_left_120fov",
|
| 72 |
+
"camera_front_wide_120fov",
|
| 73 |
+
"camera_cross_right_120fov",
|
| 74 |
+
"camera_rear_right_70fov",
|
| 75 |
+
"camera_rear_tele_30fov",
|
| 76 |
+
"camera_front_tele_30fov",
|
| 77 |
+
]
|
| 78 |
+
|
| 79 |
+
visualization_view_index_order = [camera_to_view_id[camera] for camera in visualization_camera_order]
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
class EveryNDrawSampleMultiviewVideo(EveryNDrawSample):
|
| 83 |
+
"""
|
| 84 |
+
This class is a modified version of EveryNDrawSample that saves 12 frames instead of 3.
|
| 85 |
+
"""
|
| 86 |
+
|
| 87 |
+
def __init__(
|
| 88 |
+
self,
|
| 89 |
+
*args,
|
| 90 |
+
n_view_embed=None,
|
| 91 |
+
ctrl_hint_keys=None,
|
| 92 |
+
control_weights=[1.0],
|
| 93 |
+
num_cond_frames=[0, 1],
|
| 94 |
+
fix_batch_fp=None, # For backward compatibility with transfer2 experiments
|
| 95 |
+
n_x0_level=None, # For backward compatibility with transfer2 experiments
|
| 96 |
+
show_all_frames=None, # For backward compatibility with transfer2 experiments
|
| 97 |
+
is_sample=None, # For backward compatibility with transfer2 experiments
|
| 98 |
+
**kwargs,
|
| 99 |
+
):
|
| 100 |
+
# For backward compatibility with diffusion/v2 experiments that use is_x0 instead of do_x0_prediction
|
| 101 |
+
if "is_x0" in kwargs:
|
| 102 |
+
if "do_x0_prediction" in kwargs:
|
| 103 |
+
assert kwargs["do_x0_prediction"] == kwargs["is_x0"], "do_x0_prediction and is_x0 must be the same"
|
| 104 |
+
else:
|
| 105 |
+
kwargs["do_x0_prediction"] = kwargs["is_x0"]
|
| 106 |
+
del kwargs["is_x0"]
|
| 107 |
+
|
| 108 |
+
# For backward compatibility with transfer2 experiments that use n_x0_level instead of n_sigmas_for_x0_prediction
|
| 109 |
+
if n_x0_level is not None:
|
| 110 |
+
if "n_sigmas_for_x0_prediction" in kwargs:
|
| 111 |
+
assert kwargs["n_sigmas_for_x0_prediction"] == n_x0_level, (
|
| 112 |
+
"n_sigmas_for_x0_prediction and n_x0_level must be the same"
|
| 113 |
+
)
|
| 114 |
+
else:
|
| 115 |
+
kwargs["n_sigmas_for_x0_prediction"] = n_x0_level
|
| 116 |
+
|
| 117 |
+
super().__init__(*args, **kwargs)
|
| 118 |
+
self.n_view_embed = n_view_embed
|
| 119 |
+
self.ctrl_hint_keys = ctrl_hint_keys
|
| 120 |
+
self.control_weights = control_weights
|
| 121 |
+
self.num_cond_frames = num_cond_frames
|
| 122 |
+
self.is_x0 = self.do_x0_prediction
|
| 123 |
+
if not hasattr(self, "fix_batch"):
|
| 124 |
+
self.fix_batch = None
|
| 125 |
+
if not hasattr(self, "is_sample"):
|
| 126 |
+
self.is_sample = True
|
| 127 |
+
|
| 128 |
+
def on_train_start(self, model: MultiviewVid2VidModelRectifiedFlow, iteration: int = 0) -> None:
|
| 129 |
+
return super().on_train_start(model, iteration)
|
| 130 |
+
|
| 131 |
+
def _ensure_even_dimensions(self, frame: np.ndarray) -> np.ndarray:
|
| 132 |
+
"""
|
| 133 |
+
ffmpeg (H.264) requires both H and W to be even. If either is odd we pad
|
| 134 |
+
by 1 pixel on the bottom/right using edge-replication.
|
| 135 |
+
"""
|
| 136 |
+
h, w = frame.shape[:2]
|
| 137 |
+
pad_h = h % 2
|
| 138 |
+
pad_w = w % 2
|
| 139 |
+
if pad_h or pad_w:
|
| 140 |
+
frame = cv2.copyMakeBorder(frame, 0, pad_h, 0, pad_w, cv2.BORDER_REPLICATE)
|
| 141 |
+
return frame
|
| 142 |
+
|
| 143 |
+
def save_video(self, grid, video_name, fps: int = 30):
|
| 144 |
+
grid = (grid * 255).astype(np.uint8)
|
| 145 |
+
grid = np.transpose(grid, (1, 2, 3, 0)) # (T, H, W, C)
|
| 146 |
+
|
| 147 |
+
# Convert frames to RGB format and ensure even dimensions
|
| 148 |
+
processed_frames = []
|
| 149 |
+
for frame in grid:
|
| 150 |
+
frame = self._ensure_even_dimensions(frame)
|
| 151 |
+
processed_frames.append(frame)
|
| 152 |
+
|
| 153 |
+
# Use imageio.mimsave instead of ffmpegcv.VideoWriter for better error handling
|
| 154 |
+
try:
|
| 155 |
+
if imageio is not None:
|
| 156 |
+
kwargs = {
|
| 157 |
+
"fps": fps,
|
| 158 |
+
"quality": 5, # Good quality
|
| 159 |
+
"macro_block_size": 1,
|
| 160 |
+
"ffmpeg_params": ["-c:v", "libx264", "-preset", "medium"],
|
| 161 |
+
}
|
| 162 |
+
imageio.mimsave(video_name, processed_frames, "mp4", **kwargs)
|
| 163 |
+
else:
|
| 164 |
+
raise ImportError("imageio not available")
|
| 165 |
+
except Exception as e:
|
| 166 |
+
# Fallback to ffmpegcv if imageio fails
|
| 167 |
+
if ffmpegcv is not None:
|
| 168 |
+
try:
|
| 169 |
+
with ffmpegcv.VideoWriter(video_name, "h264", fps) as writer:
|
| 170 |
+
for frame in processed_frames:
|
| 171 |
+
frame = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
|
| 172 |
+
writer.write(frame)
|
| 173 |
+
except Exception as ffmpeg_error:
|
| 174 |
+
raise RuntimeError(
|
| 175 |
+
f"Both imageio and ffmpegcv failed to save video. Imageio error: {e}, FFmpeg error: {ffmpeg_error}"
|
| 176 |
+
)
|
| 177 |
+
else:
|
| 178 |
+
raise RuntimeError(f"Neither imageio nor ffmpegcv are available. Imageio error: {e}")
|
| 179 |
+
|
| 180 |
+
def run_save(self, to_show, batch_size, n_views, base_fp_wo_ext) -> Optional[str]:
|
| 181 |
+
to_show = (1.0 + torch.stack(to_show, dim=0).clamp(-1, 1)) / 2.0 # [n, b, c, t, h, w]
|
| 182 |
+
is_single_frame = to_show.shape[3] == 1
|
| 183 |
+
n_viz_sample = min(self.n_viz_sample, batch_size)
|
| 184 |
+
|
| 185 |
+
# ! we only save first n_sample_to_save video!
|
| 186 |
+
if self.save_s3 and self.data_parallel_id < self.n_sample_to_save:
|
| 187 |
+
save_img_or_video(
|
| 188 |
+
rearrange(to_show, "n b c t h w -> c t (n h) (b w)"),
|
| 189 |
+
f"s3://rundir/{self.name}/{base_fp_wo_ext}",
|
| 190 |
+
fps=self.fps,
|
| 191 |
+
)
|
| 192 |
+
|
| 193 |
+
file_base_fp = f"{base_fp_wo_ext}_resize.jpg"
|
| 194 |
+
local_path = f"{self.local_dir}/{file_base_fp}"
|
| 195 |
+
|
| 196 |
+
file_base_fp_12frames = f"{base_fp_wo_ext}_12frames.jpg"
|
| 197 |
+
local_path_12frames = f"{self.local_dir}/{file_base_fp_12frames}"
|
| 198 |
+
|
| 199 |
+
if self.rank == 0 and wandb.run:
|
| 200 |
+
if is_single_frame: # image case
|
| 201 |
+
to_show = rearrange(
|
| 202 |
+
to_show[:, :n_viz_sample],
|
| 203 |
+
"n b c t h w -> t c (n h) (b w)",
|
| 204 |
+
)
|
| 205 |
+
image_grid = torchvision.utils.make_grid(to_show, nrow=1, padding=0, normalize=False)
|
| 206 |
+
# resize so that wandb can handle it
|
| 207 |
+
torchvision.utils.save_image(resize_image(image_grid, 1024), local_path, nrow=1, scale_each=True)
|
| 208 |
+
else:
|
| 209 |
+
to_show = to_show[:, :n_viz_sample] # [n, b, c, t, h, w]
|
| 210 |
+
# Select 12 frames for the grid
|
| 211 |
+
_T = to_show.shape[3]
|
| 212 |
+
n = 12
|
| 213 |
+
twelve_frames_list = [round(ix * (_T - 1) / (n - 1)) for ix in range(n)]
|
| 214 |
+
to_show_12frames = to_show[:, :, :, twelve_frames_list]
|
| 215 |
+
to_show_12frames = rearrange(to_show_12frames, "n b c t h w -> 1 c (n h) (b t w)")
|
| 216 |
+
image_grid_12frames = torchvision.utils.make_grid(to_show_12frames, nrow=1, padding=0, normalize=False)
|
| 217 |
+
torchvision.utils.save_image(
|
| 218 |
+
resize_image(image_grid_12frames, 1024), local_path_12frames, nrow=1, scale_each=True
|
| 219 |
+
)
|
| 220 |
+
# Create a single stacked video
|
| 221 |
+
video_tensor = rearrange(to_show, "n b c t h (v w) -> t (n h) (b v w) c", v=n_views)
|
| 222 |
+
|
| 223 |
+
# Resize width to 1024 while preserving aspect ratio (keep float to avoid quantization before resize)
|
| 224 |
+
max_w = 2048
|
| 225 |
+
T, H, W, C = video_tensor.shape
|
| 226 |
+
if W > max_w:
|
| 227 |
+
scale = max_w / W
|
| 228 |
+
new_w = max_w
|
| 229 |
+
new_h = int(H * scale)
|
| 230 |
+
# video_tensor is currently float in 0-1 range -> convert [T, H, W, C] to [T, C, H, W]
|
| 231 |
+
video_tensor_f = video_tensor.permute(0, 3, 1, 2)
|
| 232 |
+
video_tensor_f = F.interpolate(
|
| 233 |
+
video_tensor_f, size=(new_h, new_w), mode="bilinear", align_corners=False
|
| 234 |
+
)
|
| 235 |
+
video_tensor = video_tensor_f.permute(0, 2, 3, 1) # [T, H, W, C]
|
| 236 |
+
|
| 237 |
+
video_tensor = rearrange(video_tensor, "T H W C -> C T H W")
|
| 238 |
+
# Write the video
|
| 239 |
+
video_fp = f"{self.local_dir}/{base_fp_wo_ext}.mp4"
|
| 240 |
+
self.save_video(video_tensor.cpu().numpy(), video_fp, fps=self.fps)
|
| 241 |
+
|
| 242 |
+
return local_path, local_path_12frames, video_fp
|
| 243 |
+
return None
|
| 244 |
+
|
| 245 |
+
@torch.no_grad()
|
| 246 |
+
def every_n_impl(self, trainer, model, data_batch, output_batch, loss, iteration):
|
| 247 |
+
return self.every_n_impl_multiview(
|
| 248 |
+
trainer, model, None, data_batch, output_batch=output_batch, loss=loss, iteration=iteration
|
| 249 |
+
)
|
| 250 |
+
|
| 251 |
+
@torch.no_grad()
|
| 252 |
+
def every_n_impl_multiview(
|
| 253 |
+
self, trainer, model, data_batch_sample_all, data_batch_sample_n, output_batch, loss, iteration
|
| 254 |
+
):
|
| 255 |
+
if self.is_ema:
|
| 256 |
+
if not model.config.ema.enabled:
|
| 257 |
+
return
|
| 258 |
+
context = partial(model.ema_scope, "every_n_sampling")
|
| 259 |
+
else:
|
| 260 |
+
context = nullcontext
|
| 261 |
+
|
| 262 |
+
tag = "ema" if self.is_ema else "reg"
|
| 263 |
+
sample_counter = getattr(trainer, "sample_counter", iteration)
|
| 264 |
+
data_batch_for_info = data_batch_sample_all if data_batch_sample_all is not None else data_batch_sample_n
|
| 265 |
+
batch_info = {
|
| 266 |
+
"data": {
|
| 267 |
+
k: convert_to_primitive(v)
|
| 268 |
+
for k, v in data_batch_for_info.items()
|
| 269 |
+
if is_primitive(v) or isinstance(v, (list, dict))
|
| 270 |
+
},
|
| 271 |
+
"sample_counter": sample_counter,
|
| 272 |
+
"iteration": iteration,
|
| 273 |
+
"sample_n_views": data_batch_for_info["sample_n_views"].cpu().item(),
|
| 274 |
+
"n_view_embed": self.n_view_embed,
|
| 275 |
+
}
|
| 276 |
+
if is_tp_cp_pp_rank0():
|
| 277 |
+
if self.save_s3 and self.data_parallel_id < self.n_sample_to_save:
|
| 278 |
+
easy_io.dump(
|
| 279 |
+
batch_info,
|
| 280 |
+
f"s3://rundir/{self.name}/BatchInfo_ReplicateID{self.data_parallel_id:04d}_Iter{iteration:09d}.json",
|
| 281 |
+
)
|
| 282 |
+
|
| 283 |
+
samples_img_fp = []
|
| 284 |
+
with context():
|
| 285 |
+
if self.is_x0:
|
| 286 |
+
x0_img_fp, mse_loss, sigmas = self.x0_pred(
|
| 287 |
+
trainer,
|
| 288 |
+
model,
|
| 289 |
+
data_batch_for_info,
|
| 290 |
+
output_batch,
|
| 291 |
+
loss,
|
| 292 |
+
iteration,
|
| 293 |
+
)
|
| 294 |
+
if self.save_s3 and self.rank == 0:
|
| 295 |
+
easy_io.dump(
|
| 296 |
+
{
|
| 297 |
+
"mse_loss": mse_loss.tolist(),
|
| 298 |
+
"sigmas": sigmas.tolist(),
|
| 299 |
+
"iteration": iteration,
|
| 300 |
+
},
|
| 301 |
+
f"s3://rundir/{self.name}/{tag}_MSE_Iter{iteration:09d}.json",
|
| 302 |
+
)
|
| 303 |
+
if self.is_sample:
|
| 304 |
+
for data_batch in [data_batch_sample_all, data_batch_sample_n]:
|
| 305 |
+
if data_batch is None:
|
| 306 |
+
samples_img_fp.append(None)
|
| 307 |
+
continue
|
| 308 |
+
sample_img_fp = self.sample(
|
| 309 |
+
trainer,
|
| 310 |
+
model,
|
| 311 |
+
data_batch,
|
| 312 |
+
output_batch,
|
| 313 |
+
loss,
|
| 314 |
+
iteration,
|
| 315 |
+
)
|
| 316 |
+
samples_img_fp.append(sample_img_fp)
|
| 317 |
+
if self.fix_batch is not None:
|
| 318 |
+
misc.to(self.fix_batch, "cpu")
|
| 319 |
+
|
| 320 |
+
dist.barrier()
|
| 321 |
+
if wandb.run:
|
| 322 |
+
sample_counter = getattr(trainer, "sample_counter", iteration)
|
| 323 |
+
data_type = "image" if model.is_image_batch(data_batch) else "video"
|
| 324 |
+
tag += f"_{data_type}"
|
| 325 |
+
info = {
|
| 326 |
+
"trainer/global_step": iteration,
|
| 327 |
+
"sample_counter": sample_counter,
|
| 328 |
+
}
|
| 329 |
+
if self.is_x0:
|
| 330 |
+
info[f"{self.name}/{tag}_x0"] = wandb.Image(x0_img_fp, caption=f"{sample_counter}")
|
| 331 |
+
# convert mse_loss to a dict
|
| 332 |
+
mse_loss = mse_loss.tolist()
|
| 333 |
+
info.update({f"x0_pred_mse_{tag}/Sigma{sigmas[i]:0.5f}": mse_loss[i] for i in range(len(mse_loss))})
|
| 334 |
+
|
| 335 |
+
if self.is_sample:
|
| 336 |
+
sample_all_img_fp, sample_n_img_fp = samples_img_fp
|
| 337 |
+
if sample_all_img_fp is not None:
|
| 338 |
+
# info[f"{self.name}/{tag}_sample"] = wandb.Image(sample_all_img_fp[0], caption=f"{sample_counter}")
|
| 339 |
+
info[f"{self.name}/{tag}_sample_allviews_frames"] = wandb.Image(
|
| 340 |
+
sample_all_img_fp[1], caption=f"{sample_counter}"
|
| 341 |
+
)
|
| 342 |
+
info[f"{self.name}/{tag}_sample_allviews"] = wandb.Video(
|
| 343 |
+
sample_all_img_fp[2], caption=f"{sample_counter}"
|
| 344 |
+
)
|
| 345 |
+
|
| 346 |
+
# info[f"{self.name}/{tag}_sample"] = wandb.Image(sample_n_img_fp[0], caption=f"{sample_counter}")
|
| 347 |
+
info[f"{self.name}/{tag}_sample_nviews_frames"] = wandb.Image(
|
| 348 |
+
sample_n_img_fp[1], caption=f"{sample_counter}"
|
| 349 |
+
)
|
| 350 |
+
info[f"{self.name}/{tag}_sample_nviews"] = wandb.Video(sample_n_img_fp[2], caption=f"{sample_counter}")
|
| 351 |
+
wandb.log(
|
| 352 |
+
info,
|
| 353 |
+
step=iteration,
|
| 354 |
+
)
|
| 355 |
+
torch.cuda.empty_cache()
|
| 356 |
+
|
| 357 |
+
@misc.timer("EveryNDrawSample: sample")
|
| 358 |
+
def sample(self, trainer, model, data_batch, output_batch, loss, iteration):
|
| 359 |
+
"""
|
| 360 |
+
Args:
|
| 361 |
+
skip_save: to make sure FSDP can work, we run forward pass on all ranks even though we only save on rank 0 and 1
|
| 362 |
+
"""
|
| 363 |
+
n_views = len(data_batch["view_indices_selection"][0])
|
| 364 |
+
if self.fix_batch is not None:
|
| 365 |
+
data_batch = misc.to(self.fix_batch, **model.tensor_kwargs)
|
| 366 |
+
tag = "ema" if self.is_ema else "reg"
|
| 367 |
+
raw_data, x0, condition = model.get_data_and_condition(data_batch)
|
| 368 |
+
if self.use_negative_prompt:
|
| 369 |
+
batch_size = x0.shape[0]
|
| 370 |
+
data_batch["neg_t5_text_embeddings"] = misc.to(
|
| 371 |
+
repeat(
|
| 372 |
+
self.negative_prompt_data["t5_text_embeddings"],
|
| 373 |
+
"l ... -> b (v l) ...",
|
| 374 |
+
b=batch_size,
|
| 375 |
+
v=n_views,
|
| 376 |
+
),
|
| 377 |
+
**model.tensor_kwargs,
|
| 378 |
+
)
|
| 379 |
+
assert data_batch["neg_t5_text_embeddings"].shape == data_batch["t5_text_embeddings"].shape, (
|
| 380 |
+
f"{data_batch['neg_t5_text_embeddings'].shape} != {data_batch['t5_text_embeddings'].shape}"
|
| 381 |
+
)
|
| 382 |
+
data_batch["neg_t5_text_mask"] = data_batch["t5_text_mask"]
|
| 383 |
+
|
| 384 |
+
def time_to_width_dimension(mv_video):
|
| 385 |
+
"""
|
| 386 |
+
Args:
|
| 387 |
+
mv_video: (B, C, V * T, H, W)
|
| 388 |
+
Returns:
|
| 389 |
+
(B, C, T, H, V * W)
|
| 390 |
+
"""
|
| 391 |
+
current_view_index_order = [i.item() for i in data_batch["view_indices_selection"][0]]
|
| 392 |
+
expected_view_index_order = visualization_view_index_order
|
| 393 |
+
|
| 394 |
+
# Reorder views to match expected visualization order
|
| 395 |
+
if (
|
| 396 |
+
len(current_view_index_order) == len(expected_view_index_order)
|
| 397 |
+
and current_view_index_order != expected_view_index_order
|
| 398 |
+
):
|
| 399 |
+
# Create mapping from current order to expected order
|
| 400 |
+
reorder_indices = []
|
| 401 |
+
for expected_view in expected_view_index_order:
|
| 402 |
+
if expected_view in current_view_index_order:
|
| 403 |
+
reorder_indices.append(current_view_index_order.index(expected_view))
|
| 404 |
+
|
| 405 |
+
# Reshape to separate view and time dimensions
|
| 406 |
+
B, C, VT, H, W = mv_video.shape
|
| 407 |
+
T = VT // n_views
|
| 408 |
+
mv_video = rearrange(mv_video, "B C (V T) H W -> B C V T H W", V=n_views)
|
| 409 |
+
|
| 410 |
+
# Reorder views according to expected order
|
| 411 |
+
mv_video = mv_video[:, :, reorder_indices, :, :, :]
|
| 412 |
+
|
| 413 |
+
# Reshape back to original format
|
| 414 |
+
mv_video = rearrange(mv_video, "B C V T H W -> B C (V T) H W")
|
| 415 |
+
|
| 416 |
+
return rearrange(mv_video, "B C (V T) H W -> B C T H (V W)", V=n_views)
|
| 417 |
+
|
| 418 |
+
# GPU memory management before sampling to avoid OOM
|
| 419 |
+
if torch.cuda.is_available():
|
| 420 |
+
mem_allocated_before = torch.cuda.memory_allocated() / 1e9
|
| 421 |
+
mem_reserved_before = torch.cuda.memory_reserved() / 1e9
|
| 422 |
+
print(
|
| 423 |
+
f"[Rank {dist.get_rank() if dist.is_initialized() else 0}] Before sampling - "
|
| 424 |
+
f"Allocated: {mem_allocated_before:.2f}GB, Reserved: {mem_reserved_before:.2f}GB"
|
| 425 |
+
)
|
| 426 |
+
|
| 427 |
+
# Clear GPU cache to free up fragmented memory
|
| 428 |
+
torch.cuda.empty_cache()
|
| 429 |
+
|
| 430 |
+
mem_allocated_after = torch.cuda.memory_allocated() / 1e9
|
| 431 |
+
mem_reserved_after = torch.cuda.memory_reserved() / 1e9
|
| 432 |
+
print(
|
| 433 |
+
f"[Rank {dist.get_rank() if dist.is_initialized() else 0}] After cleanup - "
|
| 434 |
+
f"Allocated: {mem_allocated_after:.2f}GB, Reserved: {mem_reserved_after:.2f}GB "
|
| 435 |
+
f"(freed: Allocated={mem_allocated_before - mem_allocated_after:.2f}GB, "
|
| 436 |
+
f"Reserved={mem_reserved_before - mem_reserved_after:.2f}GB)"
|
| 437 |
+
)
|
| 438 |
+
|
| 439 |
+
to_show = []
|
| 440 |
+
# for use_apg in [False, True]:
|
| 441 |
+
for use_apg in [False]:
|
| 442 |
+
for num_cond_frames in self.num_cond_frames:
|
| 443 |
+
for control_weight in self.control_weights:
|
| 444 |
+
data_batch[NUM_CONDITIONAL_FRAMES_KEY] = num_cond_frames
|
| 445 |
+
data_batch[CONTROL_WEIGHT_KEY] = control_weight
|
| 446 |
+
for guidance in self.guidance:
|
| 447 |
+
sample = model.generate_samples_from_batch(
|
| 448 |
+
data_batch,
|
| 449 |
+
guidance=guidance,
|
| 450 |
+
# make sure no mismatch and also works for cp
|
| 451 |
+
state_shape=x0.shape[1:],
|
| 452 |
+
n_sample=x0.shape[0],
|
| 453 |
+
num_steps=self.num_sampling_step,
|
| 454 |
+
is_negative_prompt=True if self.use_negative_prompt else False,
|
| 455 |
+
)
|
| 456 |
+
if hasattr(model, "decode"):
|
| 457 |
+
sample = model.decode(sample)
|
| 458 |
+
to_show.append(sample.float().cpu())
|
| 459 |
+
|
| 460 |
+
to_show.append(raw_data.float().cpu())
|
| 461 |
+
|
| 462 |
+
# Transfer2-multiview: visualize control input
|
| 463 |
+
if self.ctrl_hint_keys:
|
| 464 |
+
# visualize input video
|
| 465 |
+
if "hint_key" in data_batch:
|
| 466 |
+
hint = data_batch[data_batch["hint_key"]]
|
| 467 |
+
for idx in range(0, hint.size(1), 3):
|
| 468 |
+
x_rgb = hint[:, idx : idx + 3]
|
| 469 |
+
to_show.append(x_rgb.float().cpu())
|
| 470 |
+
else:
|
| 471 |
+
for key in self.ctrl_hint_keys:
|
| 472 |
+
if key in data_batch and data_batch[key] is not None:
|
| 473 |
+
hint = data_batch[key]
|
| 474 |
+
log.info(f"hint: {hint.shape}")
|
| 475 |
+
to_show.append(hint.float().cpu())
|
| 476 |
+
|
| 477 |
+
to_show = [time_to_width_dimension(t) for t in to_show]
|
| 478 |
+
|
| 479 |
+
base_fp_wo_ext = f"{tag}_ReplicateID{self.data_parallel_id:04d}_Sample_Iter{iteration:09d}_{n_views}views"
|
| 480 |
+
batch_size = x0.shape[0]
|
| 481 |
+
|
| 482 |
+
if is_tp_cp_pp_rank0():
|
| 483 |
+
local_path = self.run_save(to_show, batch_size, n_views, base_fp_wo_ext)
|
| 484 |
+
return local_path
|
| 485 |
+
return None
|
cosmos_predict2/_src/predict2_multiview/callbacks/frame_loss_log.py
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
import torch
|
| 17 |
+
|
| 18 |
+
from cosmos_predict2._src.imaginaire.model import ImaginaireModel
|
| 19 |
+
from cosmos_predict2._src.imaginaire.utils.callback import Callback
|
| 20 |
+
|
| 21 |
+
"""
|
| 22 |
+
Dummy FrameLossLog callback used in multiview training with view dropout, where batches don't have the same number of views / frames.
|
| 23 |
+
"""
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
class DummyFrameLossLog(Callback):
|
| 27 |
+
def __init__(
|
| 28 |
+
self,
|
| 29 |
+
logging_iter_multipler: int = 1,
|
| 30 |
+
save_logging_iter_multipler: int = 1,
|
| 31 |
+
save_s3: bool = False,
|
| 32 |
+
) -> None:
|
| 33 |
+
pass
|
| 34 |
+
|
| 35 |
+
def on_training_step_end(
|
| 36 |
+
self,
|
| 37 |
+
model: ImaginaireModel,
|
| 38 |
+
data_batch: dict[str, torch.Tensor],
|
| 39 |
+
output_batch: dict[str, torch.Tensor],
|
| 40 |
+
loss: torch.Tensor,
|
| 41 |
+
iteration: int = 0,
|
| 42 |
+
):
|
| 43 |
+
pass
|
cosmos_predict2/_src/predict2_multiview/callbacks/log_weight.py
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
from typing import Optional
|
| 17 |
+
|
| 18 |
+
import torch
|
| 19 |
+
import wandb
|
| 20 |
+
|
| 21 |
+
from cosmos_predict2._src.imaginaire.callbacks.every_n import EveryN
|
| 22 |
+
from cosmos_predict2._src.imaginaire.model import ImaginaireModel
|
| 23 |
+
from cosmos_predict2._src.imaginaire.trainer import ImaginaireTrainer
|
| 24 |
+
from cosmos_predict2._src.imaginaire.utils import distributed, log
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
class LogWeight(EveryN):
|
| 28 |
+
def __init__(
|
| 29 |
+
self,
|
| 30 |
+
every_n: Optional[int] = 100,
|
| 31 |
+
step_size: int = 1,
|
| 32 |
+
barrier_after_run: bool = True,
|
| 33 |
+
run_at_start: bool = False,
|
| 34 |
+
):
|
| 35 |
+
super().__init__(
|
| 36 |
+
every_n=every_n,
|
| 37 |
+
step_size=step_size,
|
| 38 |
+
barrier_after_run=barrier_after_run,
|
| 39 |
+
run_at_start=run_at_start,
|
| 40 |
+
)
|
| 41 |
+
|
| 42 |
+
def every_n_impl(
|
| 43 |
+
self,
|
| 44 |
+
trainer: ImaginaireTrainer,
|
| 45 |
+
model: ImaginaireModel,
|
| 46 |
+
data_batch: dict[str, torch.Tensor],
|
| 47 |
+
output_batch: dict[str, torch.Tensor],
|
| 48 |
+
loss: torch.Tensor,
|
| 49 |
+
iteration: int,
|
| 50 |
+
) -> None:
|
| 51 |
+
if "logging_dict" in output_batch:
|
| 52 |
+
logging_dict = output_batch["logging_dict"]
|
| 53 |
+
|
| 54 |
+
if distributed.is_rank0():
|
| 55 |
+
info = {}
|
| 56 |
+
for k, v in logging_dict.items():
|
| 57 |
+
info[f"model_weight/{k}"] = v
|
| 58 |
+
|
| 59 |
+
if info and wandb.run:
|
| 60 |
+
wandb.log(info, step=iteration)
|
| 61 |
+
|
| 62 |
+
log.info(f"Log weight at iteration {iteration}: {info}")
|
cosmos_predict2/_src/predict2_multiview/callbacks/nymeria_validation_viz.py
ADDED
|
@@ -0,0 +1,335 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
"""Validation visualization for 2-actor Nymeria generation.
|
| 17 |
+
|
| 18 |
+
Every ``every_n`` iterations, generates videos for a FIXED set of ``num_train`` train + ``num_val`` val
|
| 19 |
+
samples and saves, per sample, a comparison video laid out as:
|
| 20 |
+
|
| 21 |
+
| warped cond | pose input | GT video | generated |
|
| 22 |
+
egoA | .... | .... | .... | .... | (row 1)
|
| 23 |
+
egoB | .... | .... | .... | .... | (row 2)
|
| 24 |
+
|
| 25 |
+
i.e. 4 columns (warped / pose / GT / generated) x 2 rows (actor A / actor B). Files are written under
|
| 26 |
+
``<output>/validation/`` and logged to W&B. Samples are deterministic (fixed indices) so progress is
|
| 27 |
+
comparable across iterations.
|
| 28 |
+
"""
|
| 29 |
+
|
| 30 |
+
import os
|
| 31 |
+
from contextlib import nullcontext
|
| 32 |
+
from functools import partial
|
| 33 |
+
from typing import List, Optional
|
| 34 |
+
|
| 35 |
+
import cv2
|
| 36 |
+
import numpy as np
|
| 37 |
+
import torch
|
| 38 |
+
import wandb
|
| 39 |
+
|
| 40 |
+
from cosmos_predict2._src.imaginaire.utils import distributed, log
|
| 41 |
+
from cosmos_predict2._src.imaginaire.utils.callback import Callback
|
| 42 |
+
from cosmos_predict2._src.imaginaire.visualize.video import save_img_or_video
|
| 43 |
+
|
| 44 |
+
_COLS = ["warped", "pose", "GT", "generated"]
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
class NymeriaValidationViz(Callback):
|
| 48 |
+
def __init__(
|
| 49 |
+
self,
|
| 50 |
+
every_n: int = 2000,
|
| 51 |
+
num_train: int = 5,
|
| 52 |
+
num_val: int = 5,
|
| 53 |
+
num_sampling_step: int = 35,
|
| 54 |
+
guidance: float = 7.0,
|
| 55 |
+
fps: int = 10,
|
| 56 |
+
viz_size: int = 192,
|
| 57 |
+
is_ema: bool = False,
|
| 58 |
+
run_at_start: bool = False,
|
| 59 |
+
# nymeria dataset params (match the training dataloader)
|
| 60 |
+
root: str = "/data2/nymeria_processed_longer",
|
| 61 |
+
num_video_frames: int = 77,
|
| 62 |
+
resolution_hw: tuple = (480, 480),
|
| 63 |
+
val_num_sessions: int = 4,
|
| 64 |
+
split_seed: int = 1234,
|
| 65 |
+
single_view: bool = False,
|
| 66 |
+
actor_observer: bool = False, # build the 2-view actor-observer val set (train_split/val_split.csv)
|
| 67 |
+
actor_actor: bool = False, # build the 2-view actor-ACTOR val set (longer dataset, per-view real captions)
|
| 68 |
+
comind: bool = False, # build the 2-view CoMind actor-actor val set (shared-world clip.npz L_/H_ format)
|
| 69 |
+
single_view_stage1: bool = False, # build the STAGE-1 single-view pooled ego-clip val set (real captions)
|
| 70 |
+
num_reference_frames: int = 0, # emit `reference_frames` in the val set (match the training dataloader)
|
| 71 |
+
shared_reference: bool = False, # use SHARED refs (refs_shared.npz) — match the training dataloader
|
| 72 |
+
emit_camera_poses: bool = False, # emit camera_w2c/K in the val set (for Plücker conditioning)
|
| 73 |
+
emit_depth: bool = False, # emit control_input_depth (composite depth) -> adds a "depth" viz column
|
| 74 |
+
person_pose: bool = False, # use identity-colored pose_person in the val set (match training)
|
| 75 |
+
emit_reference_pose: bool = False, # emit reference_pose (ref skeletons) for the reference-pose condition
|
| 76 |
+
train_fixed_ids: Optional[List[str]] = None, # pin the viz TRAIN samples to these pair_ids (stable
|
| 77 |
+
val_fixed_ids: Optional[List[str]] = None, # across runs / manifest-size changes). Falls back to
|
| 78 |
+
# evenly-spaced sampling when None.
|
| 79 |
+
val_manifest: str = "manifest/val_split.csv", # override the val manifest (e.g. a curated subset for
|
| 80 |
+
# checkpoint-eval sweeps: interaction-only, synthetic, ...)
|
| 81 |
+
):
|
| 82 |
+
self.train_fixed_ids = train_fixed_ids
|
| 83 |
+
self.val_fixed_ids = val_fixed_ids
|
| 84 |
+
self.val_manifest = val_manifest
|
| 85 |
+
self.actor_observer = actor_observer
|
| 86 |
+
self.actor_actor = actor_actor
|
| 87 |
+
self.comind = comind
|
| 88 |
+
self.single_view_stage1 = single_view_stage1
|
| 89 |
+
self.num_reference_frames = num_reference_frames
|
| 90 |
+
self.shared_reference = shared_reference
|
| 91 |
+
self.emit_camera_poses = emit_camera_poses
|
| 92 |
+
self.emit_depth = emit_depth
|
| 93 |
+
self.person_pose = person_pose
|
| 94 |
+
self.emit_reference_pose = emit_reference_pose
|
| 95 |
+
self.every_n = every_n
|
| 96 |
+
self.num_train = num_train
|
| 97 |
+
self.num_val = num_val
|
| 98 |
+
self.num_sampling_step = num_sampling_step
|
| 99 |
+
self.guidance = guidance
|
| 100 |
+
self.fps = fps
|
| 101 |
+
self.viz_size = viz_size
|
| 102 |
+
self.is_ema = is_ema
|
| 103 |
+
self.run_at_start = run_at_start
|
| 104 |
+
self._ds_kwargs = dict(
|
| 105 |
+
root=root, num_video_frames=num_video_frames, resolution_hw=tuple(resolution_hw),
|
| 106 |
+
val_num_sessions=val_num_sessions, split_seed=split_seed, single_view=single_view,
|
| 107 |
+
num_reference_frames=num_reference_frames, shared_reference=shared_reference,
|
| 108 |
+
emit_camera_poses=emit_camera_poses, emit_depth=emit_depth,
|
| 109 |
+
person_pose=person_pose, emit_reference_pose=emit_reference_pose,
|
| 110 |
+
)
|
| 111 |
+
self._samples: List[tuple] = [] # (tag, name, dataset, idx)
|
| 112 |
+
|
| 113 |
+
# ------------------------------------------------------------------ setup
|
| 114 |
+
def on_train_start(self, model, iteration: int = 0) -> None:
|
| 115 |
+
from cosmos_predict2._src.predict2_multiview.datasets.nymeria_pairs import (
|
| 116 |
+
NymeriaActorActorDataset,
|
| 117 |
+
NymeriaActorObserverDataset,
|
| 118 |
+
NymeriaPairsConfig,
|
| 119 |
+
NymeriaPairsDataset,
|
| 120 |
+
)
|
| 121 |
+
|
| 122 |
+
self.local_dir = f"{self.config.job.path_local}/validation"
|
| 123 |
+
if distributed.get_rank() == 0:
|
| 124 |
+
os.makedirs(self.local_dir, exist_ok=True)
|
| 125 |
+
log.info(f"NymeriaValidationViz: local_dir={self.local_dir}")
|
| 126 |
+
|
| 127 |
+
if self.comind: # 2-view CoMind actor-actor (shared-world clip.npz L_/H_ format), held-out recs for val
|
| 128 |
+
from cosmos_predict2._src.predict2_multiview.datasets.comind_pairs import (
|
| 129 |
+
ComindActorActorDataset, ComindPairsConfig,
|
| 130 |
+
)
|
| 131 |
+
cok = dict(root=self._ds_kwargs["root"], num_video_frames=self._ds_kwargs["num_video_frames"],
|
| 132 |
+
resolution_hw=self._ds_kwargs["resolution_hw"],
|
| 133 |
+
num_reference_frames=self.num_reference_frames, shared_reference=self.shared_reference,
|
| 134 |
+
emit_camera_poses=self.emit_camera_poses)
|
| 135 |
+
train_ds = ComindActorActorDataset(ComindPairsConfig(split="train", role_mix=False, **cok))
|
| 136 |
+
val_ds = ComindActorActorDataset(ComindPairsConfig(split="val", role_mix=False, **cok))
|
| 137 |
+
elif self.single_view_stage1: # STAGE-1 pooled ego clips (V=1) with real per-clip captions
|
| 138 |
+
from cosmos_predict2._src.predict2_multiview.datasets.nymeria_pairs import (
|
| 139 |
+
NymeriaSingleViewDataset, _STAGE1_SOURCES,
|
| 140 |
+
)
|
| 141 |
+
aok = {k: v for k, v in self._ds_kwargs.items()
|
| 142 |
+
if k in ("num_video_frames", "resolution_hw", "num_reference_frames", "emit_camera_poses")}
|
| 143 |
+
cfg_t = NymeriaPairsConfig(single_view=True, shared_reference=False, **aok)
|
| 144 |
+
train_ds = NymeriaSingleViewDataset(cfg_t, _STAGE1_SOURCES)
|
| 145 |
+
val_srcs = [dict(s, manifest_csv=s["manifest_csv"].replace("train_split", "val_split")) for s in _STAGE1_SOURCES]
|
| 146 |
+
val_ds = NymeriaSingleViewDataset(NymeriaPairsConfig(single_view=True, shared_reference=False, **aok), val_srcs)
|
| 147 |
+
elif self.actor_observer or self.actor_actor: # 2-view (real per-view captions, role-mix off for determinism)
|
| 148 |
+
aok = {k: v for k, v in self._ds_kwargs.items() if k in ("root", "num_video_frames", "resolution_hw", "num_reference_frames", "shared_reference", "emit_camera_poses", "emit_depth", "person_pose", "emit_reference_pose")}
|
| 149 |
+
_DS = NymeriaActorActorDataset if self.actor_actor else NymeriaActorObserverDataset
|
| 150 |
+
train_ds = _DS(NymeriaPairsConfig(manifest_csv="manifest/train_split.csv", role_mix=False, **aok))
|
| 151 |
+
val_ds = _DS(NymeriaPairsConfig(manifest_csv=self.val_manifest, role_mix=False, **aok))
|
| 152 |
+
else:
|
| 153 |
+
dummy = getattr(model.config, "text_encoder_config", None) is None
|
| 154 |
+
train_ds = NymeriaPairsDataset(NymeriaPairsConfig(split="train", dummy_text_embeddings=dummy, **self._ds_kwargs))
|
| 155 |
+
val_ds = NymeriaPairsDataset(NymeriaPairsConfig(split="val", clean_only=True, dummy_text_embeddings=dummy, **self._ds_kwargs))
|
| 156 |
+
|
| 157 |
+
# fixed, evenly-spaced indices (identical on every rank -> consistent collective generation)
|
| 158 |
+
def pick(ds, n):
|
| 159 |
+
n = min(n, len(ds))
|
| 160 |
+
return [int(i * (len(ds) / n)) for i in range(n)] if n else []
|
| 161 |
+
|
| 162 |
+
# pin by pair_id when given (stable across runs / manifest resizes); else evenly-spaced fallback.
|
| 163 |
+
def select(ds, n, fixed_ids):
|
| 164 |
+
if fixed_ids:
|
| 165 |
+
id2idx = {p.get("pair_id"): i for i, p in enumerate(ds.pairs)}
|
| 166 |
+
idxs = [id2idx[pid] for pid in fixed_ids if pid in id2idx]
|
| 167 |
+
missing = [pid for pid in fixed_ids if pid not in id2idx]
|
| 168 |
+
if missing:
|
| 169 |
+
log.warning(f"NymeriaValidationViz: {len(missing)} fixed pair_ids not in dataset: {missing[:3]}...")
|
| 170 |
+
if idxs:
|
| 171 |
+
return idxs
|
| 172 |
+
log.warning("NymeriaValidationViz: no fixed pair_ids matched; falling back to even sampling")
|
| 173 |
+
return pick(ds, n)
|
| 174 |
+
|
| 175 |
+
self._samples = [("train", f"train{i:02d}", train_ds, idx)
|
| 176 |
+
for i, idx in enumerate(select(train_ds, self.num_train, self.train_fixed_ids))]
|
| 177 |
+
self._samples += [("val", f"val{i:02d}", val_ds, idx)
|
| 178 |
+
for i, idx in enumerate(select(val_ds, self.num_val, self.val_fixed_ids))]
|
| 179 |
+
|
| 180 |
+
# run_at_start: generate immediately here (pure inference, before any training step). This lets a
|
| 181 |
+
# single-GPU checkpoint-eval run (max_iter=0) produce all viz without the training-step backward/optim
|
| 182 |
+
# memory that would OOM an unsharded model.
|
| 183 |
+
if self.run_at_start:
|
| 184 |
+
log.info(f"NymeriaValidationViz: run_at_start -> generating {len(self._samples)} samples now")
|
| 185 |
+
self._generate(model, iteration)
|
| 186 |
+
|
| 187 |
+
# ------------------------------------------------------------------ trigger
|
| 188 |
+
def on_training_step_end(self, model, data_batch, output_batch, loss, iteration: int = 0) -> None:
|
| 189 |
+
# run_at_start is handled in on_train_start (pure inference, no training step needed -> lets a
|
| 190 |
+
# single-GPU checkpoint-eval run with max_iter=0 avoid the training-step OOM). Here only periodic.
|
| 191 |
+
if iteration > 0 and iteration % self.every_n == 0:
|
| 192 |
+
self._generate(model, iteration)
|
| 193 |
+
|
| 194 |
+
def _generate(self, model, iteration: int) -> None:
|
| 195 |
+
if not self._samples:
|
| 196 |
+
return
|
| 197 |
+
if self.is_ema and not model.config.ema.enabled:
|
| 198 |
+
return
|
| 199 |
+
context = partial(model.ema_scope, "validation_viz") if self.is_ema else nullcontext
|
| 200 |
+
log_payload = {}
|
| 201 |
+
with context():
|
| 202 |
+
for tag, name, ds, idx in self._samples:
|
| 203 |
+
try:
|
| 204 |
+
path, caption = self._viz_one(model, ds, idx, iteration, name) # collective on all ranks
|
| 205 |
+
except Exception as e: # don't let viz crash training
|
| 206 |
+
log.warning(f"NymeriaValidationViz {name} failed: {type(e).__name__}: {e}")
|
| 207 |
+
continue
|
| 208 |
+
if distributed.get_rank() == 0 and path is not None and wandb.run:
|
| 209 |
+
# caption = per-view text fed to this sample (egoA / egoB), shown under the W&B video
|
| 210 |
+
log_payload[f"validation/{tag}/{name}"] = wandb.Video(path, fps=self.fps, format="mp4", caption=caption)
|
| 211 |
+
if distributed.get_rank() == 0 and (len(self._samples) > 20) and ((self._samples.index((tag, name, ds, idx)) + 1) % 25 == 0):
|
| 212 |
+
log.info(f"NymeriaValidationViz: generated {self._samples.index((tag, name, ds, idx))+1}/{len(self._samples)}")
|
| 213 |
+
if distributed.get_rank() == 0 and wandb.run and log_payload:
|
| 214 |
+
wandb.log(log_payload, step=iteration)
|
| 215 |
+
|
| 216 |
+
# ------------------------------------------------------------------ one sample
|
| 217 |
+
@torch.no_grad()
|
| 218 |
+
def _viz_one(self, model, ds, idx, iteration, name):
|
| 219 |
+
from cosmos_predict2._src.predict2_multiview.datasets.multiview import collate_fn as mv_collate
|
| 220 |
+
|
| 221 |
+
data_batch = mv_collate([ds[idx]])
|
| 222 |
+
from cosmos_predict2._src.imaginaire.utils import misc
|
| 223 |
+
|
| 224 |
+
# per-view captions actually fed to this sample (before online text-embed consumes them)
|
| 225 |
+
caps = None
|
| 226 |
+
_ai = data_batch.get("ai_caption")
|
| 227 |
+
if _ai and isinstance(_ai[0], (list, tuple)):
|
| 228 |
+
caps = list(_ai[0]) # [view0_caption, view1_caption]
|
| 229 |
+
|
| 230 |
+
data_batch = misc.to(data_batch, device="cuda")
|
| 231 |
+
|
| 232 |
+
# real-text runs provide per-view captions (ai_caption), not embeddings -> compute them online first
|
| 233 |
+
# (mirrors training_step_multiview), otherwise the conditioner has no t5_text_embeddings to read.
|
| 234 |
+
tec = getattr(model.config, "text_encoder_config", None)
|
| 235 |
+
if tec is not None and getattr(tec, "compute_online", False):
|
| 236 |
+
model.inplace_compute_text_embeddings_online(data_batch)
|
| 237 |
+
|
| 238 |
+
# generation runs the net in bf16 (training casts inputs explicitly; here we use autocast so the
|
| 239 |
+
# float32 text + sampling inputs line up with the bf16 weights, mirroring the training forward)
|
| 240 |
+
with torch.autocast("cuda", dtype=torch.bfloat16):
|
| 241 |
+
# GT + latent shape (also runs the model's pose preprocessing into data_batch)
|
| 242 |
+
raw_data, x0, _ = model.get_data_and_condition(data_batch)
|
| 243 |
+
sample = model.generate_samples_from_batch(
|
| 244 |
+
data_batch, guidance=self.guidance, state_shape=x0.shape[1:], n_sample=x0.shape[0],
|
| 245 |
+
num_steps=self.num_sampling_step, is_negative_prompt=False,
|
| 246 |
+
)
|
| 247 |
+
if hasattr(model, "decode"):
|
| 248 |
+
sample = model.decode(sample) # (B, 3, V*T, H, W) in [-1, 1]
|
| 249 |
+
|
| 250 |
+
if distributed.get_rank() != 0:
|
| 251 |
+
return None, None
|
| 252 |
+
|
| 253 |
+
def to01(x): # (B,C,V*T,H,W) -> (C,V*T,H,W) in [0,1]
|
| 254 |
+
return ((x[0].float().cpu().clamp(-1, 1) + 1.0) / 2.0)
|
| 255 |
+
|
| 256 |
+
gen = to01(sample)
|
| 257 |
+
gt = to01(raw_data)
|
| 258 |
+
warped = data_batch["control_input_warped"][0].float().cpu() / 255.0 # (3,V*T,H,W)
|
| 259 |
+
pose = data_batch["control_input_pose"][0].float().cpu() / 255.0
|
| 260 |
+
depth = None
|
| 261 |
+
if "control_input_depth" in data_batch: # composite-depth condition (warped scene + human mesh) column
|
| 262 |
+
depth = data_batch["control_input_depth"][0].float().cpu() / 255.0
|
| 263 |
+
n_views = int(data_batch["sample_n_views"][0])
|
| 264 |
+
# shared reference frames fed to this sample (3, R, H, W); normalized to [0,1] for the viz strip
|
| 265 |
+
refs = None
|
| 266 |
+
if "reference_frames" in data_batch:
|
| 267 |
+
rf = data_batch["reference_frames"][0].float().cpu()
|
| 268 |
+
refs = rf / 255.0 if rf.max() > 1.5 else rf
|
| 269 |
+
grid = self._compose(warped, pose, gt, gen, n_views, refs=refs, depth=depth) # (3, T, n_views*s + hs, ncols*s)
|
| 270 |
+
fp = f"{self.local_dir}/{name}_iter{iteration:09d}"
|
| 271 |
+
save_img_or_video(grid, fp, fps=self.fps)
|
| 272 |
+
|
| 273 |
+
# save the per-view captions locally (sidecar .txt) and return them for W&B logging
|
| 274 |
+
labels = ["egoA", "egoB", "egoC", "egoD"]
|
| 275 |
+
caption_str = None
|
| 276 |
+
if caps is not None:
|
| 277 |
+
lines = [f"{labels[v] if v < len(labels) else f'view{v}'}: {c}" for v, c in enumerate(caps)]
|
| 278 |
+
caption_str = "\n".join(lines)
|
| 279 |
+
with open(f"{fp}.txt", "w") as f:
|
| 280 |
+
f.write(caption_str + "\n")
|
| 281 |
+
return f"{fp}.mp4", caption_str
|
| 282 |
+
|
| 283 |
+
# ------------------------------------------------------------------ grid + labels
|
| 284 |
+
def _compose(self, warped, pose, gt, gen, n_views: int = 2, refs=None, depth=None):
|
| 285 |
+
s = self.viz_size
|
| 286 |
+
T = gen.shape[1] // n_views # frames per actor/view
|
| 287 |
+
titles = list(_COLS) + (["depth"] if depth is not None else []) # extra depth-condition column when present
|
| 288 |
+
ncols = len(titles)
|
| 289 |
+
|
| 290 |
+
def cell(x, v): # x:(3,V*T,H,W) -> view v -> (T,s,s,3) numpy uint8
|
| 291 |
+
a = x[:, v * T : (v + 1) * T] # (3,T,H,W)
|
| 292 |
+
a = torch.nn.functional.interpolate(a.permute(1, 0, 2, 3), size=(s, s), mode="bilinear", align_corners=False)
|
| 293 |
+
a = (a.clamp(0, 1).permute(0, 2, 3, 1).numpy() * 255).astype(np.uint8) # (T,s,s,3)
|
| 294 |
+
return a
|
| 295 |
+
|
| 296 |
+
rows = []
|
| 297 |
+
labels = ["egoA", "egoB", "egoC", "egoD"]
|
| 298 |
+
for v in range(n_views):
|
| 299 |
+
row_label = labels[v] if v < len(labels) else f"ego{v}"
|
| 300 |
+
cols = [cell(warped, v), cell(pose, v), cell(gt, v), cell(gen, v)] # each (T,s,s,3)
|
| 301 |
+
if depth is not None:
|
| 302 |
+
cols.append(cell(depth, v))
|
| 303 |
+
row = np.ascontiguousarray(np.concatenate(cols, axis=2)) # (T,s,ncols*s,3)
|
| 304 |
+
# column headers (top of actorA row) + row label (left); per-frame contiguous buffer for cv2
|
| 305 |
+
for t in range(row.shape[0]):
|
| 306 |
+
frame = np.ascontiguousarray(row[t])
|
| 307 |
+
if v == 0:
|
| 308 |
+
for c, title in enumerate(titles):
|
| 309 |
+
cv2.putText(frame, title, (c * s + 4, 16), cv2.FONT_HERSHEY_SIMPLEX, 0.5, (255, 255, 0), 1, cv2.LINE_AA)
|
| 310 |
+
cv2.putText(frame, row_label, (4, s - 8), cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 255, 255), 1, cv2.LINE_AA)
|
| 311 |
+
row[t] = frame
|
| 312 |
+
rows.append(row)
|
| 313 |
+
|
| 314 |
+
# shared-reference strip: R thumbnails side-by-side (static across all T frames), one row ABOVE the
|
| 315 |
+
# views so you can see exactly which reference images conditioned this sample.
|
| 316 |
+
if refs is not None and refs.shape[1] > 0:
|
| 317 |
+
R = refs.shape[1]
|
| 318 |
+
grid_w = ncols * s
|
| 319 |
+
hs = max(1, grid_w // R) # thumbnail size so R thumbs span the full grid width
|
| 320 |
+
thumbs = []
|
| 321 |
+
for r in range(R):
|
| 322 |
+
a = torch.nn.functional.interpolate(refs[:, r].unsqueeze(0), size=(hs, hs), mode="bilinear",
|
| 323 |
+
align_corners=False)[0]
|
| 324 |
+
a = (a.clamp(0, 1).permute(1, 2, 0).numpy() * 255).astype(np.uint8) # (hs,hs,3)
|
| 325 |
+
thumbs.append(a)
|
| 326 |
+
strip = np.ascontiguousarray(np.concatenate(thumbs, axis=1)) # (hs, R*hs, 3)
|
| 327 |
+
if strip.shape[1] != grid_w: # pad/trim to exact grid width
|
| 328 |
+
strip = np.ascontiguousarray(cv2.resize(strip, (grid_w, hs)))
|
| 329 |
+
cv2.putText(strip, f"shared refs (x{R})", (4, 14), cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 255, 0), 1, cv2.LINE_AA)
|
| 330 |
+
ref_row = np.repeat(strip[None], T, axis=0) # (T, hs, 4s, 3) static
|
| 331 |
+
rows.insert(0, ref_row) # TOP strip (above the view rows)
|
| 332 |
+
|
| 333 |
+
grid = np.concatenate(rows, axis=1) # (T, hs + n_views*s, 4s, 3)
|
| 334 |
+
grid = torch.from_numpy(grid).float().div(255).permute(3, 0, 1, 2) # (3, T, H, 4s)
|
| 335 |
+
return grid
|
cosmos_predict2/_src/predict2_multiview/callbacks/sigma_loss_analysis_per_frame.py
ADDED
|
@@ -0,0 +1,338 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
from dataclasses import dataclass
|
| 17 |
+
from typing import List, Optional, Tuple
|
| 18 |
+
|
| 19 |
+
import matplotlib
|
| 20 |
+
import matplotlib.pyplot as plt
|
| 21 |
+
import numpy as np
|
| 22 |
+
import torch
|
| 23 |
+
import torch.distributed as dist
|
| 24 |
+
import wandb
|
| 25 |
+
|
| 26 |
+
from cosmos_predict2._src.imaginaire.model import ImaginaireModel
|
| 27 |
+
from cosmos_predict2._src.imaginaire.utils import distributed, misc
|
| 28 |
+
from cosmos_predict2._src.imaginaire.utils.callback import Callback
|
| 29 |
+
from cosmos_predict2._src.imaginaire.utils.easy_io import easy_io
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
class DummySigmaLossAnalysisPerFrame(Callback):
|
| 33 |
+
def __init__(
|
| 34 |
+
self,
|
| 35 |
+
logging_iter_multipler: int = 1,
|
| 36 |
+
logging_viz_iter_multipler: int = 1,
|
| 37 |
+
save_s3: bool = False,
|
| 38 |
+
) -> None:
|
| 39 |
+
super().__init__()
|
| 40 |
+
pass
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def _get_normal_quantile_bins():
|
| 44 |
+
"""Get predefined bins based on exp(N(0,1)) distribution quantiles"""
|
| 45 |
+
# Using torch.special.ndtri (inverse of standard normal CDF)
|
| 46 |
+
# and taking exponential to get exponentially spaced bins
|
| 47 |
+
probs = torch.linspace(0.0, 1.0, 11) # 11 points gives 10 bins
|
| 48 |
+
points = torch.special.ndtri(probs)
|
| 49 |
+
# Take exponential to get exponentially spaced bins
|
| 50 |
+
points = torch.exp(points)
|
| 51 |
+
# Replace extreme values at boundaries
|
| 52 |
+
points[0] = points[1] / (points[2] / points[1]) # Extrapolate left boundary
|
| 53 |
+
points[-1] = points[-2] * (points[-2] / points[-3]) # Extrapolate right boundary
|
| 54 |
+
return points.numpy()
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
@dataclass
|
| 58 |
+
class _SigmaLossCache:
|
| 59 |
+
def __init__(self):
|
| 60 |
+
self.reset()
|
| 61 |
+
|
| 62 |
+
def reset(self):
|
| 63 |
+
self.sigma_list: List[torch.Tensor] = []
|
| 64 |
+
self.loss_list: List[torch.Tensor] = []
|
| 65 |
+
|
| 66 |
+
def add(self, sigma: torch.Tensor, loss: torch.Tensor):
|
| 67 |
+
# Convert to bf16 and store on CPU
|
| 68 |
+
self.sigma_list.append(sigma.detach().cpu().to(torch.bfloat16))
|
| 69 |
+
self.loss_list.append(loss.detach().cpu().to(torch.bfloat16))
|
| 70 |
+
|
| 71 |
+
def get_arrays(self) -> Tuple[torch.Tensor, torch.Tensor, Optional[int]]:
|
| 72 |
+
if not self.sigma_list:
|
| 73 |
+
return torch.tensor([], dtype=torch.bfloat16), torch.tensor([], dtype=torch.bfloat16), None
|
| 74 |
+
|
| 75 |
+
sigma_arr = torch.cat(self.sigma_list, dim=0) # [B*N, T] or [B*N, 1]
|
| 76 |
+
loss_arr = torch.cat(self.loss_list, dim=0) # [B*N, T]
|
| 77 |
+
|
| 78 |
+
# Handle broadcasting case where sigma is shape [B, 1]
|
| 79 |
+
if sigma_arr.shape[-1] == 1 and loss_arr.shape[-1] > 1:
|
| 80 |
+
sigma_arr = sigma_arr.expand(-1, loss_arr.shape[-1])
|
| 81 |
+
|
| 82 |
+
num_frames = loss_arr.shape[-1] if len(loss_arr.shape) > 1 else 1
|
| 83 |
+
|
| 84 |
+
assert sigma_arr.shape == loss_arr.shape, (sigma_arr.shape, loss_arr.shape)
|
| 85 |
+
|
| 86 |
+
return sigma_arr, loss_arr, num_frames
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
class SigmaLossAnalysisPerFrame(Callback):
|
| 90 |
+
def __init__(
|
| 91 |
+
self,
|
| 92 |
+
logging_iter_multipler: int = 1,
|
| 93 |
+
logging_viz_iter_multipler: int = 1,
|
| 94 |
+
save_s3: bool = False,
|
| 95 |
+
) -> None:
|
| 96 |
+
super().__init__()
|
| 97 |
+
self.save_s3 = save_s3
|
| 98 |
+
self.logging_iter_multipler = logging_iter_multipler
|
| 99 |
+
assert logging_viz_iter_multipler % logging_iter_multipler == 0
|
| 100 |
+
self.logging_viz_iter_multipler = logging_viz_iter_multipler
|
| 101 |
+
self.name = self.__class__.__name__
|
| 102 |
+
|
| 103 |
+
self.image_cache = _SigmaLossCache()
|
| 104 |
+
self.video_cache = _SigmaLossCache()
|
| 105 |
+
|
| 106 |
+
def _create_analysis_plots(
|
| 107 |
+
self,
|
| 108 |
+
sigma_arr: torch.Tensor,
|
| 109 |
+
loss_arr: torch.Tensor,
|
| 110 |
+
frame_idx: Optional[int] = None, # [N] # [N]
|
| 111 |
+
) -> Optional[wandb.Image]:
|
| 112 |
+
if len(sigma_arr) == 0:
|
| 113 |
+
return None
|
| 114 |
+
|
| 115 |
+
# Convert to numpy for plotting
|
| 116 |
+
sigma_np = sigma_arr.cpu().float().numpy()[:800]
|
| 117 |
+
loss_np = loss_arr.cpu().float().numpy()[:800]
|
| 118 |
+
|
| 119 |
+
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 5))
|
| 120 |
+
|
| 121 |
+
# Get predefined bins based on normal distribution quantiles
|
| 122 |
+
sigma_bins = _get_normal_quantile_bins()
|
| 123 |
+
|
| 124 |
+
y_tick_min, y_tick_max = 0, 1.0
|
| 125 |
+
# 2D histogram with exponential sigma bins and fixed [0,1] loss range
|
| 126 |
+
loss_bins = np.linspace(y_tick_min, y_tick_max, 20)
|
| 127 |
+
|
| 128 |
+
counts, xedges, yedges = np.histogram2d(sigma_np, loss_np, bins=(sigma_bins, loss_bins))
|
| 129 |
+
if counts.max() < 0.1:
|
| 130 |
+
return None
|
| 131 |
+
|
| 132 |
+
# Plot heatmap with exponential scale colormap
|
| 133 |
+
im = ax1.imshow(
|
| 134 |
+
counts.T,
|
| 135 |
+
origin="lower",
|
| 136 |
+
aspect="auto",
|
| 137 |
+
extent=[sigma_bins[0], sigma_bins[-1], y_tick_min, y_tick_max],
|
| 138 |
+
norm=matplotlib.colors.LogNorm(vmin=1, vmax=counts.max()),
|
| 139 |
+
)
|
| 140 |
+
plt.colorbar(im, ax=ax1)
|
| 141 |
+
|
| 142 |
+
# Set fixed loss ticks from 0 to 1
|
| 143 |
+
yticks = np.linspace(y_tick_min, y_tick_max, 6)
|
| 144 |
+
ax1.set_yticks(yticks)
|
| 145 |
+
ax1.set_yticklabels([f"{y:.1f}" for y in yticks])
|
| 146 |
+
|
| 147 |
+
ax1.set_xlabel("Sigma (Standard Normal Quantiles)")
|
| 148 |
+
ax1.set_ylabel("Loss")
|
| 149 |
+
title = "Sigma vs Loss Distribution"
|
| 150 |
+
if frame_idx is not None:
|
| 151 |
+
title += f" (Frame {frame_idx})"
|
| 152 |
+
ax1.set_title(title)
|
| 153 |
+
|
| 154 |
+
# Sigma histogram with loss statistics per bin
|
| 155 |
+
hist_counts, _ = np.histogram(sigma_np, bins=sigma_bins)
|
| 156 |
+
bin_indices = np.digitize(sigma_np, sigma_bins) - 1
|
| 157 |
+
|
| 158 |
+
# Calculate statistics per bin
|
| 159 |
+
n_bins = len(sigma_bins) - 1
|
| 160 |
+
means = np.zeros(n_bins)
|
| 161 |
+
stds = np.zeros(n_bins)
|
| 162 |
+
for i in range(n_bins):
|
| 163 |
+
bin_mask = bin_indices == i
|
| 164 |
+
if bin_mask.any():
|
| 165 |
+
means[i] = loss_np[bin_mask].mean()
|
| 166 |
+
stds[i] = loss_np[bin_mask].std()
|
| 167 |
+
else:
|
| 168 |
+
means[i] = np.nan
|
| 169 |
+
stds[i] = np.nan
|
| 170 |
+
|
| 171 |
+
# Plot histogram
|
| 172 |
+
bin_centers = (sigma_bins[:-1] + sigma_bins[1:]) / 2
|
| 173 |
+
ax2.bar(bin_centers, hist_counts, width=np.diff(sigma_bins), alpha=0.3, align="center")
|
| 174 |
+
|
| 175 |
+
# Plot loss statistics on twin axis
|
| 176 |
+
ax2_twin = ax2.twinx()
|
| 177 |
+
valid_mask = ~np.isnan(means)
|
| 178 |
+
ax2_twin.errorbar(
|
| 179 |
+
bin_centers[valid_mask], means[valid_mask], yerr=stds[valid_mask], color="red", fmt="o-", alpha=0.5
|
| 180 |
+
)
|
| 181 |
+
|
| 182 |
+
ax2.set_xlabel("Sigma (Standard Normal Quantiles)")
|
| 183 |
+
ax2.set_ylabel("Count")
|
| 184 |
+
ax2_twin.set_ylabel("Loss (mean ± std)")
|
| 185 |
+
title = "Sigma Distribution with Loss Statistics"
|
| 186 |
+
if frame_idx is not None:
|
| 187 |
+
title += f" (Frame {frame_idx})"
|
| 188 |
+
ax2.set_title(title)
|
| 189 |
+
|
| 190 |
+
# Add grid for better readability
|
| 191 |
+
ax1.grid(True, alpha=0.3)
|
| 192 |
+
ax2.grid(True, alpha=0.3)
|
| 193 |
+
|
| 194 |
+
# Add quantile labels
|
| 195 |
+
probs = torch.linspace(0.0, 1.0, 11) # 10 points for 9 internal quantiles
|
| 196 |
+
quantile_labels = [f"{p:.1%}" for p in probs]
|
| 197 |
+
ax1.set_xticks(sigma_bins[1:-1]) # Skip boundary bins
|
| 198 |
+
ax1.set_xticklabels(quantile_labels[1:-1], rotation=45)
|
| 199 |
+
ax1.set_xscale("log")
|
| 200 |
+
ax2.set_xticks(sigma_bins[1:-1])
|
| 201 |
+
ax2.set_xticklabels(quantile_labels[1:-1], rotation=45)
|
| 202 |
+
ax2.set_xscale("log")
|
| 203 |
+
|
| 204 |
+
plt.tight_layout()
|
| 205 |
+
fig_img = wandb.Image(fig)
|
| 206 |
+
plt.close(fig)
|
| 207 |
+
|
| 208 |
+
return fig_img
|
| 209 |
+
|
| 210 |
+
def _process_frame_stats(self, sigma: torch.Tensor, loss: torch.Tensor, frame_idx: int) -> dict:
|
| 211 |
+
"""Calculate statistics for a specific frame"""
|
| 212 |
+
return {
|
| 213 |
+
"sigma_log_mean": float(sigma.log().mean()),
|
| 214 |
+
"sigma_log_std": float(sigma.log().std()),
|
| 215 |
+
"loss_mean": float(loss.mean()),
|
| 216 |
+
"loss_std": float(loss.std()),
|
| 217 |
+
"loss_min": float(loss.min()),
|
| 218 |
+
"loss_max": float(loss.max()),
|
| 219 |
+
"loss_median": float(loss.median()),
|
| 220 |
+
"loss_q1": float(torch.quantile(loss.float(), 0.25)),
|
| 221 |
+
"loss_q3": float(torch.quantile(loss.float(), 0.75)),
|
| 222 |
+
}
|
| 223 |
+
|
| 224 |
+
def _gather_and_save(self, cache: _SigmaLossCache, iteration: int, prefix: str, log_viz: bool = True) -> dict:
|
| 225 |
+
info = {}
|
| 226 |
+
|
| 227 |
+
# Gather data from all ranks
|
| 228 |
+
local_sigma, local_loss, num_frames = cache.get_arrays()
|
| 229 |
+
world_size = dist.get_world_size()
|
| 230 |
+
|
| 231 |
+
if world_size > 1:
|
| 232 |
+
# Gather sizes first
|
| 233 |
+
local_size = torch.tensor([len(local_sigma)], dtype=torch.long, device="cuda")
|
| 234 |
+
sizes = [torch.zeros_like(local_size) for _ in range(world_size)]
|
| 235 |
+
dist.all_gather(sizes, local_size)
|
| 236 |
+
sizes = [s.item() for s in sizes]
|
| 237 |
+
|
| 238 |
+
# Gather data
|
| 239 |
+
max_size = max(sizes)
|
| 240 |
+
if max_size > 0:
|
| 241 |
+
# Move to GPU for gathering
|
| 242 |
+
padded_sigma = torch.zeros(max_size, num_frames or 1, dtype=torch.bfloat16, device="cuda")
|
| 243 |
+
padded_loss = torch.zeros(max_size, num_frames or 1, dtype=torch.bfloat16, device="cuda")
|
| 244 |
+
|
| 245 |
+
if len(local_sigma) > 0:
|
| 246 |
+
padded_sigma[: len(local_sigma)] = local_sigma.cuda()
|
| 247 |
+
padded_loss[: len(local_loss)] = local_loss.cuda()
|
| 248 |
+
|
| 249 |
+
all_sigma = [torch.zeros_like(padded_sigma) for _ in range(world_size)]
|
| 250 |
+
all_loss = [torch.zeros_like(padded_loss) for _ in range(world_size)]
|
| 251 |
+
|
| 252 |
+
dist.all_gather(all_sigma, padded_sigma)
|
| 253 |
+
dist.all_gather(all_loss, padded_loss)
|
| 254 |
+
|
| 255 |
+
if distributed.is_rank0():
|
| 256 |
+
# Combine data from all ranks
|
| 257 |
+
valid_sigma = []
|
| 258 |
+
valid_loss = []
|
| 259 |
+
for sigma, loss, size in zip(all_sigma, all_loss, sizes):
|
| 260 |
+
if size > 0:
|
| 261 |
+
valid_sigma.append(sigma[:size])
|
| 262 |
+
valid_loss.append(loss[:size])
|
| 263 |
+
|
| 264 |
+
if valid_sigma:
|
| 265 |
+
sigma_arr = torch.cat(valid_sigma)
|
| 266 |
+
loss_arr = torch.cat(valid_loss)
|
| 267 |
+
|
| 268 |
+
# Overall statistics
|
| 269 |
+
info[f"{prefix}/total_samples"] = sigma_arr.shape[0]
|
| 270 |
+
|
| 271 |
+
# Per-frame statistics
|
| 272 |
+
if num_frames and num_frames > 1:
|
| 273 |
+
for t in range(num_frames):
|
| 274 |
+
frame_stats = self._process_frame_stats(sigma_arr[:, t], loss_arr[:, t], t)
|
| 275 |
+
frame_prefix = f"{prefix}/frame_{t}"
|
| 276 |
+
info.update({f"{frame_prefix}/{k}": v for k, v in frame_stats.items()})
|
| 277 |
+
|
| 278 |
+
# Create per-frame visualization
|
| 279 |
+
if log_viz:
|
| 280 |
+
fig_img = self._create_analysis_plots(sigma_arr[:, t], loss_arr[:, t], t)
|
| 281 |
+
if fig_img is not None:
|
| 282 |
+
info[f"{frame_prefix}/distribution_plot"] = fig_img
|
| 283 |
+
else:
|
| 284 |
+
# Single frame case (images or single-frame stats)
|
| 285 |
+
frame_stats = self._process_frame_stats(sigma_arr.squeeze(), loss_arr.squeeze(), None)
|
| 286 |
+
info.update({f"{prefix}/{k}": v for k, v in frame_stats.items()})
|
| 287 |
+
|
| 288 |
+
# Create visualization
|
| 289 |
+
if log_viz:
|
| 290 |
+
fig_img = self._create_analysis_plots(sigma_arr.squeeze(), loss_arr.squeeze())
|
| 291 |
+
if fig_img is not None:
|
| 292 |
+
info[f"{prefix}/distribution_plot"] = fig_img
|
| 293 |
+
|
| 294 |
+
if self.save_s3:
|
| 295 |
+
save_data = {
|
| 296 |
+
"sigma": sigma_arr.cpu(),
|
| 297 |
+
"loss": loss_arr.cpu(),
|
| 298 |
+
"stats": {k: v for k, v in info.items() if not isinstance(v, wandb.Image)},
|
| 299 |
+
}
|
| 300 |
+
easy_io.dump(
|
| 301 |
+
save_data,
|
| 302 |
+
f"s3://rundir/{self.name}/{prefix}_Iter{iteration:09d}.pkl",
|
| 303 |
+
)
|
| 304 |
+
|
| 305 |
+
cache.reset()
|
| 306 |
+
return info
|
| 307 |
+
|
| 308 |
+
def on_training_step_end(
|
| 309 |
+
self,
|
| 310 |
+
model: ImaginaireModel,
|
| 311 |
+
data_batch: dict[str, torch.Tensor],
|
| 312 |
+
output_batch: dict[str, torch.Tensor],
|
| 313 |
+
loss: torch.Tensor,
|
| 314 |
+
iteration: int = 0,
|
| 315 |
+
):
|
| 316 |
+
sigma = output_batch["sigma"]
|
| 317 |
+
loss_per_frame = output_batch["edm_loss_per_frame"]
|
| 318 |
+
|
| 319 |
+
if model.is_image_batch(data_batch):
|
| 320 |
+
self.image_cache.add(sigma, loss_per_frame)
|
| 321 |
+
else:
|
| 322 |
+
self.video_cache.add(sigma, loss_per_frame)
|
| 323 |
+
|
| 324 |
+
if iteration % (self.config.trainer.logging_iter * self.logging_iter_multipler) == 0:
|
| 325 |
+
info = {}
|
| 326 |
+
|
| 327 |
+
with misc.timer("sigma_loss_analysis"):
|
| 328 |
+
log_viz = iteration % (self.config.trainer.logging_iter * self.logging_viz_iter_multipler) == 0
|
| 329 |
+
# Process image data
|
| 330 |
+
if len(self.image_cache.sigma_list) > 0:
|
| 331 |
+
info.update(self._gather_and_save(self.image_cache, iteration, "sigma_loss_image", log_viz=log_viz))
|
| 332 |
+
|
| 333 |
+
# Process video data
|
| 334 |
+
if len(self.video_cache.sigma_list) > 0:
|
| 335 |
+
info.update(self._gather_and_save(self.video_cache, iteration, "sigma_loss_video", log_viz=log_viz))
|
| 336 |
+
|
| 337 |
+
if distributed.is_rank0() and info and wandb.run:
|
| 338 |
+
wandb.log(info, step=iteration)
|
cosmos_predict2/_src/predict2_multiview/conditioner.py
ADDED
|
@@ -0,0 +1,103 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
from __future__ import annotations
|
| 17 |
+
|
| 18 |
+
from typing import List, Optional
|
| 19 |
+
|
| 20 |
+
import torch
|
| 21 |
+
|
| 22 |
+
from cosmos_predict2._src.imaginaire.utils.easy_io import easy_io
|
| 23 |
+
from cosmos_predict2._src.predict2.conditioner import AbstractEmbModel
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
class MVTextAttr(AbstractEmbModel):
|
| 27 |
+
def __init__(
|
| 28 |
+
self,
|
| 29 |
+
input_key: List[str],
|
| 30 |
+
dropout_rate: Optional[float] = 0.0,
|
| 31 |
+
use_empty_string: bool = False,
|
| 32 |
+
empty_string_embeddings_path: str = "s3://bucket/predict2_assets/reason1_empty_string_embeddings.pt",
|
| 33 |
+
credential_path: str = "credentials/s3_training.secret",
|
| 34 |
+
single_caption_length: int = 512,
|
| 35 |
+
):
|
| 36 |
+
super().__init__()
|
| 37 |
+
self._input_key = input_key
|
| 38 |
+
self._dropout_rate = dropout_rate
|
| 39 |
+
# if True, will use empty string embeddings
|
| 40 |
+
# otherwise use zero tensor embeddings
|
| 41 |
+
self.use_empty_string = use_empty_string
|
| 42 |
+
self._empty_string_embeddings_cache = None
|
| 43 |
+
self.empty_string_embeddings_path = empty_string_embeddings_path
|
| 44 |
+
self.credential_path = credential_path
|
| 45 |
+
self.single_caption_length = single_caption_length
|
| 46 |
+
|
| 47 |
+
def forward(self, token: torch.Tensor):
|
| 48 |
+
return {"crossattn_emb": token}
|
| 49 |
+
|
| 50 |
+
def _get_empty_string_embeddings(self) -> torch.Tensor:
|
| 51 |
+
"""Lazy load and cache empty string embeddings."""
|
| 52 |
+
if self._empty_string_embeddings_cache is None:
|
| 53 |
+
self._empty_string_embeddings_cache = easy_io.load(
|
| 54 |
+
self.empty_string_embeddings_path,
|
| 55 |
+
backend_args={"backend": "s3", "s3_credential_path": self.credential_path},
|
| 56 |
+
)
|
| 57 |
+
return self._empty_string_embeddings_cache
|
| 58 |
+
|
| 59 |
+
def random_dropout_input(
|
| 60 |
+
self, in_tensor: torch.Tensor, dropout_rate: Optional[float] = None, key: Optional[str] = None
|
| 61 |
+
) -> torch.Tensor:
|
| 62 |
+
if key is not None and "mask" in key:
|
| 63 |
+
return in_tensor
|
| 64 |
+
|
| 65 |
+
dropout_rate = dropout_rate if dropout_rate is not None else self._dropout_rate
|
| 66 |
+
|
| 67 |
+
B = in_tensor.shape[0] # batch size
|
| 68 |
+
S_per_view = self.single_caption_length # sequence length per view
|
| 69 |
+
if in_tensor.shape[1] % S_per_view != 0:
|
| 70 |
+
raise ValueError(
|
| 71 |
+
f"in_tensor sequence length {in_tensor.shape[1]} is not divisible by single_caption_length {S_per_view}"
|
| 72 |
+
)
|
| 73 |
+
V = in_tensor.shape[1] // S_per_view # number of views
|
| 74 |
+
|
| 75 |
+
# reshape input tensor to [B, V, S, C]
|
| 76 |
+
in_tensor_reshaped = in_tensor.view(B, V, S_per_view, -1)
|
| 77 |
+
|
| 78 |
+
# create independent dropout mask for each view: [B, V]
|
| 79 |
+
dropout_rates = torch.ones(B, V, device=in_tensor.device) * dropout_rate
|
| 80 |
+
keep_mask = torch.bernoulli(1.0 - dropout_rates).type_as(in_tensor)
|
| 81 |
+
keep_mask = keep_mask.view(B, V, 1, 1) # broadcastable shape
|
| 82 |
+
|
| 83 |
+
# prepare empty prompt data
|
| 84 |
+
if not self.use_empty_string:
|
| 85 |
+
empty_string_embeddings = torch.zeros(
|
| 86 |
+
1, 1, S_per_view, in_tensor_reshaped.shape[-1], dtype=in_tensor.dtype, device=in_tensor.device
|
| 87 |
+
)
|
| 88 |
+
else:
|
| 89 |
+
empty_string_embeddings = (
|
| 90 |
+
self._get_empty_string_embeddings().to(dtype=in_tensor.dtype, device=in_tensor.device).unsqueeze(1)
|
| 91 |
+
) # [1, 1, 512, C]
|
| 92 |
+
|
| 93 |
+
# expand empty prompt data to all views: [B, V, S, C]
|
| 94 |
+
empty_string_embeddings = empty_string_embeddings.expand(B, V, S_per_view, -1)
|
| 95 |
+
|
| 96 |
+
# apply dropout independently for each view
|
| 97 |
+
output_reshaped = keep_mask * in_tensor_reshaped + (1.0 - keep_mask) * empty_string_embeddings
|
| 98 |
+
|
| 99 |
+
# reshape back to original shape: [B, V * S, C]
|
| 100 |
+
return output_reshaped.view(B, V * S_per_view, -1)
|
| 101 |
+
|
| 102 |
+
def details(self) -> str:
|
| 103 |
+
return "Output key: [crossattn_emb]"
|
cosmos_predict2/_src/predict2_multiview/configs/__init__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/__init__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/config.py
ADDED
|
@@ -0,0 +1,47 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
from cosmos_predict2._src.imaginaire.flags import INTERNAL
|
| 17 |
+
from cosmos_predict2._src.imaginaire.utils.config_helper import import_all_modules_from_package
|
| 18 |
+
from cosmos_predict2._src.predict2.configs.video2world.config import make_config as vid2vid_make_config
|
| 19 |
+
from cosmos_predict2._src.predict2_multiview.configs.vid2vid.defaults.callbacks import register_callbacks
|
| 20 |
+
from cosmos_predict2._src.predict2_multiview.configs.vid2vid.defaults.conditioner import register_conditioner
|
| 21 |
+
from cosmos_predict2._src.predict2_multiview.configs.vid2vid.defaults.dataloader import (
|
| 22 |
+
register_multiview_dataloader,
|
| 23 |
+
)
|
| 24 |
+
from cosmos_predict2._src.predict2_multiview.configs.vid2vid.defaults.dataloader_local import register_waymo_dataloader
|
| 25 |
+
from cosmos_predict2._src.predict2_multiview.configs.vid2vid.defaults.model import register_model
|
| 26 |
+
from cosmos_predict2._src.predict2_multiview.configs.vid2vid.defaults.net import register_net
|
| 27 |
+
from cosmos_predict2._src.predict2_multiview.configs.vid2vid.defaults.optimizer import register_optimizer
|
| 28 |
+
from cosmos_predict2._src.predict2_multiview.datasets.nymeria_pairs import register_nymeria_pairs_dataloader
|
| 29 |
+
from cosmos_predict2._src.predict2_multiview.datasets.comind_pairs import register_comind_data
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def make_config():
|
| 33 |
+
c = vid2vid_make_config()
|
| 34 |
+
c.job.project = "cosmos_predict2_multiview"
|
| 35 |
+
register_conditioner()
|
| 36 |
+
register_model()
|
| 37 |
+
register_net()
|
| 38 |
+
register_multiview_dataloader()
|
| 39 |
+
register_waymo_dataloader()
|
| 40 |
+
register_nymeria_pairs_dataloader()
|
| 41 |
+
register_comind_data()
|
| 42 |
+
register_callbacks()
|
| 43 |
+
register_optimizer()
|
| 44 |
+
import_all_modules_from_package("cosmos_predict2._src.predict2_multiview.configs.vid2vid.experiment", reload=True)
|
| 45 |
+
if not INTERNAL:
|
| 46 |
+
import_all_modules_from_package("cosmos_predict2.experiments.multiview", reload=True)
|
| 47 |
+
return c
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/__init__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/callbacks.py
ADDED
|
@@ -0,0 +1,50 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
from hydra.core.config_store import ConfigStore
|
| 17 |
+
|
| 18 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 19 |
+
from cosmos_predict2._src.predict2_multiview.callbacks.log_weight import LogWeight
|
| 20 |
+
from cosmos_predict2._src.predict2_multiview.callbacks.sigma_loss_analysis_per_frame import SigmaLossAnalysisPerFrame
|
| 21 |
+
|
| 22 |
+
LOG_SIGMA_LOSS_CALLBACKS = dict(
|
| 23 |
+
sigma_loss_log=L(SigmaLossAnalysisPerFrame)(
|
| 24 |
+
save_s3="${upload_reproducible_setup}",
|
| 25 |
+
logging_iter_multipler=2,
|
| 26 |
+
logging_viz_iter_multipler=10,
|
| 27 |
+
),
|
| 28 |
+
)
|
| 29 |
+
|
| 30 |
+
LOG_WEIGHT_CALLBACKS = dict(
|
| 31 |
+
log_weight=L(LogWeight)(
|
| 32 |
+
every_n=100,
|
| 33 |
+
),
|
| 34 |
+
)
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def register_callbacks():
|
| 38 |
+
cs = ConfigStore.instance()
|
| 39 |
+
cs.store(
|
| 40 |
+
group="callbacks",
|
| 41 |
+
package="trainer.callbacks",
|
| 42 |
+
name="log_sigma_loss",
|
| 43 |
+
node=LOG_SIGMA_LOSS_CALLBACKS,
|
| 44 |
+
)
|
| 45 |
+
cs.store(
|
| 46 |
+
group="callbacks",
|
| 47 |
+
package="trainer.callbacks",
|
| 48 |
+
name="log_weight",
|
| 49 |
+
node=LOG_WEIGHT_CALLBACKS,
|
| 50 |
+
)
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/conditioner.py
ADDED
|
@@ -0,0 +1,720 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
import copy
|
| 17 |
+
import random
|
| 18 |
+
from dataclasses import dataclass, field
|
| 19 |
+
from enum import Enum
|
| 20 |
+
from typing import Any, Dict, List, Optional, Tuple, Union
|
| 21 |
+
|
| 22 |
+
import torch
|
| 23 |
+
from einops import rearrange
|
| 24 |
+
from hydra.core.config_store import ConfigStore
|
| 25 |
+
from omegaconf import ListConfig
|
| 26 |
+
|
| 27 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 28 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyDict
|
| 29 |
+
from cosmos_predict2._src.imaginaire.utils import log
|
| 30 |
+
from cosmos_predict2._src.imaginaire.utils.context_parallel import broadcast_split_tensor
|
| 31 |
+
from cosmos_predict2._src.imaginaire.utils.validator import Validator
|
| 32 |
+
from cosmos_predict2._src.predict2.conditioner import Text2WorldCondition, TextAttr
|
| 33 |
+
from cosmos_predict2._src.predict2.configs.video2world.defaults.conditioner import (
|
| 34 |
+
_SHARED_CONFIG,
|
| 35 |
+
GeneralConditioner,
|
| 36 |
+
ReMapkey,
|
| 37 |
+
Video2WorldCondition,
|
| 38 |
+
)
|
| 39 |
+
from cosmos_predict2._src.predict2_multiview.conditioner import MVTextAttr
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
class ConditionLocation(Enum):
|
| 43 |
+
"""
|
| 44 |
+
Enum representing different camera condition locations for anymulti-to-multiview video generation.
|
| 45 |
+
|
| 46 |
+
Attributes:
|
| 47 |
+
NO_CAM: Indicates no camera is used for conditioning (i.e text2world)
|
| 48 |
+
REF_CAM: Indicates a reference camera is used for conditioning. (i.e single-to-multiview-text2world)
|
| 49 |
+
ANY_CAM: Indicates any camera can be used for conditioning. (i.e any-to-multiview-text2world)
|
| 50 |
+
FIRST_RANDOM_N: Indicates a random number of frames from all cameras are used for conditioning. (i.e video2world-multiview)
|
| 51 |
+
|
| 52 |
+
Note: Multiple locations can be set together when compatible.
|
| 53 |
+
- NO_CAM cannot be set with any other location.
|
| 54 |
+
- ANY_CAM and REF_CAM cannot be set simultaneously.
|
| 55 |
+
- FIRST_RANDOM_N can be set with ANY_CAM or REF_CAM.
|
| 56 |
+
"""
|
| 57 |
+
|
| 58 |
+
NO_CAM = "no_cam"
|
| 59 |
+
REF_CAM = "ref_cam"
|
| 60 |
+
ANY_CAM = "any_cam"
|
| 61 |
+
FIRST_RANDOM_N = "first_random_n"
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
class ConditionLocationListValidator(Validator):
|
| 65 |
+
"""
|
| 66 |
+
Validator for a list of ConditionLocation objects.
|
| 67 |
+
Validates that:
|
| 68 |
+
- NO_CAM is not set with any other location
|
| 69 |
+
- ANY_CAM and REF_CAM are not set together
|
| 70 |
+
"""
|
| 71 |
+
|
| 72 |
+
def __init__(self, default: List[ConditionLocation], hidden=False, tooltip=None):
|
| 73 |
+
self.default = default
|
| 74 |
+
self.hidden = hidden
|
| 75 |
+
self.tooltip = tooltip
|
| 76 |
+
|
| 77 |
+
def validate(self, value: List[ConditionLocation]):
|
| 78 |
+
for v in value:
|
| 79 |
+
if not isinstance(v, ConditionLocation):
|
| 80 |
+
raise TypeError(f"All elements must be ConditionLocation enums, got {type(v)}: {v}")
|
| 81 |
+
if ConditionLocation.NO_CAM in value:
|
| 82 |
+
assert len(value) == 1, f"Cannot set ConditionLocation.NO_CAM and other locations together. Got {value=}"
|
| 83 |
+
elif ConditionLocation.ANY_CAM in value and ConditionLocation.REF_CAM in value:
|
| 84 |
+
raise ValueError("ConditionLocation.ANY_CAM and ConditionLocation.REF_CAM cannot be set together.")
|
| 85 |
+
return value
|
| 86 |
+
|
| 87 |
+
def __repr__(self) -> str:
|
| 88 |
+
return f"ConditionLocationValidator({self.default=}, {self.hidden=})"
|
| 89 |
+
|
| 90 |
+
def json(self):
|
| 91 |
+
return {
|
| 92 |
+
"type": ConditionLocationListValidator.__name__,
|
| 93 |
+
"default": self.default,
|
| 94 |
+
"tooltip": self.tooltip,
|
| 95 |
+
}
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
class ConditionLocationList(list):
|
| 99 |
+
def __init__(self, locations: List[ConditionLocation]):
|
| 100 |
+
enum_locations = []
|
| 101 |
+
for loc in locations:
|
| 102 |
+
if not isinstance(loc, ConditionLocation):
|
| 103 |
+
loc = ConditionLocation(loc) # Will raise ValueError if invalid
|
| 104 |
+
enum_locations.append(loc)
|
| 105 |
+
super().__init__(enum_locations)
|
| 106 |
+
self.validator = ConditionLocationListValidator(default=[])
|
| 107 |
+
self.validator.validate(self)
|
| 108 |
+
|
| 109 |
+
def __repr__(self) -> str:
|
| 110 |
+
return f"ConditionLocationList({super().__repr__()})"
|
| 111 |
+
|
| 112 |
+
def to_json(self):
|
| 113 |
+
return {
|
| 114 |
+
"type": ConditionLocationList.__name__,
|
| 115 |
+
"locations": [location.value for location in self],
|
| 116 |
+
}
|
| 117 |
+
|
| 118 |
+
|
| 119 |
+
@dataclass(frozen=True)
|
| 120 |
+
class MultiViewCondition(Video2WorldCondition):
|
| 121 |
+
state_t: Optional[int] = None
|
| 122 |
+
view_indices_B_T: Optional[torch.Tensor] = None
|
| 123 |
+
ref_cam_view_idx_sample_position: Optional[torch.Tensor] = None
|
| 124 |
+
# Pose / warped-frame conditioning for 2-actor joint generation. All laid out (B, C, V*T, H, W):
|
| 125 |
+
# pose_map_B_C_T_H_W : projected 2D pose RGB at pixel resolution (encoded by the net's pose encoder)
|
| 126 |
+
# pose_latent_B_C_T_H_W : frozen-VAE pose latent (16 ch, latent resolution) for vae_concat/vae_mlp_add
|
| 127 |
+
# warped_latent_B_C_T_H_W : warped past-frame VAE latent (16 ch, latent resolution)
|
| 128 |
+
# visibility_mask_B_C_T_H_W: warp-validity mask (1 ch, latent resolution)
|
| 129 |
+
pose_map_B_C_T_H_W: Optional[torch.Tensor] = None
|
| 130 |
+
pose_latent_B_C_T_H_W: Optional[torch.Tensor] = None
|
| 131 |
+
warped_latent_B_C_T_H_W: Optional[torch.Tensor] = None
|
| 132 |
+
visibility_mask_B_C_T_H_W: Optional[torch.Tensor] = None
|
| 133 |
+
# reference_latent_B_C_R_H_W: clean reference-frame VAE latent (16 ch, latent res), (B, C, V*R, H, W),
|
| 134 |
+
# appended as extra in-context frames inside the net for appearance conditioning
|
| 135 |
+
reference_latent_B_C_R_H_W: Optional[torch.Tensor] = None
|
| 136 |
+
# plucker_map_B_C_T_H_W: per-pixel Plücker ray map (6 ch = dir 3 + moment 3, latent res), (B, 6, V*T, H, W),
|
| 137 |
+
# in a per-pair canonical frame (view0 frame0) for cross-view shared-space grounding
|
| 138 |
+
plucker_map_B_C_T_H_W: Optional[torch.Tensor] = None
|
| 139 |
+
# reference_plucker_map_B_C_R_H_W: Plücker rays of the reference frames' past poses, same canonical frame
|
| 140 |
+
reference_plucker_map_B_C_R_H_W: Optional[torch.Tensor] = None
|
| 141 |
+
# reference_pose_latent_B_C_R_H_W: VAE latent (16 ch) of each reference frame's per-person skeleton render,
|
| 142 |
+
# (B, C, V*R, H, W), added to the reference tokens by the net's zero-init reference_pose_embedder.
|
| 143 |
+
reference_pose_latent_B_C_R_H_W: Optional[torch.Tensor] = None
|
| 144 |
+
# depth_latent_B_C_T_H_W: composite-depth VAE latent (16 ch, latent res) = warped scene depth + human mesh depth
|
| 145 |
+
depth_latent_B_C_T_H_W: Optional[torch.Tensor] = None
|
| 146 |
+
|
| 147 |
+
def set_video_condition(
|
| 148 |
+
self,
|
| 149 |
+
state_t: int,
|
| 150 |
+
gt_frames: torch.Tensor,
|
| 151 |
+
condition_locations: Union[ConditionLocationList, ListConfig] = field(
|
| 152 |
+
default_factory=lambda: ConditionLocationList([])
|
| 153 |
+
),
|
| 154 |
+
random_min_num_conditional_frames_per_view: Optional[int] = None,
|
| 155 |
+
random_max_num_conditional_frames_per_view: Optional[int] = None,
|
| 156 |
+
num_conditional_frames_per_view: Optional[int | List[int]] = None,
|
| 157 |
+
condition_cam_idx: Optional[int] = None,
|
| 158 |
+
view_condition_dropout_max: int = 0,
|
| 159 |
+
conditional_frames_probs: Optional[Dict[int, float]] = None,
|
| 160 |
+
) -> "MultiViewCondition":
|
| 161 |
+
"""
|
| 162 |
+
Sets the video conditioning frames for anymulti-to-multiview generation.
|
| 163 |
+
|
| 164 |
+
This method creates a conditioning mask for the input video frames that determines
|
| 165 |
+
which frames will be used as context frames for generating new frames. The method
|
| 166 |
+
handles video batches (T>1) and does not support images (T=1).
|
| 167 |
+
|
| 168 |
+
Args:
|
| 169 |
+
gt_frames: A tensor of ground truth frames with shape [B, C, T, H, W], where:
|
| 170 |
+
B = batch size
|
| 171 |
+
C = number of channels
|
| 172 |
+
T = number of frames per view * self.sample_n_views
|
| 173 |
+
H = height
|
| 174 |
+
W = width
|
| 175 |
+
|
| 176 |
+
random_min_num_conditional_frames_per_view: Minimum number of frames per view to use for conditioning
|
| 177 |
+
when randomly selecting a number of conditioning frames.
|
| 178 |
+
|
| 179 |
+
random_max_num_conditional_frames_per_view: Maximum number of frames per view to use for conditioning
|
| 180 |
+
when randomly selecting a number of conditioning frames.
|
| 181 |
+
|
| 182 |
+
num_conditional_frames_per_view: Optional[int | List[int]]; If provided, all examples in the batch will use
|
| 183 |
+
exactly this many frames per view for conditioning. If None, a random number of frames per view
|
| 184 |
+
between random_min_num_conditional_frames_per_view and random_max_num_conditional_frames_per_view
|
| 185 |
+
will be selected for each example in the batch. Can also be a list of integers, one for each view.
|
| 186 |
+
|
| 187 |
+
condition_cam_idx: Optional; Used only if ConditionLocation.ANY_CAM is in condition_locations.
|
| 188 |
+
If provided, all examples in the batch will use the same cam_idx for conditioning. If None,
|
| 189 |
+
a random cam_idx will be selected for each example in the batch.
|
| 190 |
+
view_condition_dropout_max: Optional; If provided and > 0, then a random number of views will be dropped from the conditioning.
|
| 191 |
+
|
| 192 |
+
conditional_frames_probs: Optional; Dictionary mapping number of frames to probabilities.
|
| 193 |
+
If provided, overrides the random_min/max_num_conditional_frames with weighted sampling.
|
| 194 |
+
Example: {0: 0.5, 1: 0.25, 2: 0.25} for 50% chance of 0 frames, 25% for 1, 25% for 2.
|
| 195 |
+
|
| 196 |
+
Returns:
|
| 197 |
+
A new MultiViewCondition object with the gt_frames and conditioning mask set.
|
| 198 |
+
The conditioning mask (condition_video_input_mask_B_C_T_H_W) is a binary tensor
|
| 199 |
+
of shape [B, 1, T, H, W] where 1 indicates frames used for conditioning and 0
|
| 200 |
+
indicates frames to be generated.
|
| 201 |
+
|
| 202 |
+
Notes:
|
| 203 |
+
- Image batches are not supported.
|
| 204 |
+
- For video batches multiple condition_locations can be provided and combined:
|
| 205 |
+
- If num_conditional_frames_per_view is provided and "random_n" is in condition_locations,
|
| 206 |
+
then all examples will use the same number of frames per view for conditioning,
|
| 207 |
+
otherwise, if num_conditional_frames_per_view is not provided,
|
| 208 |
+
then each example will randomly uses between random_min_num_conditional_frames_per_view
|
| 209 |
+
and random_max_num_conditional_frames_per_view frames per view.
|
| 210 |
+
- If "ref_cam" is in condition_locations, then for each example,
|
| 211 |
+
all frames of the first view will be used for conditioning.
|
| 212 |
+
"""
|
| 213 |
+
kwargs = self.to_dict(skip_underscore=False)
|
| 214 |
+
kwargs["state_t"] = state_t
|
| 215 |
+
kwargs["gt_frames"] = gt_frames
|
| 216 |
+
B, _, T, H, W = gt_frames.shape
|
| 217 |
+
|
| 218 |
+
if not isinstance(condition_locations, ConditionLocationList):
|
| 219 |
+
condition_locations = ConditionLocationList(condition_locations)
|
| 220 |
+
assert len(condition_locations) > 0, "condition_locations must be provided."
|
| 221 |
+
assert state_t is not None, "state_t must be provided."
|
| 222 |
+
assert T > 1, "Image batches are not supported."
|
| 223 |
+
assert T % state_t == 0, f"T must be a multiple of state_t. Got T={T} and state_t={state_t}."
|
| 224 |
+
sample_n_views = T // state_t
|
| 225 |
+
condition_video_input_mask_B_C_V_T_H_W = torch.zeros(
|
| 226 |
+
B, 1, sample_n_views, state_t, H, W, dtype=gt_frames.dtype, device=gt_frames.device
|
| 227 |
+
)
|
| 228 |
+
views_eligible_for_dropout = list(range(sample_n_views))
|
| 229 |
+
|
| 230 |
+
if ConditionLocation.REF_CAM in condition_locations:
|
| 231 |
+
ref_cam_view_idx_sample_position = kwargs["ref_cam_view_idx_sample_position"]
|
| 232 |
+
ref_cam_idx_B = (
|
| 233 |
+
torch.ones(B, dtype=torch.int32, device=ref_cam_view_idx_sample_position.device)
|
| 234 |
+
* ref_cam_view_idx_sample_position
|
| 235 |
+
)
|
| 236 |
+
condition_video_input_mask_B_C_V_T_H_W = self.enable_ref_cam_condition(
|
| 237 |
+
ref_cam_idx_B, condition_video_input_mask_B_C_V_T_H_W
|
| 238 |
+
)
|
| 239 |
+
assert (ref_cam_view_idx_sample_position == ref_cam_view_idx_sample_position[0]).all(), (
|
| 240 |
+
f"ref_cam_view_idx_sample_position must be the same for all examples. Got {ref_cam_view_idx_sample_position=}"
|
| 241 |
+
)
|
| 242 |
+
ref_cam_view_idx_sample_position_int = ref_cam_view_idx_sample_position[0].cpu().item()
|
| 243 |
+
views_eligible_for_dropout.remove(ref_cam_view_idx_sample_position_int)
|
| 244 |
+
elif ConditionLocation.ANY_CAM in condition_locations:
|
| 245 |
+
if condition_cam_idx is None:
|
| 246 |
+
assert kwargs["view_indices_B_T"].shape[-1] % sample_n_views == 0, (
|
| 247 |
+
f"view_indices_B_T last dimension must be a multiple of sample_n_views. Got view_indices_B_T.shape={kwargs['view_indices_B_T'].shape}, sample_n_views={sample_n_views}"
|
| 248 |
+
)
|
| 249 |
+
view_indices = kwargs["view_indices_B_T"]
|
| 250 |
+
selected_cam_latent_t_index = torch.randint(0, state_t, size=(B,))
|
| 251 |
+
any_cam_idx_B = view_indices[torch.arange(B), selected_cam_latent_t_index]
|
| 252 |
+
else:
|
| 253 |
+
any_cam_idx_B = torch.full((B,), condition_cam_idx, dtype=torch.int32)
|
| 254 |
+
condition_video_input_mask_B_C_V_T_H_W = self.enable_ref_cam_condition(
|
| 255 |
+
any_cam_idx_B, condition_video_input_mask_B_C_V_T_H_W
|
| 256 |
+
)
|
| 257 |
+
assert (any_cam_idx_B == any_cam_idx_B[0]).all(), (
|
| 258 |
+
f"any_cam_idx_B must be the same for all examples. Got {any_cam_idx_B=}"
|
| 259 |
+
)
|
| 260 |
+
any_cam_idx_B_int = any_cam_idx_B[0].cpu().item()
|
| 261 |
+
views_eligible_for_dropout.remove(any_cam_idx_B_int)
|
| 262 |
+
if ConditionLocation.FIRST_RANDOM_N in condition_locations:
|
| 263 |
+
if (
|
| 264 |
+
num_conditional_frames_per_view is None
|
| 265 |
+
and random_min_num_conditional_frames_per_view == random_max_num_conditional_frames_per_view
|
| 266 |
+
):
|
| 267 |
+
num_conditional_frames_per_view = random_min_num_conditional_frames_per_view
|
| 268 |
+
if num_conditional_frames_per_view is not None:
|
| 269 |
+
if isinstance(num_conditional_frames_per_view, list):
|
| 270 |
+
assert len(num_conditional_frames_per_view) == sample_n_views, (
|
| 271 |
+
f"num_conditional_frames_per_view must be a list of length {sample_n_views}. Got {num_conditional_frames_per_view=}"
|
| 272 |
+
)
|
| 273 |
+
log.info(
|
| 274 |
+
f"Setting num_conditional_frames_per_view_B_V explicitly from list: {num_conditional_frames_per_view}"
|
| 275 |
+
)
|
| 276 |
+
num_conditional_frames_per_view_B_V = torch.tensor(
|
| 277 |
+
num_conditional_frames_per_view, dtype=torch.int32
|
| 278 |
+
).repeat(B, 1)
|
| 279 |
+
else:
|
| 280 |
+
num_conditional_frames_per_view_B_V = (
|
| 281 |
+
torch.ones((B, sample_n_views), dtype=torch.int32) * num_conditional_frames_per_view
|
| 282 |
+
)
|
| 283 |
+
elif conditional_frames_probs is not None:
|
| 284 |
+
# Use weighted sampling based on provided probabilities
|
| 285 |
+
frames_options = list(conditional_frames_probs.keys())
|
| 286 |
+
weights = list(conditional_frames_probs.values())
|
| 287 |
+
num_conditional_frames_per_view_B_V = (
|
| 288 |
+
torch.tensor(random.choices(frames_options, weights=weights, k=B), dtype=torch.int32)
|
| 289 |
+
.view(B, 1)
|
| 290 |
+
.repeat(1, sample_n_views)
|
| 291 |
+
)
|
| 292 |
+
else:
|
| 293 |
+
assert (
|
| 294 |
+
random_min_num_conditional_frames_per_view is not None
|
| 295 |
+
and random_max_num_conditional_frames_per_view is not None
|
| 296 |
+
), (
|
| 297 |
+
f"random_min_num_conditional_frames_per_view and random_max_num_conditional_frames_per_view must be provided if num_conditional_frames_per_view is None. Got {random_min_num_conditional_frames_per_view=}, {random_max_num_conditional_frames_per_view=}, {num_conditional_frames_per_view=}"
|
| 298 |
+
)
|
| 299 |
+
num_conditional_frames_per_view_B_V = torch.randint(
|
| 300 |
+
random_min_num_conditional_frames_per_view,
|
| 301 |
+
random_max_num_conditional_frames_per_view + 1,
|
| 302 |
+
size=(B, 1),
|
| 303 |
+
).repeat(1, sample_n_views)
|
| 304 |
+
condition_video_input_mask_B_C_V_T_H_W = self.enable_first_random_n_condition(
|
| 305 |
+
condition_video_input_mask_B_C_V_T_H_W, num_conditional_frames_per_view_B_V
|
| 306 |
+
)
|
| 307 |
+
if view_condition_dropout_max > 0:
|
| 308 |
+
random.shuffle(views_eligible_for_dropout)
|
| 309 |
+
n_views_to_dropout = random.randint(0, view_condition_dropout_max)
|
| 310 |
+
views_to_dropout = views_eligible_for_dropout[:n_views_to_dropout]
|
| 311 |
+
for view_idx in views_to_dropout:
|
| 312 |
+
condition_video_input_mask_B_C_V_T_H_W[:, :, view_idx] = 0
|
| 313 |
+
|
| 314 |
+
condition_video_input_mask_B_C_T_H_W = rearrange(
|
| 315 |
+
condition_video_input_mask_B_C_V_T_H_W, "B C V T H W -> B C (V T) H W", V=sample_n_views
|
| 316 |
+
)
|
| 317 |
+
kwargs["condition_video_input_mask_B_C_T_H_W"] = condition_video_input_mask_B_C_T_H_W
|
| 318 |
+
return type(self)(**kwargs)
|
| 319 |
+
|
| 320 |
+
def enable_ref_cam_condition(self, cam_idx_B: torch.Tensor, condition_video_input_mask_B_C_V_T_H_W: torch.Tensor):
|
| 321 |
+
"""
|
| 322 |
+
Sets condition video input mask to 1 for all frames of the cam_idx[i] view in each example i
|
| 323 |
+
Args:
|
| 324 |
+
cam_idx_B: A tensor of shape [B]
|
| 325 |
+
condition_video_input_mask_B_C_V_T_H_W: A tensor of shape [B, 1, V, T, H, W]
|
| 326 |
+
where V is the number of views, T is the number of frames per view, H is the height, and W is the width
|
| 327 |
+
Returns:
|
| 328 |
+
A copy of the condition video input mask with the cam_idx[i] view set to 1 for example i
|
| 329 |
+
"""
|
| 330 |
+
assert condition_video_input_mask_B_C_V_T_H_W.ndim == 6, (
|
| 331 |
+
f"condition_video_input_mask_B_C_V_T_H_W must have 6 dimensions. Got {condition_video_input_mask_B_C_V_T_H_W.shape=}"
|
| 332 |
+
)
|
| 333 |
+
assert cam_idx_B.ndim == 1, f"cam_idx_B must have 1 dimension. Got {cam_idx_B.shape=}"
|
| 334 |
+
copy_condition_video_input_mask_B_C_V_T_H_W = condition_video_input_mask_B_C_V_T_H_W.clone()
|
| 335 |
+
for i in range(copy_condition_video_input_mask_B_C_V_T_H_W.shape[0]):
|
| 336 |
+
copy_condition_video_input_mask_B_C_V_T_H_W[i, :, cam_idx_B[i]] = 1
|
| 337 |
+
return copy_condition_video_input_mask_B_C_V_T_H_W
|
| 338 |
+
|
| 339 |
+
def enable_first_random_n_condition(
|
| 340 |
+
self, condition_video_input_mask_B_C_V_T_H_W: torch.Tensor, num_conditional_frames_per_view_B_V: torch.Tensor
|
| 341 |
+
):
|
| 342 |
+
"""
|
| 343 |
+
Sets condition video input mask to 1 for the first num_conditional_frames_per_view_B frames of each view
|
| 344 |
+
Args:
|
| 345 |
+
condition_video_input_mask_B_C_V_T_H_W: A tensor of shape [B, 1, V, T, H, W]
|
| 346 |
+
num_conditional_frames_per_view_B_V: A tensor of shape [B, V]
|
| 347 |
+
Returns:
|
| 348 |
+
A copy of the condition video input mask with the first num_conditional_frames_per_view_B_V frames of each view set to 1
|
| 349 |
+
"""
|
| 350 |
+
assert condition_video_input_mask_B_C_V_T_H_W.ndim == 6, (
|
| 351 |
+
"condition_video_input_mask_B_C_V_T_H_W must have 6 dimensions"
|
| 352 |
+
)
|
| 353 |
+
B, _, _, _, _, _ = condition_video_input_mask_B_C_V_T_H_W.shape
|
| 354 |
+
copy_condition_video_input_mask_B_C_V_T_H_W = condition_video_input_mask_B_C_V_T_H_W.clone()
|
| 355 |
+
for idx in range(B):
|
| 356 |
+
for view_idx in range(num_conditional_frames_per_view_B_V.shape[1]):
|
| 357 |
+
copy_condition_video_input_mask_B_C_V_T_H_W[
|
| 358 |
+
idx, :, view_idx, : num_conditional_frames_per_view_B_V[idx, view_idx]
|
| 359 |
+
] = 1
|
| 360 |
+
return copy_condition_video_input_mask_B_C_V_T_H_W
|
| 361 |
+
|
| 362 |
+
def edit_for_inference(
|
| 363 |
+
self,
|
| 364 |
+
condition_locations: Union[ConditionLocationList, ListConfig] = field(
|
| 365 |
+
default_factory=lambda: ConditionLocationList([])
|
| 366 |
+
),
|
| 367 |
+
is_cfg_conditional: bool = True,
|
| 368 |
+
num_conditional_frames_per_view: int = 1,
|
| 369 |
+
) -> "MultiViewCondition":
|
| 370 |
+
_condition = self.set_video_condition(
|
| 371 |
+
state_t=self.state_t,
|
| 372 |
+
gt_frames=self.gt_frames,
|
| 373 |
+
condition_locations=condition_locations,
|
| 374 |
+
random_min_num_conditional_frames_per_view=0,
|
| 375 |
+
random_max_num_conditional_frames_per_view=0,
|
| 376 |
+
num_conditional_frames_per_view=num_conditional_frames_per_view,
|
| 377 |
+
view_condition_dropout_max=0,
|
| 378 |
+
)
|
| 379 |
+
if not is_cfg_conditional:
|
| 380 |
+
# Do not use classifier free guidance on conditional frames.
|
| 381 |
+
# YB found that it leads to worse results.
|
| 382 |
+
_condition.use_video_condition.fill_(True)
|
| 383 |
+
return _condition
|
| 384 |
+
|
| 385 |
+
def broadcast(self, process_group: torch.distributed.ProcessGroup) -> "MultiViewCondition":
|
| 386 |
+
if self.is_broadcasted:
|
| 387 |
+
return self
|
| 388 |
+
gt_frames_B_C_T_H_W = self.gt_frames
|
| 389 |
+
view_indices_B_T = self.view_indices_B_T
|
| 390 |
+
condition_video_input_mask_B_C_T_H_W = self.condition_video_input_mask_B_C_T_H_W
|
| 391 |
+
# Pose conditioning tensors are laid out (V*T) like gt_frames. Keep them out of the generic
|
| 392 |
+
# broadcast (which would .cuda()+broadcast from rank0) and re-attach unchanged. For CP==1 (our
|
| 393 |
+
# target) this pass-through is correct; CP>1 would need the same seq-split as gt_frames below.
|
| 394 |
+
pose_map_B_C_T_H_W = self.pose_map_B_C_T_H_W
|
| 395 |
+
pose_latent_B_C_T_H_W = self.pose_latent_B_C_T_H_W
|
| 396 |
+
warped_latent_B_C_T_H_W = self.warped_latent_B_C_T_H_W
|
| 397 |
+
visibility_mask_B_C_T_H_W = self.visibility_mask_B_C_T_H_W
|
| 398 |
+
reference_latent_B_C_R_H_W = self.reference_latent_B_C_R_H_W
|
| 399 |
+
plucker_map_B_C_T_H_W = self.plucker_map_B_C_T_H_W
|
| 400 |
+
reference_plucker_map_B_C_R_H_W = self.reference_plucker_map_B_C_R_H_W
|
| 401 |
+
reference_pose_latent_B_C_R_H_W = self.reference_pose_latent_B_C_R_H_W
|
| 402 |
+
depth_latent_B_C_T_H_W = self.depth_latent_B_C_T_H_W
|
| 403 |
+
kwargs = self.to_dict(skip_underscore=False)
|
| 404 |
+
kwargs["gt_frames"] = None
|
| 405 |
+
kwargs["condition_video_input_mask_B_C_T_H_W"] = None
|
| 406 |
+
kwargs["view_indices_B_T"] = None
|
| 407 |
+
kwargs["pose_map_B_C_T_H_W"] = None
|
| 408 |
+
kwargs["pose_latent_B_C_T_H_W"] = None
|
| 409 |
+
kwargs["warped_latent_B_C_T_H_W"] = None
|
| 410 |
+
kwargs["visibility_mask_B_C_T_H_W"] = None
|
| 411 |
+
kwargs["reference_latent_B_C_R_H_W"] = None
|
| 412 |
+
kwargs["plucker_map_B_C_T_H_W"] = None
|
| 413 |
+
kwargs["reference_plucker_map_B_C_R_H_W"] = None
|
| 414 |
+
kwargs["reference_pose_latent_B_C_R_H_W"] = None
|
| 415 |
+
kwargs["depth_latent_B_C_T_H_W"] = None
|
| 416 |
+
new_condition = Text2WorldCondition.broadcast(
|
| 417 |
+
type(self)(**kwargs),
|
| 418 |
+
process_group,
|
| 419 |
+
)
|
| 420 |
+
|
| 421 |
+
kwargs = new_condition.to_dict(skip_underscore=False)
|
| 422 |
+
_, _, T, _, _ = gt_frames_B_C_T_H_W.shape
|
| 423 |
+
n_views = T // self.state_t
|
| 424 |
+
assert T % self.state_t == 0, f"T must be a multiple of state_t. Got T={T} and state_t={self.state_t}."
|
| 425 |
+
if process_group is not None:
|
| 426 |
+
if T > 1 and process_group.size() > 1:
|
| 427 |
+
log.debug(f"Broadcasting {gt_frames_B_C_T_H_W.shape=} to {n_views=} views")
|
| 428 |
+
gt_frames_B_C_V_T_H_W = rearrange(gt_frames_B_C_T_H_W, "B C (V T) H W -> B C V T H W", V=n_views)
|
| 429 |
+
condition_video_input_mask_B_C_V_T_H_W = rearrange(
|
| 430 |
+
condition_video_input_mask_B_C_T_H_W, "B C (V T) H W -> B C V T H W", V=n_views
|
| 431 |
+
)
|
| 432 |
+
view_indices_B_V_T = rearrange(view_indices_B_T, "B (V T) -> B V T", V=n_views)
|
| 433 |
+
|
| 434 |
+
gt_frames_B_C_V_T_H_W = broadcast_split_tensor(
|
| 435 |
+
gt_frames_B_C_V_T_H_W, seq_dim=3, process_group=process_group
|
| 436 |
+
)
|
| 437 |
+
condition_video_input_mask_B_C_V_T_H_W = broadcast_split_tensor(
|
| 438 |
+
condition_video_input_mask_B_C_V_T_H_W, seq_dim=3, process_group=process_group
|
| 439 |
+
)
|
| 440 |
+
view_indices_B_V_T = broadcast_split_tensor(view_indices_B_V_T, seq_dim=2, process_group=process_group)
|
| 441 |
+
|
| 442 |
+
gt_frames_B_C_T_H_W = rearrange(gt_frames_B_C_V_T_H_W, "B C V T H W -> B C (V T) H W", V=n_views)
|
| 443 |
+
condition_video_input_mask_B_C_T_H_W = rearrange(
|
| 444 |
+
condition_video_input_mask_B_C_V_T_H_W, "B C V T H W -> B C (V T) H W", V=n_views
|
| 445 |
+
)
|
| 446 |
+
view_indices_B_T = rearrange(view_indices_B_V_T, "B V T -> B (V T)", V=n_views)
|
| 447 |
+
|
| 448 |
+
kwargs["gt_frames"] = gt_frames_B_C_T_H_W
|
| 449 |
+
kwargs["condition_video_input_mask_B_C_T_H_W"] = condition_video_input_mask_B_C_T_H_W
|
| 450 |
+
kwargs["view_indices_B_T"] = view_indices_B_T
|
| 451 |
+
kwargs["pose_map_B_C_T_H_W"] = pose_map_B_C_T_H_W
|
| 452 |
+
kwargs["pose_latent_B_C_T_H_W"] = pose_latent_B_C_T_H_W
|
| 453 |
+
kwargs["warped_latent_B_C_T_H_W"] = warped_latent_B_C_T_H_W
|
| 454 |
+
kwargs["visibility_mask_B_C_T_H_W"] = visibility_mask_B_C_T_H_W
|
| 455 |
+
kwargs["reference_latent_B_C_R_H_W"] = reference_latent_B_C_R_H_W
|
| 456 |
+
kwargs["plucker_map_B_C_T_H_W"] = plucker_map_B_C_T_H_W
|
| 457 |
+
kwargs["reference_plucker_map_B_C_R_H_W"] = reference_plucker_map_B_C_R_H_W
|
| 458 |
+
kwargs["reference_pose_latent_B_C_R_H_W"] = reference_pose_latent_B_C_R_H_W
|
| 459 |
+
kwargs["depth_latent_B_C_T_H_W"] = depth_latent_B_C_T_H_W
|
| 460 |
+
return type(self)(**kwargs)
|
| 461 |
+
|
| 462 |
+
|
| 463 |
+
class MultiViewConditioner(GeneralConditioner):
|
| 464 |
+
def _forward(self, batch: Dict, override_dropout_rate: Optional[Dict[str, float]] = None) -> Dict:
|
| 465 |
+
"""Like GeneralConditioner._forward but SKIPS any embedder whose (string) input_key is absent from the
|
| 466 |
+
batch. This lets optional conditioning (reference_latent / plucker_map) be present in the shared
|
| 467 |
+
conditioner config while individual experiments produce only the subset they enable."""
|
| 468 |
+
from collections import defaultdict
|
| 469 |
+
from contextlib import nullcontext
|
| 470 |
+
|
| 471 |
+
output = defaultdict(list)
|
| 472 |
+
override_dropout_rate = override_dropout_rate or {}
|
| 473 |
+
for emb_name in override_dropout_rate.keys():
|
| 474 |
+
assert emb_name in self.embedders, f"invalid name found {emb_name}"
|
| 475 |
+
for emb_name, embedder in self.embedders.items():
|
| 476 |
+
if isinstance(embedder.input_key, str) and embedder.input_key not in batch:
|
| 477 |
+
continue # optional condition not produced by this experiment -> skip (no KeyError)
|
| 478 |
+
context = nullcontext if embedder.is_trainable else torch.no_grad
|
| 479 |
+
with context():
|
| 480 |
+
if isinstance(embedder.input_key, str):
|
| 481 |
+
emb_out = embedder(
|
| 482 |
+
embedder.random_dropout_input(
|
| 483 |
+
batch[embedder.input_key], override_dropout_rate.get(emb_name, None)
|
| 484 |
+
)
|
| 485 |
+
)
|
| 486 |
+
else:
|
| 487 |
+
emb_out = embedder(
|
| 488 |
+
*[
|
| 489 |
+
embedder.random_dropout_input(batch.get(k), override_dropout_rate.get(emb_name, None), k)
|
| 490 |
+
for k in embedder.input_key
|
| 491 |
+
]
|
| 492 |
+
)
|
| 493 |
+
for k, v in emb_out.items():
|
| 494 |
+
output[k].append(v)
|
| 495 |
+
return {k: torch.cat(v, dim=self.KEY2DIM.get(k, -1)) for k, v in output.items()}
|
| 496 |
+
|
| 497 |
+
def forward(self, batch: Dict, override_dropout_rate: Optional[Dict[str, float]] = None) -> MultiViewCondition:
|
| 498 |
+
output = self._forward(batch, override_dropout_rate)
|
| 499 |
+
return MultiViewCondition(**output)
|
| 500 |
+
|
| 501 |
+
def get_condition_with_negative_prompt(
|
| 502 |
+
self,
|
| 503 |
+
data_batch: Dict,
|
| 504 |
+
) -> Tuple[Any, Any]:
|
| 505 |
+
"""
|
| 506 |
+
Similar functionality as get_condition_uncondition
|
| 507 |
+
But use negative prompts for unconditon
|
| 508 |
+
"""
|
| 509 |
+
cond_dropout_rates, uncond_dropout_rates = {}, {}
|
| 510 |
+
for emb_name, embedder in self.embedders.items():
|
| 511 |
+
cond_dropout_rates[emb_name] = 0.0
|
| 512 |
+
if isinstance(embedder, TextAttr) or isinstance(embedder, MVTextAttr):
|
| 513 |
+
uncond_dropout_rates[emb_name] = 0.0
|
| 514 |
+
else:
|
| 515 |
+
uncond_dropout_rates[emb_name] = 1.0 if embedder.dropout_rate > 1e-4 else 0.0
|
| 516 |
+
|
| 517 |
+
data_batch_neg_prompt = copy.deepcopy(data_batch)
|
| 518 |
+
if "neg_t5_text_embeddings" in data_batch_neg_prompt:
|
| 519 |
+
if isinstance(data_batch_neg_prompt["neg_t5_text_embeddings"], torch.Tensor):
|
| 520 |
+
data_batch_neg_prompt["t5_text_embeddings"] = data_batch_neg_prompt["neg_t5_text_embeddings"]
|
| 521 |
+
|
| 522 |
+
condition: Any = self(data_batch, override_dropout_rate=cond_dropout_rates)
|
| 523 |
+
un_condition: Any = self(data_batch_neg_prompt, override_dropout_rate=uncond_dropout_rates)
|
| 524 |
+
|
| 525 |
+
return condition, un_condition
|
| 526 |
+
|
| 527 |
+
|
| 528 |
+
MultiViewConditionerConfig: LazyDict = L(MultiViewConditioner)(
|
| 529 |
+
**_SHARED_CONFIG,
|
| 530 |
+
view_indices_B_T=L(ReMapkey)(
|
| 531 |
+
input_key="latent_view_indices_B_T",
|
| 532 |
+
output_key="view_indices_B_T",
|
| 533 |
+
dropout_rate=0.0,
|
| 534 |
+
dtype=None,
|
| 535 |
+
),
|
| 536 |
+
ref_cam_view_idx_sample_position=L(ReMapkey)(
|
| 537 |
+
input_key="ref_cam_view_idx_sample_position",
|
| 538 |
+
output_key="ref_cam_view_idx_sample_position",
|
| 539 |
+
dropout_rate=0.0,
|
| 540 |
+
dtype=None,
|
| 541 |
+
),
|
| 542 |
+
)
|
| 543 |
+
|
| 544 |
+
|
| 545 |
+
class TextAttrEmptyStringDropout(TextAttr):
|
| 546 |
+
def __init__(
|
| 547 |
+
self,
|
| 548 |
+
input_key: str,
|
| 549 |
+
pos_input_key: str,
|
| 550 |
+
dropout_input_key: str,
|
| 551 |
+
dropout_rate: Optional[float] = 0.0,
|
| 552 |
+
use_empty_string: bool = False,
|
| 553 |
+
**kwargs,
|
| 554 |
+
):
|
| 555 |
+
self._input_key = input_key
|
| 556 |
+
self._pos_input_key = pos_input_key
|
| 557 |
+
self._dropout_input_key = dropout_input_key
|
| 558 |
+
self._dropout_rate = dropout_rate
|
| 559 |
+
self._use_empty_string = use_empty_string
|
| 560 |
+
super().__init__(input_key, dropout_rate)
|
| 561 |
+
|
| 562 |
+
def forward(self, tensor: torch.Tensor):
|
| 563 |
+
return {"crossattn_emb": tensor}
|
| 564 |
+
|
| 565 |
+
def random_dropout_input(
|
| 566 |
+
self,
|
| 567 |
+
in_tensor_dict: torch.Tensor | Dict[str, torch.Tensor],
|
| 568 |
+
dropout_rate: Optional[float] = None,
|
| 569 |
+
key: Optional[str] = None,
|
| 570 |
+
) -> torch.Tensor:
|
| 571 |
+
if key is not None and "mask" in key:
|
| 572 |
+
return in_tensor_dict
|
| 573 |
+
del key
|
| 574 |
+
assert isinstance(in_tensor_dict, dict), f"in_tensor_dict must be a dict. Got {type(in_tensor_dict)}"
|
| 575 |
+
in_tensor = in_tensor_dict[self._pos_input_key]
|
| 576 |
+
B = in_tensor.shape[0]
|
| 577 |
+
dropout_rate = dropout_rate if dropout_rate is not None else self.dropout_rate
|
| 578 |
+
keep_mask = torch.bernoulli((1.0 - dropout_rate) * torch.ones(B)).type_as(in_tensor)
|
| 579 |
+
if self._use_empty_string:
|
| 580 |
+
empty_prompt = in_tensor_dict[self._dropout_input_key]
|
| 581 |
+
if empty_prompt.shape[0] != B:
|
| 582 |
+
empty_prompt = empty_prompt.repeat(B, 1, 1)
|
| 583 |
+
else:
|
| 584 |
+
empty_prompt = torch.zeros_like(in_tensor)
|
| 585 |
+
|
| 586 |
+
return keep_mask * in_tensor + (1 - keep_mask) * empty_prompt
|
| 587 |
+
|
| 588 |
+
def details(self) -> str:
|
| 589 |
+
return "Output key: [crossattn_emb]"
|
| 590 |
+
|
| 591 |
+
|
| 592 |
+
_SHARED_CONFIG_PER_VIEW_DROPOUT = copy.deepcopy(_SHARED_CONFIG)
|
| 593 |
+
_SHARED_CONFIG_PER_VIEW_DROPOUT["text"] = L(MVTextAttr)(
|
| 594 |
+
input_key=["t5_text_embeddings"],
|
| 595 |
+
dropout_rate=0.2,
|
| 596 |
+
use_empty_string=False,
|
| 597 |
+
)
|
| 598 |
+
|
| 599 |
+
MultiViewConditionerPerViewDropoutConfig: LazyDict = L(MultiViewConditioner)(
|
| 600 |
+
**_SHARED_CONFIG_PER_VIEW_DROPOUT,
|
| 601 |
+
view_indices_B_T=L(ReMapkey)(
|
| 602 |
+
input_key="latent_view_indices_B_T",
|
| 603 |
+
output_key="view_indices_B_T",
|
| 604 |
+
dropout_rate=0.0,
|
| 605 |
+
dtype=None,
|
| 606 |
+
),
|
| 607 |
+
ref_cam_view_idx_sample_position=L(ReMapkey)(
|
| 608 |
+
input_key="ref_cam_view_idx_sample_position",
|
| 609 |
+
output_key="ref_cam_view_idx_sample_position",
|
| 610 |
+
dropout_rate=0.0,
|
| 611 |
+
dtype=None,
|
| 612 |
+
),
|
| 613 |
+
)
|
| 614 |
+
|
| 615 |
+
|
| 616 |
+
# Conditioner for 2-actor pose-conditioned joint generation. Adds three ReMapkey embedders that thread the
|
| 617 |
+
# preprocessed pose / warped-latent / visibility tensors (written into the data_batch by the model's
|
| 618 |
+
# get_data_and_condition) into the MultiViewCondition fields consumed by MultiViewPoseDiT.forward.
|
| 619 |
+
# - pose_map : dropout 0.1 -> classifier-free-guidance on pose (zeroed in the uncond branch).
|
| 620 |
+
# - warped_latent : dropout 0.0 -> geometric conditioning always present (no guidance).
|
| 621 |
+
# - visibility_mask : dropout 0.0 -> always present.
|
| 622 |
+
MultiViewPoseConditionerConfig: LazyDict = L(MultiViewConditioner)(
|
| 623 |
+
**_SHARED_CONFIG,
|
| 624 |
+
view_indices_B_T=L(ReMapkey)(
|
| 625 |
+
input_key="latent_view_indices_B_T",
|
| 626 |
+
output_key="view_indices_B_T",
|
| 627 |
+
dropout_rate=0.0,
|
| 628 |
+
dtype=None,
|
| 629 |
+
),
|
| 630 |
+
ref_cam_view_idx_sample_position=L(ReMapkey)(
|
| 631 |
+
input_key="ref_cam_view_idx_sample_position",
|
| 632 |
+
output_key="ref_cam_view_idx_sample_position",
|
| 633 |
+
dropout_rate=0.0,
|
| 634 |
+
dtype=None,
|
| 635 |
+
),
|
| 636 |
+
pose_map=L(ReMapkey)(
|
| 637 |
+
input_key="pose_map",
|
| 638 |
+
output_key="pose_map_B_C_T_H_W",
|
| 639 |
+
dropout_rate=0.1,
|
| 640 |
+
dtype=None,
|
| 641 |
+
),
|
| 642 |
+
warped_latent=L(ReMapkey)(
|
| 643 |
+
input_key="warped_latent",
|
| 644 |
+
output_key="warped_latent_B_C_T_H_W",
|
| 645 |
+
dropout_rate=0.0,
|
| 646 |
+
dtype=None,
|
| 647 |
+
),
|
| 648 |
+
visibility_mask=L(ReMapkey)(
|
| 649 |
+
input_key="visibility_mask",
|
| 650 |
+
output_key="visibility_mask_B_C_T_H_W",
|
| 651 |
+
dropout_rate=0.0,
|
| 652 |
+
dtype=None,
|
| 653 |
+
),
|
| 654 |
+
)
|
| 655 |
+
|
| 656 |
+
# Variant that ALSO remaps the frozen-VAE pose latent (for net pose_mode = vae_concat / vae_mlp_add).
|
| 657 |
+
# pose_latent carries the classifier-free-guidance dropout (0.1) for pose; pose_map is still present but the
|
| 658 |
+
# net ignores it in the VAE pose modes.
|
| 659 |
+
MultiViewPoseLatentConditionerConfig: LazyDict = L(MultiViewConditioner)(
|
| 660 |
+
**_SHARED_CONFIG,
|
| 661 |
+
view_indices_B_T=L(ReMapkey)(
|
| 662 |
+
input_key="latent_view_indices_B_T", output_key="view_indices_B_T", dropout_rate=0.0, dtype=None,
|
| 663 |
+
),
|
| 664 |
+
ref_cam_view_idx_sample_position=L(ReMapkey)(
|
| 665 |
+
input_key="ref_cam_view_idx_sample_position", output_key="ref_cam_view_idx_sample_position",
|
| 666 |
+
dropout_rate=0.0, dtype=None,
|
| 667 |
+
),
|
| 668 |
+
pose_latent=L(ReMapkey)(
|
| 669 |
+
input_key="pose_latent", output_key="pose_latent_B_C_T_H_W", dropout_rate=0.1, dtype=None,
|
| 670 |
+
),
|
| 671 |
+
warped_latent=L(ReMapkey)(
|
| 672 |
+
input_key="warped_latent", output_key="warped_latent_B_C_T_H_W", dropout_rate=0.0, dtype=None,
|
| 673 |
+
),
|
| 674 |
+
visibility_mask=L(ReMapkey)(
|
| 675 |
+
input_key="visibility_mask", output_key="visibility_mask_B_C_T_H_W", dropout_rate=0.0, dtype=None,
|
| 676 |
+
),
|
| 677 |
+
reference_latent=L(ReMapkey)(
|
| 678 |
+
input_key="reference_latent", output_key="reference_latent_B_C_R_H_W", dropout_rate=0.0, dtype=None,
|
| 679 |
+
),
|
| 680 |
+
plucker_map=L(ReMapkey)(
|
| 681 |
+
input_key="plucker_map", output_key="plucker_map_B_C_T_H_W", dropout_rate=0.0, dtype=None,
|
| 682 |
+
),
|
| 683 |
+
reference_plucker_map=L(ReMapkey)(
|
| 684 |
+
input_key="reference_plucker_map", output_key="reference_plucker_map_B_C_R_H_W", dropout_rate=0.0, dtype=None,
|
| 685 |
+
),
|
| 686 |
+
reference_pose_latent=L(ReMapkey)(
|
| 687 |
+
input_key="reference_pose_latent", output_key="reference_pose_latent_B_C_R_H_W", dropout_rate=0.0, dtype=None,
|
| 688 |
+
),
|
| 689 |
+
depth_latent=L(ReMapkey)(
|
| 690 |
+
input_key="depth_latent", output_key="depth_latent_B_C_T_H_W", dropout_rate=0.0, dtype=None,
|
| 691 |
+
),
|
| 692 |
+
)
|
| 693 |
+
|
| 694 |
+
|
| 695 |
+
def register_conditioner():
|
| 696 |
+
cs = ConfigStore.instance()
|
| 697 |
+
cs.store(
|
| 698 |
+
group="conditioner",
|
| 699 |
+
package="model.config.conditioner",
|
| 700 |
+
name="video_prediction_multiview_conditioner",
|
| 701 |
+
node=MultiViewConditionerConfig,
|
| 702 |
+
)
|
| 703 |
+
cs.store(
|
| 704 |
+
group="conditioner",
|
| 705 |
+
package="model.config.conditioner",
|
| 706 |
+
name="video_prediction_multiview_pose_conditioner",
|
| 707 |
+
node=MultiViewPoseConditionerConfig,
|
| 708 |
+
)
|
| 709 |
+
cs.store(
|
| 710 |
+
group="conditioner",
|
| 711 |
+
package="model.config.conditioner",
|
| 712 |
+
name="video_prediction_multiview_pose_latent_conditioner",
|
| 713 |
+
node=MultiViewPoseLatentConditionerConfig,
|
| 714 |
+
)
|
| 715 |
+
cs.store(
|
| 716 |
+
group="conditioner",
|
| 717 |
+
package="model.config.conditioner",
|
| 718 |
+
name="video_prediction_multiview_conditioner_per_view_dropout",
|
| 719 |
+
node=MultiViewConditionerPerViewDropoutConfig,
|
| 720 |
+
)
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/dataloader.py
ADDED
|
@@ -0,0 +1,101 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
from hydra.core.config_store import ConfigStore
|
| 18 |
+
|
| 19 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 20 |
+
from cosmos_predict2._src.predict2_multiview.datasets.multiview import (
|
| 21 |
+
DEFAULT_CAMERAS,
|
| 22 |
+
AugmentationConfig,
|
| 23 |
+
get_multiview_video_loader,
|
| 24 |
+
)
|
| 25 |
+
|
| 26 |
+
DEFAULT_CAMERA_VIEW_CONFIGS = {
|
| 27 |
+
"7views": DEFAULT_CAMERAS,
|
| 28 |
+
"4views": [
|
| 29 |
+
"camera_front_wide_120fov",
|
| 30 |
+
"camera_cross_right_120fov",
|
| 31 |
+
"camera_rear_tele_30fov",
|
| 32 |
+
"camera_cross_left_120fov",
|
| 33 |
+
],
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def register_multiview_dataloader() -> None:
|
| 38 |
+
"""Register multiview video dataloader configurations."""
|
| 39 |
+
|
| 40 |
+
cs = ConfigStore.instance()
|
| 41 |
+
|
| 42 |
+
# alpamayo
|
| 43 |
+
datasets = ["alpamayo_dec2024"]
|
| 44 |
+
object_stores = ["gcs", "s3"]
|
| 45 |
+
resolutions = [
|
| 46 |
+
("480p", (480, 832)),
|
| 47 |
+
("720p", (720, 1280)),
|
| 48 |
+
("1080p", (1080, 1920)),
|
| 49 |
+
]
|
| 50 |
+
fps = [
|
| 51 |
+
("10fps", 3),
|
| 52 |
+
("15fps", 2),
|
| 53 |
+
("30fps", 1),
|
| 54 |
+
]
|
| 55 |
+
num_video_frames = [
|
| 56 |
+
("29frames", 29),
|
| 57 |
+
("61frames", 61),
|
| 58 |
+
("93frames", 93),
|
| 59 |
+
]
|
| 60 |
+
cs.store(
|
| 61 |
+
group="data_val",
|
| 62 |
+
package="dataloader_val",
|
| 63 |
+
name="mock",
|
| 64 |
+
node=L(get_multiview_video_loader)(
|
| 65 |
+
dataset_name=datasets[0],
|
| 66 |
+
is_train=False,
|
| 67 |
+
object_store="s3",
|
| 68 |
+
augmentation_config=L(AugmentationConfig)(
|
| 69 |
+
resolution_hw=resolutions[0][1],
|
| 70 |
+
fps_downsample_factor=fps[0][1],
|
| 71 |
+
num_video_frames=num_video_frames[0][1],
|
| 72 |
+
camera_keys=DEFAULT_CAMERA_VIEW_CONFIGS["7views"],
|
| 73 |
+
),
|
| 74 |
+
batch_size=1,
|
| 75 |
+
num_workers=2,
|
| 76 |
+
),
|
| 77 |
+
)
|
| 78 |
+
|
| 79 |
+
for dataset in datasets:
|
| 80 |
+
for object_store in object_stores:
|
| 81 |
+
for resolution_str, resolution_hw in resolutions:
|
| 82 |
+
for fps_str, downsample_factor in fps:
|
| 83 |
+
for num_video_frames_str, num_frames in num_video_frames:
|
| 84 |
+
for views_str, camera_keys in DEFAULT_CAMERA_VIEW_CONFIGS.items():
|
| 85 |
+
name = f"video_{dataset}_{object_store}_{resolution_str}_{fps_str}_{num_video_frames_str}_{views_str}"
|
| 86 |
+
cs.store(
|
| 87 |
+
group="data_train",
|
| 88 |
+
package="dataloader_train",
|
| 89 |
+
name=name,
|
| 90 |
+
node=L(get_multiview_video_loader)(
|
| 91 |
+
dataset_name=dataset,
|
| 92 |
+
is_train=True,
|
| 93 |
+
object_store=object_store,
|
| 94 |
+
augmentation_config=L(AugmentationConfig)(
|
| 95 |
+
resolution_hw=resolution_hw,
|
| 96 |
+
fps_downsample_factor=downsample_factor,
|
| 97 |
+
num_video_frames=num_frames,
|
| 98 |
+
camera_keys=camera_keys,
|
| 99 |
+
),
|
| 100 |
+
),
|
| 101 |
+
)
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/dataloader_local.py
ADDED
|
@@ -0,0 +1,111 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
"""Local file-based dataloader configurations."""
|
| 17 |
+
|
| 18 |
+
import torch.distributed as dist
|
| 19 |
+
from hydra.core.config_store import ConfigStore
|
| 20 |
+
|
| 21 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 22 |
+
from cosmos_predict2._src.predict2.datasets.local_datasets.dataset_video import get_generic_dataloader, get_sampler
|
| 23 |
+
from cosmos_predict2._src.predict2_multiview.datasets.local import WaymoLocalDataset
|
| 24 |
+
from cosmos_predict2._src.predict2_multiview.datasets.multiview import AugmentationConfig, collate_fn
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def register_waymo_dataloader() -> None:
|
| 28 |
+
"""Register local file-based dataloader configurations."""
|
| 29 |
+
|
| 30 |
+
cs = ConfigStore.instance()
|
| 31 |
+
|
| 32 |
+
waymo_dataset = L(WaymoLocalDataset)(
|
| 33 |
+
video_file_dirs=["datasets/multiview/waymo/input"],
|
| 34 |
+
augmentation_config=L(AugmentationConfig)(
|
| 35 |
+
resolution_hw=(720, 1280),
|
| 36 |
+
fps_downsample_factor=1,
|
| 37 |
+
num_video_frames=29,
|
| 38 |
+
camera_keys=[
|
| 39 |
+
"pinhole_front",
|
| 40 |
+
"pinhole_front_right",
|
| 41 |
+
"pinhole_side_right",
|
| 42 |
+
"pinhole_side_left",
|
| 43 |
+
"pinhole_front_left",
|
| 44 |
+
],
|
| 45 |
+
camera_view_mapping={
|
| 46 |
+
"pinhole_front": 0,
|
| 47 |
+
"pinhole_front_right": 1,
|
| 48 |
+
"pinhole_side_right": 2,
|
| 49 |
+
# no pinhole_back in the dataset, so skip ID 3
|
| 50 |
+
"pinhole_side_left": 4,
|
| 51 |
+
"pinhole_front_left": 5,
|
| 52 |
+
# no pinehole_front_tele in the dataset skip ID 6
|
| 53 |
+
},
|
| 54 |
+
camera_video_key_mapping={
|
| 55 |
+
"pinhole_front": "video_pinhole_front",
|
| 56 |
+
"pinhole_front_right": "video_pinhole_front_right",
|
| 57 |
+
"pinhole_side_right": "video_pinhole_side_right",
|
| 58 |
+
"pinhole_side_left": "video_pinhole_side_left",
|
| 59 |
+
"pinhole_front_left": "video_pinhole_front_left",
|
| 60 |
+
},
|
| 61 |
+
camera_caption_key_mapping={
|
| 62 |
+
"pinhole_front": "caption_pinhole_front",
|
| 63 |
+
"pinhole_front_right": "caption_pinhole_front_right",
|
| 64 |
+
"pinhole_side_right": "caption_pinhole_side_right",
|
| 65 |
+
"pinhole_side_left": "caption_pinhole_side_left",
|
| 66 |
+
"pinhole_front_left": "caption_pinhole_front_left",
|
| 67 |
+
},
|
| 68 |
+
caption_probability={
|
| 69 |
+
"long": 1.0,
|
| 70 |
+
},
|
| 71 |
+
single_caption_camera_name="pinhole_front",
|
| 72 |
+
add_view_prefix_to_caption=True,
|
| 73 |
+
camera_prefix_mapping={
|
| 74 |
+
"pinhole_front": "The video is captured from a camera mounted on a car. The camera is facing forward.",
|
| 75 |
+
"pinhole_front_right": "The video is captured from a camera mounted on a car. The camera is facing to the front right.",
|
| 76 |
+
"pinhole_side_right": "The video is captured from a camera mounted on a car. The camera is facing to the side right.",
|
| 77 |
+
"pinhole_side_left": "The video is captured from a camera mounted on a car. The camera is facing to the side left.",
|
| 78 |
+
"pinhole_front_left": "The video is captured from a camera mounted on a car. The camera is facing to the front left.",
|
| 79 |
+
},
|
| 80 |
+
),
|
| 81 |
+
)
|
| 82 |
+
|
| 83 |
+
cs.store(
|
| 84 |
+
group="data_train",
|
| 85 |
+
package="dataloader_train",
|
| 86 |
+
name=f"waymo",
|
| 87 |
+
node=L(get_generic_dataloader)(
|
| 88 |
+
dataset=waymo_dataset,
|
| 89 |
+
sampler=L(get_sampler)(dataset=waymo_dataset) if dist.is_initialized() else None,
|
| 90 |
+
collate_fn=collate_fn,
|
| 91 |
+
batch_size=1,
|
| 92 |
+
drop_last=True,
|
| 93 |
+
num_workers=4,
|
| 94 |
+
pin_memory=True,
|
| 95 |
+
),
|
| 96 |
+
)
|
| 97 |
+
|
| 98 |
+
cs.store(
|
| 99 |
+
group="data_val",
|
| 100 |
+
package="dataloader_val",
|
| 101 |
+
name=f"waymo",
|
| 102 |
+
node=L(get_generic_dataloader)(
|
| 103 |
+
dataset=waymo_dataset,
|
| 104 |
+
sampler=L(get_sampler)(dataset=waymo_dataset) if dist.is_initialized() else None,
|
| 105 |
+
collate_fn=collate_fn,
|
| 106 |
+
batch_size=1,
|
| 107 |
+
drop_last=True,
|
| 108 |
+
num_workers=4,
|
| 109 |
+
pin_memory=True,
|
| 110 |
+
),
|
| 111 |
+
)
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/model.py
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
from hydra.core.config_store import ConfigStore
|
| 17 |
+
|
| 18 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 19 |
+
from cosmos_predict2._src.predict2_multiview.models.multiview_pose_model_rectified_flow import (
|
| 20 |
+
MultiviewVid2VidPoseModelRectifiedFlow,
|
| 21 |
+
MultiviewVid2VidPoseModelRectifiedFlowConfig,
|
| 22 |
+
)
|
| 23 |
+
from cosmos_predict2._src.predict2_multiview.models.multiview_vid2vid_model_rectified_flow import (
|
| 24 |
+
MultiviewVid2VidModelRectifiedFlow,
|
| 25 |
+
MultiviewVid2VidModelRectifiedFlowConfig,
|
| 26 |
+
)
|
| 27 |
+
|
| 28 |
+
FSDP_RECTIFIED_FLOW_CONFIG = dict(
|
| 29 |
+
trainer=dict(
|
| 30 |
+
distributed_parallelism="fsdp",
|
| 31 |
+
),
|
| 32 |
+
model=L(MultiviewVid2VidModelRectifiedFlow)(
|
| 33 |
+
config=MultiviewVid2VidModelRectifiedFlowConfig(),
|
| 34 |
+
_recursive_=False,
|
| 35 |
+
),
|
| 36 |
+
)
|
| 37 |
+
|
| 38 |
+
FSDP_RECTIFIED_FLOW_POSE_CONFIG = dict(
|
| 39 |
+
trainer=dict(
|
| 40 |
+
distributed_parallelism="fsdp",
|
| 41 |
+
),
|
| 42 |
+
model=L(MultiviewVid2VidPoseModelRectifiedFlow)(
|
| 43 |
+
config=MultiviewVid2VidPoseModelRectifiedFlowConfig(),
|
| 44 |
+
_recursive_=False,
|
| 45 |
+
),
|
| 46 |
+
)
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def register_model():
|
| 50 |
+
cs = ConfigStore.instance()
|
| 51 |
+
cs.store(group="model", package="_global_", name="fsdp_rectified_flow_multiview", node=FSDP_RECTIFIED_FLOW_CONFIG)
|
| 52 |
+
cs.store(
|
| 53 |
+
group="model",
|
| 54 |
+
package="_global_",
|
| 55 |
+
name="fsdp_rectified_flow_multiview_pose",
|
| 56 |
+
node=FSDP_RECTIFIED_FLOW_POSE_CONFIG,
|
| 57 |
+
)
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/net.py
ADDED
|
@@ -0,0 +1,190 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
import copy
|
| 17 |
+
|
| 18 |
+
from hydra.core.config_store import ConfigStore
|
| 19 |
+
|
| 20 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 21 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyDict
|
| 22 |
+
from cosmos_predict2._src.predict2.networks.minimal_v4_dit import SACConfig
|
| 23 |
+
from cosmos_predict2._src.predict2_multiview.networks.multiview_cross_dit import MultiViewCrossDiT, MultiViewSACConfig
|
| 24 |
+
from cosmos_predict2._src.predict2_multiview.networks.multiview_dit import MultiViewDiT
|
| 25 |
+
from cosmos_predict2._src.predict2_multiview.networks.multiview_pose_dit import MultiViewPoseDiT
|
| 26 |
+
|
| 27 |
+
COSMOS_V1_7B_MULTIVIEW_NET: LazyDict = L(MultiViewDiT)(
|
| 28 |
+
max_img_h=240,
|
| 29 |
+
max_img_w=240,
|
| 30 |
+
max_frames=128,
|
| 31 |
+
in_channels=16,
|
| 32 |
+
out_channels=16,
|
| 33 |
+
patch_spatial=2,
|
| 34 |
+
patch_temporal=1,
|
| 35 |
+
model_channels=4096,
|
| 36 |
+
num_blocks=28,
|
| 37 |
+
num_heads=32,
|
| 38 |
+
concat_padding_mask=True,
|
| 39 |
+
pos_emb_cls="rope3d",
|
| 40 |
+
pos_emb_learnable=True,
|
| 41 |
+
pos_emb_interpolation="crop",
|
| 42 |
+
use_adaln_lora=True,
|
| 43 |
+
adaln_lora_dim=256,
|
| 44 |
+
atten_backend="minimal_a2a",
|
| 45 |
+
extra_per_block_abs_pos_emb=True,
|
| 46 |
+
rope_h_extrapolation_ratio=1.0,
|
| 47 |
+
rope_w_extrapolation_ratio=1.0,
|
| 48 |
+
rope_t_extrapolation_ratio=2.0,
|
| 49 |
+
sac_config=SACConfig(),
|
| 50 |
+
n_cameras_emb=7,
|
| 51 |
+
view_condition_dim=6,
|
| 52 |
+
concat_view_embedding=True,
|
| 53 |
+
use_wan_fp32_strategy=False,
|
| 54 |
+
layer_mask=None,
|
| 55 |
+
)
|
| 56 |
+
|
| 57 |
+
COSMOS_V1_2B_MULTIVIEW_NET = copy.deepcopy(COSMOS_V1_7B_MULTIVIEW_NET)
|
| 58 |
+
COSMOS_V1_2B_MULTIVIEW_NET.model_channels = 2048
|
| 59 |
+
COSMOS_V1_2B_MULTIVIEW_NET.num_blocks = 28
|
| 60 |
+
COSMOS_V1_2B_MULTIVIEW_NET.num_heads = 16
|
| 61 |
+
COSMOS_V1_2B_MULTIVIEW_NET.extra_per_block_abs_pos_emb = False
|
| 62 |
+
COSMOS_V1_2B_MULTIVIEW_NET.rope_t_extrapolation_ratio = 1.0
|
| 63 |
+
|
| 64 |
+
COSMOS_V1_14B_MULTIVIEW_NET = copy.deepcopy(COSMOS_V1_7B_MULTIVIEW_NET)
|
| 65 |
+
COSMOS_V1_14B_MULTIVIEW_NET.model_channels = 5120
|
| 66 |
+
COSMOS_V1_14B_MULTIVIEW_NET.num_blocks = 36
|
| 67 |
+
COSMOS_V1_14B_MULTIVIEW_NET.num_heads = 40
|
| 68 |
+
COSMOS_V1_14B_MULTIVIEW_NET.extra_per_block_abs_pos_emb = False
|
| 69 |
+
COSMOS_V1_14B_MULTIVIEW_NET.rope_t_extrapolation_ratio = 1.0
|
| 70 |
+
|
| 71 |
+
mini_net = copy.deepcopy(COSMOS_V1_7B_MULTIVIEW_NET)
|
| 72 |
+
mini_net.model_channels = 1024
|
| 73 |
+
mini_net.num_heads = 8
|
| 74 |
+
mini_net.num_blocks = 2
|
| 75 |
+
mini_net.rope_t_extrapolation_ratio = 1.0
|
| 76 |
+
|
| 77 |
+
# 2B multi-view network + pose encoder + warped/visibility conditioning (2-actor joint generation).
|
| 78 |
+
# Derived from the 2B multiview net; only the target class and the pose-specific kwargs change.
|
| 79 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET = copy.deepcopy(COSMOS_V1_2B_MULTIVIEW_NET)
|
| 80 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET._target_ = MultiViewPoseDiT
|
| 81 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.warped_latent_channels = 16
|
| 82 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.visibility_channels = 1
|
| 83 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.pose_in_channels = 3
|
| 84 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.pose_hidden_dim = 16
|
| 85 |
+
# in-context reference-frame appearance conditioning (default OFF -> other experiments unchanged)
|
| 86 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.enable_reference_frames = False
|
| 87 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.num_reference_frames = 0
|
| 88 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.ref_rope_offset = 50
|
| 89 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.ref_rope_stride = 5
|
| 90 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.shared_reference = False # refs = one shared set, no per-view view-embedding
|
| 91 |
+
# per-pixel Plücker ray conditioning (default OFF)
|
| 92 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.enable_plucker = False
|
| 93 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.plucker_channels = 6
|
| 94 |
+
COSMOS_V1_2B_MULTIVIEW_POSE_NET.enable_reference_plucker = False
|
| 95 |
+
|
| 96 |
+
# tiny variant for fast shape/smoke tests
|
| 97 |
+
mini_pose_net = copy.deepcopy(COSMOS_V1_2B_MULTIVIEW_POSE_NET)
|
| 98 |
+
mini_pose_net.model_channels = 1024
|
| 99 |
+
mini_pose_net.num_heads = 8
|
| 100 |
+
mini_pose_net.num_blocks = 2
|
| 101 |
+
|
| 102 |
+
# modified according to COSMOS_V1_2B_MULTIVIEW_NET
|
| 103 |
+
COSMOS_V1_2B_MULTIVIEW_CROSSVIEW_NET = LazyDict = L(MultiViewCrossDiT)(
|
| 104 |
+
max_img_h=240,
|
| 105 |
+
max_img_w=240,
|
| 106 |
+
max_frames=128,
|
| 107 |
+
in_channels=16,
|
| 108 |
+
out_channels=16,
|
| 109 |
+
patch_spatial=2,
|
| 110 |
+
patch_temporal=1,
|
| 111 |
+
model_channels=2048,
|
| 112 |
+
num_blocks=28,
|
| 113 |
+
num_heads=16,
|
| 114 |
+
concat_padding_mask=True,
|
| 115 |
+
pos_emb_cls="rope3d",
|
| 116 |
+
pos_emb_learnable=True,
|
| 117 |
+
pos_emb_interpolation="crop",
|
| 118 |
+
use_adaln_lora=True,
|
| 119 |
+
adaln_lora_dim=256,
|
| 120 |
+
atten_backend="minimal_a2a",
|
| 121 |
+
extra_per_block_abs_pos_emb=False,
|
| 122 |
+
rope_h_extrapolation_ratio=1.0,
|
| 123 |
+
rope_w_extrapolation_ratio=1.0,
|
| 124 |
+
rope_t_extrapolation_ratio=1.0,
|
| 125 |
+
sac_config=MultiViewSACConfig(),
|
| 126 |
+
n_cameras_emb=7,
|
| 127 |
+
view_condition_dim=6,
|
| 128 |
+
concat_view_embedding=False,
|
| 129 |
+
adaln_view_embedding=True,
|
| 130 |
+
enable_cross_view_attn=True,
|
| 131 |
+
use_wan_fp32_strategy=False,
|
| 132 |
+
layer_mask=None,
|
| 133 |
+
)
|
| 134 |
+
|
| 135 |
+
# modify according to COSMOS_V1_14B_MULTIVIEW_NET
|
| 136 |
+
COSMOS_V1_14B_MULTIVIEW_CROSSVIEW_NET = LazyDict = L(MultiViewCrossDiT)(
|
| 137 |
+
max_img_h=240,
|
| 138 |
+
max_img_w=240,
|
| 139 |
+
max_frames=128,
|
| 140 |
+
in_channels=16,
|
| 141 |
+
out_channels=16,
|
| 142 |
+
patch_spatial=2,
|
| 143 |
+
patch_temporal=1,
|
| 144 |
+
model_channels=5120,
|
| 145 |
+
num_blocks=36,
|
| 146 |
+
num_heads=40,
|
| 147 |
+
concat_padding_mask=True,
|
| 148 |
+
pos_emb_cls="rope3d",
|
| 149 |
+
pos_emb_learnable=True,
|
| 150 |
+
pos_emb_interpolation="crop",
|
| 151 |
+
use_adaln_lora=True,
|
| 152 |
+
adaln_lora_dim=256,
|
| 153 |
+
atten_backend="minimal_a2a",
|
| 154 |
+
extra_per_block_abs_pos_emb=False,
|
| 155 |
+
rope_h_extrapolation_ratio=1.0,
|
| 156 |
+
rope_w_extrapolation_ratio=1.0,
|
| 157 |
+
rope_t_extrapolation_ratio=1.0,
|
| 158 |
+
sac_config=MultiViewSACConfig(),
|
| 159 |
+
n_cameras_emb=7,
|
| 160 |
+
view_condition_dim=6,
|
| 161 |
+
concat_view_embedding=False,
|
| 162 |
+
adaln_view_embedding=True,
|
| 163 |
+
enable_cross_view_attn=True,
|
| 164 |
+
use_wan_fp32_strategy=False,
|
| 165 |
+
layer_mask=None,
|
| 166 |
+
)
|
| 167 |
+
|
| 168 |
+
|
| 169 |
+
def register_net():
|
| 170 |
+
cs = ConfigStore.instance()
|
| 171 |
+
cs.store(group="net", package="model.config.net", name="mini_net", node=mini_net)
|
| 172 |
+
cs.store(group="net", package="model.config.net", name="cosmos_v1_2B_multiview", node=COSMOS_V1_2B_MULTIVIEW_NET)
|
| 173 |
+
cs.store(group="net", package="model.config.net", name="cosmos_v1_7B_multiview", node=COSMOS_V1_7B_MULTIVIEW_NET)
|
| 174 |
+
cs.store(group="net", package="model.config.net", name="cosmos_v1_14B_multiview", node=COSMOS_V1_14B_MULTIVIEW_NET)
|
| 175 |
+
cs.store(
|
| 176 |
+
group="net",
|
| 177 |
+
package="model.config.net",
|
| 178 |
+
name="cosmos_v1_2B_multiview_crossview",
|
| 179 |
+
node=COSMOS_V1_2B_MULTIVIEW_CROSSVIEW_NET,
|
| 180 |
+
)
|
| 181 |
+
cs.store(
|
| 182 |
+
group="net",
|
| 183 |
+
package="model.config.net",
|
| 184 |
+
name="cosmos_v1_14B_multiview_crossview",
|
| 185 |
+
node=COSMOS_V1_14B_MULTIVIEW_CROSSVIEW_NET,
|
| 186 |
+
)
|
| 187 |
+
cs.store(
|
| 188 |
+
group="net", package="model.config.net", name="cosmos_v1_2B_multiview_pose", node=COSMOS_V1_2B_MULTIVIEW_POSE_NET
|
| 189 |
+
)
|
| 190 |
+
cs.store(group="net", package="model.config.net", name="mini_pose_net", node=mini_pose_net)
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/optimizer.py
ADDED
|
@@ -0,0 +1,72 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
from hydra.core.config_store import ConfigStore
|
| 17 |
+
|
| 18 |
+
from cosmos_predict2._src.imaginaire.lazy_config import PLACEHOLDER, LazyDict
|
| 19 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 20 |
+
from cosmos_predict2._src.predict2.utils.optim_instantiate import get_base_optimizer
|
| 21 |
+
from cosmos_predict2._src.predict2_multiview.utils.optim_instantiate import get_multiple_optimizer
|
| 22 |
+
|
| 23 |
+
AdamWConfig = L(get_base_optimizer)(
|
| 24 |
+
model=PLACEHOLDER,
|
| 25 |
+
lr=1e-4,
|
| 26 |
+
weight_decay=0.1,
|
| 27 |
+
betas=[0.9, 0.99],
|
| 28 |
+
optim_type="adamw",
|
| 29 |
+
eps=1e-8,
|
| 30 |
+
fused=True,
|
| 31 |
+
)
|
| 32 |
+
|
| 33 |
+
FusedAdamWConfig: LazyDict = L(get_base_optimizer)(
|
| 34 |
+
model=PLACEHOLDER,
|
| 35 |
+
lr=1e-4,
|
| 36 |
+
weight_decay=0.1,
|
| 37 |
+
betas=[0.9, 0.99],
|
| 38 |
+
optim_type="fusedadam",
|
| 39 |
+
eps=1e-8,
|
| 40 |
+
master_weights=True,
|
| 41 |
+
capturable=True,
|
| 42 |
+
)
|
| 43 |
+
|
| 44 |
+
MultipleAdamWConfig = L(get_multiple_optimizer)(
|
| 45 |
+
model=PLACEHOLDER,
|
| 46 |
+
lr=1e-4,
|
| 47 |
+
weight_decay=1e-3,
|
| 48 |
+
betas=[0.9, 0.999],
|
| 49 |
+
optim_type="adamw",
|
| 50 |
+
eps=1e-8,
|
| 51 |
+
fused=True,
|
| 52 |
+
lr_overrides=[], # New format: list of dicts with 'pattern', 'lr', and optional 'match_type'
|
| 53 |
+
)
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
MultipleFusedAdamWConfig = L(get_multiple_optimizer)(
|
| 57 |
+
model=PLACEHOLDER,
|
| 58 |
+
lr=1e-4,
|
| 59 |
+
weight_decay=1e-3,
|
| 60 |
+
betas=[0.9, 0.999],
|
| 61 |
+
optim_type="fusedadam",
|
| 62 |
+
eps=1e-8,
|
| 63 |
+
lr_overrides=[], # New format: list of dicts with 'pattern', 'lr', and optional 'match_type'
|
| 64 |
+
)
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
def register_optimizer():
|
| 68 |
+
cs = ConfigStore.instance()
|
| 69 |
+
cs.store(group="optimizer", package="optimizer", name="fusedadamw", node=FusedAdamWConfig)
|
| 70 |
+
cs.store(group="optimizer", package="optimizer", name="adamw", node=AdamWConfig)
|
| 71 |
+
cs.store(group="optimizer", package="optimizer", name="multipleadamw", node=MultipleAdamWConfig)
|
| 72 |
+
cs.store(group="optimizer", package="optimizer", name="multiplefusedadamw", node=MultipleFusedAdamWConfig)
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/__init__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/buttercup2p5_rectified_flow.py
ADDED
|
@@ -0,0 +1,229 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
"""Experiments for vanilla attention training on 2B model.
|
| 17 |
+
|
| 18 |
+
All experiments are using datasets located on AWS S3 unless specified.
|
| 19 |
+
If you want to use datasets on GCS, change the name of the override /data_train from e.g.
|
| 20 |
+
video_alpamayo_dec2024_s3_720p_10fps_93frames_7views -> video_alpamayo_dec2024_gcs_720p_10fps_93frames_7views
|
| 21 |
+
"""
|
| 22 |
+
|
| 23 |
+
from hydra.core.config_store import ConfigStore
|
| 24 |
+
|
| 25 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 26 |
+
from cosmos_predict2._src.predict2.text_encoders.text_encoder import EmbeddingConcatStrategy
|
| 27 |
+
from cosmos_predict2._src.predict2_multiview.callbacks.every_n_draw_sample_multiviewvideo import (
|
| 28 |
+
EveryNDrawSampleMultiviewVideo,
|
| 29 |
+
)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def buttercup_predict2p5_2b_7views_res720p_fps30_t8_joint_alpamayo1capviewprefix_allcapsviewprefix_29frames_nofps_uniform_dropoutt0() -> (
|
| 33 |
+
dict
|
| 34 |
+
):
|
| 35 |
+
state_t = 8
|
| 36 |
+
return dict(
|
| 37 |
+
defaults=[
|
| 38 |
+
{"override /ckpt_type": "dcp"},
|
| 39 |
+
{"override /optimizer": "adamw"},
|
| 40 |
+
{"override /callbacks": ["basic", "viz_online_sampling", "wandb", "cluster_speed"]},
|
| 41 |
+
{"override /checkpoint": "s3"},
|
| 42 |
+
{"override /tokenizer": "wan2pt1_tokenizer"},
|
| 43 |
+
{"override /data_train": "video_alpamayo_dec2024_s3_720p_30fps_29frames_7views"},
|
| 44 |
+
{"override /conditioner": "video_prediction_multiview_conditioner"},
|
| 45 |
+
{"override /model": "fsdp_rectified_flow_multiview"},
|
| 46 |
+
{"override /net": "cosmos_v1_2B_multiview"},
|
| 47 |
+
"_self_",
|
| 48 |
+
],
|
| 49 |
+
job=dict(
|
| 50 |
+
group="cosmos2_mv",
|
| 51 |
+
name="buttercup_predict2p5_2b_7views_res720p_fps30_t8_joint_alpamayo1capviewprefix_allcapsviewprefix_29frames_nofps_uniform_dropoutt0",
|
| 52 |
+
),
|
| 53 |
+
optimizer=dict(
|
| 54 |
+
lr=3e-5, # 2**(-14.5) = 3.0517578125e-05
|
| 55 |
+
weight_decay=1e-3,
|
| 56 |
+
betas=[0.9, 0.999],
|
| 57 |
+
),
|
| 58 |
+
scheduler=dict(
|
| 59 |
+
f_max=[0.99],
|
| 60 |
+
f_min=[0.4],
|
| 61 |
+
warm_up_steps=[100],
|
| 62 |
+
cycle_lengths=[400_000],
|
| 63 |
+
),
|
| 64 |
+
checkpoint=dict(
|
| 65 |
+
load_from_object_store=dict(
|
| 66 |
+
enabled=True,
|
| 67 |
+
),
|
| 68 |
+
save_to_object_store=dict(
|
| 69 |
+
enabled=True,
|
| 70 |
+
),
|
| 71 |
+
save_iter=500,
|
| 72 |
+
load_path="cosmos_predict2_multiview/cosmos2_mv/buttercup_predict2p5_2b_7views_res720p_fps30_t8_from48kfps30mv_condprobs0442_joint_alpamayo1capnoviewprefix_allcapsviewprefix_29frames_nofps-0/checkpoints/iter_000005000",
|
| 73 |
+
),
|
| 74 |
+
trainer=dict(
|
| 75 |
+
callbacks=dict(
|
| 76 |
+
compile_tokenizer=dict(
|
| 77 |
+
enabled=False,
|
| 78 |
+
),
|
| 79 |
+
iter_speed=dict(
|
| 80 |
+
hit_thres=50,
|
| 81 |
+
every_n=100,
|
| 82 |
+
),
|
| 83 |
+
every_n_sample_reg=L(EveryNDrawSampleMultiviewVideo)(
|
| 84 |
+
every_n=2_000,
|
| 85 |
+
do_x0_prediction=False,
|
| 86 |
+
is_ema=False,
|
| 87 |
+
num_sampling_step=35,
|
| 88 |
+
guidance=[0, 3, 7],
|
| 89 |
+
fps=30,
|
| 90 |
+
),
|
| 91 |
+
every_n_sample_ema=L(EveryNDrawSampleMultiviewVideo)(
|
| 92 |
+
every_n=2_000,
|
| 93 |
+
do_x0_prediction=False,
|
| 94 |
+
is_ema=True,
|
| 95 |
+
num_sampling_step=35,
|
| 96 |
+
guidance=[0, 3, 7],
|
| 97 |
+
fps=30,
|
| 98 |
+
),
|
| 99 |
+
),
|
| 100 |
+
),
|
| 101 |
+
model_parallel=dict(
|
| 102 |
+
context_parallel_size=8,
|
| 103 |
+
),
|
| 104 |
+
model=dict(
|
| 105 |
+
config=dict(
|
| 106 |
+
min_num_conditional_frames=0,
|
| 107 |
+
max_num_conditional_frames=2,
|
| 108 |
+
conditional_frames_probs={0: 0.5, 1: 0.25, 2: 0.25},
|
| 109 |
+
condition_locations=["first_random_n"],
|
| 110 |
+
fsdp_shard_size=8,
|
| 111 |
+
resolution="720p",
|
| 112 |
+
state_t=state_t,
|
| 113 |
+
shift=5,
|
| 114 |
+
use_dynamic_shift=False,
|
| 115 |
+
train_time_weight="uniform",
|
| 116 |
+
train_time_distribution="logitnormal",
|
| 117 |
+
online_text_embeddings_as_dict=False,
|
| 118 |
+
net=dict(
|
| 119 |
+
concat_view_embedding=True,
|
| 120 |
+
view_condition_dim=7,
|
| 121 |
+
state_t=8,
|
| 122 |
+
n_cameras_emb=7,
|
| 123 |
+
rope_enable_fps_modulation=False,
|
| 124 |
+
rope_h_extrapolation_ratio=3.0,
|
| 125 |
+
rope_w_extrapolation_ratio=3.0,
|
| 126 |
+
rope_t_extrapolation_ratio=float(state_t) / 24,
|
| 127 |
+
timestep_scale=0.001,
|
| 128 |
+
sac_config=dict(
|
| 129 |
+
mode="predict2_2b_720",
|
| 130 |
+
),
|
| 131 |
+
use_crossattn_projection=True,
|
| 132 |
+
crossattn_proj_in_channels=100352,
|
| 133 |
+
crossattn_emb_channels=1024,
|
| 134 |
+
use_wan_fp32_strategy=True,
|
| 135 |
+
),
|
| 136 |
+
conditioner=dict(
|
| 137 |
+
use_video_condition=dict(
|
| 138 |
+
dropout_rate=0.0,
|
| 139 |
+
),
|
| 140 |
+
text=dict(
|
| 141 |
+
dropout_rate=0.0,
|
| 142 |
+
use_empty_string=False,
|
| 143 |
+
),
|
| 144 |
+
),
|
| 145 |
+
tokenizer=dict(
|
| 146 |
+
temporal_window=16,
|
| 147 |
+
),
|
| 148 |
+
text_encoder_class="reason1p1_7B",
|
| 149 |
+
text_encoder_config=dict(
|
| 150 |
+
embedding_concat_strategy=str(EmbeddingConcatStrategy.FULL_CONCAT),
|
| 151 |
+
compute_online=True,
|
| 152 |
+
ckpt_path="s3://bucket/cosmos_reasoning1/sft_exp700/sft_exp721-1_qwen7b_tl_721_5vs5_s3_balanced_n32_resume_16k/checkpoints/iter_000016000/model/",
|
| 153 |
+
),
|
| 154 |
+
),
|
| 155 |
+
),
|
| 156 |
+
dataloader_train=dict(
|
| 157 |
+
augmentation_config=dict(
|
| 158 |
+
single_caption_camera_name="camera_front_wide_120fov",
|
| 159 |
+
add_view_prefix_to_caption=True,
|
| 160 |
+
)
|
| 161 |
+
),
|
| 162 |
+
)
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
# Note, use GB200 cluster for this config, as it OOMs on H100 clusters.
|
| 166 |
+
# For context parallelism of 24, use multiples of 24 GPUs or 6 nodes (since 4 GB200 per node).
|
| 167 |
+
def buttercup_predict2p5_2b_mv_7views_res720p_fps30_t24_joint_alpamayo1capviewprefix_allcapsviewprefix_93frames_nofps_uniform_dropoutt0() -> (
|
| 168 |
+
dict
|
| 169 |
+
):
|
| 170 |
+
state_t = 24
|
| 171 |
+
return dict(
|
| 172 |
+
defaults=[
|
| 173 |
+
"/experiment/buttercup_predict2p5_2b_7views_res720p_fps30_t8_joint_alpamayo1capviewprefix_allcapsviewprefix_29frames_nofps_uniform_dropoutt0",
|
| 174 |
+
{"override /data_train": "video_alpamayo_dec2024_gcs_720p_30fps_93frames_7views"},
|
| 175 |
+
"_self_",
|
| 176 |
+
],
|
| 177 |
+
job=dict(
|
| 178 |
+
group="cosmos2_mv",
|
| 179 |
+
name="buttercup_predict2p5_2b_mv_7views_res720p_fps30_t24_joint_alpamayo1capviewprefix_allcapsviewprefix_93frames_nofps_uniform_dropoutt0",
|
| 180 |
+
),
|
| 181 |
+
checkpoint=dict(
|
| 182 |
+
load_path="cosmos_predict2_multiview/cosmos2_mv/buttercup_predict2p5_2b_7views_res720p_fps30_t8_from48kfps30mv_condprobs0442_joint_alpamayo1capnoviewprefix_allcapsviewprefix_29frames_nofps-0/checkpoints/iter_000005000",
|
| 183 |
+
),
|
| 184 |
+
model_parallel=dict(
|
| 185 |
+
context_parallel_size=8,
|
| 186 |
+
),
|
| 187 |
+
model=dict(
|
| 188 |
+
config=dict(
|
| 189 |
+
state_t=state_t,
|
| 190 |
+
net=dict(
|
| 191 |
+
state_t=state_t,
|
| 192 |
+
rope_enable_fps_modulation=False,
|
| 193 |
+
rope_h_extrapolation_ratio=3.0,
|
| 194 |
+
rope_w_extrapolation_ratio=3.0,
|
| 195 |
+
rope_t_extrapolation_ratio=state_t / 24.0,
|
| 196 |
+
sac_config=dict(
|
| 197 |
+
mode="predict2_2b_720",
|
| 198 |
+
),
|
| 199 |
+
),
|
| 200 |
+
),
|
| 201 |
+
),
|
| 202 |
+
trainer=dict(
|
| 203 |
+
straggler_detection=dict(
|
| 204 |
+
enabled=False,
|
| 205 |
+
),
|
| 206 |
+
),
|
| 207 |
+
dataloader_train=dict(
|
| 208 |
+
augmentation_config=dict(
|
| 209 |
+
single_caption_camera_name="camera_front_wide_120fov",
|
| 210 |
+
add_view_prefix_to_caption=True,
|
| 211 |
+
),
|
| 212 |
+
),
|
| 213 |
+
)
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
experiments = [
|
| 217 |
+
buttercup_predict2p5_2b_7views_res720p_fps30_t8_joint_alpamayo1capviewprefix_allcapsviewprefix_29frames_nofps_uniform_dropoutt0(),
|
| 218 |
+
buttercup_predict2p5_2b_mv_7views_res720p_fps30_t24_joint_alpamayo1capviewprefix_allcapsviewprefix_93frames_nofps_uniform_dropoutt0(),
|
| 219 |
+
]
|
| 220 |
+
|
| 221 |
+
cs = ConfigStore.instance()
|
| 222 |
+
|
| 223 |
+
for _item in experiments:
|
| 224 |
+
cs.store(
|
| 225 |
+
group="experiment",
|
| 226 |
+
package="_global_",
|
| 227 |
+
name=_item["job"]["name"],
|
| 228 |
+
node=_item,
|
| 229 |
+
)
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/buttercup2p5_rectified_flow_14b.py
ADDED
|
@@ -0,0 +1,243 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
from hydra.core.config_store import ConfigStore
|
| 18 |
+
|
| 19 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 20 |
+
from cosmos_predict2._src.predict2.text_encoders.text_encoder import EmbeddingConcatStrategy
|
| 21 |
+
from cosmos_predict2._src.predict2_multiview.callbacks.every_n_draw_sample_multiviewvideo import (
|
| 22 |
+
EveryNDrawSampleMultiviewVideo,
|
| 23 |
+
)
|
| 24 |
+
from cosmos_predict2._src.predict2_multiview.callbacks.log_weight import LogWeight
|
| 25 |
+
|
| 26 |
+
"""
|
| 27 |
+
torchrun --nproc_per_node=8 --master_port=12341 -m scripts.train --config=cosmos_predict2/_src/predict2_multiview/configs/vid2vid/config.py -- experiment=buttercup_predict2p5_14b_7views_res720p_fps30_t8_frombase2p5_condprobs0442_joint_alpamayo1capnoviewprefix_allcapsviewprefix_29frames_nofps_uniform
|
| 28 |
+
"""
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def buttercup_predict2p5_14b_7views_res720p_fps10_t8_fromavfinetune_allcaps_29frames_nofps_uniform() -> dict:
|
| 32 |
+
state_t = 8
|
| 33 |
+
return dict(
|
| 34 |
+
defaults=[
|
| 35 |
+
{"override /data_train": "video_alpamayo_dec2024_gcs_720p_10fps_29frames_7views"},
|
| 36 |
+
{"override /conditioner": "video_prediction_multiview_conditioner"},
|
| 37 |
+
{"override /model": "fsdp_rectified_flow_multiview"},
|
| 38 |
+
{"override /net": "cosmos_v1_14B_multiview"},
|
| 39 |
+
{"override /ckpt_type": "dcp"},
|
| 40 |
+
{"override /optimizer": "adamw"},
|
| 41 |
+
{
|
| 42 |
+
"override /callbacks": [
|
| 43 |
+
"basic",
|
| 44 |
+
"viz_online_sampling",
|
| 45 |
+
"wandb",
|
| 46 |
+
"cluster_speed",
|
| 47 |
+
]
|
| 48 |
+
},
|
| 49 |
+
{"override /checkpoint": "s3"},
|
| 50 |
+
{"override /tokenizer": "wan2pt1_tokenizer"},
|
| 51 |
+
"_self_",
|
| 52 |
+
],
|
| 53 |
+
job=dict(
|
| 54 |
+
group="cosmos2p5_mv",
|
| 55 |
+
name="buttercup_predict2p5_14b_7views_res720p_fps10_t8_fromavfinetune_allcaps_29frames_nofps_uniform",
|
| 56 |
+
),
|
| 57 |
+
checkpoint=dict(
|
| 58 |
+
save_to_object_store=dict(
|
| 59 |
+
enabled=True,
|
| 60 |
+
),
|
| 61 |
+
load_from_object_store=dict(
|
| 62 |
+
enabled=True,
|
| 63 |
+
),
|
| 64 |
+
save_iter=250,
|
| 65 |
+
load_path="cosmos_diffusion_v2/official_runs_vid2vid/Stage-c_pt_4-Index-43-Size-14B-Res-720-Fps-16-Note-rf_av_high_sigma_uniform/checkpoints/iter_000022000/",
|
| 66 |
+
),
|
| 67 |
+
model_parallel=dict(
|
| 68 |
+
context_parallel_size=8,
|
| 69 |
+
),
|
| 70 |
+
optimizer=dict(
|
| 71 |
+
lr=2 ** (-14.5),
|
| 72 |
+
weight_decay=0.001,
|
| 73 |
+
betas=[0.9, 0.999],
|
| 74 |
+
),
|
| 75 |
+
scheduler=dict(
|
| 76 |
+
f_max=[0.3],
|
| 77 |
+
f_min=[0.1],
|
| 78 |
+
warm_up_steps=[2_000],
|
| 79 |
+
cycle_lengths=[200_000],
|
| 80 |
+
),
|
| 81 |
+
model=dict(
|
| 82 |
+
config=dict(
|
| 83 |
+
min_num_conditional_frames=0,
|
| 84 |
+
max_num_conditional_frames=2,
|
| 85 |
+
conditional_frames_probs={0: 0.6, 1: 0.2, 2: 0.2},
|
| 86 |
+
condition_locations=["first_random_n"],
|
| 87 |
+
fsdp_shard_size=32,
|
| 88 |
+
resolution="720p",
|
| 89 |
+
online_text_embeddings_as_dict=False,
|
| 90 |
+
state_t=state_t,
|
| 91 |
+
shift=5,
|
| 92 |
+
use_dynamic_shift=False,
|
| 93 |
+
train_time_weight="uniform",
|
| 94 |
+
train_time_distribution="logitnormal",
|
| 95 |
+
use_high_sigma_strategy=True,
|
| 96 |
+
net=dict(
|
| 97 |
+
concat_view_embedding=True,
|
| 98 |
+
view_condition_dim=7,
|
| 99 |
+
n_cameras_emb=7,
|
| 100 |
+
state_t=state_t,
|
| 101 |
+
rope_enable_fps_modulation=False,
|
| 102 |
+
rope_h_extrapolation_ratio=3.0,
|
| 103 |
+
rope_w_extrapolation_ratio=3.0,
|
| 104 |
+
rope_t_extrapolation_ratio=float(state_t) / 24,
|
| 105 |
+
timestep_scale=0.001,
|
| 106 |
+
sac_config=dict(
|
| 107 |
+
mode="predict2_14b_720_aggressive",
|
| 108 |
+
),
|
| 109 |
+
use_crossattn_projection=True,
|
| 110 |
+
crossattn_proj_in_channels=100352,
|
| 111 |
+
crossattn_emb_channels=1024,
|
| 112 |
+
use_wan_fp32_strategy=True,
|
| 113 |
+
),
|
| 114 |
+
conditioner=dict(
|
| 115 |
+
use_video_condition=dict(
|
| 116 |
+
dropout_rate=0.0,
|
| 117 |
+
),
|
| 118 |
+
text=dict(
|
| 119 |
+
dropout_rate=0.2,
|
| 120 |
+
use_empty_string=False,
|
| 121 |
+
),
|
| 122 |
+
),
|
| 123 |
+
tokenizer=dict(
|
| 124 |
+
temporal_window=16,
|
| 125 |
+
),
|
| 126 |
+
text_encoder_class="reason1p1_7B",
|
| 127 |
+
text_encoder_config=dict(
|
| 128 |
+
embedding_concat_strategy=str(EmbeddingConcatStrategy.FULL_CONCAT),
|
| 129 |
+
compute_online=True,
|
| 130 |
+
ckpt_path="s3://bucket/cosmos_reasoning1/sft_exp700/sft_exp721-1_qwen7b_tl_721_5vs5_s3_balanced_n32_resume_16k/checkpoints/iter_000016000/model/",
|
| 131 |
+
s3_credential_path="credentials/s3_checkpoint.secret",
|
| 132 |
+
),
|
| 133 |
+
)
|
| 134 |
+
),
|
| 135 |
+
trainer=dict(
|
| 136 |
+
straggler_detection=dict(
|
| 137 |
+
enabled=False,
|
| 138 |
+
),
|
| 139 |
+
callbacks=dict(
|
| 140 |
+
compile_tokenizer=dict(
|
| 141 |
+
enabled=False,
|
| 142 |
+
),
|
| 143 |
+
iter_speed=dict(
|
| 144 |
+
hit_thres=50,
|
| 145 |
+
every_n=100,
|
| 146 |
+
),
|
| 147 |
+
every_n_sample_ema=L(EveryNDrawSampleMultiviewVideo)(
|
| 148 |
+
every_n=2_000,
|
| 149 |
+
do_x0_prediction=False,
|
| 150 |
+
is_ema=True,
|
| 151 |
+
num_sampling_step=35,
|
| 152 |
+
guidance=[0, 3],
|
| 153 |
+
fps=10,
|
| 154 |
+
run_at_start=True,
|
| 155 |
+
),
|
| 156 |
+
),
|
| 157 |
+
),
|
| 158 |
+
dataloader_train=dict(
|
| 159 |
+
augmentation_config=dict(
|
| 160 |
+
caption_probability={
|
| 161 |
+
"qwen2p5_7b_caption": 0.7,
|
| 162 |
+
"qwen2p5_7b_caption_medium": 0.2,
|
| 163 |
+
"qwen2p5_7b_caption_short": 0.1,
|
| 164 |
+
},
|
| 165 |
+
)
|
| 166 |
+
),
|
| 167 |
+
)
|
| 168 |
+
|
| 169 |
+
|
| 170 |
+
def buttercup_predict2p5_14b_crossview_7views_res720p_fps10_t8_fromavfinetune_allcaps_29frames_nofps_uniform():
|
| 171 |
+
return dict(
|
| 172 |
+
defaults=[
|
| 173 |
+
"/experiment/buttercup_predict2p5_14b_7views_res720p_fps10_t8_fromavfinetune_allcaps_29frames_nofps_uniform",
|
| 174 |
+
{"override /net": "cosmos_v1_14B_multiview_crossview"},
|
| 175 |
+
{"override /optimizer": "multiplefusedadamw"},
|
| 176 |
+
"_self_",
|
| 177 |
+
],
|
| 178 |
+
job=dict(
|
| 179 |
+
group="cosmos2_mv2",
|
| 180 |
+
name="buttercup_predict2p5_14b_crossview_7views_res720p_fps10_t8_fromavfinetune_allcaps_29frames_nofps_uniform",
|
| 181 |
+
),
|
| 182 |
+
model=dict(
|
| 183 |
+
config=dict(
|
| 184 |
+
net=dict(
|
| 185 |
+
cross_view_attn_map_str={
|
| 186 |
+
"camera_front_wide_120fov": [
|
| 187 |
+
"camera_cross_left_120fov",
|
| 188 |
+
"camera_cross_right_120fov",
|
| 189 |
+
"camera_front_tele_30fov",
|
| 190 |
+
],
|
| 191 |
+
"camera_cross_left_120fov": ["camera_front_wide_120fov", "camera_rear_left_70fov"],
|
| 192 |
+
"camera_cross_right_120fov": ["camera_front_wide_120fov", "camera_rear_right_70fov"],
|
| 193 |
+
"camera_rear_left_70fov": ["camera_cross_left_120fov", "camera_rear_tele_30fov"],
|
| 194 |
+
"camera_rear_right_70fov": ["camera_cross_right_120fov", "camera_rear_tele_30fov"],
|
| 195 |
+
"camera_rear_tele_30fov": ["camera_rear_left_70fov", "camera_rear_right_70fov"],
|
| 196 |
+
"camera_front_tele_30fov": ["camera_front_wide_120fov"],
|
| 197 |
+
},
|
| 198 |
+
camera_to_view_id={
|
| 199 |
+
"camera_front_wide_120fov": 0,
|
| 200 |
+
"camera_cross_left_120fov": 5,
|
| 201 |
+
"camera_cross_right_120fov": 1,
|
| 202 |
+
"camera_rear_left_70fov": 4,
|
| 203 |
+
"camera_rear_right_70fov": 2,
|
| 204 |
+
"camera_rear_tele_30fov": 3,
|
| 205 |
+
"camera_front_tele_30fov": 6,
|
| 206 |
+
},
|
| 207 |
+
),
|
| 208 |
+
),
|
| 209 |
+
),
|
| 210 |
+
optimizer=dict(
|
| 211 |
+
lr=3e-5,
|
| 212 |
+
lr_overrides=[
|
| 213 |
+
{"pattern": "cross_view_attn", "lr": 1e-4, "match_type": "contains"},
|
| 214 |
+
],
|
| 215 |
+
),
|
| 216 |
+
trainer=dict(
|
| 217 |
+
logging_iter=50,
|
| 218 |
+
callbacks=dict(
|
| 219 |
+
log_weight=L(LogWeight)(
|
| 220 |
+
every_n=50,
|
| 221 |
+
),
|
| 222 |
+
every_n_sample_reg=dict(
|
| 223 |
+
every_n=1500,
|
| 224 |
+
),
|
| 225 |
+
),
|
| 226 |
+
),
|
| 227 |
+
)
|
| 228 |
+
|
| 229 |
+
|
| 230 |
+
experiments = [
|
| 231 |
+
buttercup_predict2p5_14b_7views_res720p_fps10_t8_fromavfinetune_allcaps_29frames_nofps_uniform(),
|
| 232 |
+
buttercup_predict2p5_14b_crossview_7views_res720p_fps10_t8_fromavfinetune_allcaps_29frames_nofps_uniform(),
|
| 233 |
+
]
|
| 234 |
+
|
| 235 |
+
cs = ConfigStore.instance()
|
| 236 |
+
|
| 237 |
+
for _item in experiments:
|
| 238 |
+
cs.store(
|
| 239 |
+
group="experiment",
|
| 240 |
+
package="_global_",
|
| 241 |
+
name=_item["job"]["name"],
|
| 242 |
+
node=_item,
|
| 243 |
+
)
|
cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/nymeria_pose_2actor.py
ADDED
|
@@ -0,0 +1,1005 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
"""2-actor joint video generation on Nymeria, full fine-tune of Cosmos Predict 2.5 (2B) on 4 GPUs with FSDP.
|
| 17 |
+
|
| 18 |
+
Actor 1 / Actor 2 are the two "views" (V=2). Per actor: frame 0 is given, frames 1.. are generated; the
|
| 19 |
+
warped past-frame (pixel -> VAE latent) + visibility mask are channel-concat conditioning (additive zero-init
|
| 20 |
+
PatchEmbed), and the projected 2D pose RGB map is fed through the 3D-conv pose encoder and added to the patch
|
| 21 |
+
tokens. VAE + text encoder are frozen; the DiT backbone + pose encoder + view embeddings are fully trained.
|
| 22 |
+
|
| 23 |
+
Architecture mirrors the pretrained ``buttercup_predict2p5_2b ... t8`` 2B multiview net so the base checkpoint
|
| 24 |
+
warm-starts cleanly (the new pose/cond modules load as missing keys).
|
| 25 |
+
|
| 26 |
+
state_t = 9 latent frames (= 1 + (33 - 1) // 4 for the 33-frame nymeria_processed clips).
|
| 27 |
+
|
| 28 |
+
Launch (4 GPUs):
|
| 29 |
+
torchrun --nproc_per_node=4 -m scripts.train \
|
| 30 |
+
--config=cosmos_predict2/_src/predict2_multiview/configs/vid2vid/config.py -- \
|
| 31 |
+
experiment=nymeria_pose_2actor_2b_t9_fsdp4
|
| 32 |
+
|
| 33 |
+
NOTE (deployment): set ``checkpoint.load_path`` to your local 2B multiview checkpoint dir for warm start, and
|
| 34 |
+
ensure the frozen text-encoder checkpoint (``model.config.text_encoder_config.ckpt_path``) is reachable. Both
|
| 35 |
+
are inherited from the buttercup recipe below and may point at S3 by default.
|
| 36 |
+
"""
|
| 37 |
+
|
| 38 |
+
from hydra.core.config_store import ConfigStore
|
| 39 |
+
|
| 40 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 41 |
+
from cosmos_predict2._src.predict2.text_encoders.text_encoder import EmbeddingConcatStrategy
|
| 42 |
+
from cosmos_predict2._src.predict2_multiview.callbacks.every_n_draw_sample_multiviewvideo import (
|
| 43 |
+
EveryNDrawSampleMultiviewVideo,
|
| 44 |
+
)
|
| 45 |
+
from cosmos_predict2._src.predict2_multiview.callbacks.nymeria_validation_viz import NymeriaValidationViz
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def nymeria_pose_2actor_2b_t9_fsdp4() -> dict:
|
| 49 |
+
state_t = 9 # 1 + (33 - 1) // 4 ; matches 33 target frames per actor @10fps
|
| 50 |
+
return dict(
|
| 51 |
+
defaults=[
|
| 52 |
+
{"override /ckpt_type": "dcp"},
|
| 53 |
+
{"override /optimizer": "adamw"},
|
| 54 |
+
{"override /callbacks": ["basic", "viz_online_sampling"]},
|
| 55 |
+
{"override /checkpoint": "s3"}, # used with object-store disabled below -> writes to local job dir
|
| 56 |
+
{"override /tokenizer": "wan2pt1_tokenizer"},
|
| 57 |
+
{"override /data_train": "nymeria_pairs_480p_10fps_33frames"},
|
| 58 |
+
{"override /data_val": "nymeria_pairs_480p_10fps_33frames"},
|
| 59 |
+
{"override /conditioner": "video_prediction_multiview_pose_conditioner"},
|
| 60 |
+
{"override /model": "fsdp_rectified_flow_multiview_pose"},
|
| 61 |
+
{"override /net": "cosmos_v1_2B_multiview_pose"},
|
| 62 |
+
"_self_",
|
| 63 |
+
],
|
| 64 |
+
job=dict(
|
| 65 |
+
group="nymeria_pose",
|
| 66 |
+
name="nymeria_pose_2actor_2b_t9_fsdp4",
|
| 67 |
+
),
|
| 68 |
+
optimizer=dict(
|
| 69 |
+
lr=3e-5,
|
| 70 |
+
weight_decay=1e-3,
|
| 71 |
+
betas=[0.9, 0.999],
|
| 72 |
+
),
|
| 73 |
+
scheduler=dict(
|
| 74 |
+
f_max=[0.99],
|
| 75 |
+
f_min=[0.4],
|
| 76 |
+
warm_up_steps=[100],
|
| 77 |
+
cycle_lengths=[400_000],
|
| 78 |
+
),
|
| 79 |
+
checkpoint=dict(
|
| 80 |
+
# local training: keep checkpoints on the local filesystem (no object store).
|
| 81 |
+
load_from_object_store=dict(enabled=False),
|
| 82 |
+
save_to_object_store=dict(enabled=False),
|
| 83 |
+
save_iter=500,
|
| 84 |
+
# TODO(deploy): point at your local 2B multiview checkpoint dir to warm-start the DiT backbone.
|
| 85 |
+
# Leave as "" to train the backbone from its random init (pose/cond modules are always fresh).
|
| 86 |
+
load_path="",
|
| 87 |
+
strict_resume=False, # tolerate the new pose/cond parameters missing from the base checkpoint
|
| 88 |
+
),
|
| 89 |
+
trainer=dict(
|
| 90 |
+
max_iter=400_000,
|
| 91 |
+
callbacks=dict(
|
| 92 |
+
every_n_sample_reg=L(EveryNDrawSampleMultiviewVideo)(
|
| 93 |
+
every_n=2_000,
|
| 94 |
+
do_x0_prediction=False,
|
| 95 |
+
is_ema=False,
|
| 96 |
+
num_sampling_step=35,
|
| 97 |
+
guidance=[0, 3, 7],
|
| 98 |
+
fps=10,
|
| 99 |
+
),
|
| 100 |
+
every_n_sample_ema=L(EveryNDrawSampleMultiviewVideo)(
|
| 101 |
+
every_n=2_000,
|
| 102 |
+
do_x0_prediction=False,
|
| 103 |
+
is_ema=True,
|
| 104 |
+
num_sampling_step=35,
|
| 105 |
+
guidance=[0, 3, 7],
|
| 106 |
+
fps=10,
|
| 107 |
+
),
|
| 108 |
+
),
|
| 109 |
+
),
|
| 110 |
+
model_parallel=dict(
|
| 111 |
+
context_parallel_size=1,
|
| 112 |
+
),
|
| 113 |
+
model=dict(
|
| 114 |
+
config=dict(
|
| 115 |
+
# frame 0 given, the rest generated (1 conditional frame per actor)
|
| 116 |
+
min_num_conditional_frames_per_view=1,
|
| 117 |
+
max_num_conditional_frames_per_view=1,
|
| 118 |
+
conditional_frames_probs={1: 1.0},
|
| 119 |
+
condition_locations=["first_random_n"],
|
| 120 |
+
fsdp_shard_size=4, # FSDP across the 4 GPUs
|
| 121 |
+
resolution="480",
|
| 122 |
+
state_t=state_t,
|
| 123 |
+
shift=5,
|
| 124 |
+
use_dynamic_shift=False,
|
| 125 |
+
train_time_weight="uniform",
|
| 126 |
+
train_time_distribution="logitnormal",
|
| 127 |
+
online_text_embeddings_as_dict=False,
|
| 128 |
+
net=dict(
|
| 129 |
+
concat_view_embedding=True,
|
| 130 |
+
view_condition_dim=7,
|
| 131 |
+
state_t=state_t,
|
| 132 |
+
n_cameras_emb=7,
|
| 133 |
+
rope_enable_fps_modulation=False,
|
| 134 |
+
rope_h_extrapolation_ratio=3.0,
|
| 135 |
+
rope_w_extrapolation_ratio=3.0,
|
| 136 |
+
rope_t_extrapolation_ratio=float(state_t) / 24.0,
|
| 137 |
+
timestep_scale=0.001,
|
| 138 |
+
sac_config=dict(mode="predict2_2b_720"),
|
| 139 |
+
use_crossattn_projection=True,
|
| 140 |
+
crossattn_proj_in_channels=100352,
|
| 141 |
+
crossattn_emb_channels=1024,
|
| 142 |
+
use_wan_fp32_strategy=True,
|
| 143 |
+
# pose / warped / visibility conditioning channels
|
| 144 |
+
warped_latent_channels=16,
|
| 145 |
+
visibility_channels=1,
|
| 146 |
+
pose_in_channels=3,
|
| 147 |
+
pose_hidden_dim=16,
|
| 148 |
+
),
|
| 149 |
+
conditioner=dict(
|
| 150 |
+
use_video_condition=dict(dropout_rate=0.0),
|
| 151 |
+
text=dict(dropout_rate=0.0, use_empty_string=False),
|
| 152 |
+
),
|
| 153 |
+
tokenizer=dict(temporal_window=16),
|
| 154 |
+
text_encoder_class="reason1p1_7B",
|
| 155 |
+
text_encoder_config=dict(
|
| 156 |
+
embedding_concat_strategy=str(EmbeddingConcatStrategy.FULL_CONCAT),
|
| 157 |
+
compute_online=True,
|
| 158 |
+
ckpt_path="s3://bucket/cosmos_reasoning1/sft_exp700/sft_exp721-1_qwen7b_tl_721_5vs5_s3_balanced_n32_resume_16k/checkpoints/iter_000016000/model/",
|
| 159 |
+
),
|
| 160 |
+
),
|
| 161 |
+
),
|
| 162 |
+
)
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
def _local_wan_vae_path() -> str:
|
| 166 |
+
"""Resolve the locally-cached Wan2.1 VAE (downloaded from HF); used for offline runs."""
|
| 167 |
+
from huggingface_hub import hf_hub_download
|
| 168 |
+
|
| 169 |
+
return hf_hub_download("Wan-AI/Wan2.1-T2V-1.3B", "Wan2.1_VAE.pth")
|
| 170 |
+
|
| 171 |
+
|
| 172 |
+
def _base_2b_multiview_ckpt() -> str:
|
| 173 |
+
"""Local path to the pretrained 2B multiview checkpoint (HF mirror of MODEL_CHECKPOINTS auto/multiview).
|
| 174 |
+
|
| 175 |
+
Returns "" if it cannot be resolved/downloaded (then the backbone trains from random init).
|
| 176 |
+
"""
|
| 177 |
+
try:
|
| 178 |
+
from huggingface_hub import hf_hub_download
|
| 179 |
+
|
| 180 |
+
return hf_hub_download(
|
| 181 |
+
"nvidia/Cosmos-Predict2.5-2B",
|
| 182 |
+
"auto/multiview/524af350-2e43-496c-8590-3646ae1325da_ema_bf16.pt",
|
| 183 |
+
revision="865baf084d4c9e850eac59a021277d5a9b9e8b63",
|
| 184 |
+
)
|
| 185 |
+
except Exception:
|
| 186 |
+
return ""
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
def nymeria_pose_2actor_2b_smoke() -> dict:
|
| 190 |
+
"""Offline smoke test: real Wan VAE (local) + zero dummy text (no S3 text encoder), random-init backbone.
|
| 191 |
+
|
| 192 |
+
Exercises the full pipeline end-to-end: dataloader -> VAE encode (video + warped) -> pose preprocessing
|
| 193 |
+
-> MultiViewPoseDiT forward -> rectified-flow loss -> backward -> FSDP -> optimizer step -> checkpoint.
|
| 194 |
+
|
| 195 |
+
Single GPU: torchrun --nproc_per_node=1 -m scripts.train --config=<cfg> -- experiment=nymeria_pose_2actor_2b_smoke model.config.fsdp_shard_size=1
|
| 196 |
+
4-GPU FSDP: torchrun --nproc_per_node=4 -m scripts.train --config=<cfg> -- experiment=nymeria_pose_2actor_2b_smoke
|
| 197 |
+
"""
|
| 198 |
+
state_t = 9
|
| 199 |
+
return dict(
|
| 200 |
+
defaults=[
|
| 201 |
+
{"override /ckpt_type": "dcp"},
|
| 202 |
+
{"override /optimizer": "adamw"},
|
| 203 |
+
{"override /callbacks": ["basic"]},
|
| 204 |
+
{"override /checkpoint": "s3"},
|
| 205 |
+
{"override /tokenizer": "wan2pt1_tokenizer"},
|
| 206 |
+
{"override /data_train": "nymeria_pairs_smoke"},
|
| 207 |
+
{"override /data_val": "nymeria_pairs_smoke"},
|
| 208 |
+
{"override /conditioner": "video_prediction_multiview_pose_conditioner"},
|
| 209 |
+
{"override /model": "fsdp_rectified_flow_multiview_pose"},
|
| 210 |
+
{"override /net": "cosmos_v1_2B_multiview_pose"},
|
| 211 |
+
"_self_",
|
| 212 |
+
],
|
| 213 |
+
job=dict(group="nymeria_pose", name="nymeria_pose_2actor_2b_smoke"),
|
| 214 |
+
optimizer=dict(lr=3e-5, weight_decay=1e-3, betas=[0.9, 0.999]),
|
| 215 |
+
scheduler=dict(f_max=[0.99], f_min=[0.4], warm_up_steps=[1], cycle_lengths=[1000]),
|
| 216 |
+
checkpoint=dict(
|
| 217 |
+
load_from_object_store=dict(enabled=False),
|
| 218 |
+
save_to_object_store=dict(enabled=False),
|
| 219 |
+
save_iter=1000,
|
| 220 |
+
load_path="",
|
| 221 |
+
strict_resume=False,
|
| 222 |
+
),
|
| 223 |
+
trainer=dict(max_iter=3, logging_iter=1),
|
| 224 |
+
model_parallel=dict(context_parallel_size=1),
|
| 225 |
+
model=dict(
|
| 226 |
+
config=dict(
|
| 227 |
+
min_num_conditional_frames_per_view=1,
|
| 228 |
+
max_num_conditional_frames_per_view=1,
|
| 229 |
+
conditional_frames_probs={1: 1.0},
|
| 230 |
+
condition_locations=["first_random_n"],
|
| 231 |
+
fsdp_shard_size=4,
|
| 232 |
+
resolution="480",
|
| 233 |
+
state_t=state_t,
|
| 234 |
+
shift=5,
|
| 235 |
+
use_dynamic_shift=False,
|
| 236 |
+
train_time_weight="uniform",
|
| 237 |
+
train_time_distribution="logitnormal",
|
| 238 |
+
online_text_embeddings_as_dict=False,
|
| 239 |
+
ema=dict(enabled=False),
|
| 240 |
+
net=dict(
|
| 241 |
+
concat_view_embedding=True,
|
| 242 |
+
view_condition_dim=7,
|
| 243 |
+
state_t=state_t,
|
| 244 |
+
n_cameras_emb=7,
|
| 245 |
+
rope_enable_fps_modulation=False,
|
| 246 |
+
rope_h_extrapolation_ratio=3.0,
|
| 247 |
+
rope_w_extrapolation_ratio=3.0,
|
| 248 |
+
rope_t_extrapolation_ratio=float(state_t) / 24.0,
|
| 249 |
+
timestep_scale=0.001,
|
| 250 |
+
sac_config=dict(mode="predict2_2b_720"),
|
| 251 |
+
use_crossattn_projection=False, # consume dummy crossattn (1024-d) directly
|
| 252 |
+
crossattn_emb_channels=1024,
|
| 253 |
+
use_wan_fp32_strategy=True,
|
| 254 |
+
warped_latent_channels=16,
|
| 255 |
+
visibility_channels=1,
|
| 256 |
+
pose_in_channels=3,
|
| 257 |
+
pose_hidden_dim=16,
|
| 258 |
+
),
|
| 259 |
+
conditioner=dict(
|
| 260 |
+
use_video_condition=dict(dropout_rate=0.0),
|
| 261 |
+
text=dict(dropout_rate=0.0, use_empty_string=False),
|
| 262 |
+
),
|
| 263 |
+
tokenizer=dict(vae_pth=_local_wan_vae_path(), temporal_window=16),
|
| 264 |
+
text_encoder_class="reason1p1_7B", # valid label; encoder is NOT built (config below is None)
|
| 265 |
+
text_encoder_config=None, # disable the (S3-only) text encoder; dataset provides dummy embeddings
|
| 266 |
+
),
|
| 267 |
+
),
|
| 268 |
+
)
|
| 269 |
+
|
| 270 |
+
|
| 271 |
+
def nymeria_pose_2actor_2b_longer() -> dict:
|
| 272 |
+
"""Production run on the LONGER dataset (200 past + 77 target frames), 4-GPU FSDP full fine-tune.
|
| 273 |
+
|
| 274 |
+
- Warm-starts the 2B multiview backbone from the HF checkpoint (auto/multiview ema_bf16.pt); the new
|
| 275 |
+
pose encoder + warped/visibility embedders load as missing keys (kept at zero-init) under strict_resume=False.
|
| 276 |
+
- VAE = local Wan2.1 VAE. Text is disabled (Nymeria has no captions) -> dummy zero crossattn + no crossattn
|
| 277 |
+
projection. To use real text instead: set data_train/val to `nymeria_pairs_longer_77frames`, restore
|
| 278 |
+
`use_crossattn_projection=True` / `crossattn_proj_in_channels=100352`, and provide the reason1 text encoder.
|
| 279 |
+
- Outputs (checkpoints + validation videos) go under $IMAGINAIRE_OUTPUT_ROOT/multi-ego/nymeria_pose/<name>,
|
| 280 |
+
and W&B logs to entity/project from $WANDB_ENTITY / job.project with run name == <name> (== the folder).
|
| 281 |
+
"""
|
| 282 |
+
state_t = 20 # 1 + (77 - 1) // 4 ; matches 77 target frames per actor @10fps
|
| 283 |
+
return dict(
|
| 284 |
+
defaults=[
|
| 285 |
+
{"override /ckpt_type": "dcp"},
|
| 286 |
+
{"override /optimizer": "adamw"},
|
| 287 |
+
{"override /callbacks": ["basic", "wandb"]},
|
| 288 |
+
{"override /checkpoint": "s3"},
|
| 289 |
+
{"override /tokenizer": "wan2pt1_tokenizer"},
|
| 290 |
+
{"override /data_train": "nymeria_train_ready_smoke"},
|
| 291 |
+
{"override /data_val": "nymeria_train_ready_smoke"},
|
| 292 |
+
{"override /conditioner": "video_prediction_multiview_pose_conditioner"},
|
| 293 |
+
{"override /model": "fsdp_rectified_flow_multiview_pose"},
|
| 294 |
+
{"override /net": "cosmos_v1_2B_multiview_pose"},
|
| 295 |
+
"_self_",
|
| 296 |
+
],
|
| 297 |
+
# folder -> $IMAGINAIRE_OUTPUT_ROOT/train/<name> (project="train", empty group). W&B project/entity are
|
| 298 |
+
# decoupled via $WANDB_PROJECT / $WANDB_ENTITY (set in sh/train_nymeria_longer.sh).
|
| 299 |
+
job=dict(group="", name="nymeria_pose_2actor_2b_longer", project="train"),
|
| 300 |
+
optimizer=dict(lr=3e-5, weight_decay=1e-3, betas=[0.9, 0.999]),
|
| 301 |
+
scheduler=dict(f_max=[0.99], f_min=[0.4], warm_up_steps=[100], cycle_lengths=[400_000]),
|
| 302 |
+
checkpoint=dict(
|
| 303 |
+
load_from_object_store=dict(enabled=False),
|
| 304 |
+
save_to_object_store=dict(enabled=False),
|
| 305 |
+
save_iter=500,
|
| 306 |
+
load_path=_base_2b_multiview_ckpt(), # warm-start backbone; "" -> random init
|
| 307 |
+
strict_resume=False,
|
| 308 |
+
),
|
| 309 |
+
trainer=dict(
|
| 310 |
+
max_iter=400_000,
|
| 311 |
+
logging_iter=50,
|
| 312 |
+
callbacks=dict(
|
| 313 |
+
# custom 2-actor validation viz: [warped | pose | GT | generated] x [egoA; egoB], 5 train + 5 val.
|
| 314 |
+
validation_viz=L(NymeriaValidationViz)(
|
| 315 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 316 |
+
num_video_frames=77,
|
| 317 |
+
),
|
| 318 |
+
),
|
| 319 |
+
),
|
| 320 |
+
model_parallel=dict(context_parallel_size=1),
|
| 321 |
+
model=dict(
|
| 322 |
+
config=dict(
|
| 323 |
+
min_num_conditional_frames_per_view=1,
|
| 324 |
+
max_num_conditional_frames_per_view=1,
|
| 325 |
+
conditional_frames_probs={1: 1.0},
|
| 326 |
+
condition_locations=["first_random_n"],
|
| 327 |
+
fsdp_shard_size=4,
|
| 328 |
+
resolution="480",
|
| 329 |
+
state_t=state_t,
|
| 330 |
+
shift=5,
|
| 331 |
+
use_dynamic_shift=False,
|
| 332 |
+
train_time_weight="uniform",
|
| 333 |
+
train_time_distribution="logitnormal",
|
| 334 |
+
online_text_embeddings_as_dict=False,
|
| 335 |
+
net=dict(
|
| 336 |
+
concat_view_embedding=True,
|
| 337 |
+
view_condition_dim=7,
|
| 338 |
+
state_t=state_t,
|
| 339 |
+
n_cameras_emb=7,
|
| 340 |
+
rope_enable_fps_modulation=False,
|
| 341 |
+
rope_h_extrapolation_ratio=3.0,
|
| 342 |
+
rope_w_extrapolation_ratio=3.0,
|
| 343 |
+
rope_t_extrapolation_ratio=float(state_t) / 24.0,
|
| 344 |
+
timestep_scale=0.001,
|
| 345 |
+
sac_config=dict(mode="predict2_2b_720"),
|
| 346 |
+
use_crossattn_projection=False,
|
| 347 |
+
crossattn_emb_channels=1024,
|
| 348 |
+
use_wan_fp32_strategy=True,
|
| 349 |
+
warped_latent_channels=16,
|
| 350 |
+
visibility_channels=1,
|
| 351 |
+
pose_in_channels=3,
|
| 352 |
+
pose_hidden_dim=16,
|
| 353 |
+
),
|
| 354 |
+
conditioner=dict(
|
| 355 |
+
use_video_condition=dict(dropout_rate=0.0),
|
| 356 |
+
text=dict(dropout_rate=0.0, use_empty_string=False),
|
| 357 |
+
),
|
| 358 |
+
tokenizer=dict(vae_pth=_local_wan_vae_path(), temporal_window=16),
|
| 359 |
+
text_encoder_class="reason1p1_7B",
|
| 360 |
+
text_encoder_config=None,
|
| 361 |
+
),
|
| 362 |
+
),
|
| 363 |
+
)
|
| 364 |
+
|
| 365 |
+
|
| 366 |
+
def nymeria_pose_2b_stage1_singleview() -> dict:
|
| 367 |
+
"""Stage 1: single-view (V=1) pose-encoder pretraining, 2-GPU FSDP.
|
| 368 |
+
|
| 369 |
+
Each actor is an independent V=1 sample (simpler + 2x data + faster iters). Warm-starts from base 2B.
|
| 370 |
+
The pose encoder uses a SMALL-RANDOM final-proj init (pose_proj_init_std>0) instead of zero-init so its
|
| 371 |
+
conv feature extractor gets gradient from step 0 (zero-init starves the convs until the proj grows, which
|
| 372 |
+
is a key cause of slow pose learning). Warped conditioning is kept as-is. After this warms up the pose
|
| 373 |
+
encoder, fine-tune the joint 2-actor model (nymeria_pose_2actor_2b_longer) warm-started from this stage.
|
| 374 |
+
"""
|
| 375 |
+
return dict(
|
| 376 |
+
defaults=[
|
| 377 |
+
"/experiment/nymeria_pose_2actor_2b_longer",
|
| 378 |
+
{"override /data_train": "nymeria_single_view"},
|
| 379 |
+
{"override /data_val": "nymeria_single_view"},
|
| 380 |
+
"_self_",
|
| 381 |
+
],
|
| 382 |
+
job=dict(group="", name="nymeria_pose_2b_stage1_singleview", project="train"),
|
| 383 |
+
model=dict(
|
| 384 |
+
config=dict(
|
| 385 |
+
fsdp_shard_size=2, # 2-GPU run
|
| 386 |
+
net=dict(
|
| 387 |
+
pose_proj_init_std=0.02, # small-random pose proj init -> faster pose learning
|
| 388 |
+
freeze_view_embedding=True, # view embedding is meaningless with V=1 -> don't train it
|
| 389 |
+
),
|
| 390 |
+
)
|
| 391 |
+
),
|
| 392 |
+
trainer=dict(
|
| 393 |
+
callbacks=dict(
|
| 394 |
+
validation_viz=L(NymeriaValidationViz)(
|
| 395 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 396 |
+
num_video_frames=77, single_view=True,
|
| 397 |
+
),
|
| 398 |
+
),
|
| 399 |
+
),
|
| 400 |
+
)
|
| 401 |
+
|
| 402 |
+
|
| 403 |
+
def nymeria_pose_2b_stage1_sv_vaeconcat() -> dict:
|
| 404 |
+
"""Stage-1 single-view, VAE pose encoder + CHANNEL CONCAT (Setting 1).
|
| 405 |
+
|
| 406 |
+
The pixel pose RGB is passed through the FROZEN Wan2.1 VAE -> a 16ch latent-resolution pose latent, which
|
| 407 |
+
is CHANNEL-CONCATENATED with the warped latent + visibility mask into the single zero-init cond_embedder
|
| 408 |
+
(one linear patch projection mixes all conditioning latents). The Conv3d PoseEncoder is NOT built.
|
| 409 |
+
Warm-starts from base 2B; the cond_embedder loads as a (zero-init) missing key. Run on GPU 0,1.
|
| 410 |
+
"""
|
| 411 |
+
return dict(
|
| 412 |
+
defaults=[
|
| 413 |
+
"/experiment/nymeria_pose_2b_stage1_singleview",
|
| 414 |
+
{"override /conditioner": "video_prediction_multiview_pose_latent_conditioner"},
|
| 415 |
+
"_self_",
|
| 416 |
+
],
|
| 417 |
+
job=dict(group="", name="nymeria_pose_2b_stage1_sv_vaeconcat", project="train"),
|
| 418 |
+
model=dict(
|
| 419 |
+
config=dict(
|
| 420 |
+
pose_via_vae=True, # preprocess: pose RGB -> frozen VAE -> pose_latent
|
| 421 |
+
net=dict(pose_mode="vae_concat"),
|
| 422 |
+
)
|
| 423 |
+
),
|
| 424 |
+
)
|
| 425 |
+
|
| 426 |
+
|
| 427 |
+
def nymeria_pose_2b_stage1_sv_vaemlp() -> dict:
|
| 428 |
+
"""Stage-1 single-view, VAE pose encoder + trainable MLP + ADD (Setting 2).
|
| 429 |
+
|
| 430 |
+
Pixel pose RGB -> FROZEN Wan2.1 VAE -> 16ch pose latent -> trainable PatchEmbed + MLP -> ADDED to the
|
| 431 |
+
patch tokens (MLP last layer zero-init -> day-0 == base). The Conv3d PoseEncoder is NOT built; the VAE
|
| 432 |
+
stays frozen and only the small pose MLP/embedder (+ backbone) train. Run on GPU 2,3.
|
| 433 |
+
"""
|
| 434 |
+
return dict(
|
| 435 |
+
defaults=[
|
| 436 |
+
"/experiment/nymeria_pose_2b_stage1_singleview",
|
| 437 |
+
{"override /conditioner": "video_prediction_multiview_pose_latent_conditioner"},
|
| 438 |
+
"_self_",
|
| 439 |
+
],
|
| 440 |
+
job=dict(group="", name="nymeria_pose_2b_stage1_sv_vaemlp", project="train"),
|
| 441 |
+
model=dict(
|
| 442 |
+
config=dict(
|
| 443 |
+
pose_via_vae=True,
|
| 444 |
+
net=dict(pose_mode="vae_mlp_add"),
|
| 445 |
+
)
|
| 446 |
+
),
|
| 447 |
+
)
|
| 448 |
+
|
| 449 |
+
|
| 450 |
+
def nymeria_pose_2view_actorobs() -> dict:
|
| 451 |
+
"""2-view actor-observer joint generation: vae_concat pose + REAL per-view text (reason1), 4-GPU FSDP.
|
| 452 |
+
|
| 453 |
+
- Data: /data2/nymeria_processed_single actor+observer pairs (train_split/val_split.csv), role-mix 50% so
|
| 454 |
+
the view embedding can't memorize actor/observer roles.
|
| 455 |
+
- Text: HF Cosmos-Reason1-7B (FULL_CONCAT 100352) online, per-view captions (actor caption on the actor
|
| 456 |
+
view; "C is observing the partner. The partner ..." on the observer view). crossattn projection restored
|
| 457 |
+
(100352->1024) -> warm-starts from base 2B.
|
| 458 |
+
- Pose: vae_concat (pose RGB -> frozen VAE -> pose_latent, channel-concat with warped+vis in cond_embedder).
|
| 459 |
+
Uses the pose_latent conditioner. View embedding TRAINABLE (role-mix handles actor/observer symmetry).
|
| 460 |
+
- Warm-start: base 2B multiview. Requires env COSMOS_QWEN_TOKENIZER_DIR=/data/cosmos_reason1_7b.
|
| 461 |
+
"""
|
| 462 |
+
state_t = 20
|
| 463 |
+
return dict(
|
| 464 |
+
defaults=[
|
| 465 |
+
{"override /ckpt_type": "dcp"},
|
| 466 |
+
{"override /optimizer": "adamw"},
|
| 467 |
+
{"override /callbacks": ["basic", "wandb"]},
|
| 468 |
+
{"override /checkpoint": "s3"},
|
| 469 |
+
{"override /tokenizer": "wan2pt1_tokenizer"},
|
| 470 |
+
{"override /data_train": "nymeria_actor_observer"},
|
| 471 |
+
{"override /data_val": "nymeria_actor_observer"},
|
| 472 |
+
{"override /conditioner": "video_prediction_multiview_pose_latent_conditioner"},
|
| 473 |
+
{"override /model": "fsdp_rectified_flow_multiview_pose"},
|
| 474 |
+
{"override /net": "cosmos_v1_2B_multiview_pose"},
|
| 475 |
+
"_self_",
|
| 476 |
+
],
|
| 477 |
+
job=dict(group="", name="nymeria_pose_2view_actorobs", project="train"),
|
| 478 |
+
optimizer=dict(lr=3e-5, weight_decay=1e-3, betas=[0.9, 0.999]),
|
| 479 |
+
scheduler=dict(f_max=[0.99], f_min=[0.4], warm_up_steps=[100], cycle_lengths=[400_000]),
|
| 480 |
+
checkpoint=dict(
|
| 481 |
+
load_from_object_store=dict(enabled=False),
|
| 482 |
+
save_to_object_store=dict(enabled=False),
|
| 483 |
+
save_iter=500,
|
| 484 |
+
load_path=_base_2b_multiview_ckpt(),
|
| 485 |
+
strict_resume=False,
|
| 486 |
+
),
|
| 487 |
+
trainer=dict(
|
| 488 |
+
max_iter=400_000,
|
| 489 |
+
logging_iter=50,
|
| 490 |
+
callbacks=dict(
|
| 491 |
+
validation_viz=L(NymeriaValidationViz)(
|
| 492 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 493 |
+
num_video_frames=77, root="/data2/nymeria_processed_single", actor_observer=True,
|
| 494 |
+
),
|
| 495 |
+
),
|
| 496 |
+
),
|
| 497 |
+
model_parallel=dict(context_parallel_size=1),
|
| 498 |
+
model=dict(
|
| 499 |
+
config=dict(
|
| 500 |
+
min_num_conditional_frames_per_view=1,
|
| 501 |
+
max_num_conditional_frames_per_view=1,
|
| 502 |
+
conditional_frames_probs={1: 1.0},
|
| 503 |
+
condition_locations=["first_random_n"],
|
| 504 |
+
fsdp_shard_size=4,
|
| 505 |
+
resolution="480",
|
| 506 |
+
state_t=state_t,
|
| 507 |
+
shift=5,
|
| 508 |
+
use_dynamic_shift=False,
|
| 509 |
+
train_time_weight="uniform",
|
| 510 |
+
train_time_distribution="logitnormal",
|
| 511 |
+
online_text_embeddings_as_dict=False,
|
| 512 |
+
pose_via_vae=True, # pose RGB -> frozen VAE -> pose_latent (for vae_concat)
|
| 513 |
+
net=dict(
|
| 514 |
+
concat_view_embedding=True,
|
| 515 |
+
view_condition_dim=7,
|
| 516 |
+
state_t=state_t,
|
| 517 |
+
n_cameras_emb=7,
|
| 518 |
+
rope_enable_fps_modulation=False,
|
| 519 |
+
rope_h_extrapolation_ratio=3.0,
|
| 520 |
+
rope_w_extrapolation_ratio=3.0,
|
| 521 |
+
rope_t_extrapolation_ratio=float(state_t) / 24.0,
|
| 522 |
+
timestep_scale=0.001,
|
| 523 |
+
sac_config=dict(mode="predict2_2b_720"),
|
| 524 |
+
use_crossattn_projection=True,
|
| 525 |
+
crossattn_proj_in_channels=100352,
|
| 526 |
+
crossattn_emb_channels=1024,
|
| 527 |
+
use_wan_fp32_strategy=True,
|
| 528 |
+
warped_latent_channels=16,
|
| 529 |
+
visibility_channels=1,
|
| 530 |
+
pose_in_channels=3,
|
| 531 |
+
pose_hidden_dim=16,
|
| 532 |
+
pose_mode="vae_concat",
|
| 533 |
+
),
|
| 534 |
+
conditioner=dict(
|
| 535 |
+
use_video_condition=dict(dropout_rate=0.0),
|
| 536 |
+
text=dict(dropout_rate=0.2, use_empty_string=False),
|
| 537 |
+
),
|
| 538 |
+
tokenizer=dict(vae_pth=_local_wan_vae_path(), temporal_window=16),
|
| 539 |
+
text_encoder_class="reason1p1_7B",
|
| 540 |
+
text_encoder_config=dict(
|
| 541 |
+
embedding_concat_strategy=str(EmbeddingConcatStrategy.FULL_CONCAT),
|
| 542 |
+
compute_online=True,
|
| 543 |
+
ckpt_path="/data/cosmos_reason1_7b",
|
| 544 |
+
),
|
| 545 |
+
),
|
| 546 |
+
),
|
| 547 |
+
)
|
| 548 |
+
|
| 549 |
+
|
| 550 |
+
def nymeria_pose_2view_actorobs_refs() -> dict:
|
| 551 |
+
"""nymeria_pose_2view_actorobs + in-context REFERENCE-FRAME appearance conditioning.
|
| 552 |
+
|
| 553 |
+
Appends R=6 clean source frames per view as extra self-attention tokens at fixed non-contiguous temporal
|
| 554 |
+
RoPE positions [50,55,60,65,70,75], to sharpen appearance detail (complements the geometry-only warped_cond).
|
| 555 |
+
Only new parameter is a zero-init ref gate -> day-0 output ~= base; warm-starts from base 2B. Fresh run.
|
| 556 |
+
"""
|
| 557 |
+
cfg = nymeria_pose_2view_actorobs()
|
| 558 |
+
cfg["job"]["name"] = "nymeria_pose_2view_actorobs_refs"
|
| 559 |
+
# dataset variants that emit `reference_frames` (R=6 clean source frames/view)
|
| 560 |
+
for d in cfg["defaults"]:
|
| 561 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 562 |
+
d["override /data_train"] = "nymeria_actor_observer_refs"
|
| 563 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 564 |
+
d["override /data_val"] = "nymeria_actor_observer_refs"
|
| 565 |
+
# model: VAE-encode the reference frames
|
| 566 |
+
cfg["model"]["config"]["num_reference_frames"] = 4
|
| 567 |
+
# net: enable the in-context reference tokens
|
| 568 |
+
cfg["model"]["config"]["net"].update(
|
| 569 |
+
enable_reference_frames=True, num_reference_frames=4, ref_rope_offset=50, ref_rope_stride=5,
|
| 570 |
+
)
|
| 571 |
+
# validation viz must also emit reference frames so generation sees them
|
| 572 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 573 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 574 |
+
num_video_frames=77, root="/data2/nymeria_processed_single", actor_observer=True,
|
| 575 |
+
num_reference_frames=4,
|
| 576 |
+
)
|
| 577 |
+
return cfg
|
| 578 |
+
|
| 579 |
+
|
| 580 |
+
def nymeria_pose_2view_actorobs_refs_campose() -> dict:
|
| 581 |
+
"""refs (R=4) + per-pixel Plücker camera-ray conditioning, for cross-view shared-space awareness.
|
| 582 |
+
|
| 583 |
+
Adds a 6-ch Plücker ray map (dir+moment, per-pair canonical frame = view0 frame0) via a zero-init PatchEmbed
|
| 584 |
+
added to the tokens (same additive pattern as warped/pose; day-0 == base). Targets the two views not
|
| 585 |
+
recognizing they're the same space / cross-view consistency. Fresh from base 2B.
|
| 586 |
+
"""
|
| 587 |
+
cfg = nymeria_pose_2view_actorobs_refs() # inherits refs (R=4) config
|
| 588 |
+
cfg["job"]["name"] = "nymeria_pose_2view_actorobs_refs_campose"
|
| 589 |
+
# dataloaders that also emit camera poses (campose.npz)
|
| 590 |
+
for d in cfg["defaults"]:
|
| 591 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 592 |
+
d["override /data_train"] = "nymeria_actor_observer_refs_campose"
|
| 593 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 594 |
+
d["override /data_val"] = "nymeria_actor_observer_refs_campose"
|
| 595 |
+
# model: compute Plücker maps
|
| 596 |
+
cfg["model"]["config"]["enable_plucker"] = True
|
| 597 |
+
# net: enable the Plücker embedder
|
| 598 |
+
cfg["model"]["config"]["net"].update(enable_plucker=True, plucker_channels=6)
|
| 599 |
+
# validation viz must emit camera poses too
|
| 600 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 601 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 602 |
+
num_video_frames=77, root="/data2/nymeria_processed_single", actor_observer=True,
|
| 603 |
+
num_reference_frames=4, emit_camera_poses=True,
|
| 604 |
+
)
|
| 605 |
+
return cfg
|
| 606 |
+
|
| 607 |
+
|
| 608 |
+
def nymeria_pose_2view_actorobs_campose() -> dict:
|
| 609 |
+
"""Ablation of `..._refs_campose` with the REFERENCE-frame condition REMOVED: Plücker camera conditioning
|
| 610 |
+
only (no in-context reference frames). A/B partner to isolate what the reference frames contribute vs the
|
| 611 |
+
Plücker shared-space signal. Same split/val, fresh from base 2B."""
|
| 612 |
+
cfg = nymeria_pose_2view_actorobs_refs_campose()
|
| 613 |
+
cfg["job"]["name"] = "nymeria_pose_2view_actorobs_campose"
|
| 614 |
+
# camera-pose-only dataloaders (no reference frames)
|
| 615 |
+
for d in cfg["defaults"]:
|
| 616 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 617 |
+
d["override /data_train"] = "nymeria_actor_observer_campose"
|
| 618 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 619 |
+
d["override /data_val"] = "nymeria_actor_observer_campose"
|
| 620 |
+
# disable reference frames (keep Plücker on)
|
| 621 |
+
cfg["model"]["config"]["num_reference_frames"] = 0
|
| 622 |
+
cfg["model"]["config"]["net"].update(enable_reference_frames=False, num_reference_frames=0)
|
| 623 |
+
# validation viz: no reference frames, keep camera poses
|
| 624 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 625 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 626 |
+
num_video_frames=77, root="/data2/nymeria_processed_single", actor_observer=True,
|
| 627 |
+
num_reference_frames=0, emit_camera_poses=True,
|
| 628 |
+
)
|
| 629 |
+
return cfg
|
| 630 |
+
|
| 631 |
+
|
| 632 |
+
def nymeria_pose_2view_actorobs_refs_campose_refpose() -> dict:
|
| 633 |
+
"""`..._refs_campose` + POSED references: the in-context reference frames also get a Plücker ray map from
|
| 634 |
+
their OWN past camera poses (in the same per-pair canonical frame), so references are geometrically grounded
|
| 635 |
+
(posed reference views), not just floating appearance. Fresh from base 2B."""
|
| 636 |
+
cfg = nymeria_pose_2view_actorobs_refs_campose()
|
| 637 |
+
cfg["job"]["name"] = "nymeria_pose_2view_actorobs_refs_campose_refpose"
|
| 638 |
+
cfg["model"]["config"]["enable_reference_plucker"] = True
|
| 639 |
+
cfg["model"]["config"]["net"].update(enable_reference_plucker=True)
|
| 640 |
+
# (dataloaders already emit reference frames + camera poses; the dataset also emits reference_cam_w2c)
|
| 641 |
+
return cfg
|
| 642 |
+
|
| 643 |
+
|
| 644 |
+
def nymeria_pose_2view_actorobs_refs_campose_refpose_shared() -> dict:
|
| 645 |
+
"""`..._refs_campose_refpose` but with SHARED references: instead of per-view reference pools, one greedy
|
| 646 |
+
set-cover selection over the COMBINED pool (both views' past) is fed to both view slots; the per-view
|
| 647 |
+
view-embedding on reference tokens is DROPPED (shared_reference=True). Refs are geometrically identified
|
| 648 |
+
only by their posed Plücker; both views attend all of them via cross-view self-attention. Uses
|
| 649 |
+
refs_shared.npz (backfill_shared_refs.py). Fresh from base 2B."""
|
| 650 |
+
cfg = nymeria_pose_2view_actorobs_refs_campose_refpose()
|
| 651 |
+
cfg["job"]["name"] = "nymeria_pose_2view_actorobs_refs_campose_refpose_shared"
|
| 652 |
+
# net + (net-nested) share flag: drop the ref view-embedding
|
| 653 |
+
cfg["model"]["config"]["net"].update(shared_reference=True)
|
| 654 |
+
# shared-reference dataloaders (refs_shared.npz)
|
| 655 |
+
for d in cfg["defaults"]:
|
| 656 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 657 |
+
d["override /data_train"] = "nymeria_actor_observer_refs_campose_shared"
|
| 658 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 659 |
+
d["override /data_val"] = "nymeria_actor_observer_refs_campose_shared"
|
| 660 |
+
# validation viz: shared refs on the actor-observer val set
|
| 661 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 662 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 663 |
+
num_video_frames=77, root="/data2/nymeria_processed_single", actor_observer=True,
|
| 664 |
+
num_reference_frames=4, shared_reference=True, emit_camera_poses=True,
|
| 665 |
+
)
|
| 666 |
+
return cfg
|
| 667 |
+
|
| 668 |
+
|
| 669 |
+
def nymeria_pose_2view_actoractor_refs_campose_refpose() -> dict:
|
| 670 |
+
"""Same structure as `..._refs_campose_refpose` (refs + Plücker + posed refs) but on the ACTOR-ACTOR longer
|
| 671 |
+
dataset (/data2/nymeria_processed_longer): two ego actors generated jointly, each with its own real action
|
| 672 |
+
caption. Fresh from base 2B."""
|
| 673 |
+
cfg = nymeria_pose_2view_actorobs_refs_campose_refpose()
|
| 674 |
+
cfg["job"]["name"] = "nymeria_pose_2view_actoractor_refs_campose_refpose"
|
| 675 |
+
for d in cfg["defaults"]:
|
| 676 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 677 |
+
d["override /data_train"] = "nymeria_actor_actor_refs_campose"
|
| 678 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 679 |
+
d["override /data_val"] = "nymeria_actor_actor_refs_campose"
|
| 680 |
+
# validation viz on the actor-actor longer val set
|
| 681 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 682 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 683 |
+
num_video_frames=77, root="/data2/nymeria_processed_longer", actor_actor=True,
|
| 684 |
+
num_reference_frames=4, emit_camera_poses=True,
|
| 685 |
+
)
|
| 686 |
+
return cfg
|
| 687 |
+
|
| 688 |
+
|
| 689 |
+
# fixed viz samples (real longer): 5 val (Loc_46+Loc_BX held-out) + 5 train (6 train locations) -> stable
|
| 690 |
+
# wandb comparison across all actor-actor runs regardless of manifest size. See manifest/viz_fixed_ids.json.
|
| 691 |
+
_AA_VAL_FIXED_IDS = [
|
| 692 |
+
"paul_act1_t2fpwg_003927__justin_act1_q1vpto_003941",
|
| 693 |
+
"paul_act2_pryhwf_031416__justin_act2_bw0jz0_031395",
|
| 694 |
+
"patricia_act2_c993gv_003234__greg_act2_jc6wnc_003243",
|
| 695 |
+
"patricia_act4_ifc9zd_005775__greg_act4_chhp0x_004682",
|
| 696 |
+
"patricia_act5_swxtds_038808__greg_act5_hu4348_040953",
|
| 697 |
+
]
|
| 698 |
+
_AA_TRAIN_FIXED_IDS = [
|
| 699 |
+
"carly_act0_et4azi_025410__andrew_act0_us908q_025439",
|
| 700 |
+
"sarah_act0_v24gb2_034881__alexis_act0_0e9t84_034200",
|
| 701 |
+
"jodi_act4_iplp4x_018018__katie_act4_ie44mj_017665",
|
| 702 |
+
"shawn_act7_onxbbz_019404__brenda_act7_mnqj9y_019417",
|
| 703 |
+
"randall_act2_nwllyf_016632__erica_act2_nbdcde_016609",
|
| 704 |
+
]
|
| 705 |
+
|
| 706 |
+
|
| 707 |
+
def nymeria_pose_2view_actoractor_refs_campose_refpose_shared() -> dict:
|
| 708 |
+
"""SHARED-reference ACTOR-ACTOR, stage-2 of the stagewise plan: warm-start from the actor-observer SHARED
|
| 709 |
+
checkpoint, then fine-tune on the REAL longer (location-holdout) + SYNTHETIC combined set. Both views are
|
| 710 |
+
ego actors with their own captions; shared refs (no view-embedding) + Plücker + posed refs. Validation runs
|
| 711 |
+
on the REAL longer location-holdout val with FIXED viz pair_ids (stable wandb comparison)."""
|
| 712 |
+
cfg = nymeria_pose_2view_actorobs_refs_campose_refpose_shared() # shared net + refs+plucker+refpose
|
| 713 |
+
cfg["job"]["name"] = "nymeria_pose_2view_actoractor_refs_campose_refpose_shared"
|
| 714 |
+
# warm-start from the actor-observer SHARED checkpoint (model-only staged dir -> fresh optimizer, iter 0)
|
| 715 |
+
cfg["checkpoint"]["load_path"] = "/data/model_output/warmstart/actorobs_shared_iter5500"
|
| 716 |
+
cfg["checkpoint"]["strict_resume"] = False
|
| 717 |
+
# data: REAL longer + SYNTHETIC combined (train); val = REAL longer location-holdout only
|
| 718 |
+
for d in cfg["defaults"]:
|
| 719 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 720 |
+
d["override /data_train"] = "nymeria_actor_actor_refs_campose_shared_plus_synth"
|
| 721 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 722 |
+
d["override /data_val"] = "nymeria_actor_actor_refs_campose_shared"
|
| 723 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 724 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 725 |
+
num_video_frames=77, root="/data2/nymeria_processed_longer", actor_actor=True,
|
| 726 |
+
num_reference_frames=4, shared_reference=True, emit_camera_poses=True,
|
| 727 |
+
train_fixed_ids=_AA_TRAIN_FIXED_IDS, val_fixed_ids=_AA_VAL_FIXED_IDS,
|
| 728 |
+
)
|
| 729 |
+
return cfg
|
| 730 |
+
|
| 731 |
+
|
| 732 |
+
def nymeria_pose_2b_stage1_lora() -> dict:
|
| 733 |
+
"""STAGE-1 of the motion curriculum: single-view (V=1) LoRA pretraining on the MAX pooled ego-clip set
|
| 734 |
+
(real single actor + longer + synthetic, ~41k clips). Full conditioning minus cross-view: real per-clip
|
| 735 |
+
caption + pose + warping + plucker(self-frame0) + OWN refs (R=4). Backbone = LoRA(rank 64, prior-preserving);
|
| 736 |
+
new zero-init conditioning modules (pose/plucker/ref embedders, gates) stay fully trainable. Fresh from base
|
| 737 |
+
2B; warm-starts the 2-view stages 2/3. 2-GPU."""
|
| 738 |
+
cfg = nymeria_pose_2view_actorobs_refs_campose_refpose() # refs(4)+plucker+refpose, per-view (own refs)
|
| 739 |
+
cfg["job"]["name"] = "nymeria_pose_2b_stage1_lora"
|
| 740 |
+
cfg["checkpoint"]["load_path"] = _base_2b_multiview_ckpt() # start of the curriculum chain
|
| 741 |
+
cfg["checkpoint"]["strict_resume"] = False
|
| 742 |
+
for d in cfg["defaults"]:
|
| 743 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 744 |
+
d["override /data_train"] = "nymeria_single_view_stage1"
|
| 745 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 746 |
+
d["override /data_val"] = "nymeria_single_view_stage1"
|
| 747 |
+
mc = cfg["model"]["config"]
|
| 748 |
+
mc["fsdp_shard_size"] = 2
|
| 749 |
+
mc["use_lora"] = True
|
| 750 |
+
mc["lora_rank"] = 64
|
| 751 |
+
mc["lora_alpha"] = 64
|
| 752 |
+
mc["net"].update(freeze_view_embedding=True) # V=1 -> view embedding meaningless
|
| 753 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 754 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 755 |
+
num_video_frames=77, single_view_stage1=True,
|
| 756 |
+
num_reference_frames=4, shared_reference=False, emit_camera_poses=True,
|
| 757 |
+
)
|
| 758 |
+
return cfg
|
| 759 |
+
|
| 760 |
+
|
| 761 |
+
def nymeria_pose_2view_actoractor_refs_campose_refpose_shared_synthonly() -> dict:
|
| 762 |
+
"""DIAGNOSTIC: SYNTHETIC-ONLY actor-actor (fresh from base 2B, no real-data influence). Isolates whether the
|
| 763 |
+
black/broken human appearance seen in synthetic val of the real+synth run comes from the REAL dataset
|
| 764 |
+
(masking prior) or is inherent to the synthetic data/masking. Val viz on synthetic val (no real fixed ids)."""
|
| 765 |
+
cfg = nymeria_pose_2view_actoractor_refs_campose_refpose_shared()
|
| 766 |
+
cfg["job"]["name"] = "nymeria_pose_2view_actoractor_refs_campose_refpose_shared_synthonly"
|
| 767 |
+
cfg["checkpoint"]["load_path"] = _base_2b_multiview_ckpt() # fresh from base -> zero real-data influence
|
| 768 |
+
for d in cfg["defaults"]:
|
| 769 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 770 |
+
d["override /data_train"] = "nymeria_actor_actor_refs_campose_shared_synthonly"
|
| 771 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 772 |
+
d["override /data_val"] = "nymeria_actor_actor_refs_campose_shared_synthonly"
|
| 773 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 774 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 775 |
+
num_video_frames=77, root="/data3/synthetic_processed_multi", actor_actor=True,
|
| 776 |
+
num_reference_frames=4, shared_reference=True, emit_camera_poses=True,
|
| 777 |
+
train_fixed_ids=None, val_fixed_ids=None, # synthetic viz (real longer fixed ids don't apply)
|
| 778 |
+
)
|
| 779 |
+
return cfg
|
| 780 |
+
|
| 781 |
+
|
| 782 |
+
# CoMind val handover/collaboration clips (from val holdout recs 1326e688/313483e1) — pinned so every validation
|
| 783 |
+
# shows whether the model generates person-to-person object transfer. Verified loadable (refs+campose present).
|
| 784 |
+
_COMIND_HANDOVER_VAL_IDS = [
|
| 785 |
+
"313483e1_000831", # partner washes zucchini -> HANDS IT TO the wearer, who dries it
|
| 786 |
+
"1326e688_003624", # partner rinses carrot; wearer's hands enter to TAKE THE CARROT FROM the partner
|
| 787 |
+
"1326e688_052008", # partner washes dishes -> HANDS THE WEARER A PAIR OF CHOPSTICKS
|
| 788 |
+
"1326e688_040200", # partner stirs pan; wearer takes ... (near GT handover: bowl @frame 40292)
|
| 789 |
+
"1326e688_028392", # partner holds bowl+spoon, places in sink; wearer reaches (bowl transfer)
|
| 790 |
+
"1326e688_055608", # pan handling near GT handover (pan @frame 55781)
|
| 791 |
+
]
|
| 792 |
+
|
| 793 |
+
|
| 794 |
+
def comind_actoractor_personpose_shared() -> dict:
|
| 795 |
+
"""CoMind (Aria 2-person kitchen) 2-view actor-actor. Dataset = CoMind ONLY (32 multislam shared-world recs,
|
| 796 |
+
~12k pairs). Skeleton pose is IDENTITY-colored (person_pose: leader=EGO / helper=PARTNER, view-invariant),
|
| 797 |
+
NOT wearer/partner role-colored. Conditioning = combined-pool warp RGB + person pose + GREEDY shared refs +
|
| 798 |
+
cross-view Plücker (campose/refpose). DEPTH channels are stored in clip.npz but NOT loaded (no depth condition
|
| 799 |
+
for now). Warm-start from the Nymeria actor-observer SHARED checkpoint. Val = CoMind held-out recs."""
|
| 800 |
+
cfg = nymeria_pose_2view_actoractor_refs_campose_refpose_shared() # shared net + refs+plucker+refpose
|
| 801 |
+
cfg["job"]["name"] = "comind_actoractor_personpose_shared"
|
| 802 |
+
cfg["checkpoint"]["load_path"] = "/data/model_output/warmstart/actorobs_shared_iter5500" # nymeria actor-observer
|
| 803 |
+
cfg["checkpoint"]["strict_resume"] = False
|
| 804 |
+
for d in cfg["defaults"]:
|
| 805 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 806 |
+
d["override /data_train"] = "comind_actoractor_personpose_shared"
|
| 807 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 808 |
+
d["override /data_val"] = "comind_actoractor_personpose_shared"
|
| 809 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 810 |
+
every_n=500, num_train=5, num_val=len(_COMIND_HANDOVER_VAL_IDS), num_sampling_step=35, guidance=7.0, fps=10,
|
| 811 |
+
num_video_frames=77, root="/data4/comind_dataset", comind=True,
|
| 812 |
+
num_reference_frames=4, shared_reference=True, emit_camera_poses=True,
|
| 813 |
+
train_fixed_ids=None, val_fixed_ids=_COMIND_HANDOVER_VAL_IDS, # pin handover/collab clips
|
| 814 |
+
)
|
| 815 |
+
return cfg
|
| 816 |
+
|
| 817 |
+
|
| 818 |
+
def nymeria_synth_actoractor_depth_shared() -> dict:
|
| 819 |
+
"""Nymeria longer + SYNTHETIC actor-actor + composite DEPTH condition. Depth = warped scene depth + human
|
| 820 |
+
mesh depth (mesh_cond depth_comp, RGB) -> frozen VAE -> depth_latent(16) -> zero-init depth_embedder ADDED to
|
| 821 |
+
tokens (== VAE-channel-concat, warm-start clean). Same net (refs+plucker+refpose) as the actor-actor shared
|
| 822 |
+
run; warm-start from the Nymeria actor-observer SHARED checkpoint. Val = nymeria longer holdout.
|
| 823 |
+
NOTE: synthetic depth is 100% ready; nymeria longer mesh Stage-B is in progress -> missing clips get ZERO
|
| 824 |
+
depth (coverage grows as Stage-B completes)."""
|
| 825 |
+
cfg = nymeria_pose_2view_actoractor_refs_campose_refpose_shared() # refs+plucker+refpose, warm-start actorobs
|
| 826 |
+
cfg["job"]["name"] = "nymeria_synth_actoractor_depth_shared"
|
| 827 |
+
# warm-start from the DEPTH actor-observer run (nymeria_actorobs_depth_shared iter 5500) — pose+refs+plucker
|
| 828 |
+
# +DEPTH embedder are ALL already trained there (net matches exactly), so this continues from a depth-aware
|
| 829 |
+
# actor-observer rather than the depth-less actorobs_shared. Model-only staged dir -> fresh optimizer.
|
| 830 |
+
cfg["checkpoint"]["load_path"] = "/data/model_output/warmstart/actorobs_depth_iter5500"
|
| 831 |
+
cfg["checkpoint"]["strict_resume"] = False
|
| 832 |
+
cfg["model"]["config"]["depth_via_vae"] = True # RGB depth -> frozen VAE -> depth_latent
|
| 833 |
+
cfg["model"]["config"]["net"].update(enable_depth=True) # zero-init depth_embedder (additive)
|
| 834 |
+
cfg["trainer"]["callbacks"]["validation_viz"]["emit_depth"] = True # add a "depth" column to the viz grid
|
| 835 |
+
for d in cfg["defaults"]:
|
| 836 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 837 |
+
d["override /data_train"] = "nymeria_actor_actor_refs_campose_shared_plus_synth_depth"
|
| 838 |
+
return cfg
|
| 839 |
+
|
| 840 |
+
|
| 841 |
+
def nymeria_actorobs_depth_shared() -> dict:
|
| 842 |
+
"""Nymeria ACTOR-OBSERVER (single) + composite DEPTH condition — the stage-1 depth run (train this FIRST,
|
| 843 |
+
then warm-start the actor-actor depth run from it). Actor-observer SHARED model (refs+plucker+refpose) with
|
| 844 |
+
warped-scene + human-mesh depth (mesh_cond depth_comp) -> frozen VAE -> depth_latent(16) -> zero-init
|
| 845 |
+
depth_embedder ADDED to tokens (== VAE-channel-concat, day-0 contribution 0 so warm-start is clean).
|
| 846 |
+
Warm-start from the trained actor-observer SHARED checkpoint (actorobs_shared_iter5500) so only the depth
|
| 847 |
+
branch is new. single depth (mesh Stage-B) ~100% ready. Val = nymeria single holdout + a depth viz column."""
|
| 848 |
+
cfg = nymeria_pose_2view_actorobs_refs_campose_refpose_shared() # actor-observer shared: refs+plucker+refpose
|
| 849 |
+
cfg["job"]["name"] = "nymeria_actorobs_depth_shared"
|
| 850 |
+
# FRESH from Cosmos base 2B (NOT warm-started from a trained actor-observer ckpt): the pose/refs/plucker/
|
| 851 |
+
# refpose/depth modules are all new (zero-init or missing keys) and train from scratch on top of the base
|
| 852 |
+
# video backbone. So iter-1 loss starts HIGH (~base video prior), unlike a warm-start.
|
| 853 |
+
cfg["checkpoint"]["load_path"] = _base_2b_multiview_ckpt()
|
| 854 |
+
cfg["checkpoint"]["strict_resume"] = False
|
| 855 |
+
cfg["model"]["config"]["depth_via_vae"] = True # RGB depth -> frozen VAE -> depth_latent
|
| 856 |
+
cfg["model"]["config"]["net"].update(enable_depth=True) # zero-init depth_embedder (additive)
|
| 857 |
+
cfg["trainer"]["callbacks"]["validation_viz"]["emit_depth"] = True # add a "depth" column to the viz grid
|
| 858 |
+
for d in cfg["defaults"]:
|
| 859 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 860 |
+
d["override /data_train"] = "nymeria_actor_observer_refs_campose_shared_depth"
|
| 861 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 862 |
+
d["override /data_val"] = "nymeria_actor_observer_refs_campose_shared_depth"
|
| 863 |
+
return cfg
|
| 864 |
+
|
| 865 |
+
|
| 866 |
+
def comind_synth_actoractor_personpose_shared() -> dict:
|
| 867 |
+
"""CoMind + SYNTHETIC combined 2-view actor-actor (both person-pose). Train = CoMind (32 shared-world recs,
|
| 868 |
+
~10.8k) + synthetic (/data3/synthetic_processed_multi, ~17.4k) pooled. Validation stays on the CoMind val
|
| 869 |
+
(handover/collab clips pinned, same as comind_actoractor_personpose_shared). Warm-start from the Nymeria
|
| 870 |
+
actor-observer SHARED checkpoint. Depth NOT loaded."""
|
| 871 |
+
cfg = comind_actoractor_personpose_shared() # comind net/conditioning/viz + actor-observer warm-start
|
| 872 |
+
cfg["job"]["name"] = "comind_synth_actoractor_personpose_shared"
|
| 873 |
+
for d in cfg["defaults"]:
|
| 874 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 875 |
+
d["override /data_train"] = "comind_synth_actoractor_personpose_shared" # comind + synthetic
|
| 876 |
+
# data_val stays "comind_actoractor_personpose_shared" (CoMind val) — set by the parent
|
| 877 |
+
return cfg
|
| 878 |
+
|
| 879 |
+
|
| 880 |
+
def vroid_actoractor_refpose() -> dict:
|
| 881 |
+
"""Synthetic vroid_batch (/data4/vroid_batch) 2-view actor-actor, FRESH from Cosmos base 2B (no warm-start).
|
| 882 |
+
Same net as the shared actor-actor runs (person-pose + shared refs + Plücker + posed refs) PLUS a new
|
| 883 |
+
REFERENCE-POSE condition: each reference frame's per-person skeleton render (refs_shared `reference_pose`)
|
| 884 |
+
-> frozen VAE -> zero-init `reference_pose_embedder` ADDED to the reference tokens (mirrors reference-Plücker,
|
| 885 |
+
so the in-context references now carry BOTH RGB appearance and pose). 4-GPU."""
|
| 886 |
+
cfg = nymeria_pose_2view_actoractor_refs_campose_refpose_shared() # shared refs + plucker + refpose net
|
| 887 |
+
cfg["job"]["name"] = "vroid_actoractor_refpose_ar_generaltext" # AR recipe (2-cond + noise) + general captions
|
| 888 |
+
cfg["checkpoint"]["load_path"] = _base_2b_multiview_ckpt() # FRESH from Cosmos base 2B
|
| 889 |
+
cfg["checkpoint"]["strict_resume"] = False
|
| 890 |
+
# NEW: reference-pose conditioning (model VAE-encodes reference_pose -> net zero-init reference_pose_embedder)
|
| 891 |
+
cfg["model"]["config"]["enable_reference_pose"] = True
|
| 892 |
+
cfg["model"]["config"]["net"].update(enable_reference_pose=True)
|
| 893 |
+
# AR training recipe (from the handoff): condition on 1 OR 2 latents (50/50 per view) + noise the cond
|
| 894 |
+
# frames (prob 0.5) so the model is robust to imperfect (self-generated) cond frames at AR inference.
|
| 895 |
+
# `0: 0.0` is load-bearing (explicit so it can't be picked). See AR_HANDOFF/README.
|
| 896 |
+
cfg["model"]["config"]["conditional_frames_probs"] = {0: 0.0, 1: 0.5, 2: 0.5}
|
| 897 |
+
cfg["model"]["config"]["max_num_conditional_frames_per_view"] = 2
|
| 898 |
+
cfg["model"]["config"]["noisy_conditioning_prob"] = 0.5
|
| 899 |
+
cfg["model"]["config"]["noisy_conditioning_scale"] = 0.1
|
| 900 |
+
cfg["model"]["config"]["noisy_conditioning_min_frames"] = 2
|
| 901 |
+
# data: vroid synthetic (person-pose + reference-pose emitted)
|
| 902 |
+
for d in cfg["defaults"]:
|
| 903 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 904 |
+
d["override /data_train"] = "vroid_actoractor_refpose"
|
| 905 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 906 |
+
d["override /data_val"] = "vroid_actoractor_refpose"
|
| 907 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 908 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 909 |
+
num_video_frames=77, root="/data4/vroid_batch", actor_actor=True,
|
| 910 |
+
num_reference_frames=4, shared_reference=True, emit_camera_poses=True,
|
| 911 |
+
person_pose=True, emit_reference_pose=True, val_manifest="manifest/val_avatar_disjoint.csv",
|
| 912 |
+
)
|
| 913 |
+
return cfg
|
| 914 |
+
|
| 915 |
+
|
| 916 |
+
def vroid_overfit5() -> dict:
|
| 917 |
+
"""OVERFIT test: warm-start from the vroid ar_generaltext iter9000 checkpoint and deliberately overfit on
|
| 918 |
+
5 hand-picked val pairs (manifest/overfit5.csv). Train AND validate on the SAME 5 pairs so we can watch the
|
| 919 |
+
model memorize them. Same net/recipe as vroid_actoractor_refpose_ar_generaltext (AR 2-cond+noise, refpose)."""
|
| 920 |
+
cfg = vroid_actoractor_refpose() # inherits AR recipe + refpose + general captions
|
| 921 |
+
cfg["job"]["name"] = "vroid_overfit5"
|
| 922 |
+
# warm-start from the iter9000 checkpoint of the main run (NOT base 2B)
|
| 923 |
+
cfg["checkpoint"]["load_path"] = (
|
| 924 |
+
"/data/model_output/train/vroid_actoractor_refpose_ar_generaltext/checkpoints/iter_000009000"
|
| 925 |
+
)
|
| 926 |
+
cfg["checkpoint"]["strict_resume"] = False # load model weights; fresh optimizer/iteration for the overfit
|
| 927 |
+
cfg["checkpoint"]["save_iter"] = 200 # save often so we can inspect the memorization curve
|
| 928 |
+
cfg["trainer"]["max_iter"] = 4000
|
| 929 |
+
# data: the 5-pair overfit set for BOTH train and val
|
| 930 |
+
for d in cfg["defaults"]:
|
| 931 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 932 |
+
d["override /data_train"] = "vroid_overfit5"
|
| 933 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 934 |
+
d["override /data_val"] = "vroid_overfit5"
|
| 935 |
+
# validation viz on the SAME 5 pairs, frequently (num_val=5 -> all of them)
|
| 936 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 937 |
+
every_n=200, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 938 |
+
num_video_frames=77, root="/data4/vroid_batch", actor_actor=True,
|
| 939 |
+
num_reference_frames=4, shared_reference=True, emit_camera_poses=True,
|
| 940 |
+
person_pose=True, emit_reference_pose=True, val_manifest="manifest/overfit5.csv",
|
| 941 |
+
)
|
| 942 |
+
return cfg
|
| 943 |
+
|
| 944 |
+
|
| 945 |
+
def nymeria_vroid_mixed_refpose_pretrain() -> dict:
|
| 946 |
+
"""PRETRAIN on a MIX of nymeria actor-OBSERVER (real /data2/nymeria_processed_single) + vroid actor-ACTOR
|
| 947 |
+
(synthetic /data4/vroid_batch), pooled uniformly. Full conditioning structure INCLUDING reference-HUMAN-POSE
|
| 948 |
+
(skeleton): shared refs + Plücker campose + person-pose + reference_pose. Fresh from Cosmos base 2B. Standard
|
| 949 |
+
training recipe (no AR teacher-forcing/noise). 4-GPU."""
|
| 950 |
+
cfg = nymeria_pose_2view_actoractor_refs_campose_refpose_shared() # shared refs + plucker + refpose net
|
| 951 |
+
cfg["job"]["name"] = "nymeria_vroid_mixed_refpose_pretrain"
|
| 952 |
+
cfg["checkpoint"]["load_path"] = _base_2b_multiview_ckpt() # FRESH from Cosmos base 2B
|
| 953 |
+
cfg["checkpoint"]["strict_resume"] = False
|
| 954 |
+
# reference-HUMAN-POSE conditioning (VAE-encode reference_pose -> net zero-init reference_pose_embedder)
|
| 955 |
+
cfg["model"]["config"]["enable_reference_pose"] = True
|
| 956 |
+
cfg["model"]["config"]["net"].update(enable_reference_pose=True)
|
| 957 |
+
# data: the mixed obs + vroid loader (person-pose + reference-pose emitted for both)
|
| 958 |
+
for d in cfg["defaults"]:
|
| 959 |
+
if isinstance(d, dict) and "override /data_train" in d:
|
| 960 |
+
d["override /data_train"] = "nymeria_obs_vroid_mixed_refpose"
|
| 961 |
+
if isinstance(d, dict) and "override /data_val" in d:
|
| 962 |
+
d["override /data_val"] = "nymeria_obs_vroid_mixed_refpose"
|
| 963 |
+
# validation viz on the vroid actor-actor val set (person-pose + reference-pose)
|
| 964 |
+
cfg["trainer"]["callbacks"]["validation_viz"] = L(NymeriaValidationViz)(
|
| 965 |
+
every_n=500, num_train=5, num_val=5, num_sampling_step=35, guidance=7.0, fps=10,
|
| 966 |
+
num_video_frames=77, root="/data4/vroid_batch", actor_actor=True,
|
| 967 |
+
num_reference_frames=4, shared_reference=True, emit_camera_poses=True,
|
| 968 |
+
person_pose=True, emit_reference_pose=True, val_manifest="manifest/val_avatar_disjoint.csv",
|
| 969 |
+
)
|
| 970 |
+
return cfg
|
| 971 |
+
|
| 972 |
+
|
| 973 |
+
experiments = [
|
| 974 |
+
nymeria_pose_2actor_2b_smoke(),
|
| 975 |
+
nymeria_pose_2actor_2b_longer(),
|
| 976 |
+
nymeria_pose_2b_stage1_singleview(),
|
| 977 |
+
nymeria_pose_2b_stage1_sv_vaeconcat(),
|
| 978 |
+
nymeria_pose_2b_stage1_sv_vaemlp(),
|
| 979 |
+
nymeria_pose_2view_actorobs(),
|
| 980 |
+
nymeria_pose_2view_actorobs_refs(),
|
| 981 |
+
nymeria_pose_2view_actorobs_refs_campose(),
|
| 982 |
+
nymeria_pose_2view_actorobs_campose(),
|
| 983 |
+
nymeria_pose_2view_actorobs_refs_campose_refpose(),
|
| 984 |
+
nymeria_pose_2view_actorobs_refs_campose_refpose_shared(),
|
| 985 |
+
nymeria_pose_2view_actoractor_refs_campose_refpose(),
|
| 986 |
+
nymeria_pose_2view_actoractor_refs_campose_refpose_shared(),
|
| 987 |
+
nymeria_pose_2view_actoractor_refs_campose_refpose_shared_synthonly(),
|
| 988 |
+
nymeria_pose_2b_stage1_lora(),
|
| 989 |
+
comind_actoractor_personpose_shared(),
|
| 990 |
+
comind_synth_actoractor_personpose_shared(),
|
| 991 |
+
nymeria_synth_actoractor_depth_shared(),
|
| 992 |
+
nymeria_actorobs_depth_shared(),
|
| 993 |
+
vroid_actoractor_refpose(),
|
| 994 |
+
vroid_overfit5(),
|
| 995 |
+
nymeria_vroid_mixed_refpose_pretrain(),
|
| 996 |
+
]
|
| 997 |
+
|
| 998 |
+
cs = ConfigStore.instance()
|
| 999 |
+
for _item in experiments:
|
| 1000 |
+
cs.store(
|
| 1001 |
+
group="experiment",
|
| 1002 |
+
package="_global_",
|
| 1003 |
+
name=_item["job"]["name"],
|
| 1004 |
+
node=_item,
|
| 1005 |
+
)
|
cosmos_predict2/_src/predict2_multiview/datasets/__init__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
cosmos_predict2/_src/predict2_multiview/datasets/comind_pairs.py
ADDED
|
@@ -0,0 +1,302 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""CoMind 2-view actor-actor dataset (shared-world, multislam-eligible 32 recs).
|
| 2 |
+
|
| 3 |
+
Unlike Nymeria (two separate per-clip-dir clip.npz per pair), ONE CoMind clip.npz holds BOTH ego views —
|
| 4 |
+
``L_`` (leader) and ``H_`` (helper) — as JPEG byte-array streams (warped_cond / pose / target) + packed
|
| 5 |
+
visibility + Sapiens keypoints. Sidecars produced by the shared-world rebuild:
|
| 6 |
+
* ``<clip>.refs.npz`` — 8 GREEDY shared reference frames (combined leader+helper pool) + shared-world w2c/K.
|
| 7 |
+
* ``<clip>.campose.npz`` — both views' TARGET-frame camera poses (77) + K, in the multislam shared world.
|
| 8 |
+
Captions: ``manifest/captions/<clip_id>.json`` -> {"leader": ..., "helper": ...}.
|
| 9 |
+
|
| 10 |
+
Emits the SAME sample dict as ``NymeriaActorActorDataset`` (shared refs / camera poses / role-mix) so it drops
|
| 11 |
+
straight into the existing 2-view experiment + net. Combine with synthetic via ``get_comind_actor_actor_combined_loader``.
|
| 12 |
+
"""
|
| 13 |
+
import csv
|
| 14 |
+
import json
|
| 15 |
+
import os
|
| 16 |
+
import random
|
| 17 |
+
from typing import Any, Optional
|
| 18 |
+
|
| 19 |
+
import attrs
|
| 20 |
+
import cv2
|
| 21 |
+
import numpy as np
|
| 22 |
+
import torch
|
| 23 |
+
from hydra.core.config_store import ConfigStore
|
| 24 |
+
|
| 25 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 26 |
+
from cosmos_predict2._src.imaginaire.utils import log
|
| 27 |
+
from cosmos_predict2._src.predict2.datasets.cached_replay_dataloader import get_cached_replay_dataloader
|
| 28 |
+
from cosmos_predict2._src.predict2_multiview.datasets.multiview import collate_fn
|
| 29 |
+
from cosmos_predict2._src.predict2_multiview.datasets.nymeria_pairs import (
|
| 30 |
+
NymeriaActorActorDataset,
|
| 31 |
+
NymeriaPairsConfig,
|
| 32 |
+
_resize_thwc_uint8,
|
| 33 |
+
)
|
| 34 |
+
|
| 35 |
+
cv2.setNumThreads(1)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
@attrs.define(slots=False)
|
| 39 |
+
class ComindPairsConfig:
|
| 40 |
+
root: str = "/data4/comind_dataset"
|
| 41 |
+
resolution_hw: tuple[int, int] = (480, 480)
|
| 42 |
+
num_video_frames: int = 77
|
| 43 |
+
fps: float = 10.0
|
| 44 |
+
num_reference_frames: int = 4 # per-view slot; shared set = 2*R (=8) greedy refs
|
| 45 |
+
shared_reference: bool = True # refs carry no view identity (net drops ref view-emb)
|
| 46 |
+
emit_camera_poses: bool = True # cross-view Plücker (campose.npz) + posed refs (refs.npz w2c)
|
| 47 |
+
role_mix: bool = True # swap view0<->view1 (+captions) 50%
|
| 48 |
+
person_pose: bool = True # use IDENTITY-colored pose (L/H_pose_person: leader=EGO, helper=PARTNER
|
| 49 |
+
# view-invariant) instead of role-colored pose (wearer/partner)
|
| 50 |
+
reference_pose: bool = True # emit reference_pose (Sapiens skeleton on each reference frame, person-ID
|
| 51 |
+
# palette) alongside reference_frames — human-pose condition for the refs
|
| 52 |
+
split: str = "all" # "train" | "val" | "all" — held out by RECORDING
|
| 53 |
+
val_num_sessions: int = 3 # recordings held out for val
|
| 54 |
+
split_seed: int = 1234
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def _jdec(b: bytes) -> np.ndarray:
|
| 58 |
+
return cv2.cvtColor(cv2.imdecode(np.frombuffer(b, np.uint8), cv2.IMREAD_COLOR), cv2.COLOR_BGR2RGB)
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def _stack_jpeg(arr: np.ndarray) -> np.ndarray: # object array of JPEG bytes -> (T,H,W,3) uint8
|
| 62 |
+
return np.stack([_jdec(b) for b in arr])
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
def _fit_poses(w2c: np.ndarray, n: int) -> np.ndarray:
|
| 66 |
+
m = w2c.shape[0]
|
| 67 |
+
if m < n:
|
| 68 |
+
return np.concatenate([w2c, np.repeat(w2c[-1:], n - m, 0)], 0)
|
| 69 |
+
return w2c[:n]
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
class ComindActorActorDataset(torch.utils.data.Dataset):
|
| 73 |
+
def __init__(self, config: ComindPairsConfig):
|
| 74 |
+
self.config = config
|
| 75 |
+
self.root = config.root
|
| 76 |
+
self.clips_dir = os.path.join(self.root, "clips")
|
| 77 |
+
self.resolution_hw = tuple(config.resolution_hw)
|
| 78 |
+
self.num_frames = int(config.num_video_frames)
|
| 79 |
+
self.R = int(config.num_reference_frames)
|
| 80 |
+
self.emit_cam = bool(config.emit_camera_poses)
|
| 81 |
+
self.cap_dir = os.path.join(self.root, "manifest", "captions")
|
| 82 |
+
|
| 83 |
+
elig = sorted(json.load(open(os.path.join(self.root, "manifest", "multislam_paths.json"))))
|
| 84 |
+
rng = random.Random(config.split_seed)
|
| 85 |
+
shuffled = list(elig)
|
| 86 |
+
rng.shuffle(shuffled)
|
| 87 |
+
val_recs = set(shuffled[: config.val_num_sessions])
|
| 88 |
+
elig = set(elig)
|
| 89 |
+
|
| 90 |
+
self.pairs: list[dict[str, str]] = []
|
| 91 |
+
with open(os.path.join(self.root, "manifest", "clip_pairs.csv"), newline="") as f:
|
| 92 |
+
for row in csv.DictReader(f):
|
| 93 |
+
rec8 = row["recording"][:8]
|
| 94 |
+
if rec8 not in elig:
|
| 95 |
+
continue
|
| 96 |
+
cid = row["clip_id"]
|
| 97 |
+
npz = os.path.join(self.clips_dir, rec8, f"{cid}.npz")
|
| 98 |
+
refs = npz[:-4] + ".refs.npz"
|
| 99 |
+
camp = npz[:-4] + ".campose.npz"
|
| 100 |
+
if not (os.path.isfile(npz) and os.path.isfile(refs)):
|
| 101 |
+
continue
|
| 102 |
+
if self.emit_cam and not os.path.isfile(camp):
|
| 103 |
+
continue
|
| 104 |
+
in_val = rec8 in val_recs
|
| 105 |
+
if config.split == "train" and in_val:
|
| 106 |
+
continue
|
| 107 |
+
if config.split == "val" and not in_val:
|
| 108 |
+
continue
|
| 109 |
+
self.pairs.append({"pair_id": cid, "rec8": rec8, "npz": npz, "refs": refs, "camp": camp})
|
| 110 |
+
log.info(
|
| 111 |
+
f"ComindActorActorDataset[{config.split}]: {len(self.pairs)} clips "
|
| 112 |
+
f"({len(elig)} eligible recs, {len(val_recs)} held out for val, R={self.R}, emit_cam={self.emit_cam})"
|
| 113 |
+
)
|
| 114 |
+
if not self.pairs:
|
| 115 |
+
raise RuntimeError(f"No usable CoMind clips for split={config.split} under {self.clips_dir}")
|
| 116 |
+
|
| 117 |
+
def __len__(self) -> int:
|
| 118 |
+
return len(self.pairs)
|
| 119 |
+
|
| 120 |
+
def _resize(self, thwc: np.ndarray, n: int) -> torch.Tensor:
|
| 121 |
+
return _resize_thwc_uint8(thwc, n, self.resolution_hw)
|
| 122 |
+
|
| 123 |
+
def _load_view(self, z, refs, camp, side: str, slot: int) -> dict[str, torch.Tensor]:
|
| 124 |
+
"""side: 'L'(leader) | 'H'(helper). slot: final view position (0/1) -> shared-ref R-slice."""
|
| 125 |
+
out: dict[str, torch.Tensor] = {}
|
| 126 |
+
out["video"] = self._resize(_stack_jpeg(z[f"{side}_target"]), self.num_frames)
|
| 127 |
+
pose_key = f"{side}_pose_person" if (self.config.person_pose and f"{side}_pose_person" in z.files) else f"{side}_pose"
|
| 128 |
+
out["control_input_pose"] = self._resize(_stack_jpeg(z[pose_key]), self.num_frames)
|
| 129 |
+
out["control_input_warped"] = self._resize(_stack_jpeg(z[f"{side}_warped"]), self.num_frames)
|
| 130 |
+
vt, vh, vw = (int(x) for x in z["vis_shape"])
|
| 131 |
+
vis = np.unpackbits(z[f"{side}_vis_packed"])[: vt * vh * vw].reshape(vt, vh, vw)
|
| 132 |
+
vis = (vis.astype(np.uint8) * 255)[..., None] # (T,H,W,1)
|
| 133 |
+
out["control_input_visibility"] = self._resize(vis, self.num_frames)
|
| 134 |
+
if self.R > 0:
|
| 135 |
+
allf = refs["reference_frames"] # (2R, h, h, 3) greedy shared set (top-gain first)
|
| 136 |
+
lo = slot * self.R
|
| 137 |
+
out["reference_frames"] = self._resize(allf[lo : lo + self.R], self.R)
|
| 138 |
+
if self.config.reference_pose and "reference_pose" in refs.files: # Sapiens skeleton per ref frame (person-ID palette)
|
| 139 |
+
out["reference_pose"] = self._resize(refs["reference_pose"][lo : lo + self.R], self.R)
|
| 140 |
+
if self.emit_cam:
|
| 141 |
+
out["cam_ref_w2c"] = torch.from_numpy(refs["reference_w2c"].astype(np.float32)[lo : lo + self.R])
|
| 142 |
+
if self.emit_cam:
|
| 143 |
+
out["cam_w2c"] = torch.from_numpy(_fit_poses(camp[f"{side}_w2c"].astype(np.float32), self.num_frames))
|
| 144 |
+
out["cam_K"] = torch.from_numpy(camp[f"K_{'leader' if side == 'L' else 'helper'}"].astype(np.float32))
|
| 145 |
+
out["cam_src_res"] = torch.tensor(504, dtype=torch.int64)
|
| 146 |
+
return out
|
| 147 |
+
|
| 148 |
+
def _caption(self, cid: str) -> tuple[str, str]:
|
| 149 |
+
p = os.path.join(self.cap_dir, f"{cid}.json")
|
| 150 |
+
if os.path.isfile(p):
|
| 151 |
+
d = json.load(open(p))
|
| 152 |
+
return d.get("leader", ""), d.get("helper", "")
|
| 153 |
+
return "", ""
|
| 154 |
+
|
| 155 |
+
def __getitem__(self, idx: int) -> dict[str, Any]:
|
| 156 |
+
for _ in range(16):
|
| 157 |
+
try:
|
| 158 |
+
return self._build_sample(idx)
|
| 159 |
+
except Exception as e:
|
| 160 |
+
log.warning(f"ComindActorActorDataset: bad sample {self.pairs[idx]['pair_id']} ({type(e).__name__}: {str(e)[:80]}); retrying")
|
| 161 |
+
idx = random.randrange(len(self.pairs))
|
| 162 |
+
raise RuntimeError("ComindActorActorDataset: too many unreadable samples in a row")
|
| 163 |
+
|
| 164 |
+
def _build_sample(self, idx: int) -> dict[str, Any]:
|
| 165 |
+
p = self.pairs[idx]
|
| 166 |
+
z = np.load(p["npz"], allow_pickle=True)
|
| 167 |
+
refs = np.load(p["refs"], allow_pickle=True)
|
| 168 |
+
camp = np.load(p["camp"], allow_pickle=True) if self.emit_cam else None
|
| 169 |
+
cap_l, cap_h = self._caption(p["pair_id"])
|
| 170 |
+
|
| 171 |
+
views = [("L", cap_l), ("H", cap_h)] # view0=leader, view1=helper
|
| 172 |
+
if self.config.role_mix and random.random() < 0.5:
|
| 173 |
+
views = views[::-1]
|
| 174 |
+
actors = [self._load_view(z, refs, camp, side, slot=i) for i, (side, _) in enumerate(views)]
|
| 175 |
+
captions = [c for _, c in views]
|
| 176 |
+
n_views = 2
|
| 177 |
+
|
| 178 |
+
sample: dict[str, Any] = {}
|
| 179 |
+
for key in actors[0]:
|
| 180 |
+
if key.startswith("cam_"):
|
| 181 |
+
continue
|
| 182 |
+
sample[key] = torch.cat([a[key] for a in actors], dim=1).contiguous()
|
| 183 |
+
if "cam_w2c" in actors[0]:
|
| 184 |
+
sample["camera_w2c"] = torch.stack([a["cam_w2c"] for a in actors], dim=0).contiguous()
|
| 185 |
+
sample["camera_K"] = torch.stack([a["cam_K"] for a in actors], dim=0).contiguous()
|
| 186 |
+
sample["camera_src_res"] = torch.stack([a["cam_src_res"] for a in actors], dim=0).contiguous()
|
| 187 |
+
if "cam_ref_w2c" in actors[0]:
|
| 188 |
+
sample["reference_cam_w2c"] = torch.stack([a["cam_ref_w2c"] for a in actors], dim=0).contiguous()
|
| 189 |
+
|
| 190 |
+
T = self.num_frames
|
| 191 |
+
sample.update(
|
| 192 |
+
{
|
| 193 |
+
"__key__": p["pair_id"],
|
| 194 |
+
"__url__": p["pair_id"],
|
| 195 |
+
"ai_caption": captions,
|
| 196 |
+
"view_indices": torch.tensor([0] * T + [1] * T, dtype=torch.int64),
|
| 197 |
+
"fps": torch.tensor(self.config.fps, dtype=torch.float64),
|
| 198 |
+
"num_video_frames_per_view": torch.tensor(T, dtype=torch.int64),
|
| 199 |
+
"view_indices_selection": torch.tensor([0, 1], dtype=torch.int64),
|
| 200 |
+
"camera_keys_selection": ["view0", "view1"],
|
| 201 |
+
"sample_n_views": torch.tensor(n_views, dtype=torch.int64),
|
| 202 |
+
"padding_mask": torch.zeros((1, *self.resolution_hw), dtype=torch.float32),
|
| 203 |
+
"ref_cam_view_idx_sample_position": torch.tensor(-1, dtype=torch.int64),
|
| 204 |
+
"front_cam_view_idx_sample_position": torch.tensor(0, dtype=torch.int64),
|
| 205 |
+
}
|
| 206 |
+
)
|
| 207 |
+
return sample
|
| 208 |
+
|
| 209 |
+
|
| 210 |
+
def get_comind_actor_actor_combined_loader(
|
| 211 |
+
*,
|
| 212 |
+
comind_root: str = "/data4/comind_dataset",
|
| 213 |
+
synth_root: str = "/data3/synthetic_processed_multi",
|
| 214 |
+
synth_manifest: str = "manifest/train_split.csv",
|
| 215 |
+
resolution_hw: tuple[int, int] = (480, 480),
|
| 216 |
+
num_video_frames: int = 77,
|
| 217 |
+
fps: float = 10.0,
|
| 218 |
+
role_mix: bool = True,
|
| 219 |
+
num_reference_frames: int = 4,
|
| 220 |
+
shared_reference: bool = True,
|
| 221 |
+
emit_camera_poses: bool = True,
|
| 222 |
+
person_pose: bool = True,
|
| 223 |
+
split: str = "train",
|
| 224 |
+
batch_size: int = 1,
|
| 225 |
+
num_workers: int = 4,
|
| 226 |
+
prefetch_factor: Optional[int] = 2,
|
| 227 |
+
is_train: bool = True,
|
| 228 |
+
**kwargs: Any,
|
| 229 |
+
):
|
| 230 |
+
"""ConcatDataset over CoMind (shared-world 32 recs) [+ synthetic]. Both emit the same sample dict; pooled
|
| 231 |
+
uniformly. synth_root="" -> CoMind-only. Val (is_train=False) uses each set's held-out split."""
|
| 232 |
+
parts: list[torch.utils.data.Dataset] = [
|
| 233 |
+
ComindActorActorDataset(
|
| 234 |
+
ComindPairsConfig(
|
| 235 |
+
root=comind_root, resolution_hw=tuple(resolution_hw), num_video_frames=num_video_frames, fps=fps,
|
| 236 |
+
num_reference_frames=num_reference_frames, shared_reference=shared_reference,
|
| 237 |
+
emit_camera_poses=emit_camera_poses, role_mix=role_mix, person_pose=person_pose, split=split,
|
| 238 |
+
)
|
| 239 |
+
)
|
| 240 |
+
]
|
| 241 |
+
if synth_root:
|
| 242 |
+
parts.append(
|
| 243 |
+
NymeriaActorActorDataset(
|
| 244 |
+
NymeriaPairsConfig(
|
| 245 |
+
root=synth_root, manifest_csv=synth_manifest, resolution_hw=tuple(resolution_hw),
|
| 246 |
+
num_video_frames=num_video_frames, fps=fps, role_mix=role_mix,
|
| 247 |
+
num_reference_frames=num_reference_frames, shared_reference=shared_reference,
|
| 248 |
+
emit_camera_poses=emit_camera_poses, person_pose=person_pose,
|
| 249 |
+
)
|
| 250 |
+
)
|
| 251 |
+
)
|
| 252 |
+
dataset = torch.utils.data.ConcatDataset(parts) if len(parts) > 1 else parts[0]
|
| 253 |
+
log.info(f"comind+synth combined[{split}]: {[len(p) for p in parts]} -> {len(dataset)} clips")
|
| 254 |
+
sampler = None
|
| 255 |
+
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
| 256 |
+
sampler = torch.utils.data.distributed.DistributedSampler(dataset, shuffle=is_train, drop_last=True)
|
| 257 |
+
return get_cached_replay_dataloader(
|
| 258 |
+
webdataset=False, use_cache=False, dataset=dataset, batch_size=batch_size, num_workers=num_workers,
|
| 259 |
+
sampler=sampler, shuffle=(sampler is None and is_train), drop_last=True,
|
| 260 |
+
prefetch_factor=prefetch_factor if num_workers > 0 else None, persistent_workers=num_workers > 0,
|
| 261 |
+
pin_memory=False, collate_fn=collate_fn, cache_replay_name="comind_actor_actor_combined_dataloader",
|
| 262 |
+
)
|
| 263 |
+
|
| 264 |
+
|
| 265 |
+
def register_comind_data() -> None:
|
| 266 |
+
cs = ConfigStore.instance()
|
| 267 |
+
# CoMind-ONLY (shared-world 32 recs), 2-view actor-actor, IDENTITY-colored pose (person_pose), greedy shared
|
| 268 |
+
# refs + cross-view Plücker. Depth channels stored but NOT loaded (person_pose pose, no depth condition).
|
| 269 |
+
cs.store(
|
| 270 |
+
group="data_train", package="dataloader_train", name="comind_actoractor_personpose_shared",
|
| 271 |
+
node=L(get_comind_actor_actor_combined_loader)(
|
| 272 |
+
synth_root="", split="train", num_video_frames=77, role_mix=True, num_reference_frames=4,
|
| 273 |
+
shared_reference=True, emit_camera_poses=True, person_pose=True, is_train=True, batch_size=1, num_workers=4),
|
| 274 |
+
)
|
| 275 |
+
cs.store(
|
| 276 |
+
group="data_val", package="dataloader_val", name="comind_actoractor_personpose_shared",
|
| 277 |
+
node=L(get_comind_actor_actor_combined_loader)(
|
| 278 |
+
synth_root="", split="val", num_video_frames=77, role_mix=False, num_reference_frames=4,
|
| 279 |
+
shared_reference=True, emit_camera_poses=True, person_pose=True, is_train=False, batch_size=1, num_workers=2),
|
| 280 |
+
)
|
| 281 |
+
# CoMind + SYNTHETIC combined, 2-view actor-actor, person-pose (both use pose_person). Train mixes both;
|
| 282 |
+
# validation stays on the CoMind val (data_val below = comind-only).
|
| 283 |
+
cs.store(
|
| 284 |
+
group="data_train", package="dataloader_train", name="comind_synth_actoractor_personpose_shared",
|
| 285 |
+
node=L(get_comind_actor_actor_combined_loader)(
|
| 286 |
+
synth_root="/data3/synthetic_processed_multi", split="train", num_video_frames=77, role_mix=True,
|
| 287 |
+
num_reference_frames=4, shared_reference=True, emit_camera_poses=True, person_pose=True,
|
| 288 |
+
is_train=True, batch_size=1, num_workers=4),
|
| 289 |
+
)
|
| 290 |
+
# (kept) CoMind + synthetic combined variant, role pose — earlier config.
|
| 291 |
+
cs.store(
|
| 292 |
+
group="data_train", package="dataloader_train", name="comind_actor_actor_refs_campose_shared_plus_synth",
|
| 293 |
+
node=L(get_comind_actor_actor_combined_loader)(
|
| 294 |
+
split="train", num_video_frames=77, role_mix=True, num_reference_frames=4,
|
| 295 |
+
shared_reference=True, emit_camera_poses=True, person_pose=True, is_train=True, batch_size=1, num_workers=4),
|
| 296 |
+
)
|
| 297 |
+
cs.store(
|
| 298 |
+
group="data_val", package="dataloader_val", name="comind_actor_actor_refs_campose_shared",
|
| 299 |
+
node=L(get_comind_actor_actor_combined_loader)(
|
| 300 |
+
synth_root="", split="val", num_video_frames=77, role_mix=False, num_reference_frames=4,
|
| 301 |
+
shared_reference=True, emit_camera_poses=True, person_pose=True, is_train=False, batch_size=1, num_workers=2),
|
| 302 |
+
)
|
cosmos_predict2/_src/predict2_multiview/datasets/local.py
ADDED
|
@@ -0,0 +1,173 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
"""
|
| 5 |
+
Local file-based datasets.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
import glob
|
| 9 |
+
import os
|
| 10 |
+
import random
|
| 11 |
+
from pathlib import Path
|
| 12 |
+
|
| 13 |
+
import pandas as pd
|
| 14 |
+
from torch.utils.data import Dataset
|
| 15 |
+
|
| 16 |
+
from cosmos_predict2._src.predict2_multiview.datasets.multiview import (
|
| 17 |
+
AugmentationConfig,
|
| 18 |
+
make_augmentations,
|
| 19 |
+
)
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
class WaymoLocalDataset(Dataset):
|
| 23 |
+
def __init__(
|
| 24 |
+
self,
|
| 25 |
+
video_file_dirs: list[str],
|
| 26 |
+
augmentation_config: AugmentationConfig,
|
| 27 |
+
shuffle: bool = True,
|
| 28 |
+
gc_every_n: int = 100,
|
| 29 |
+
) -> None:
|
| 30 |
+
self.video_file_dirs = video_file_dirs
|
| 31 |
+
self.augmentation_config = augmentation_config
|
| 32 |
+
self.shuffle = shuffle
|
| 33 |
+
self.gc_every_n = gc_every_n
|
| 34 |
+
|
| 35 |
+
self.augmentations, self.dataset_keys = make_augmentations(augmentation_config)
|
| 36 |
+
self.sample_dirs = self.build_sample_dirs(self.video_file_dirs)
|
| 37 |
+
|
| 38 |
+
if self.shuffle:
|
| 39 |
+
random.shuffle(self.sample_dirs)
|
| 40 |
+
|
| 41 |
+
def build_sample_dirs(self, video_file_dirs: list[str]) -> list[str]:
|
| 42 |
+
sample_dirs = []
|
| 43 |
+
for video_file_dir in video_file_dirs:
|
| 44 |
+
for sample_dir in glob.glob(os.path.join(video_file_dir, "**")):
|
| 45 |
+
if os.path.isdir(sample_dir):
|
| 46 |
+
sample_dirs.append(sample_dir)
|
| 47 |
+
return sample_dirs
|
| 48 |
+
|
| 49 |
+
def load_data(self, sample_dir: str) -> dict:
|
| 50 |
+
sample_id = sample_dir.split("/")[-1]
|
| 51 |
+
data_dict = dict()
|
| 52 |
+
for filename in glob.glob(os.path.join(sample_dir, "*.mp4")):
|
| 53 |
+
with open(filename, "rb") as f:
|
| 54 |
+
view = filename.split("/")[-1].split(".")[0]
|
| 55 |
+
data_dict[f"video_{view}"] = f.read()
|
| 56 |
+
return data_dict
|
| 57 |
+
|
| 58 |
+
def load_caption(self, sample_dir: str) -> dict:
|
| 59 |
+
caption_path = os.path.join(sample_dir, "caption.jsonl")
|
| 60 |
+
with open(caption_path, "r") as f:
|
| 61 |
+
caption_df = pd.read_json(f, lines=True, orient="records")
|
| 62 |
+
|
| 63 |
+
caption_dict = dict()
|
| 64 |
+
for view_name, view_df in caption_df.groupby("view"):
|
| 65 |
+
caption_styles = dict()
|
| 66 |
+
for row in view_df.itertuples():
|
| 67 |
+
caption = row.caption
|
| 68 |
+
tag = None if pd.isna(row.tag) else row.tag
|
| 69 |
+
caption_styles[tag or "long"] = caption
|
| 70 |
+
caption_dict[f"caption_{view_name}"] = {
|
| 71 |
+
"t2w_windows": [
|
| 72 |
+
{"start_frame": 0, "end_frame": self.augmentation_config.num_video_frames, **caption_styles}
|
| 73 |
+
]
|
| 74 |
+
}
|
| 75 |
+
return caption_dict
|
| 76 |
+
|
| 77 |
+
def __len__(self) -> int:
|
| 78 |
+
return len(self.sample_dirs)
|
| 79 |
+
|
| 80 |
+
def __getitem__(self, index: int) -> dict:
|
| 81 |
+
sample_dir = self.sample_dirs[index]
|
| 82 |
+
data_dict = {
|
| 83 |
+
"__key__": str(index),
|
| 84 |
+
"__url__": str(sample_dir[index]),
|
| 85 |
+
}
|
| 86 |
+
data_dict.update(self.load_data(sample_dir))
|
| 87 |
+
data_dict.update(self.load_caption(sample_dir))
|
| 88 |
+
for k, aug in self.augmentations.items():
|
| 89 |
+
data_dict = aug(data_dict)
|
| 90 |
+
return data_dict
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
class LocalMultiViewDataset(Dataset):
|
| 94 |
+
"""Dataset wrapper for local multiview sample."""
|
| 95 |
+
|
| 96 |
+
def __init__(
|
| 97 |
+
self,
|
| 98 |
+
video_file_dicts: list[dict[str, bytes | Path | None]],
|
| 99 |
+
prompts: list[str],
|
| 100 |
+
augmentation_config: AugmentationConfig,
|
| 101 |
+
camera_key_adapter: dict[str, str] | None = None,
|
| 102 |
+
control_file_dicts: list[dict[str, bytes | Path | None]] | None = None,
|
| 103 |
+
) -> None:
|
| 104 |
+
self.video_file_dicts = video_file_dicts
|
| 105 |
+
self.prompts = prompts
|
| 106 |
+
self.augmentation_config = augmentation_config
|
| 107 |
+
self.camera_key_adapter = camera_key_adapter
|
| 108 |
+
self.control_file_dicts = control_file_dicts
|
| 109 |
+
|
| 110 |
+
if self.control_file_dicts is not None and len(self.video_file_dicts) != len(self.control_file_dicts):
|
| 111 |
+
raise ValueError("Number of video file dicts and control file dicts must be the same!")
|
| 112 |
+
|
| 113 |
+
if len(self.prompts) != len(self.video_file_dicts):
|
| 114 |
+
raise ValueError("Number of prompts and video file dicts must be the same!")
|
| 115 |
+
|
| 116 |
+
if self.augmentation_config.single_caption_camera_name is None:
|
| 117 |
+
raise ValueError(
|
| 118 |
+
"`single_caption_camera_name` must be set since only single prompt is provided by dataset!"
|
| 119 |
+
)
|
| 120 |
+
|
| 121 |
+
self.augmentations, self.dataset_keys = make_augmentations(augmentation_config)
|
| 122 |
+
|
| 123 |
+
def __len__(self) -> int:
|
| 124 |
+
return len(self.video_file_dicts)
|
| 125 |
+
|
| 126 |
+
def __getitem__(self, index: int) -> dict:
|
| 127 |
+
data_dict = {
|
| 128 |
+
"__key__": str(index),
|
| 129 |
+
"__url__": "local_dataset",
|
| 130 |
+
}
|
| 131 |
+
|
| 132 |
+
for view_key, filepath in self.video_file_dicts[index].items():
|
| 133 |
+
if filepath is None:
|
| 134 |
+
raise ValueError(f"view_key {view_key} has null filepath!")
|
| 135 |
+
default_key = self.camera_key_adapter[view_key] if self.camera_key_adapter else view_key
|
| 136 |
+
video_key = self.augmentation_config.camera_video_key_mapping[default_key]
|
| 137 |
+
if isinstance(filepath, bytes):
|
| 138 |
+
data_dict[video_key] = filepath
|
| 139 |
+
else:
|
| 140 |
+
with open(filepath, "rb") as f:
|
| 141 |
+
data_dict[video_key] = f.read()
|
| 142 |
+
|
| 143 |
+
if self.control_file_dicts is not None:
|
| 144 |
+
for view_key, filepath in self.control_file_dicts[index].items():
|
| 145 |
+
if filepath is None:
|
| 146 |
+
raise ValueError(f"view_key {view_key} has null filepath!")
|
| 147 |
+
default_key = self.camera_key_adapter[view_key] if self.camera_key_adapter else view_key
|
| 148 |
+
control_key = self.augmentation_config.camera_control_key_mapping[default_key]
|
| 149 |
+
if isinstance(filepath, bytes):
|
| 150 |
+
data_dict[control_key] = filepath
|
| 151 |
+
else:
|
| 152 |
+
with open(filepath, "rb") as f:
|
| 153 |
+
data_dict[control_key] = f.read()
|
| 154 |
+
|
| 155 |
+
caption_styles = dict(
|
| 156 |
+
zip(
|
| 157 |
+
self.augmentation_config.caption_probability.keys(),
|
| 158 |
+
[self.prompts[index] for _ in range(len(self.augmentation_config.caption_probability))],
|
| 159 |
+
)
|
| 160 |
+
)
|
| 161 |
+
|
| 162 |
+
caption_key = self.augmentation_config.camera_caption_key_mapping[
|
| 163 |
+
self.augmentation_config.single_caption_camera_name
|
| 164 |
+
]
|
| 165 |
+
data_dict[caption_key] = {
|
| 166 |
+
"t2w_windows": [
|
| 167 |
+
{"start_frame": 0, "end_frame": self.augmentation_config.num_video_frames, **caption_styles}
|
| 168 |
+
]
|
| 169 |
+
}
|
| 170 |
+
|
| 171 |
+
for k, aug in self.augmentations.items():
|
| 172 |
+
data_dict = aug(data_dict)
|
| 173 |
+
return data_dict
|
cosmos_predict2/_src/predict2_multiview/datasets/multiview.py
ADDED
|
@@ -0,0 +1,547 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
"""
|
| 17 |
+
Webloaders of datasets and augmentations for visual-text multiview dataset for AV
|
| 18 |
+
"""
|
| 19 |
+
|
| 20 |
+
try:
|
| 21 |
+
from megatron.core import parallel_state
|
| 22 |
+
|
| 23 |
+
USE_MEGATRON = True
|
| 24 |
+
except ImportError:
|
| 25 |
+
USE_MEGATRON = False
|
| 26 |
+
import io
|
| 27 |
+
import random
|
| 28 |
+
from typing import Any, Final, Literal, Optional, TypeAlias
|
| 29 |
+
|
| 30 |
+
import attrs
|
| 31 |
+
import torch
|
| 32 |
+
import webdataset as wds
|
| 33 |
+
from einops import rearrange
|
| 34 |
+
from torchvision.transforms import InterpolationMode, Resize
|
| 35 |
+
|
| 36 |
+
import cosmos_predict2._src.predict2.datasets.distributor.parallel_sync_multi_aspect_ratio as parallel_sync_multi_aspect_ratio
|
| 37 |
+
from cosmos_predict2._src.imaginaire.datasets.decoders.json_loader import json_decoder
|
| 38 |
+
from cosmos_predict2._src.imaginaire.datasets.decoders.video_decoder import video_naive_bytes
|
| 39 |
+
from cosmos_predict2._src.imaginaire.datasets.webdataset.augmentors.augmentor import Augmentor
|
| 40 |
+
from cosmos_predict2._src.imaginaire.datasets.webdataset.config.schema import DatasetConfig
|
| 41 |
+
from cosmos_predict2._src.imaginaire.datasets.webdataset.distributors import ShardlistBasic
|
| 42 |
+
from cosmos_predict2._src.imaginaire.datasets.webdataset.webdataset_ext import Dataset
|
| 43 |
+
from cosmos_predict2._src.imaginaire.utils import log
|
| 44 |
+
from cosmos_predict2._src.predict2.datasets.cached_replay_dataloader import get_cached_replay_dataloader
|
| 45 |
+
from cosmos_predict2._src.predict2_multiview.datasets.wdinfo_utils import DEFAULT_CATALOG, get_video_dataset_info
|
| 46 |
+
|
| 47 |
+
CameraKeyType: TypeAlias = str
|
| 48 |
+
|
| 49 |
+
# for Autonomous Driving Dataset (Alpamayo and Mads)
|
| 50 |
+
|
| 51 |
+
DEFAULT_CAMERAS: Final[tuple[CameraKeyType, ...]] = (
|
| 52 |
+
"camera_front_wide_120fov",
|
| 53 |
+
"camera_cross_right_120fov",
|
| 54 |
+
"camera_rear_right_70fov",
|
| 55 |
+
"camera_rear_tele_30fov",
|
| 56 |
+
"camera_rear_left_70fov",
|
| 57 |
+
"camera_cross_left_120fov",
|
| 58 |
+
"camera_front_tele_30fov",
|
| 59 |
+
)
|
| 60 |
+
|
| 61 |
+
DEFAULT_CAMERA_VIEW_MAPPING: Final = dict(zip(DEFAULT_CAMERAS, range(len(DEFAULT_CAMERAS))))
|
| 62 |
+
|
| 63 |
+
DEFAULT_CAPTION_PREFIXES: Final = {
|
| 64 |
+
"camera_front_wide_120fov": "The video is captured from a camera mounted on a car. The camera is facing forward.",
|
| 65 |
+
"camera_cross_right_120fov": "The video is captured from a camera mounted on a car. The camera is facing to the right.",
|
| 66 |
+
"camera_rear_right_70fov": "The video is captured from a camera mounted on a car. The camera is facing the rear right side.",
|
| 67 |
+
"camera_rear_tele_30fov": "The video is captured from a camera mounted on a car. The camera is facing backwards.",
|
| 68 |
+
"camera_rear_left_70fov": "The video is captured from a camera mounted on a car. The camera is facing the rear left side.",
|
| 69 |
+
"camera_cross_left_120fov": "The video is captured from a camera mounted on a car. The camera is facing to the left.",
|
| 70 |
+
"camera_front_tele_30fov": "The video is captured from a telephoto camera mounted on a car. The camera is facing forward.",
|
| 71 |
+
}
|
| 72 |
+
|
| 73 |
+
DEFAULT_CAPTION_KEY_MAPPING: Final = dict(
|
| 74 |
+
zip(DEFAULT_CAMERAS, [f"metas_{k}_10s_chunks_qwen2p5_vl_32b" for k in DEFAULT_CAMERAS])
|
| 75 |
+
)
|
| 76 |
+
DEFAULT_VIDEO_KEY_MAPPING: Final = dict(zip(DEFAULT_CAMERAS, [f"video_{k}" for k in DEFAULT_CAMERAS]))
|
| 77 |
+
|
| 78 |
+
# Agibot 3-view (head_color, hand_left, hand_right) multiview multicontrol
|
| 79 |
+
AGIBOT_VIEWS: Final[tuple[CameraKeyType, ...]] = ("head_color", "hand_left", "hand_right")
|
| 80 |
+
AGIBOT_VIEW_MAPPING: Final = dict(zip(AGIBOT_VIEWS, range(len(AGIBOT_VIEWS))))
|
| 81 |
+
AGIBOT_VIDEO_KEY_MAPPING: Final = dict(zip(AGIBOT_VIEWS, [f"video_{k}" for k in AGIBOT_VIEWS]))
|
| 82 |
+
AGIBOT_CAPTION_KEY_MAPPING: Final = dict(zip(AGIBOT_VIEWS, [f"metas_{k}" for k in AGIBOT_VIEWS]))
|
| 83 |
+
AGIBOT_CONTROL_KEY_MAPPING: Final = dict(zip(AGIBOT_VIEWS, [f"control_{k}" for k in AGIBOT_VIEWS]))
|
| 84 |
+
AGIBOT_CAPTION_PREFIXES: Final = {
|
| 85 |
+
"head_color": "The video is captured from a camera mounted on the head of the subject, facing forward.",
|
| 86 |
+
"hand_left": "The video is captured from a camera mounted on the left hand of the subject.",
|
| 87 |
+
"hand_right": "The video is captured from a camera mounted on the right hand of the subject.",
|
| 88 |
+
}
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
class UnpackMetas(Augmentor):
|
| 92 |
+
"""Unpack metas from single meta dicts list into per-camera meta dicts."""
|
| 93 |
+
|
| 94 |
+
def __init__(
|
| 95 |
+
self, position_to_camera_mapping: dict[int, str], input_key: str = "metas", output_prefix: str = "metas_"
|
| 96 |
+
) -> None:
|
| 97 |
+
super().__init__([], {})
|
| 98 |
+
self.position_to_camera_mapping = position_to_camera_mapping
|
| 99 |
+
self.input_key = input_key
|
| 100 |
+
self.output_prefix = output_prefix
|
| 101 |
+
|
| 102 |
+
def __call__(self, data: dict[str, Any]) -> dict[str, Any]:
|
| 103 |
+
metas = data.pop(self.input_key)
|
| 104 |
+
for i, meta in enumerate(metas):
|
| 105 |
+
camera_name = self.position_to_camera_mapping[i]
|
| 106 |
+
data[f"{self.output_prefix}{camera_name}"] = meta
|
| 107 |
+
return data
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
class ExtractFramesAndCaptions(Augmentor):
|
| 111 |
+
"""Extract frames from a videos."""
|
| 112 |
+
|
| 113 |
+
def __init__(
|
| 114 |
+
self,
|
| 115 |
+
camera_order: list[CameraKeyType],
|
| 116 |
+
num_frames: int,
|
| 117 |
+
resolution_hw: tuple[int, int],
|
| 118 |
+
fps_downsample_factor: int,
|
| 119 |
+
caption_probability: dict[str, float],
|
| 120 |
+
camera_view_mapping: dict[CameraKeyType, int],
|
| 121 |
+
camera_caption_key_mapping: dict[CameraKeyType, str],
|
| 122 |
+
camera_video_key_mapping: dict[CameraKeyType, str],
|
| 123 |
+
camera_control_key_mapping: Optional[dict[CameraKeyType, str]] = None,
|
| 124 |
+
add_view_prefix_to_caption: bool = False,
|
| 125 |
+
camera_prefix_mapping: Optional[dict[CameraKeyType, str]] = None,
|
| 126 |
+
single_caption_camera_name: Optional[CameraKeyType] = None,
|
| 127 |
+
window_random_frame_offset_range: Optional[tuple[int, int]] = None,
|
| 128 |
+
) -> None:
|
| 129 |
+
"""Extracts frames and captions from video/metadata dicts.
|
| 130 |
+
|
| 131 |
+
Args:
|
| 132 |
+
camera_order: Order of cameras to extract
|
| 133 |
+
num_frames: Number of frames to extract
|
| 134 |
+
resolution_hw: Resolution of the extracted frames
|
| 135 |
+
fps_downsample_factor: FPS downsample factor
|
| 136 |
+
caption_probability: Probability of each caption type in t2w window
|
| 137 |
+
camera_view_mapping: Mapping of camera keys to view indices
|
| 138 |
+
camera_caption_key_mapping: Mapping of camera keys to caption keys
|
| 139 |
+
camera_video_key_mapping: Mapping of camera keys to video keys
|
| 140 |
+
camera_control_key_mapping: Mapping of camera keys to control keys
|
| 141 |
+
add_view_prefix_to_caption: Whether to add caption prefix for all views
|
| 142 |
+
camera_prefix_mapping: Mapping of camera keys to prefixes
|
| 143 |
+
single_caption_camera_name: Name of the camera key to use for single caption conditioning.
|
| 144 |
+
If `add_view_prefix_to_caption` is True, will still provide prefixes for other views.
|
| 145 |
+
window_random_frame_offset_range: Optional range of random offset to add to the start frame of the extracted window.
|
| 146 |
+
|
| 147 |
+
Returns:
|
| 148 |
+
data: Dictionary with resized tensors of frames and captions
|
| 149 |
+
"""
|
| 150 |
+
super().__init__([], {})
|
| 151 |
+
self.camera_order = camera_order
|
| 152 |
+
self.num_frames = num_frames
|
| 153 |
+
self.resolution_hw = resolution_hw
|
| 154 |
+
self.fps_downsample_factor = fps_downsample_factor
|
| 155 |
+
self.caption_probability = caption_probability
|
| 156 |
+
self.camera_view_mapping = camera_view_mapping
|
| 157 |
+
self.camera_caption_key_mapping = camera_caption_key_mapping
|
| 158 |
+
self.camera_video_key_mapping = camera_video_key_mapping
|
| 159 |
+
self.camera_control_key_mapping = camera_control_key_mapping
|
| 160 |
+
self.add_view_prefix_to_caption = add_view_prefix_to_caption
|
| 161 |
+
self.camera_prefix_mapping = camera_prefix_mapping
|
| 162 |
+
self.single_caption_camera_name = single_caption_camera_name
|
| 163 |
+
self.window_random_frame_offset_range = window_random_frame_offset_range
|
| 164 |
+
|
| 165 |
+
if self.add_view_prefix_to_caption and self.camera_prefix_mapping is None:
|
| 166 |
+
raise ValueError("camera_prefix_mapping is required when add_view_prefix_to_caption is True")
|
| 167 |
+
|
| 168 |
+
if set(self.camera_caption_key_mapping.keys()) != set(self.camera_video_key_mapping.keys()):
|
| 169 |
+
raise ValueError(
|
| 170 |
+
f"Mismatching keys {set(self.camera_caption_key_mapping.keys())} != {set(self.camera_video_key_mapping.keys())}"
|
| 171 |
+
)
|
| 172 |
+
if self.camera_control_key_mapping is not None:
|
| 173 |
+
if set(self.camera_control_key_mapping.keys()) != set(self.camera_caption_key_mapping.keys()):
|
| 174 |
+
raise ValueError(
|
| 175 |
+
f"Mismatching keys {set(self.camera_control_key_mapping.keys())} != {set(self.camera_caption_key_mapping.keys())}"
|
| 176 |
+
)
|
| 177 |
+
for camera_name in self.camera_caption_key_mapping.keys():
|
| 178 |
+
if camera_name not in self.camera_view_mapping:
|
| 179 |
+
raise ValueError(f"Camera name {camera_name} not found in camera view mapping")
|
| 180 |
+
if self.single_caption_camera_name and self.single_caption_camera_name not in self.camera_order:
|
| 181 |
+
raise ValueError(
|
| 182 |
+
f"Single caption camera name {self.single_caption_camera_name} must appear in selected cameras"
|
| 183 |
+
)
|
| 184 |
+
|
| 185 |
+
if self.window_random_frame_offset_range is not None:
|
| 186 |
+
start_range, end_range = self.window_random_frame_offset_range
|
| 187 |
+
if start_range < 0 or end_range < 0:
|
| 188 |
+
raise ValueError("`window_random_frame_offset_range` must be non-negative")
|
| 189 |
+
if start_range > end_range:
|
| 190 |
+
raise ValueError("`window_random_frame_offset_range` start must be less than end")
|
| 191 |
+
|
| 192 |
+
def __call__(self, data: dict[str, Any]) -> dict[str, Any] | None:
|
| 193 |
+
"""Extract frames from a video."""
|
| 194 |
+
|
| 195 |
+
chunk_index, extracted_frame_ids, video_fps = None, None, None
|
| 196 |
+
(
|
| 197 |
+
captions,
|
| 198 |
+
multiview_frames,
|
| 199 |
+
multiview_control,
|
| 200 |
+
view_indices,
|
| 201 |
+
view_indices_selection,
|
| 202 |
+
camera_keys_selection,
|
| 203 |
+
original_sizes,
|
| 204 |
+
) = (
|
| 205 |
+
[],
|
| 206 |
+
[],
|
| 207 |
+
[],
|
| 208 |
+
[],
|
| 209 |
+
[],
|
| 210 |
+
[],
|
| 211 |
+
[],
|
| 212 |
+
)
|
| 213 |
+
|
| 214 |
+
for camera_name in self.camera_order:
|
| 215 |
+
video_key = self.camera_video_key_mapping[camera_name]
|
| 216 |
+
if self.single_caption_camera_name:
|
| 217 |
+
meta_key = self.camera_caption_key_mapping[self.single_caption_camera_name]
|
| 218 |
+
else:
|
| 219 |
+
meta_key = self.camera_caption_key_mapping[camera_name]
|
| 220 |
+
|
| 221 |
+
t2w_windows = data[meta_key]["t2w_windows"]
|
| 222 |
+
if chunk_index is None:
|
| 223 |
+
chunk_index = random.choice(list(range(len(t2w_windows))))
|
| 224 |
+
window = t2w_windows[chunk_index]
|
| 225 |
+
|
| 226 |
+
# extract caption
|
| 227 |
+
choices = list(self.caption_probability.keys())
|
| 228 |
+
weights = list(self.caption_probability.values())
|
| 229 |
+
caption_style = random.choices(choices, weights=weights)[0]
|
| 230 |
+
caption = ""
|
| 231 |
+
if self.single_caption_camera_name:
|
| 232 |
+
if camera_name == self.single_caption_camera_name:
|
| 233 |
+
caption = window[caption_style]
|
| 234 |
+
else:
|
| 235 |
+
caption = window[caption_style]
|
| 236 |
+
|
| 237 |
+
assert isinstance(caption, str), f"Caption is not a string: {caption}"
|
| 238 |
+
if self.add_view_prefix_to_caption:
|
| 239 |
+
caption = f"{self.camera_prefix_mapping[camera_name]} {caption}"
|
| 240 |
+
captions.append(caption)
|
| 241 |
+
|
| 242 |
+
# extract frames
|
| 243 |
+
random_offset = 0
|
| 244 |
+
if self.window_random_frame_offset_range is not None:
|
| 245 |
+
random_offset = random.randint(*self.window_random_frame_offset_range)
|
| 246 |
+
frame_start = window["start_frame"] + random_offset
|
| 247 |
+
frame_end = frame_start + self.num_frames * self.fps_downsample_factor
|
| 248 |
+
frame_indices = list(range(frame_start, frame_end, self.fps_downsample_factor))
|
| 249 |
+
try:
|
| 250 |
+
frames, original_fps, original_hw = self.extract_frames(
|
| 251 |
+
data[video_key], frame_indices, self.resolution_hw
|
| 252 |
+
)
|
| 253 |
+
except Exception as e:
|
| 254 |
+
log.error(f"Error extracting frames for camera {camera_name}: {e}")
|
| 255 |
+
return None
|
| 256 |
+
assert len(frames) == self.num_frames, f"Expected {self.num_frames} frames, got {len(frames)}"
|
| 257 |
+
multiview_frames.append(frames)
|
| 258 |
+
|
| 259 |
+
# check consistency between videos
|
| 260 |
+
if extracted_frame_ids is None:
|
| 261 |
+
extracted_frame_ids = frame_indices
|
| 262 |
+
elif frame_indices != extracted_frame_ids:
|
| 263 |
+
raise ValueError("Extracted frame IDs do not match")
|
| 264 |
+
|
| 265 |
+
if video_fps is None:
|
| 266 |
+
video_fps = original_fps
|
| 267 |
+
elif video_fps != original_fps:
|
| 268 |
+
raise ValueError("Video FPS does not match")
|
| 269 |
+
original_sizes.append(list(original_hw))
|
| 270 |
+
|
| 271 |
+
# extract control frames if available
|
| 272 |
+
if self.camera_control_key_mapping is not None:
|
| 273 |
+
control_key = self.camera_control_key_mapping[camera_name]
|
| 274 |
+
try:
|
| 275 |
+
control_frames, control_fps, _ = self.extract_frames(
|
| 276 |
+
data[control_key], frame_indices, self.resolution_hw
|
| 277 |
+
)
|
| 278 |
+
except Exception as e:
|
| 279 |
+
log.error(f"Error extracting control frames for camera {camera_name}: {e}")
|
| 280 |
+
return None
|
| 281 |
+
if len(control_frames) != self.num_frames:
|
| 282 |
+
raise ValueError(f"Expected {self.num_frames} frames, got {len(control_frames)}")
|
| 283 |
+
if control_fps != original_fps:
|
| 284 |
+
raise ValueError(f"Control FPS {control_fps} does not match video FPS {original_fps}")
|
| 285 |
+
multiview_control.append(control_frames)
|
| 286 |
+
|
| 287 |
+
view_indices.extend([self.camera_view_mapping[camera_name]] * self.num_frames)
|
| 288 |
+
view_indices_selection.append(self.camera_view_mapping[camera_name])
|
| 289 |
+
camera_keys_selection.append(camera_name)
|
| 290 |
+
|
| 291 |
+
front_cam_view_idx_sample_position = (
|
| 292 |
+
torch.tensor(self.camera_order.index(self.single_caption_camera_name), dtype=torch.int64)
|
| 293 |
+
if self.single_caption_camera_name
|
| 294 |
+
else None
|
| 295 |
+
)
|
| 296 |
+
if self.single_caption_camera_name and not self.add_view_prefix_to_caption:
|
| 297 |
+
captions = [captions[front_cam_view_idx_sample_position]]
|
| 298 |
+
|
| 299 |
+
if video_fps % self.fps_downsample_factor != 0:
|
| 300 |
+
raise ValueError("Original FPS is not divisible by FPS downsample factor")
|
| 301 |
+
fps = video_fps / self.fps_downsample_factor
|
| 302 |
+
|
| 303 |
+
sample = {
|
| 304 |
+
"__key__": data["__key__"],
|
| 305 |
+
"__url__": data["__url__"],
|
| 306 |
+
"video": rearrange(torch.cat(multiview_frames, dim=0), "t c h w -> c t h w"),
|
| 307 |
+
"ai_caption": captions,
|
| 308 |
+
"view_indices": torch.tensor(view_indices, dtype=torch.int64),
|
| 309 |
+
"fps": torch.tensor(fps, dtype=torch.float64),
|
| 310 |
+
"chunk_index": torch.tensor(chunk_index, dtype=torch.int64),
|
| 311 |
+
"frame_indices": torch.tensor(extracted_frame_ids, dtype=torch.int64),
|
| 312 |
+
"num_video_frames_per_view": torch.tensor(len(extracted_frame_ids), dtype=torch.int64),
|
| 313 |
+
"view_indices_selection": torch.tensor(view_indices_selection, dtype=torch.int64),
|
| 314 |
+
"camera_keys_selection": camera_keys_selection,
|
| 315 |
+
"sample_n_views": torch.tensor(len(camera_keys_selection), dtype=torch.int64),
|
| 316 |
+
"padding_mask": torch.zeros((1, *self.resolution_hw), dtype=torch.float32),
|
| 317 |
+
"ref_cam_view_idx_sample_position": torch.tensor(-1, dtype=torch.int64),
|
| 318 |
+
"front_cam_view_idx_sample_position": front_cam_view_idx_sample_position,
|
| 319 |
+
"original_hw": torch.tensor(original_sizes, dtype=torch.int64),
|
| 320 |
+
}
|
| 321 |
+
if self.camera_control_key_mapping is not None:
|
| 322 |
+
sample["control_input_hdmap_bbox"] = rearrange(torch.cat(multiview_control, dim=0), "t c h w -> c t h w")
|
| 323 |
+
return sample
|
| 324 |
+
|
| 325 |
+
@staticmethod
|
| 326 |
+
def extract_frames(
|
| 327 |
+
video: bytes, frame_indices: list[int], resolution_hw: tuple[int, int]
|
| 328 |
+
) -> tuple[torch.Tensor, float, tuple[int, int]]:
|
| 329 |
+
"""Extract frames from a video given start and end frame range."""
|
| 330 |
+
|
| 331 |
+
from decord import VideoReader
|
| 332 |
+
|
| 333 |
+
video_reader = VideoReader(io.BytesIO(video))
|
| 334 |
+
fps = video_reader.get_avg_fps()
|
| 335 |
+
frames = video_reader.get_batch(frame_indices).asnumpy()
|
| 336 |
+
frames = rearrange(torch.from_numpy(frames), "t h w c -> t c h w")
|
| 337 |
+
original_h, original_w = frames.shape[-2:]
|
| 338 |
+
return (
|
| 339 |
+
Resize(resolution_hw, interpolation=InterpolationMode.BILINEAR, antialias=True)(frames),
|
| 340 |
+
fps,
|
| 341 |
+
(original_h, original_w),
|
| 342 |
+
)
|
| 343 |
+
|
| 344 |
+
|
| 345 |
+
def get_multiview_dataset(
|
| 346 |
+
*,
|
| 347 |
+
dataset_name: str,
|
| 348 |
+
is_train: bool,
|
| 349 |
+
object_store: Literal["gcs", "s3"],
|
| 350 |
+
dataset_keys: list[str],
|
| 351 |
+
augmentations: dict[str, Augmentor],
|
| 352 |
+
dataset_catalog: dict[str, dict[str, list[str]]],
|
| 353 |
+
) -> Dataset:
|
| 354 |
+
"""Get video-text dataset with optional custom augmentation factory.
|
| 355 |
+
|
| 356 |
+
Args:
|
| 357 |
+
is_train: Whether this is for training
|
| 358 |
+
dataset_name: Name of dataset to use for loading wdinfo files
|
| 359 |
+
object_store: Object store to use ("gcs" or "s3")
|
| 360 |
+
dataset_keys: List of keys to use for loading dataset
|
| 361 |
+
augmentations: Augmentations map to apply to dataset
|
| 362 |
+
dataset_catalog: Dataset catalog to use for loading dataset
|
| 363 |
+
"""
|
| 364 |
+
|
| 365 |
+
dataset_info = get_video_dataset_info(
|
| 366 |
+
dataset_name,
|
| 367 |
+
object_store=object_store,
|
| 368 |
+
dataset_keys=dataset_keys,
|
| 369 |
+
dataset_catalog=dataset_catalog,
|
| 370 |
+
)
|
| 371 |
+
|
| 372 |
+
if (
|
| 373 |
+
USE_MEGATRON
|
| 374 |
+
and parallel_state.is_initialized()
|
| 375 |
+
and (
|
| 376 |
+
parallel_state.get_context_parallel_world_size() > 1
|
| 377 |
+
or parallel_state.get_tensor_model_parallel_world_size() > 1
|
| 378 |
+
)
|
| 379 |
+
):
|
| 380 |
+
distributor_fn = parallel_sync_multi_aspect_ratio.ShardlistMultiAspectRatioParallelSync
|
| 381 |
+
else:
|
| 382 |
+
distributor_fn = ShardlistBasic
|
| 383 |
+
|
| 384 |
+
video_data_config = DatasetConfig(
|
| 385 |
+
keys=[], # keys are defined per dataset
|
| 386 |
+
buffer_size=1,
|
| 387 |
+
streaming_download=True,
|
| 388 |
+
dataset_info=dataset_info,
|
| 389 |
+
distributor=distributor_fn(
|
| 390 |
+
shuffle=is_train,
|
| 391 |
+
split_by_node=True,
|
| 392 |
+
split_by_worker=True,
|
| 393 |
+
resume_flag=True,
|
| 394 |
+
verbose=False,
|
| 395 |
+
is_infinite_loader=is_train,
|
| 396 |
+
),
|
| 397 |
+
decoders=[
|
| 398 |
+
video_naive_bytes(),
|
| 399 |
+
json_decoder,
|
| 400 |
+
],
|
| 401 |
+
augmentation=augmentations,
|
| 402 |
+
remove_extension_from_keys=True,
|
| 403 |
+
)
|
| 404 |
+
|
| 405 |
+
return Dataset(
|
| 406 |
+
config=video_data_config,
|
| 407 |
+
handler=wds.warn_and_continue,
|
| 408 |
+
decoder_handler=wds.warn_and_continue,
|
| 409 |
+
detshuffle=False,
|
| 410 |
+
)
|
| 411 |
+
|
| 412 |
+
|
| 413 |
+
def collate_fn(batch: list[dict[str, Any]]) -> dict[str, Any]:
|
| 414 |
+
merged = dict()
|
| 415 |
+
is_tensor = dict()
|
| 416 |
+
for row in batch:
|
| 417 |
+
for key, value in row.items():
|
| 418 |
+
if key not in merged:
|
| 419 |
+
merged[key] = []
|
| 420 |
+
if isinstance(value, torch.Tensor):
|
| 421 |
+
is_tensor[key] = True
|
| 422 |
+
merged[key].append(value)
|
| 423 |
+
for key, value in merged.items():
|
| 424 |
+
if is_tensor.get(key, False):
|
| 425 |
+
merged[key] = torch.stack(value, dim=0)
|
| 426 |
+
return merged
|
| 427 |
+
|
| 428 |
+
|
| 429 |
+
@attrs.define(slots=False)
|
| 430 |
+
class AugmentationConfig:
|
| 431 |
+
"""Configuration for video augmentation."""
|
| 432 |
+
|
| 433 |
+
resolution_hw: tuple[int, int] = (1080, 1920)
|
| 434 |
+
fps_downsample_factor: int = 1
|
| 435 |
+
num_video_frames: int = 93
|
| 436 |
+
caption_probability: dict[str, float] = {
|
| 437 |
+
"qwen2p5_7b_caption": 0.7,
|
| 438 |
+
"qwen2p5_7b_caption_medium": 0.2,
|
| 439 |
+
"qwen2p5_7b_caption_short": 0.1,
|
| 440 |
+
}
|
| 441 |
+
camera_keys: tuple[CameraKeyType, ...] = DEFAULT_CAMERAS
|
| 442 |
+
camera_view_mapping: dict[CameraKeyType, int] = DEFAULT_CAMERA_VIEW_MAPPING
|
| 443 |
+
camera_caption_key_mapping: dict[CameraKeyType, str] = DEFAULT_CAPTION_KEY_MAPPING
|
| 444 |
+
camera_video_key_mapping: dict[CameraKeyType, str] = DEFAULT_VIDEO_KEY_MAPPING
|
| 445 |
+
camera_control_key_mapping: Optional[dict[CameraKeyType, str]] = None
|
| 446 |
+
position_to_camera_mapping: Optional[dict[int, CameraKeyType]] = None
|
| 447 |
+
add_view_prefix_to_caption: bool = False
|
| 448 |
+
camera_prefix_mapping: Optional[dict[CameraKeyType, str]] = DEFAULT_CAPTION_PREFIXES
|
| 449 |
+
single_caption_camera_name: Optional[CameraKeyType] = None
|
| 450 |
+
window_random_frame_offset_range: Optional[tuple[int, int]] = None
|
| 451 |
+
|
| 452 |
+
def __attrs_post_init__(self) -> None:
|
| 453 |
+
"""Post initialization checks for camera keys consistency."""
|
| 454 |
+
|
| 455 |
+
for camera_key in self.camera_keys:
|
| 456 |
+
for attr_name in [
|
| 457 |
+
"camera_view_mapping",
|
| 458 |
+
"camera_caption_key_mapping",
|
| 459 |
+
"camera_video_key_mapping",
|
| 460 |
+
"camera_control_key_mapping",
|
| 461 |
+
"camera_prefix_mapping",
|
| 462 |
+
]:
|
| 463 |
+
attr = getattr(self, attr_name)
|
| 464 |
+
if attr is not None:
|
| 465 |
+
if camera_key not in attr:
|
| 466 |
+
raise ValueError(f"Camera key {camera_key} not found in `{attr_name}` mapping!")
|
| 467 |
+
if self.single_caption_camera_name is not None:
|
| 468 |
+
if self.single_caption_camera_name not in self.camera_keys:
|
| 469 |
+
raise ValueError(
|
| 470 |
+
f"Single caption camera key {self.single_caption_camera_name} not found in camera keys!"
|
| 471 |
+
)
|
| 472 |
+
|
| 473 |
+
|
| 474 |
+
def make_augmentations(augmentation_config: AugmentationConfig) -> tuple[dict[str, Augmentor], list[str]]:
|
| 475 |
+
"""Make augmentations for multiview video dataset."""
|
| 476 |
+
|
| 477 |
+
augmentations = dict()
|
| 478 |
+
if augmentation_config.position_to_camera_mapping is not None:
|
| 479 |
+
augmentations["unpack_metas"] = UnpackMetas(
|
| 480 |
+
position_to_camera_mapping=augmentation_config.position_to_camera_mapping
|
| 481 |
+
)
|
| 482 |
+
|
| 483 |
+
augmentations["extract_frames_and_captions"] = ExtractFramesAndCaptions(
|
| 484 |
+
camera_order=augmentation_config.camera_keys,
|
| 485 |
+
num_frames=augmentation_config.num_video_frames,
|
| 486 |
+
resolution_hw=augmentation_config.resolution_hw,
|
| 487 |
+
fps_downsample_factor=augmentation_config.fps_downsample_factor,
|
| 488 |
+
caption_probability=augmentation_config.caption_probability,
|
| 489 |
+
camera_view_mapping=augmentation_config.camera_view_mapping,
|
| 490 |
+
camera_caption_key_mapping=augmentation_config.camera_caption_key_mapping,
|
| 491 |
+
camera_video_key_mapping=augmentation_config.camera_video_key_mapping,
|
| 492 |
+
camera_control_key_mapping=augmentation_config.camera_control_key_mapping,
|
| 493 |
+
add_view_prefix_to_caption=augmentation_config.add_view_prefix_to_caption,
|
| 494 |
+
camera_prefix_mapping=augmentation_config.camera_prefix_mapping,
|
| 495 |
+
single_caption_camera_name=augmentation_config.single_caption_camera_name,
|
| 496 |
+
window_random_frame_offset_range=augmentation_config.window_random_frame_offset_range,
|
| 497 |
+
)
|
| 498 |
+
|
| 499 |
+
# define dataset keys to load
|
| 500 |
+
dataset_keys = list(augmentation_config.camera_video_key_mapping.values())
|
| 501 |
+
if augmentation_config.position_to_camera_mapping is not None:
|
| 502 |
+
dataset_keys.append("metas")
|
| 503 |
+
else:
|
| 504 |
+
dataset_keys.extend(augmentation_config.camera_caption_key_mapping.values())
|
| 505 |
+
if augmentation_config.camera_control_key_mapping is not None:
|
| 506 |
+
dataset_keys.extend(augmentation_config.camera_control_key_mapping.values())
|
| 507 |
+
|
| 508 |
+
return augmentations, dataset_keys
|
| 509 |
+
|
| 510 |
+
|
| 511 |
+
def get_multiview_video_loader(
|
| 512 |
+
*,
|
| 513 |
+
dataset_name: str,
|
| 514 |
+
is_train: bool,
|
| 515 |
+
object_store: Literal["gcs", "s3"] = "s3",
|
| 516 |
+
augmentation_config: AugmentationConfig = AugmentationConfig(),
|
| 517 |
+
batch_size: int = 1,
|
| 518 |
+
num_workers: int = 4,
|
| 519 |
+
prefetch_factor: int | None = 1,
|
| 520 |
+
**kwargs: Any,
|
| 521 |
+
):
|
| 522 |
+
"""Get video loader for alpamayo multiview dataset
|
| 523 |
+
pass kwargs to tolerate `dataloaders` from inheritance
|
| 524 |
+
"""
|
| 525 |
+
|
| 526 |
+
# make augmentations
|
| 527 |
+
augmentations, dataset_keys = make_augmentations(augmentation_config)
|
| 528 |
+
|
| 529 |
+
# get dataloader
|
| 530 |
+
return get_cached_replay_dataloader(
|
| 531 |
+
dataset=get_multiview_dataset(
|
| 532 |
+
is_train=is_train,
|
| 533 |
+
object_store=object_store,
|
| 534 |
+
dataset_name=dataset_name,
|
| 535 |
+
dataset_keys=dataset_keys,
|
| 536 |
+
dataset_catalog=DEFAULT_CATALOG,
|
| 537 |
+
augmentations=augmentations,
|
| 538 |
+
),
|
| 539 |
+
num_workers=num_workers,
|
| 540 |
+
batch_size=batch_size,
|
| 541 |
+
sampler=None,
|
| 542 |
+
prefetch_factor=prefetch_factor if num_workers > 0 else None,
|
| 543 |
+
persistent_workers=num_workers > 0,
|
| 544 |
+
pin_memory=False,
|
| 545 |
+
collate_fn=collate_fn,
|
| 546 |
+
cache_replay_name="video_dataloader",
|
| 547 |
+
)
|
cosmos_predict2/_src/predict2_multiview/datasets/nymeria_pairs.py
ADDED
|
@@ -0,0 +1,1074 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
"""Nymeria 2-actor paired clip dataset for joint multi-view video generation.
|
| 17 |
+
|
| 18 |
+
Driven by ``manifest/train_ready.csv`` (only pairs whose BOTH clips finished person+hands warping). Each row
|
| 19 |
+
references two per-actor ``clip.npz`` files (``actor1_npz`` / ``actor2_npz``); each npz holds raw arrays:
|
| 20 |
+
* ``target_rgb`` (T, H, W, 3) uint8 -- ground-truth video to generate
|
| 21 |
+
* ``pose`` (T, H, W, 3) uint8 -- projected 2D pose RGB (driving pose)
|
| 22 |
+
* ``warped_cond`` (T, H, W, 3) uint8 -- warped past-frame conditioning (pixel RGB)
|
| 23 |
+
* ``visibility_packed`` (packbits) + ``vis_shape`` (T, H, W) -- warp-validity mask (bool -> 1ch)
|
| 24 |
+
* ``keypoints`` (T, 2, 23, 3), ``src_frame`` (T,) -- unused here
|
| 25 |
+
|
| 26 |
+
(Falls back to per-clip mp4 files if a CSV row has no ``*_npz`` columns, e.g. the old ``clip_pairs.csv``.)
|
| 27 |
+
|
| 28 |
+
Actor 1 / Actor 2 are emitted as the two "views" (V=2): every tensor is laid out ``(C, V*T, H, W)`` with the
|
| 29 |
+
two actors concatenated along the temporal axis. ``actor{1,2}_n_extra`` count untracked 3rd parties; set
|
| 30 |
+
``clean_only=True`` to keep only pairs with no 3rd party (n_extra==0 for both actors).
|
| 31 |
+
"""
|
| 32 |
+
|
| 33 |
+
import csv
|
| 34 |
+
import os
|
| 35 |
+
import random
|
| 36 |
+
import re
|
| 37 |
+
from typing import Any, Optional
|
| 38 |
+
|
| 39 |
+
import attrs
|
| 40 |
+
import numpy as np
|
| 41 |
+
import torch
|
| 42 |
+
from einops import rearrange
|
| 43 |
+
from hydra.core.config_store import ConfigStore
|
| 44 |
+
from torchvision.transforms import InterpolationMode, Resize
|
| 45 |
+
|
| 46 |
+
from cosmos_predict2._src.imaginaire.lazy_config import LazyCall as L
|
| 47 |
+
from cosmos_predict2._src.imaginaire.utils import log
|
| 48 |
+
from cosmos_predict2._src.predict2.datasets.cached_replay_dataloader import get_cached_replay_dataloader
|
| 49 |
+
from cosmos_predict2._src.predict2_multiview.datasets.multiview import collate_fn
|
| 50 |
+
|
| 51 |
+
# npz array keys
|
| 52 |
+
_NPZ_KEYS = {"video": "target_rgb", "control_input_pose": "pose", "control_input_warped": "warped_cond"}
|
| 53 |
+
# mp4 fallback files
|
| 54 |
+
_VIDEO_FILES = {
|
| 55 |
+
"video": "target.mp4",
|
| 56 |
+
"control_input_pose": "pose.mp4",
|
| 57 |
+
"control_input_warped": "warped_cond.mp4",
|
| 58 |
+
"control_input_visibility": "visibility.mp4",
|
| 59 |
+
}
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def _resize_thwc_uint8(frames_thwc: np.ndarray, num_frames: int, resolution_hw: tuple[int, int]) -> torch.Tensor:
|
| 63 |
+
"""(T, H, W, C) uint8 -> resized (C, T, H', W') uint8, padded/truncated to num_frames."""
|
| 64 |
+
t = torch.from_numpy(np.ascontiguousarray(frames_thwc)) # (T,H,W,C)
|
| 65 |
+
t = rearrange(t, "t h w c -> t c h w")
|
| 66 |
+
n = t.shape[0]
|
| 67 |
+
if n < num_frames:
|
| 68 |
+
t = torch.cat([t, t[-1:].repeat(num_frames - n, 1, 1, 1)], dim=0)
|
| 69 |
+
elif n > num_frames:
|
| 70 |
+
t = t[:num_frames]
|
| 71 |
+
t = Resize(resolution_hw, interpolation=InterpolationMode.BILINEAR, antialias=True)(t.float()).to(torch.uint8)
|
| 72 |
+
return rearrange(t, "t c h w -> c t h w").contiguous()
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
def _read_video_frames(path: str, num_frames: int, resolution_hw: tuple[int, int]) -> torch.Tensor:
|
| 76 |
+
"""mp4 fallback: first num_frames frames -> (C, T, H', W') uint8."""
|
| 77 |
+
from decord import VideoReader
|
| 78 |
+
|
| 79 |
+
vr = VideoReader(path)
|
| 80 |
+
n = min(num_frames, len(vr))
|
| 81 |
+
frames = vr.get_batch(list(range(n))).asnumpy() # (n,H,W,C)
|
| 82 |
+
return _resize_thwc_uint8(frames, num_frames, resolution_hw)
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
@attrs.define(slots=False)
|
| 86 |
+
class NymeriaPairsConfig:
|
| 87 |
+
root: str = "/data2/nymeria_processed_longer"
|
| 88 |
+
manifest_csv: str = "manifest/train_ready.csv"
|
| 89 |
+
resolution_hw: tuple[int, int] = (480, 480)
|
| 90 |
+
num_video_frames: int = 77 # target frames per actor (77 for longer/train_ready)
|
| 91 |
+
fps: float = 10.0
|
| 92 |
+
caption: str = "a person performing an everyday activity" # placeholder; text encoder is frozen
|
| 93 |
+
clean_only: bool = False # keep only pairs with no untracked 3rd party (n_extra==0 for both actors)
|
| 94 |
+
single_view: bool = False # stage-1 pose pretraining: emit each actor as an independent V=1 sample
|
| 95 |
+
num_reference_frames: int = 0 # >0: emit `reference_frames` (R clean source frames/view) for appearance
|
| 96 |
+
# conditioning (in-context reference tokens). Loaded from clip.npz
|
| 97 |
+
# `reference_frames` or a sibling refs.npz (see backfill_refs.py).
|
| 98 |
+
shared_reference: bool = False # SHARED refs: instead of per-view pools, load ONE greedy-selected set from
|
| 99 |
+
# refs_shared.npz (identical in both clip dirs; V*R frames + w2c) and hand
|
| 100 |
+
# view-slot v its contiguous R-slice (top-gain first). Refs carry no view
|
| 101 |
+
# identity (net drops the ref view-embedding); both views attend all via
|
| 102 |
+
# cross-view self-attn. Pairs with net.shared_reference=True.
|
| 103 |
+
emit_camera_poses: bool = False # emit per-view camera_w2c/K for Plücker ray conditioning (campose.npz)
|
| 104 |
+
emit_depth: bool = False # emit `control_input_depth` = composite depth RGB (mesh_cond.npz `depth_comp`:
|
| 105 |
+
# warped scene depth + human mesh depth) for the net's VAE depth conditioning
|
| 106 |
+
person_pose: bool = False # use IDENTITY-colored pose (npz `pose_person`) instead of role-colored `pose`
|
| 107 |
+
emit_reference_pose: bool = False # emit `reference_pose` = per-person skeleton of each shared ref (refs_shared.npz)
|
| 108 |
+
role_mix: bool = False # actor-observer 2-view: randomly swap view0<->view1 (+captions) 50% so the
|
| 109 |
+
# view embedding can't memorize actor/observer roles
|
| 110 |
+
# train/val split — held out by SESSION (=(actor1_seq, actor2_seq)) so adjacent tiled windows of the same
|
| 111 |
+
# recording never leak across the split. "all" uses everything (no split).
|
| 112 |
+
split: str = "all" # "train" | "val" | "all"
|
| 113 |
+
val_num_sessions: int = 4 # number of sessions held out for validation (~10% of 36 sessions)
|
| 114 |
+
split_seed: int = 1234 # deterministic session assignment
|
| 115 |
+
val_max_pairs: int = 40 # cap the val set size for fast periodic validation (0 = uncapped)
|
| 116 |
+
# smoke/offline mode: emit zero text embeddings so the (S3-only) text encoder can be disabled.
|
| 117 |
+
# text_tokens_per_view MUST be 512 (MultiViewCrossAttention derives n_views as context_len // 512).
|
| 118 |
+
dummy_text_embeddings: bool = False
|
| 119 |
+
text_embed_dim: int = 1024
|
| 120 |
+
text_tokens_per_view: int = 512
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
class NymeriaPairsDataset(torch.utils.data.Dataset):
|
| 124 |
+
def __init__(self, config: NymeriaPairsConfig):
|
| 125 |
+
self.config = config
|
| 126 |
+
self.root = config.root
|
| 127 |
+
self.clips_dir = os.path.join(self.root, "clips")
|
| 128 |
+
self.resolution_hw = tuple(config.resolution_hw)
|
| 129 |
+
self.num_frames = int(config.num_video_frames)
|
| 130 |
+
manifest_path = os.path.join(self.root, config.manifest_csv)
|
| 131 |
+
|
| 132 |
+
# 1) collect ALL usable candidate rows (with session + clean flag). The clean filter is applied AFTER
|
| 133 |
+
# the session split so val_sessions are computed from the same full universe for train and val.
|
| 134 |
+
candidates: list[dict[str, str]] = []
|
| 135 |
+
with open(manifest_path, newline="") as f:
|
| 136 |
+
for row in csv.DictReader(f):
|
| 137 |
+
session = f"{row.get('actor1_seq', '')}__{row.get('actor2_seq', '')}"
|
| 138 |
+
clean = int(row.get("actor1_n_extra", 0)) == 0 and int(row.get("actor2_n_extra", 0)) == 0
|
| 139 |
+
a1_npz, a2_npz = row.get("actor1_npz", ""), row.get("actor2_npz", "")
|
| 140 |
+
if a1_npz and a2_npz: # npz path (train_ready.csv)
|
| 141 |
+
if os.path.isfile(a1_npz) and os.path.isfile(a2_npz):
|
| 142 |
+
candidates.append({"pair_id": row["pair_id"], "a1": a1_npz, "a2": a2_npz, "mode": "npz", "session": session, "clean": clean})
|
| 143 |
+
else: # mp4 fallback (clip_pairs.csv)
|
| 144 |
+
a1 = os.path.join(self.clips_dir, row["actor1_clip"])
|
| 145 |
+
a2 = os.path.join(self.clips_dir, row["actor2_clip"])
|
| 146 |
+
if self._mp4_complete(a1) and self._mp4_complete(a2):
|
| 147 |
+
candidates.append({"pair_id": row["pair_id"], "a1": a1, "a2": a2, "mode": "mp4", "session": session, "clean": clean})
|
| 148 |
+
|
| 149 |
+
# 2) deterministic session-level holdout from the FULL session universe (no temporal-window leakage,
|
| 150 |
+
# and train/val are guaranteed disjoint by session regardless of clean_only)
|
| 151 |
+
import random as _random
|
| 152 |
+
|
| 153 |
+
sessions = sorted({c["session"] for c in candidates})
|
| 154 |
+
rng = _random.Random(config.split_seed)
|
| 155 |
+
rng.shuffle(sessions)
|
| 156 |
+
val_sessions = set(sessions[: config.val_num_sessions])
|
| 157 |
+
if config.split == "train":
|
| 158 |
+
candidates = [c for c in candidates if c["session"] not in val_sessions]
|
| 159 |
+
elif config.split == "val":
|
| 160 |
+
candidates = [c for c in candidates if c["session"] in val_sessions]
|
| 161 |
+
|
| 162 |
+
# 3) clean filter (after split)
|
| 163 |
+
n_extra_dropped = 0
|
| 164 |
+
if config.clean_only:
|
| 165 |
+
before = len(candidates)
|
| 166 |
+
candidates = [c for c in candidates if c["clean"]]
|
| 167 |
+
n_extra_dropped = before - len(candidates)
|
| 168 |
+
|
| 169 |
+
# 3b) single-view (stage-1): expand each pair into two independent V=1 samples (one per actor)
|
| 170 |
+
if config.single_view:
|
| 171 |
+
singles = []
|
| 172 |
+
for c in candidates:
|
| 173 |
+
for k in ("a1", "a2"):
|
| 174 |
+
singles.append({"pair_id": f"{c['pair_id']}__{k}", "a1": c[k], "mode": c["mode"], "session": c["session"]})
|
| 175 |
+
candidates = singles
|
| 176 |
+
|
| 177 |
+
# 4) val cap (after expansion)
|
| 178 |
+
if config.split == "val" and config.val_max_pairs and len(candidates) > config.val_max_pairs:
|
| 179 |
+
candidates.sort(key=lambda c: c["pair_id"])
|
| 180 |
+
step = len(candidates) / config.val_max_pairs
|
| 181 |
+
candidates = [candidates[int(i * step)] for i in range(config.val_max_pairs)]
|
| 182 |
+
|
| 183 |
+
self.pairs = candidates
|
| 184 |
+
log.info(
|
| 185 |
+
f"NymeriaPairsDataset[{config.split}]: {len(self.pairs)} pairs from {manifest_path} "
|
| 186 |
+
f"(clean_only={config.clean_only}, {len(sessions)} sessions, {len(val_sessions)} held out for val, "
|
| 187 |
+
f"dropped {n_extra_dropped} for 3rd-party)"
|
| 188 |
+
)
|
| 189 |
+
if len(self.pairs) == 0:
|
| 190 |
+
raise RuntimeError(f"No usable clip pairs for split={config.split} in {manifest_path}")
|
| 191 |
+
|
| 192 |
+
@staticmethod
|
| 193 |
+
def _mp4_complete(clip_dir: str) -> bool:
|
| 194 |
+
return os.path.isdir(clip_dir) and all(
|
| 195 |
+
os.path.isfile(os.path.join(clip_dir, fn)) for fn in _VIDEO_FILES.values()
|
| 196 |
+
)
|
| 197 |
+
|
| 198 |
+
def __len__(self) -> int:
|
| 199 |
+
return len(self.pairs)
|
| 200 |
+
|
| 201 |
+
def _load_actor(self, ref: str, mode: str, slot: int = 0) -> dict[str, torch.Tensor]:
|
| 202 |
+
out: dict[str, torch.Tensor] = {}
|
| 203 |
+
if mode == "npz":
|
| 204 |
+
d = np.load(ref, allow_pickle=True)
|
| 205 |
+
for key, npz_key in _NPZ_KEYS.items():
|
| 206 |
+
out[key] = _resize_thwc_uint8(d[npz_key], self.num_frames, self.resolution_hw)
|
| 207 |
+
if getattr(self.config, "person_pose", False) and "pose_person" in d.files: # identity-colored pose
|
| 208 |
+
out["control_input_pose"] = _resize_thwc_uint8(d["pose_person"], self.num_frames, self.resolution_hw)
|
| 209 |
+
if getattr(self.config, "emit_depth", False): # composite depth RGB (mesh_cond.npz depth_comp)
|
| 210 |
+
mc = os.path.join(os.path.dirname(ref), "mesh_cond.npz")
|
| 211 |
+
if os.path.isfile(mc):
|
| 212 |
+
dc = np.load(mc)["depth_comp"]
|
| 213 |
+
else: # clips still lacking mesh_cond (Stage-B in progress) -> zero depth (shape-consistent for concat)
|
| 214 |
+
dc = np.zeros((self.num_frames, self.resolution_hw[0], self.resolution_hw[1], 3), np.uint8)
|
| 215 |
+
out["control_input_depth"] = _resize_thwc_uint8(dc, self.num_frames, self.resolution_hw)
|
| 216 |
+
# visibility: unpack bits -> (T,H,W) -> (T,H,W,1)
|
| 217 |
+
vt, vh, vw = (int(x) for x in d["vis_shape"])
|
| 218 |
+
vis = np.unpackbits(d["visibility_packed"])[: vt * vh * vw].reshape(vt, vh, vw)
|
| 219 |
+
vis = (vis.astype(np.uint8) * 255)[..., None] # (T,H,W,1) in {0,255}
|
| 220 |
+
out["control_input_visibility"] = _resize_thwc_uint8(vis, self.num_frames, self.resolution_hw)
|
| 221 |
+
# reference frames (clean source frames) for appearance conditioning -> (3, R, H, W) uint8
|
| 222 |
+
R = int(getattr(self.config, "num_reference_frames", 0))
|
| 223 |
+
if R > 0 and getattr(self.config, "shared_reference", False):
|
| 224 |
+
# SHARED refs: one greedy set (V*R frames, top-gain first) identical in both clip dirs; hand
|
| 225 |
+
# THIS view-slot its contiguous R-slice. Refs carry no view identity (net drops ref view-emb).
|
| 226 |
+
sib = os.path.join(os.path.dirname(ref), "refs_shared.npz")
|
| 227 |
+
if not os.path.isfile(sib):
|
| 228 |
+
raise KeyError(f"no refs_shared for {ref} (run backfill_shared_refs.py)")
|
| 229 |
+
z = np.load(sib)
|
| 230 |
+
allf = z["reference_frames"] # (V*R, h, h, 3)
|
| 231 |
+
lo = slot * R
|
| 232 |
+
if lo + R > allf.shape[0]:
|
| 233 |
+
raise KeyError(f"refs_shared has {allf.shape[0]} < {lo+R} frames for slot {slot} R={R}")
|
| 234 |
+
out["reference_frames"] = _resize_thwc_uint8(allf[lo : lo + R], R, self.resolution_hw)
|
| 235 |
+
if getattr(self.config, "emit_camera_poses", False):
|
| 236 |
+
rw2c = z["reference_w2c"].astype(np.float32)[lo : lo + R] # (R,4,4) source-frame world->cam
|
| 237 |
+
out["cam_ref_w2c"] = torch.from_numpy(rw2c)
|
| 238 |
+
if getattr(self.config, "emit_reference_pose", False) and "reference_pose" in z.files:
|
| 239 |
+
# per-person skeleton render of each shared ref (this view-slot's R-slice), same layout as refs
|
| 240 |
+
out["reference_pose"] = _resize_thwc_uint8(z["reference_pose"][lo : lo + R], R, self.resolution_hw)
|
| 241 |
+
elif R > 0:
|
| 242 |
+
if "reference_frames" in d.files: # clip.npz native (new warps)
|
| 243 |
+
ref_frames = d["reference_frames"]
|
| 244 |
+
else: # sibling refs.npz (backfilled)
|
| 245 |
+
sib = os.path.join(os.path.dirname(ref), "refs.npz")
|
| 246 |
+
if not os.path.isfile(sib):
|
| 247 |
+
raise KeyError(f"no reference_frames for {ref} (run backfill_refs.py)")
|
| 248 |
+
ref_frames = np.load(sib)["reference_frames"]
|
| 249 |
+
# stored refs may hold != R frames; pick R (coverage-spaced if more, pad-last if fewer). The
|
| 250 |
+
# SAME `sel` is reused for the reference camera poses so frames <-> poses stay aligned.
|
| 251 |
+
k_stored = ref_frames.shape[0]
|
| 252 |
+
if k_stored >= R:
|
| 253 |
+
sel = np.linspace(0, k_stored - 1, R).round().astype(int)
|
| 254 |
+
else:
|
| 255 |
+
sel = np.concatenate([np.arange(k_stored), np.full(R - k_stored, k_stored - 1, dtype=int)])
|
| 256 |
+
out["reference_frames"] = _resize_thwc_uint8(ref_frames[sel], R, self.resolution_hw)
|
| 257 |
+
# posed-reference Plücker: past camera poses of the reference frames (aligned via `sel`)
|
| 258 |
+
if getattr(self.config, "emit_camera_poses", False):
|
| 259 |
+
rp = os.path.join(os.path.dirname(ref), "refpose.npz")
|
| 260 |
+
if not os.path.isfile(rp):
|
| 261 |
+
raise KeyError(f"no refpose for {ref} (run backfill_refpose.py)")
|
| 262 |
+
rw2c = np.load(rp)["reference_w2c"].astype(np.float32)[sel] # (R,4,4)
|
| 263 |
+
out["cam_ref_w2c"] = torch.from_numpy(rw2c)
|
| 264 |
+
# camera poses (target-frame w2c + K, shared SLAM world) for Plücker ray conditioning
|
| 265 |
+
if getattr(self.config, "emit_camera_poses", False):
|
| 266 |
+
cp = os.path.join(os.path.dirname(ref), "campose.npz")
|
| 267 |
+
if not os.path.isfile(cp):
|
| 268 |
+
raise KeyError(f"no campose for {ref} (run backfill_campose.py)")
|
| 269 |
+
cpz = np.load(cp)
|
| 270 |
+
w2c = cpz["w2c"].astype(np.float32) # (T,4,4) target-frame world->cam
|
| 271 |
+
n = w2c.shape[0]
|
| 272 |
+
if n < self.num_frames: # pad by repeating last frame's pose
|
| 273 |
+
w2c = np.concatenate([w2c, np.repeat(w2c[-1:], self.num_frames - n, 0)], 0)
|
| 274 |
+
elif n > self.num_frames:
|
| 275 |
+
w2c = w2c[: self.num_frames]
|
| 276 |
+
out["cam_w2c"] = torch.from_numpy(w2c) # (T,4,4)
|
| 277 |
+
out["cam_K"] = torch.from_numpy(cpz["K"].astype(np.float32)) # (3,3)
|
| 278 |
+
out["cam_src_res"] = torch.tensor(int(cpz["src_res"]), dtype=torch.int64)
|
| 279 |
+
else:
|
| 280 |
+
for key, fn in _VIDEO_FILES.items():
|
| 281 |
+
out[key] = _read_video_frames(os.path.join(ref, fn), self.num_frames, self.resolution_hw)
|
| 282 |
+
return out
|
| 283 |
+
|
| 284 |
+
def __getitem__(self, idx: int) -> dict[str, Any]:
|
| 285 |
+
# robust to corrupt / mid-write clip.npz (warping may still be running): skip + retry another sample
|
| 286 |
+
for _ in range(16):
|
| 287 |
+
try:
|
| 288 |
+
return self._build_sample(idx)
|
| 289 |
+
except Exception as e:
|
| 290 |
+
log.warning(f"NymeriaPairsDataset: bad sample {self.pairs[idx]['pair_id']} ({type(e).__name__}); retrying")
|
| 291 |
+
idx = random.randrange(len(self.pairs))
|
| 292 |
+
raise RuntimeError("NymeriaPairsDataset: too many unreadable samples in a row")
|
| 293 |
+
|
| 294 |
+
def _build_sample(self, idx: int) -> dict[str, Any]:
|
| 295 |
+
pair = self.pairs[idx]
|
| 296 |
+
# single-view: one actor (V=1); paired: both actors stacked on the temporal axis (V=2)
|
| 297 |
+
actors = [self._load_actor(pair["a1"], pair["mode"], slot=0)]
|
| 298 |
+
if not self.config.single_view:
|
| 299 |
+
actors.append(self._load_actor(pair["a2"], pair["mode"], slot=1))
|
| 300 |
+
n_views = len(actors)
|
| 301 |
+
|
| 302 |
+
sample: dict[str, Any] = {}
|
| 303 |
+
for key in actors[0]: # concat actors along temporal axis -> (C, V*T, H, W)
|
| 304 |
+
if key.startswith("cam_"): # per-view metadata, stacked on the view dim below
|
| 305 |
+
continue
|
| 306 |
+
sample[key] = torch.cat([a[key] for a in actors], dim=1).contiguous()
|
| 307 |
+
if "cam_w2c" in actors[0]:
|
| 308 |
+
sample["camera_w2c"] = torch.stack([a["cam_w2c"] for a in actors], dim=0).contiguous()
|
| 309 |
+
sample["camera_K"] = torch.stack([a["cam_K"] for a in actors], dim=0).contiguous()
|
| 310 |
+
sample["camera_src_res"] = torch.stack([a["cam_src_res"] for a in actors], dim=0).contiguous()
|
| 311 |
+
if "cam_ref_w2c" in actors[0]: # reference-frame past poses: (V, R, 4, 4)
|
| 312 |
+
sample["reference_cam_w2c"] = torch.stack([a["cam_ref_w2c"] for a in actors], dim=0).contiguous()
|
| 313 |
+
|
| 314 |
+
T = self.num_frames
|
| 315 |
+
view_indices = [v for v in range(n_views) for _ in range(T)]
|
| 316 |
+
sample.update(
|
| 317 |
+
{
|
| 318 |
+
"__key__": pair["pair_id"],
|
| 319 |
+
"__url__": pair["pair_id"],
|
| 320 |
+
"ai_caption": [self.config.caption for _ in range(n_views)],
|
| 321 |
+
"view_indices": torch.tensor(view_indices, dtype=torch.int64),
|
| 322 |
+
"fps": torch.tensor(self.config.fps, dtype=torch.float64),
|
| 323 |
+
"num_video_frames_per_view": torch.tensor(T, dtype=torch.int64),
|
| 324 |
+
"view_indices_selection": torch.tensor(list(range(n_views)), dtype=torch.int64),
|
| 325 |
+
"camera_keys_selection": [f"actor{v + 1}" for v in range(n_views)],
|
| 326 |
+
"sample_n_views": torch.tensor(n_views, dtype=torch.int64),
|
| 327 |
+
"padding_mask": torch.zeros((1, *self.resolution_hw), dtype=torch.float32),
|
| 328 |
+
"ref_cam_view_idx_sample_position": torch.tensor(-1, dtype=torch.int64),
|
| 329 |
+
"front_cam_view_idx_sample_position": torch.tensor(0, dtype=torch.int64),
|
| 330 |
+
}
|
| 331 |
+
)
|
| 332 |
+
if self.config.dummy_text_embeddings:
|
| 333 |
+
n_tok = n_views * self.config.text_tokens_per_view
|
| 334 |
+
sample["t5_text_embeddings"] = torch.zeros(n_tok, self.config.text_embed_dim, dtype=torch.float32)
|
| 335 |
+
sample["t5_text_mask"] = torch.ones(n_tok, dtype=torch.float32)
|
| 336 |
+
return sample
|
| 337 |
+
|
| 338 |
+
|
| 339 |
+
class NymeriaActorObserverDataset(NymeriaPairsDataset):
|
| 340 |
+
"""2-view actor-observer dataset (e.g. /data2/nymeria_processed_single).
|
| 341 |
+
|
| 342 |
+
Each row of the manifest = one actor(head) clip + one observer clip + a `text` caption of the actor's
|
| 343 |
+
action. Reads a pre-split manifest (train_split.csv / val_split.csv) directly. Emits V=2 with:
|
| 344 |
+
* role mixing (``role_mix``): 50% of the time view0/view1 (and their captions) are swapped, so the
|
| 345 |
+
view embedding cannot encode "view0=actor / view1=observer".
|
| 346 |
+
* per-view captions for per-view text cross-attention:
|
| 347 |
+
actor view -> the CSV `text`
|
| 348 |
+
observer view -> "C is observing the partner. The partner <text with 'C' removed>"
|
| 349 |
+
Real captions are emitted in ``ai_caption`` (no dummy embeddings): the model's online multiview text
|
| 350 |
+
encoder computes one 512-token embedding per view.
|
| 351 |
+
"""
|
| 352 |
+
|
| 353 |
+
def __init__(self, config: NymeriaPairsConfig):
|
| 354 |
+
self.config = config
|
| 355 |
+
self.root = config.root
|
| 356 |
+
self.clips_dir = os.path.join(self.root, "clips")
|
| 357 |
+
self.resolution_hw = tuple(config.resolution_hw)
|
| 358 |
+
self.num_frames = int(config.num_video_frames)
|
| 359 |
+
manifest_path = os.path.join(self.root, config.manifest_csv)
|
| 360 |
+
self.pairs: list[dict[str, str]] = []
|
| 361 |
+
with open(manifest_path, newline="") as f:
|
| 362 |
+
for row in csv.DictReader(f):
|
| 363 |
+
actor_npz = os.path.join(self.clips_dir, row["actor1_clip"], "clip.npz") # head/actor
|
| 364 |
+
obs_npz = os.path.join(self.clips_dir, row["actor2_clip"], "clip.npz") # observer
|
| 365 |
+
if os.path.isfile(actor_npz) and os.path.isfile(obs_npz):
|
| 366 |
+
self.pairs.append(
|
| 367 |
+
{"pair_id": row["pair_id"], "actor": actor_npz, "observer": obs_npz, "text": row.get("text", "")}
|
| 368 |
+
)
|
| 369 |
+
log.info(f"NymeriaActorObserverDataset: {len(self.pairs)} pairs from {manifest_path} (role_mix={config.role_mix})")
|
| 370 |
+
if len(self.pairs) == 0:
|
| 371 |
+
raise RuntimeError(f"No usable actor-observer pairs in {manifest_path}")
|
| 372 |
+
|
| 373 |
+
@staticmethod
|
| 374 |
+
def _observer_caption(actor_text: str) -> str:
|
| 375 |
+
# frame the actor action from the observer's view: subsequent "C" subjects become "the partner"
|
| 376 |
+
# (kept as subjects so multi-sentence actions stay grammatical); the first is dropped since the
|
| 377 |
+
# template already prepends "The partner".
|
| 378 |
+
action = re.sub(r"\bC\b", "the partner", actor_text).strip()
|
| 379 |
+
action = re.sub(r"^the partner\s+", "", action) # drop the leading one (prepended below)
|
| 380 |
+
caption = f"C is observing the partner. The partner {action}"
|
| 381 |
+
caption = re.sub(r"([.!?]\s+)the partner", r"\1The partner", caption) # capitalize after sentence end
|
| 382 |
+
return caption
|
| 383 |
+
|
| 384 |
+
def __getitem__(self, idx: int) -> dict[str, Any]:
|
| 385 |
+
# robust to corrupt / mid-write clip.npz (warping may still be running): skip + retry another sample
|
| 386 |
+
for _ in range(16):
|
| 387 |
+
try:
|
| 388 |
+
return self._build_sample(idx)
|
| 389 |
+
except Exception as e:
|
| 390 |
+
log.warning(
|
| 391 |
+
f"NymeriaActorObserverDataset: bad sample {self.pairs[idx]['pair_id']} "
|
| 392 |
+
f"({type(e).__name__}: {str(e)[:80]}); retrying another"
|
| 393 |
+
)
|
| 394 |
+
idx = random.randrange(len(self.pairs))
|
| 395 |
+
raise RuntimeError("NymeriaActorObserverDataset: too many unreadable samples in a row")
|
| 396 |
+
|
| 397 |
+
def _get_roles(self, p: dict) -> list:
|
| 398 |
+
"""(clip_npz, caption) per view. Actor-observer: view0=actor(real text), view1=observer(transformed)."""
|
| 399 |
+
return [(p["actor"], p["text"]), (p["observer"], self._observer_caption(p["text"]))]
|
| 400 |
+
|
| 401 |
+
def _build_sample(self, idx: int) -> dict[str, Any]:
|
| 402 |
+
p = self.pairs[idx]
|
| 403 |
+
roles = self._get_roles(p)
|
| 404 |
+
if self.config.role_mix and random.random() < 0.5:
|
| 405 |
+
roles = roles[::-1] # swap the two view slots (caption follows the clip)
|
| 406 |
+
|
| 407 |
+
actors = [self._load_actor(npz, "npz", slot=i) for i, (npz, _) in enumerate(roles)]
|
| 408 |
+
captions = [cap for _, cap in roles]
|
| 409 |
+
n_views = 2
|
| 410 |
+
|
| 411 |
+
sample: dict[str, Any] = {}
|
| 412 |
+
for key in actors[0]: # concat the two views along the temporal axis -> (C, V*T, H, W)
|
| 413 |
+
if key.startswith("cam_"): # per-view metadata, stacked on the view dim below (not temporal-concat)
|
| 414 |
+
continue
|
| 415 |
+
sample[key] = torch.cat([a[key] for a in actors], dim=1).contiguous()
|
| 416 |
+
if "cam_w2c" in actors[0]: # camera poses: (V, T, 4, 4) / (V, 3, 3) / (V,)
|
| 417 |
+
sample["camera_w2c"] = torch.stack([a["cam_w2c"] for a in actors], dim=0).contiguous()
|
| 418 |
+
sample["camera_K"] = torch.stack([a["cam_K"] for a in actors], dim=0).contiguous()
|
| 419 |
+
sample["camera_src_res"] = torch.stack([a["cam_src_res"] for a in actors], dim=0).contiguous()
|
| 420 |
+
if "cam_ref_w2c" in actors[0]: # reference-frame past poses: (V, R, 4, 4)
|
| 421 |
+
sample["reference_cam_w2c"] = torch.stack([a["cam_ref_w2c"] for a in actors], dim=0).contiguous()
|
| 422 |
+
|
| 423 |
+
T = self.num_frames
|
| 424 |
+
sample.update(
|
| 425 |
+
{
|
| 426 |
+
"__key__": p["pair_id"],
|
| 427 |
+
"__url__": p["pair_id"],
|
| 428 |
+
"ai_caption": captions, # [view0_caption, view1_caption] -> per-view cross-attention
|
| 429 |
+
"view_indices": torch.tensor([0] * T + [1] * T, dtype=torch.int64),
|
| 430 |
+
"fps": torch.tensor(self.config.fps, dtype=torch.float64),
|
| 431 |
+
"num_video_frames_per_view": torch.tensor(T, dtype=torch.int64),
|
| 432 |
+
"view_indices_selection": torch.tensor([0, 1], dtype=torch.int64),
|
| 433 |
+
"camera_keys_selection": ["view0", "view1"],
|
| 434 |
+
"sample_n_views": torch.tensor(n_views, dtype=torch.int64),
|
| 435 |
+
"padding_mask": torch.zeros((1, *self.resolution_hw), dtype=torch.float32),
|
| 436 |
+
"ref_cam_view_idx_sample_position": torch.tensor(-1, dtype=torch.int64),
|
| 437 |
+
"front_cam_view_idx_sample_position": torch.tensor(0, dtype=torch.int64),
|
| 438 |
+
}
|
| 439 |
+
)
|
| 440 |
+
return sample
|
| 441 |
+
|
| 442 |
+
|
| 443 |
+
class NymeriaActorActorDataset(NymeriaActorObserverDataset):
|
| 444 |
+
"""2-view actor-ACTOR dataset (/data2/nymeria_processed_longer): two ego actors generated jointly, each with
|
| 445 |
+
its OWN real action caption (from captions/<seq>.csv, keyed by clip_id). Same conditioning structure as the
|
| 446 |
+
actor-observer setup (refs / camera poses / role-mix) but BOTH views are actors (no observer transform).
|
| 447 |
+
Manifest columns: actor1_clip / actor2_clip (clip dir names)."""
|
| 448 |
+
|
| 449 |
+
def __init__(self, config: NymeriaPairsConfig):
|
| 450 |
+
import glob as _glob
|
| 451 |
+
|
| 452 |
+
self.config = config
|
| 453 |
+
self.root = config.root
|
| 454 |
+
self.clips_dir = os.path.join(self.root, "clips")
|
| 455 |
+
self.resolution_hw = tuple(config.resolution_hw)
|
| 456 |
+
self.num_frames = int(config.num_video_frames)
|
| 457 |
+
# caption lookup: clip_id -> caption (per-seq CSVs)
|
| 458 |
+
captions: dict[str, str] = {}
|
| 459 |
+
for f in _glob.glob(os.path.join(self.root, "captions", "*.csv")):
|
| 460 |
+
for row in csv.DictReader(open(f)):
|
| 461 |
+
captions[row["clip_id"]] = row.get("caption", "")
|
| 462 |
+
manifest_path = os.path.join(self.root, config.manifest_csv)
|
| 463 |
+
self.pairs = []
|
| 464 |
+
with open(manifest_path, newline="") as fp:
|
| 465 |
+
for row in csv.DictReader(fp):
|
| 466 |
+
a1, a2 = row["actor1_clip"], row["actor2_clip"]
|
| 467 |
+
# prefer the manifest's absolute npz path (clips may live outside {root}/clips, e.g. vroid groups/)
|
| 468 |
+
n1 = row.get("actor1_npz") or os.path.join(self.clips_dir, a1, "clip.npz")
|
| 469 |
+
n2 = row.get("actor2_npz") or os.path.join(self.clips_dir, a2, "clip.npz")
|
| 470 |
+
c1, c2 = captions.get(a1, ""), captions.get(a2, "")
|
| 471 |
+
if os.path.isfile(n1) and os.path.isfile(n2) and c1 and c2:
|
| 472 |
+
self.pairs.append({"pair_id": row["pair_id"], "actor": n1, "observer": n2, "cap0": c1, "cap1": c2})
|
| 473 |
+
log.info(f"NymeriaActorActorDataset: {len(self.pairs)} pairs from {manifest_path} (role_mix={config.role_mix})")
|
| 474 |
+
if not self.pairs:
|
| 475 |
+
raise RuntimeError(f"No usable actor-actor pairs in {manifest_path}")
|
| 476 |
+
|
| 477 |
+
def _get_roles(self, p: dict) -> list:
|
| 478 |
+
# both views are actors -> each gets its own real caption (no observer transform)
|
| 479 |
+
return [(p["actor"], p["cap0"]), (p["observer"], p["cap1"])]
|
| 480 |
+
|
| 481 |
+
|
| 482 |
+
def get_nymeria_actor_actor_loader(
|
| 483 |
+
*,
|
| 484 |
+
root: str = "/data2/nymeria_processed_longer",
|
| 485 |
+
manifest_csv: str = "manifest/train_split.csv",
|
| 486 |
+
resolution_hw: tuple[int, int] = (480, 480),
|
| 487 |
+
num_video_frames: int = 77,
|
| 488 |
+
fps: float = 10.0,
|
| 489 |
+
role_mix: bool = True,
|
| 490 |
+
num_reference_frames: int = 0,
|
| 491 |
+
shared_reference: bool = False,
|
| 492 |
+
emit_camera_poses: bool = False,
|
| 493 |
+
person_pose: bool = False,
|
| 494 |
+
emit_reference_pose: bool = False,
|
| 495 |
+
batch_size: int = 1,
|
| 496 |
+
num_workers: int = 4,
|
| 497 |
+
prefetch_factor: Optional[int] = 2,
|
| 498 |
+
is_train: bool = True,
|
| 499 |
+
**kwargs: Any,
|
| 500 |
+
):
|
| 501 |
+
dataset = NymeriaActorActorDataset(
|
| 502 |
+
NymeriaPairsConfig(
|
| 503 |
+
root=root, manifest_csv=manifest_csv, resolution_hw=tuple(resolution_hw),
|
| 504 |
+
num_video_frames=num_video_frames, fps=fps, role_mix=role_mix,
|
| 505 |
+
num_reference_frames=num_reference_frames, shared_reference=shared_reference,
|
| 506 |
+
emit_camera_poses=emit_camera_poses, person_pose=person_pose,
|
| 507 |
+
emit_reference_pose=emit_reference_pose,
|
| 508 |
+
)
|
| 509 |
+
)
|
| 510 |
+
sampler = None
|
| 511 |
+
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
| 512 |
+
sampler = torch.utils.data.distributed.DistributedSampler(dataset, shuffle=is_train, drop_last=True)
|
| 513 |
+
return get_cached_replay_dataloader(
|
| 514 |
+
webdataset=False, use_cache=False, dataset=dataset, batch_size=batch_size, num_workers=num_workers,
|
| 515 |
+
sampler=sampler, shuffle=(sampler is None and is_train), drop_last=True,
|
| 516 |
+
prefetch_factor=prefetch_factor if num_workers > 0 else None, persistent_workers=num_workers > 0,
|
| 517 |
+
pin_memory=False, collate_fn=collate_fn, cache_replay_name="nymeria_actor_actor_dataloader",
|
| 518 |
+
)
|
| 519 |
+
|
| 520 |
+
|
| 521 |
+
class NymeriaSingleViewDataset(NymeriaPairsDataset):
|
| 522 |
+
"""STAGE-1 single-view: pool individual EGO clips across roots (real single actor, longer both egos,
|
| 523 |
+
synthetic both egos) as independent V=1 samples, each with its OWN real action caption. Conditions =
|
| 524 |
+
pose + warping + plucker(self-frame0) + own refs (shared_reference=False -> per-clip refs.npz). Motion-
|
| 525 |
+
focused pretraining before the 2-view stages. `sources` = list of dicts {root, manifest_csv, ego_cols,
|
| 526 |
+
caption_kind ("text_col"|"csv_glob")}."""
|
| 527 |
+
|
| 528 |
+
def __init__(self, config: NymeriaPairsConfig, sources: list):
|
| 529 |
+
import glob as _glob
|
| 530 |
+
self.config = config
|
| 531 |
+
self.resolution_hw = tuple(config.resolution_hw)
|
| 532 |
+
self.num_frames = int(config.num_video_frames)
|
| 533 |
+
self.pairs = []
|
| 534 |
+
for src in sources:
|
| 535 |
+
root = src["root"]
|
| 536 |
+
caps = {}
|
| 537 |
+
if src["caption_kind"] == "csv_glob":
|
| 538 |
+
for f in _glob.glob(os.path.join(root, "captions", "*.csv")):
|
| 539 |
+
for r in csv.DictReader(open(f)):
|
| 540 |
+
caps[r["clip_id"]] = r.get("caption", "")
|
| 541 |
+
with open(os.path.join(root, src["manifest_csv"]), newline="") as fp:
|
| 542 |
+
for row in csv.DictReader(fp):
|
| 543 |
+
for col in src["ego_cols"]:
|
| 544 |
+
clip = row.get(col, "")
|
| 545 |
+
if not clip:
|
| 546 |
+
continue
|
| 547 |
+
npz = os.path.join(root, "clips", clip, "clip.npz")
|
| 548 |
+
cap = row.get("text", "") if src["caption_kind"] == "text_col" else caps.get(clip, "")
|
| 549 |
+
if cap and os.path.isfile(npz):
|
| 550 |
+
self.pairs.append({"a1": npz, "mode": "npz", "pair_id": clip, "caption": cap,
|
| 551 |
+
"session": row.get("actor1_seq", clip)})
|
| 552 |
+
log.info(f"NymeriaSingleViewDataset: {len(self.pairs)} ego clips from {len(sources)} sources")
|
| 553 |
+
if not self.pairs:
|
| 554 |
+
raise RuntimeError("NymeriaSingleViewDataset: no ego clips found")
|
| 555 |
+
|
| 556 |
+
def _build_sample(self, idx: int) -> dict[str, Any]:
|
| 557 |
+
sample = super()._build_sample(idx) # V=1 (config.single_view=True); placeholder caption
|
| 558 |
+
sample["ai_caption"] = [self.pairs[idx]["caption"]] # real per-clip action caption
|
| 559 |
+
return sample
|
| 560 |
+
|
| 561 |
+
|
| 562 |
+
_STAGE1_SOURCES = [
|
| 563 |
+
dict(root="/data2/nymeria_processed_single", manifest_csv="manifest/train_split.csv",
|
| 564 |
+
ego_cols=["actor1_clip"], caption_kind="text_col"), # actor(head)=ego; observer(exo) excluded
|
| 565 |
+
dict(root="/data2/nymeria_processed_longer", manifest_csv="manifest/train_split.csv",
|
| 566 |
+
ego_cols=["actor1_clip", "actor2_clip"], caption_kind="csv_glob"),
|
| 567 |
+
dict(root="/data3/synthetic_processed_multi", manifest_csv="manifest/train_split.csv",
|
| 568 |
+
ego_cols=["actor1_clip", "actor2_clip"], caption_kind="csv_glob"),
|
| 569 |
+
]
|
| 570 |
+
|
| 571 |
+
|
| 572 |
+
def get_nymeria_single_view_loader(
|
| 573 |
+
*, sources=None, val: bool = False, resolution_hw: tuple[int, int] = (480, 480), num_video_frames: int = 77,
|
| 574 |
+
fps: float = 10.0, num_reference_frames: int = 4, emit_camera_poses: bool = True,
|
| 575 |
+
batch_size: int = 1, num_workers: int = 4, prefetch_factor: Optional[int] = 2, is_train: bool = True,
|
| 576 |
+
**kwargs: Any,
|
| 577 |
+
):
|
| 578 |
+
srcs = sources if sources is not None else _STAGE1_SOURCES
|
| 579 |
+
if val: # swap to each source's val split
|
| 580 |
+
srcs = [dict(s, manifest_csv=s["manifest_csv"].replace("train_split", "val_split")) for s in srcs]
|
| 581 |
+
cfg = NymeriaPairsConfig(
|
| 582 |
+
resolution_hw=tuple(resolution_hw), num_video_frames=num_video_frames, fps=fps, single_view=True,
|
| 583 |
+
num_reference_frames=num_reference_frames, shared_reference=False, emit_camera_poses=emit_camera_poses,
|
| 584 |
+
)
|
| 585 |
+
dataset = NymeriaSingleViewDataset(cfg, srcs)
|
| 586 |
+
sampler = None
|
| 587 |
+
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
| 588 |
+
sampler = torch.utils.data.distributed.DistributedSampler(dataset, shuffle=is_train, drop_last=True)
|
| 589 |
+
return get_cached_replay_dataloader(
|
| 590 |
+
webdataset=False, use_cache=False, dataset=dataset, batch_size=batch_size, num_workers=num_workers,
|
| 591 |
+
sampler=sampler, shuffle=(sampler is None and is_train), drop_last=True,
|
| 592 |
+
prefetch_factor=prefetch_factor if num_workers > 0 else None, persistent_workers=num_workers > 0,
|
| 593 |
+
pin_memory=False, collate_fn=collate_fn, cache_replay_name="nymeria_single_view_dataloader",
|
| 594 |
+
)
|
| 595 |
+
|
| 596 |
+
|
| 597 |
+
def get_nymeria_actor_actor_combined_loader(
|
| 598 |
+
*,
|
| 599 |
+
roots: tuple = ("/data2/nymeria_processed_longer", "/data3/synthetic_processed_multi"),
|
| 600 |
+
manifest_csv: str = "manifest/train_split.csv",
|
| 601 |
+
resolution_hw: tuple[int, int] = (480, 480),
|
| 602 |
+
num_video_frames: int = 77,
|
| 603 |
+
fps: float = 10.0,
|
| 604 |
+
role_mix: bool = True,
|
| 605 |
+
num_reference_frames: int = 0,
|
| 606 |
+
shared_reference: bool = False,
|
| 607 |
+
emit_camera_poses: bool = False,
|
| 608 |
+
emit_depth: bool = False,
|
| 609 |
+
person_pose: bool = False,
|
| 610 |
+
batch_size: int = 1,
|
| 611 |
+
num_workers: int = 4,
|
| 612 |
+
prefetch_factor: Optional[int] = 2,
|
| 613 |
+
is_train: bool = True,
|
| 614 |
+
**kwargs: Any,
|
| 615 |
+
):
|
| 616 |
+
"""ConcatDataset over MULTIPLE actor-actor roots (e.g. real longer + synthetic), all sharing the same
|
| 617 |
+
conditioning config. Each root is a NymeriaActorActorDataset; samples are pooled uniformly (no reweighting)."""
|
| 618 |
+
parts = [
|
| 619 |
+
NymeriaActorActorDataset(
|
| 620 |
+
NymeriaPairsConfig(
|
| 621 |
+
root=r, manifest_csv=manifest_csv, resolution_hw=tuple(resolution_hw),
|
| 622 |
+
num_video_frames=num_video_frames, fps=fps, role_mix=role_mix,
|
| 623 |
+
num_reference_frames=num_reference_frames, shared_reference=shared_reference,
|
| 624 |
+
emit_camera_poses=emit_camera_poses, emit_depth=emit_depth, person_pose=person_pose,
|
| 625 |
+
)
|
| 626 |
+
)
|
| 627 |
+
for r in roots
|
| 628 |
+
]
|
| 629 |
+
dataset = torch.utils.data.ConcatDataset(parts)
|
| 630 |
+
log.info(f"actor-actor combined: {[len(p) for p in parts]} -> {len(dataset)} pairs from {list(roots)}")
|
| 631 |
+
sampler = None
|
| 632 |
+
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
| 633 |
+
sampler = torch.utils.data.distributed.DistributedSampler(dataset, shuffle=is_train, drop_last=True)
|
| 634 |
+
return get_cached_replay_dataloader(
|
| 635 |
+
webdataset=False, use_cache=False, dataset=dataset, batch_size=batch_size, num_workers=num_workers,
|
| 636 |
+
sampler=sampler, shuffle=(sampler is None and is_train), drop_last=True,
|
| 637 |
+
prefetch_factor=prefetch_factor if num_workers > 0 else None, persistent_workers=num_workers > 0,
|
| 638 |
+
pin_memory=False, collate_fn=collate_fn, cache_replay_name="nymeria_actor_actor_combined_dataloader",
|
| 639 |
+
)
|
| 640 |
+
|
| 641 |
+
|
| 642 |
+
def get_nymeria_obs_vroid_mixed_loader(
|
| 643 |
+
*,
|
| 644 |
+
obs_root: str = "/data2/nymeria_processed_single",
|
| 645 |
+
vroid_root: str = "/data4/vroid_batch",
|
| 646 |
+
obs_manifest: str = "manifest/train_split.csv",
|
| 647 |
+
vroid_manifest: str = "manifest/train_split.csv",
|
| 648 |
+
resolution_hw: tuple[int, int] = (480, 480),
|
| 649 |
+
num_video_frames: int = 77,
|
| 650 |
+
fps: float = 10.0,
|
| 651 |
+
role_mix: bool = True,
|
| 652 |
+
num_reference_frames: int = 4,
|
| 653 |
+
shared_reference: bool = True,
|
| 654 |
+
emit_camera_poses: bool = True,
|
| 655 |
+
person_pose: bool = True,
|
| 656 |
+
emit_reference_pose: bool = True,
|
| 657 |
+
batch_size: int = 1,
|
| 658 |
+
num_workers: int = 4,
|
| 659 |
+
prefetch_factor: Optional[int] = 2,
|
| 660 |
+
is_train: bool = True,
|
| 661 |
+
**kwargs: Any,
|
| 662 |
+
):
|
| 663 |
+
"""PRETRAIN mix: nymeria actor-OBSERVER (real) + vroid actor-ACTOR (synthetic), pooled uniformly via a
|
| 664 |
+
ConcatDataset. Both share the SAME conditioning structure (shared refs + Plücker campose + person-pose +
|
| 665 |
+
REFERENCE-POSE), so batches are interchangeable to the model; each dataset applies its own role handling
|
| 666 |
+
(observer transform vs actor-actor) and loads its own captions."""
|
| 667 |
+
common = dict(
|
| 668 |
+
resolution_hw=tuple(resolution_hw), num_video_frames=num_video_frames, fps=fps, role_mix=role_mix,
|
| 669 |
+
num_reference_frames=num_reference_frames, shared_reference=shared_reference,
|
| 670 |
+
emit_camera_poses=emit_camera_poses, person_pose=person_pose, emit_reference_pose=emit_reference_pose,
|
| 671 |
+
)
|
| 672 |
+
obs = NymeriaActorObserverDataset(NymeriaPairsConfig(root=obs_root, manifest_csv=obs_manifest, **common))
|
| 673 |
+
vro = NymeriaActorActorDataset(NymeriaPairsConfig(root=vroid_root, manifest_csv=vroid_manifest, **common))
|
| 674 |
+
dataset = torch.utils.data.ConcatDataset([obs, vro])
|
| 675 |
+
log.info(f"obs+vroid mixed: obs={len(obs)} + vroid={len(vro)} -> {len(dataset)} pairs")
|
| 676 |
+
sampler = None
|
| 677 |
+
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
| 678 |
+
sampler = torch.utils.data.distributed.DistributedSampler(dataset, shuffle=is_train, drop_last=True)
|
| 679 |
+
return get_cached_replay_dataloader(
|
| 680 |
+
webdataset=False, use_cache=False, dataset=dataset, batch_size=batch_size, num_workers=num_workers,
|
| 681 |
+
sampler=sampler, shuffle=(sampler is None and is_train), drop_last=True,
|
| 682 |
+
prefetch_factor=prefetch_factor if num_workers > 0 else None, persistent_workers=num_workers > 0,
|
| 683 |
+
pin_memory=False, collate_fn=collate_fn, cache_replay_name="nymeria_obs_vroid_mixed_dataloader",
|
| 684 |
+
)
|
| 685 |
+
|
| 686 |
+
|
| 687 |
+
def get_nymeria_actor_observer_loader(
|
| 688 |
+
*,
|
| 689 |
+
root: str = "/data2/nymeria_processed_single",
|
| 690 |
+
manifest_csv: str = "manifest/train_split.csv",
|
| 691 |
+
resolution_hw: tuple[int, int] = (480, 480),
|
| 692 |
+
num_video_frames: int = 77,
|
| 693 |
+
fps: float = 10.0,
|
| 694 |
+
role_mix: bool = True,
|
| 695 |
+
num_reference_frames: int = 0,
|
| 696 |
+
shared_reference: bool = False,
|
| 697 |
+
emit_camera_poses: bool = False,
|
| 698 |
+
emit_depth: bool = False,
|
| 699 |
+
person_pose: bool = False,
|
| 700 |
+
emit_reference_pose: bool = False,
|
| 701 |
+
batch_size: int = 1,
|
| 702 |
+
num_workers: int = 4,
|
| 703 |
+
prefetch_factor: Optional[int] = 2,
|
| 704 |
+
is_train: bool = True,
|
| 705 |
+
**kwargs: Any,
|
| 706 |
+
):
|
| 707 |
+
dataset = NymeriaActorObserverDataset(
|
| 708 |
+
NymeriaPairsConfig(
|
| 709 |
+
root=root,
|
| 710 |
+
manifest_csv=manifest_csv,
|
| 711 |
+
resolution_hw=tuple(resolution_hw),
|
| 712 |
+
num_video_frames=num_video_frames,
|
| 713 |
+
fps=fps,
|
| 714 |
+
role_mix=role_mix,
|
| 715 |
+
num_reference_frames=num_reference_frames,
|
| 716 |
+
shared_reference=shared_reference,
|
| 717 |
+
emit_camera_poses=emit_camera_poses,
|
| 718 |
+
emit_depth=emit_depth,
|
| 719 |
+
person_pose=person_pose,
|
| 720 |
+
emit_reference_pose=emit_reference_pose,
|
| 721 |
+
)
|
| 722 |
+
)
|
| 723 |
+
sampler = None
|
| 724 |
+
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
| 725 |
+
sampler = torch.utils.data.distributed.DistributedSampler(dataset, shuffle=is_train, drop_last=True)
|
| 726 |
+
return get_cached_replay_dataloader(
|
| 727 |
+
webdataset=False,
|
| 728 |
+
use_cache=False,
|
| 729 |
+
dataset=dataset,
|
| 730 |
+
batch_size=batch_size,
|
| 731 |
+
num_workers=num_workers,
|
| 732 |
+
sampler=sampler,
|
| 733 |
+
shuffle=(sampler is None and is_train),
|
| 734 |
+
drop_last=True,
|
| 735 |
+
prefetch_factor=prefetch_factor if num_workers > 0 else None,
|
| 736 |
+
persistent_workers=num_workers > 0,
|
| 737 |
+
pin_memory=False,
|
| 738 |
+
collate_fn=collate_fn,
|
| 739 |
+
cache_replay_name="nymeria_actor_observer_dataloader",
|
| 740 |
+
)
|
| 741 |
+
|
| 742 |
+
|
| 743 |
+
def get_nymeria_pairs_loader(
|
| 744 |
+
*,
|
| 745 |
+
root: str = "/data2/nymeria_processed_longer",
|
| 746 |
+
manifest_csv: str = "manifest/train_ready.csv",
|
| 747 |
+
resolution_hw: tuple[int, int] = (480, 480),
|
| 748 |
+
num_video_frames: int = 77,
|
| 749 |
+
fps: float = 10.0,
|
| 750 |
+
caption: str = "a person performing an everyday activity",
|
| 751 |
+
clean_only: bool = False,
|
| 752 |
+
single_view: bool = False,
|
| 753 |
+
split: str = "all",
|
| 754 |
+
val_num_sessions: int = 4,
|
| 755 |
+
split_seed: int = 1234,
|
| 756 |
+
val_max_pairs: int = 40,
|
| 757 |
+
dummy_text_embeddings: bool = False,
|
| 758 |
+
batch_size: int = 1,
|
| 759 |
+
num_workers: int = 4,
|
| 760 |
+
prefetch_factor: Optional[int] = 2,
|
| 761 |
+
is_train: bool = True,
|
| 762 |
+
**kwargs: Any,
|
| 763 |
+
):
|
| 764 |
+
dataset = NymeriaPairsDataset(
|
| 765 |
+
NymeriaPairsConfig(
|
| 766 |
+
root=root,
|
| 767 |
+
manifest_csv=manifest_csv,
|
| 768 |
+
resolution_hw=tuple(resolution_hw),
|
| 769 |
+
num_video_frames=num_video_frames,
|
| 770 |
+
fps=fps,
|
| 771 |
+
caption=caption,
|
| 772 |
+
clean_only=clean_only,
|
| 773 |
+
single_view=single_view,
|
| 774 |
+
split=split,
|
| 775 |
+
val_num_sessions=val_num_sessions,
|
| 776 |
+
split_seed=split_seed,
|
| 777 |
+
val_max_pairs=val_max_pairs,
|
| 778 |
+
dummy_text_embeddings=dummy_text_embeddings,
|
| 779 |
+
)
|
| 780 |
+
)
|
| 781 |
+
|
| 782 |
+
sampler = None
|
| 783 |
+
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
| 784 |
+
sampler = torch.utils.data.distributed.DistributedSampler(dataset, shuffle=is_train, drop_last=True)
|
| 785 |
+
|
| 786 |
+
return get_cached_replay_dataloader(
|
| 787 |
+
webdataset=False,
|
| 788 |
+
use_cache=False,
|
| 789 |
+
dataset=dataset,
|
| 790 |
+
batch_size=batch_size,
|
| 791 |
+
num_workers=num_workers,
|
| 792 |
+
sampler=sampler,
|
| 793 |
+
shuffle=(sampler is None and is_train),
|
| 794 |
+
drop_last=True,
|
| 795 |
+
prefetch_factor=prefetch_factor if num_workers > 0 else None,
|
| 796 |
+
persistent_workers=num_workers > 0,
|
| 797 |
+
pin_memory=False,
|
| 798 |
+
collate_fn=collate_fn,
|
| 799 |
+
cache_replay_name="nymeria_pairs_dataloader",
|
| 800 |
+
)
|
| 801 |
+
|
| 802 |
+
|
| 803 |
+
def register_nymeria_pairs_dataloader() -> None:
|
| 804 |
+
cs = ConfigStore.instance()
|
| 805 |
+
|
| 806 |
+
def _store(name: str, **kw):
|
| 807 |
+
"""Register a name with a session-held-out train split (data_train) and val split (data_val).
|
| 808 |
+
|
| 809 |
+
data_train -> split="train" (all sessions except the held-out ones).
|
| 810 |
+
data_val -> split="val" (only held-out sessions, clean subset, capped to val_max_pairs).
|
| 811 |
+
"""
|
| 812 |
+
cs.store(
|
| 813 |
+
group="data_train", package="dataloader_train", name=name,
|
| 814 |
+
node=L(get_nymeria_pairs_loader)(is_train=True, split="train", **kw),
|
| 815 |
+
)
|
| 816 |
+
val_kw = dict(kw)
|
| 817 |
+
val_kw["clean_only"] = True # cleaner validation regardless of the train clean setting
|
| 818 |
+
val_kw["num_workers"] = 2 # validation set is small
|
| 819 |
+
cs.store(
|
| 820 |
+
group="data_val", package="dataloader_val", name=name,
|
| 821 |
+
node=L(get_nymeria_pairs_loader)(is_train=False, split="val", **val_kw),
|
| 822 |
+
)
|
| 823 |
+
|
| 824 |
+
# LONGER (200+77, state_t=20) from train_ready.csv (clip.npz). 77 target frames -> state_t = 1+(77-1)//4 = 20.
|
| 825 |
+
_store("nymeria_train_ready", root="/data2/nymeria_processed_longer", num_video_frames=77, batch_size=1, num_workers=4)
|
| 826 |
+
_store("nymeria_train_ready_smoke", root="/data2/nymeria_processed_longer", num_video_frames=77,
|
| 827 |
+
dummy_text_embeddings=True, batch_size=1, num_workers=4)
|
| 828 |
+
_store("nymeria_train_ready_clean_smoke", root="/data2/nymeria_processed_longer", num_video_frames=77,
|
| 829 |
+
clean_only=True, dummy_text_embeddings=True, batch_size=1, num_workers=4)
|
| 830 |
+
# stage-1 single-view (V=1) pose pretraining: each actor an independent sample
|
| 831 |
+
_store("nymeria_single_view", root="/data2/nymeria_processed_longer", num_video_frames=77,
|
| 832 |
+
single_view=True, dummy_text_embeddings=True, batch_size=1, num_workers=4)
|
| 833 |
+
|
| 834 |
+
# 2-view actor-observer (real per-view text). train: role-mix on; val: role-mix off (deterministic).
|
| 835 |
+
cs.store(
|
| 836 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_observer",
|
| 837 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 838 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/train_split.csv",
|
| 839 |
+
num_video_frames=77, role_mix=True, is_train=True, batch_size=1, num_workers=4),
|
| 840 |
+
)
|
| 841 |
+
cs.store(
|
| 842 |
+
group="data_val", package="dataloader_val", name="nymeria_actor_observer",
|
| 843 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 844 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/val_split.csv",
|
| 845 |
+
num_video_frames=77, role_mix=False, is_train=False, batch_size=1, num_workers=2),
|
| 846 |
+
)
|
| 847 |
+
# 2-view actor-observer WITH reference-frame appearance conditioning (R=6 clean source frames/view)
|
| 848 |
+
cs.store(
|
| 849 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_observer_refs",
|
| 850 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 851 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/train_split.csv",
|
| 852 |
+
num_video_frames=77, role_mix=True, num_reference_frames=4, is_train=True, batch_size=1, num_workers=4),
|
| 853 |
+
)
|
| 854 |
+
cs.store(
|
| 855 |
+
group="data_val", package="dataloader_val", name="nymeria_actor_observer_refs",
|
| 856 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 857 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/val_split.csv",
|
| 858 |
+
num_video_frames=77, role_mix=False, num_reference_frames=4, is_train=False, batch_size=1, num_workers=2),
|
| 859 |
+
)
|
| 860 |
+
# 2-view actor-observer WITH reference frames (R=4) + camera-pose (Plücker) conditioning
|
| 861 |
+
cs.store(
|
| 862 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_observer_refs_campose",
|
| 863 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 864 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/train_split.csv",
|
| 865 |
+
num_video_frames=77, role_mix=True, num_reference_frames=4, emit_camera_poses=True,
|
| 866 |
+
is_train=True, batch_size=1, num_workers=4),
|
| 867 |
+
)
|
| 868 |
+
cs.store(
|
| 869 |
+
group="data_val", package="dataloader_val", name="nymeria_actor_observer_refs_campose",
|
| 870 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 871 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/val_split.csv",
|
| 872 |
+
num_video_frames=77, role_mix=False, num_reference_frames=4, emit_camera_poses=True,
|
| 873 |
+
is_train=False, batch_size=1, num_workers=2),
|
| 874 |
+
)
|
| 875 |
+
# 2-view ACTOR-ACTOR (longer dataset) with reference frames (R=4) + camera-pose (Plücker) conditioning.
|
| 876 |
+
# Both views are ego actors, each with its own real caption (captions/<seq>.csv).
|
| 877 |
+
cs.store(
|
| 878 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_actor_refs_campose",
|
| 879 |
+
node=L(get_nymeria_actor_actor_loader)(
|
| 880 |
+
root="/data2/nymeria_processed_longer", manifest_csv="manifest/train_split.csv",
|
| 881 |
+
num_video_frames=77, role_mix=True, num_reference_frames=4, emit_camera_poses=True,
|
| 882 |
+
is_train=True, batch_size=1, num_workers=4),
|
| 883 |
+
)
|
| 884 |
+
cs.store(
|
| 885 |
+
group="data_val", package="dataloader_val", name="nymeria_actor_actor_refs_campose",
|
| 886 |
+
node=L(get_nymeria_actor_actor_loader)(
|
| 887 |
+
root="/data2/nymeria_processed_longer", manifest_csv="manifest/val_split.csv",
|
| 888 |
+
num_video_frames=77, role_mix=False, num_reference_frames=4, emit_camera_poses=True,
|
| 889 |
+
is_train=False, batch_size=1, num_workers=2),
|
| 890 |
+
)
|
| 891 |
+
|
| 892 |
+
# 2-view actor-observer with SHARED greedy references (R=4/slot -> 8 shared) + camera-pose Plücker.
|
| 893 |
+
# Pairs with net/model shared_reference=True (ref view-embedding dropped). Uses refs_shared.npz.
|
| 894 |
+
cs.store(
|
| 895 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_observer_refs_campose_shared",
|
| 896 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 897 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/train_split.csv",
|
| 898 |
+
num_video_frames=77, role_mix=True, num_reference_frames=4, shared_reference=True,
|
| 899 |
+
emit_camera_poses=True, is_train=True, batch_size=1, num_workers=4),
|
| 900 |
+
)
|
| 901 |
+
cs.store(
|
| 902 |
+
group="data_val", package="dataloader_val", name="nymeria_actor_observer_refs_campose_shared",
|
| 903 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 904 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/val_split.csv",
|
| 905 |
+
num_video_frames=77, role_mix=False, num_reference_frames=4, shared_reference=True,
|
| 906 |
+
emit_camera_poses=True, is_train=False, batch_size=1, num_workers=2),
|
| 907 |
+
)
|
| 908 |
+
# + DEPTH condition: actor-observer SHARED refs with composite-depth (mesh_cond depth_comp) emitted as
|
| 909 |
+
# control_input_depth -> net VAE depth_embedder. single depth (mesh Stage-B) is ~100% ready.
|
| 910 |
+
cs.store(
|
| 911 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_observer_refs_campose_shared_depth",
|
| 912 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 913 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/train_split.csv",
|
| 914 |
+
num_video_frames=77, role_mix=True, num_reference_frames=4, shared_reference=True,
|
| 915 |
+
emit_camera_poses=True, emit_depth=True, is_train=True, batch_size=1, num_workers=4),
|
| 916 |
+
)
|
| 917 |
+
cs.store(
|
| 918 |
+
group="data_val", package="dataloader_val", name="nymeria_actor_observer_refs_campose_shared_depth",
|
| 919 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 920 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/val_split.csv",
|
| 921 |
+
num_video_frames=77, role_mix=False, num_reference_frames=4, shared_reference=True,
|
| 922 |
+
emit_camera_poses=True, emit_depth=True, is_train=False, batch_size=1, num_workers=2),
|
| 923 |
+
)
|
| 924 |
+
# 2-view ACTOR-ACTOR with SHARED greedy references + camera-pose Plücker.
|
| 925 |
+
cs.store(
|
| 926 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_actor_refs_campose_shared",
|
| 927 |
+
node=L(get_nymeria_actor_actor_loader)(
|
| 928 |
+
root="/data2/nymeria_processed_longer", manifest_csv="manifest/train_split.csv",
|
| 929 |
+
num_video_frames=77, role_mix=True, num_reference_frames=4, shared_reference=True,
|
| 930 |
+
emit_camera_poses=True, is_train=True, batch_size=1, num_workers=4),
|
| 931 |
+
)
|
| 932 |
+
cs.store(
|
| 933 |
+
group="data_val", package="dataloader_val", name="nymeria_actor_actor_refs_campose_shared",
|
| 934 |
+
node=L(get_nymeria_actor_actor_loader)(
|
| 935 |
+
root="/data2/nymeria_processed_longer", manifest_csv="manifest/val_split.csv",
|
| 936 |
+
num_video_frames=77, role_mix=False, num_reference_frames=4, shared_reference=True,
|
| 937 |
+
emit_camera_poses=True, is_train=False, batch_size=1, num_workers=2),
|
| 938 |
+
)
|
| 939 |
+
|
| 940 |
+
# VROID synthetic actor-actor: person-pose + shared refs + Plücker + REFERENCE-POSE condition
|
| 941 |
+
# (refs_shared `reference_pose`). /data4/vroid_batch. Captions from captions/core4d_all.csv (general:
|
| 942 |
+
# "A first-person view of C interacting with a partner."). detail_all.csv (per-action) is kept as .bak.
|
| 943 |
+
cs.store(
|
| 944 |
+
group="data_train", package="dataloader_train", name="vroid_actoractor_refpose",
|
| 945 |
+
node=L(get_nymeria_actor_actor_loader)(
|
| 946 |
+
root="/data4/vroid_batch", manifest_csv="manifest/train_split.csv",
|
| 947 |
+
num_video_frames=77, role_mix=True, num_reference_frames=4, shared_reference=True,
|
| 948 |
+
emit_camera_poses=True, person_pose=True, emit_reference_pose=True,
|
| 949 |
+
is_train=True, batch_size=1, num_workers=4),
|
| 950 |
+
)
|
| 951 |
+
cs.store(
|
| 952 |
+
group="data_val", package="dataloader_val", name="vroid_actoractor_refpose",
|
| 953 |
+
node=L(get_nymeria_actor_actor_loader)(
|
| 954 |
+
root="/data4/vroid_batch", manifest_csv="manifest/val_avatar_disjoint.csv",
|
| 955 |
+
num_video_frames=77, role_mix=False, num_reference_frames=4, shared_reference=True,
|
| 956 |
+
emit_camera_poses=True, person_pose=True, emit_reference_pose=True,
|
| 957 |
+
is_train=False, batch_size=1, num_workers=2),
|
| 958 |
+
)
|
| 959 |
+
|
| 960 |
+
# PRETRAIN MIX: nymeria actor-OBSERVER (real /data2/nymeria_processed_single) + vroid actor-ACTOR
|
| 961 |
+
# (synthetic /data4/vroid_batch), pooled uniformly. Shared refs + Plücker campose + person-pose + refpose.
|
| 962 |
+
cs.store(
|
| 963 |
+
group="data_train", package="dataloader_train", name="nymeria_obs_vroid_mixed_refpose",
|
| 964 |
+
node=L(get_nymeria_obs_vroid_mixed_loader)(
|
| 965 |
+
obs_root="/data2/nymeria_processed_single", vroid_root="/data4/vroid_batch",
|
| 966 |
+
obs_manifest="manifest/train_split.csv", vroid_manifest="manifest/train_split.csv",
|
| 967 |
+
num_video_frames=77, role_mix=True, num_reference_frames=4, shared_reference=True,
|
| 968 |
+
emit_camera_poses=True, person_pose=True, emit_reference_pose=True,
|
| 969 |
+
is_train=True, batch_size=1, num_workers=4),
|
| 970 |
+
)
|
| 971 |
+
cs.store(
|
| 972 |
+
group="data_val", package="dataloader_val", name="nymeria_obs_vroid_mixed_refpose",
|
| 973 |
+
node=L(get_nymeria_obs_vroid_mixed_loader)(
|
| 974 |
+
obs_root="/data2/nymeria_processed_single", vroid_root="/data4/vroid_batch",
|
| 975 |
+
obs_manifest="manifest/val_split.csv", vroid_manifest="manifest/val_avatar_disjoint.csv",
|
| 976 |
+
num_video_frames=77, role_mix=False, num_reference_frames=4, shared_reference=True,
|
| 977 |
+
emit_camera_poses=True, person_pose=True, emit_reference_pose=True,
|
| 978 |
+
is_train=False, batch_size=1, num_workers=2),
|
| 979 |
+
)
|
| 980 |
+
|
| 981 |
+
# OVERFIT: train AND val on the SAME 5 hand-picked pairs (manifest/overfit5.csv) — deliberate overfit test.
|
| 982 |
+
# role_mix=False on both so train and val see the identical pairs/orientation.
|
| 983 |
+
for grp, pkg in [("data_train", "dataloader_train"), ("data_val", "dataloader_val")]:
|
| 984 |
+
cs.store(
|
| 985 |
+
group=grp, package=pkg, name="vroid_overfit5",
|
| 986 |
+
node=L(get_nymeria_actor_actor_loader)(
|
| 987 |
+
root="/data4/vroid_batch", manifest_csv="manifest/overfit5.csv",
|
| 988 |
+
num_video_frames=77, role_mix=False, num_reference_frames=4, shared_reference=True,
|
| 989 |
+
emit_camera_poses=True, person_pose=True, emit_reference_pose=True,
|
| 990 |
+
is_train=(grp == "data_train"), batch_size=1, num_workers=2),
|
| 991 |
+
)
|
| 992 |
+
|
| 993 |
+
# 2-view ACTOR-ACTOR SHARED refs, REAL longer + SYNTHETIC combined (all pairs, no reweighting). Val stays
|
| 994 |
+
# on the REAL longer location-holdout set only (nymeria_actor_actor_refs_campose_shared val).
|
| 995 |
+
cs.store(
|
| 996 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_actor_refs_campose_shared_plus_synth",
|
| 997 |
+
node=L(get_nymeria_actor_actor_combined_loader)(
|
| 998 |
+
roots=["/data2/nymeria_processed_longer", "/data3/synthetic_processed_multi"],
|
| 999 |
+
manifest_csv="manifest/train_split.csv", num_video_frames=77, role_mix=True,
|
| 1000 |
+
num_reference_frames=4, shared_reference=True, emit_camera_poses=True,
|
| 1001 |
+
is_train=True, batch_size=1, num_workers=4),
|
| 1002 |
+
)
|
| 1003 |
+
|
| 1004 |
+
# + DEPTH condition: REAL longer + SYNTHETIC combined actor-actor with composite-depth (mesh_cond depth_comp)
|
| 1005 |
+
# emitted as control_input_depth -> net VAE depth_embedder. Synthetic depth 100% ready; nymeria mesh Stage-B
|
| 1006 |
+
# in progress (missing -> zero depth, coverage grows). Val stays on nymeria longer holdout.
|
| 1007 |
+
cs.store(
|
| 1008 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_actor_refs_campose_shared_plus_synth_depth",
|
| 1009 |
+
node=L(get_nymeria_actor_actor_combined_loader)(
|
| 1010 |
+
roots=["/data2/nymeria_processed_longer", "/data3/synthetic_processed_multi"],
|
| 1011 |
+
manifest_csv="manifest/train_split.csv", num_video_frames=77, role_mix=True,
|
| 1012 |
+
num_reference_frames=4, shared_reference=True, emit_camera_poses=True, emit_depth=True,
|
| 1013 |
+
is_train=True, batch_size=1, num_workers=4),
|
| 1014 |
+
)
|
| 1015 |
+
|
| 1016 |
+
# STAGE-1 SINGLE-VIEW (motion pretraining): pooled ego clips (real single actor + longer + synthetic),
|
| 1017 |
+
# per-clip real caption + pose + warping + plucker + own refs (R=4). LoRA-friendly, max data.
|
| 1018 |
+
cs.store(
|
| 1019 |
+
group="data_train", package="dataloader_train", name="nymeria_single_view_stage1",
|
| 1020 |
+
node=L(get_nymeria_single_view_loader)(
|
| 1021 |
+
val=False, num_video_frames=77, num_reference_frames=4, emit_camera_poses=True,
|
| 1022 |
+
is_train=True, batch_size=1, num_workers=4),
|
| 1023 |
+
)
|
| 1024 |
+
cs.store(
|
| 1025 |
+
group="data_val", package="dataloader_val", name="nymeria_single_view_stage1",
|
| 1026 |
+
node=L(get_nymeria_single_view_loader)(
|
| 1027 |
+
val=True, num_video_frames=77, num_reference_frames=4, emit_camera_poses=True,
|
| 1028 |
+
is_train=False, batch_size=1, num_workers=2),
|
| 1029 |
+
)
|
| 1030 |
+
|
| 1031 |
+
# 2-view ACTOR-ACTOR SHARED refs, SYNTHETIC ONLY (diagnostic: isolate synthetic human appearance / masking
|
| 1032 |
+
# from the real-data influence in the combined run). root = synthetic; val = synthetic val_split.
|
| 1033 |
+
cs.store(
|
| 1034 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_actor_refs_campose_shared_synthonly",
|
| 1035 |
+
node=L(get_nymeria_actor_actor_loader)(
|
| 1036 |
+
root="/data3/synthetic_processed_multi", manifest_csv="manifest/train_split.csv",
|
| 1037 |
+
num_video_frames=77, role_mix=True, num_reference_frames=4, shared_reference=True,
|
| 1038 |
+
emit_camera_poses=True, is_train=True, batch_size=1, num_workers=4),
|
| 1039 |
+
)
|
| 1040 |
+
cs.store(
|
| 1041 |
+
group="data_val", package="dataloader_val", name="nymeria_actor_actor_refs_campose_shared_synthonly",
|
| 1042 |
+
node=L(get_nymeria_actor_actor_loader)(
|
| 1043 |
+
root="/data3/synthetic_processed_multi", manifest_csv="manifest/val_split.csv",
|
| 1044 |
+
num_video_frames=77, role_mix=False, num_reference_frames=4, shared_reference=True,
|
| 1045 |
+
emit_camera_poses=True, is_train=False, batch_size=1, num_workers=2),
|
| 1046 |
+
)
|
| 1047 |
+
|
| 1048 |
+
# 2-view actor-observer with camera-pose (Plücker) conditioning ONLY (no reference frames) — ablation
|
| 1049 |
+
cs.store(
|
| 1050 |
+
group="data_train", package="dataloader_train", name="nymeria_actor_observer_campose",
|
| 1051 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 1052 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/train_split.csv",
|
| 1053 |
+
num_video_frames=77, role_mix=True, num_reference_frames=0, emit_camera_poses=True,
|
| 1054 |
+
is_train=True, batch_size=1, num_workers=4),
|
| 1055 |
+
)
|
| 1056 |
+
cs.store(
|
| 1057 |
+
group="data_val", package="dataloader_val", name="nymeria_actor_observer_campose",
|
| 1058 |
+
node=L(get_nymeria_actor_observer_loader)(
|
| 1059 |
+
root="/data2/nymeria_processed_single", manifest_csv="manifest/val_split.csv",
|
| 1060 |
+
num_video_frames=77, role_mix=False, num_reference_frames=0, emit_camera_poses=True,
|
| 1061 |
+
is_train=False, batch_size=1, num_workers=2),
|
| 1062 |
+
)
|
| 1063 |
+
|
| 1064 |
+
# legacy 33-frame (mp4) smoke kept for quick tests (no split)
|
| 1065 |
+
cs.store(
|
| 1066 |
+
group="data_train", package="dataloader_train", name="nymeria_pairs_smoke",
|
| 1067 |
+
node=L(get_nymeria_pairs_loader)(root="/data2/nymeria_processed", manifest_csv="manifest/clip_pairs.csv",
|
| 1068 |
+
num_video_frames=33, dummy_text_embeddings=True, batch_size=1, num_workers=2, is_train=True),
|
| 1069 |
+
)
|
| 1070 |
+
cs.store(
|
| 1071 |
+
group="data_val", package="dataloader_val", name="nymeria_pairs_smoke",
|
| 1072 |
+
node=L(get_nymeria_pairs_loader)(root="/data2/nymeria_processed", manifest_csv="manifest/clip_pairs.csv",
|
| 1073 |
+
num_video_frames=33, dummy_text_embeddings=True, batch_size=1, num_workers=2, is_train=False),
|
| 1074 |
+
)
|
cosmos_predict2/_src/predict2_multiview/datasets/wdinfo_utils.py
ADDED
|
@@ -0,0 +1,79 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
"""Utility functions for handling wdinfo files."""
|
| 17 |
+
|
| 18 |
+
from typing import Literal, Mapping
|
| 19 |
+
|
| 20 |
+
from cosmos_predict2._src.imaginaire import config
|
| 21 |
+
from cosmos_predict2._src.imaginaire.datasets.webdataset.config.schema import DatasetInfo
|
| 22 |
+
|
| 23 |
+
DEFAULT_CATALOG: Mapping = {
|
| 24 |
+
"alpamayo_dec2024": {
|
| 25 |
+
"sensitive": [
|
| 26 |
+
"wdinfo/alpamayo_dec2024/v0/resolution_1080/aspect_ratio_16_9/duration_10_30/wdinfo_test.json",
|
| 27 |
+
],
|
| 28 |
+
},
|
| 29 |
+
"mads_multiview_0823": {
|
| 30 |
+
"sensitive": [
|
| 31 |
+
"wdinfo/mads/cosmos-mads-dataset-transfer2-multiview-0823/v0/driving/resolution_720/aspect_ratio_16_9/duration_5_10/wdinfo_08232025.json",
|
| 32 |
+
],
|
| 33 |
+
},
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def get_video_dataset_info(
|
| 38 |
+
source_name: str,
|
| 39 |
+
*,
|
| 40 |
+
dataset_keys: list[str] | None = None,
|
| 41 |
+
object_store: Literal["gcs", "s3"] = "gcs",
|
| 42 |
+
dataset_catalog: dict = DEFAULT_CATALOG,
|
| 43 |
+
) -> list[DatasetInfo]:
|
| 44 |
+
if source_name not in dataset_catalog:
|
| 45 |
+
raise KeyError(
|
| 46 |
+
f"Source {source_name} not found in dataset catalog. Available keys are {dataset_catalog.keys()}"
|
| 47 |
+
)
|
| 48 |
+
|
| 49 |
+
# Create the wdinfo files here
|
| 50 |
+
dataset_infos = []
|
| 51 |
+
for sensitive_type, wdinfos in dataset_catalog[source_name].items():
|
| 52 |
+
if object_store == "gcs":
|
| 53 |
+
bucket = "bucket" if sensitive_type == "nonsensitive" else "bucket-s"
|
| 54 |
+
elif object_store == "s3":
|
| 55 |
+
bucket = "bucket" if sensitive_type == "nonsensitive" else "bucket-sensitive"
|
| 56 |
+
else:
|
| 57 |
+
raise ValueError("Cosmos data: only support gcs or s3 for object store")
|
| 58 |
+
|
| 59 |
+
if not wdinfos:
|
| 60 |
+
continue
|
| 61 |
+
|
| 62 |
+
dataset_infos.append(
|
| 63 |
+
DatasetInfo(
|
| 64 |
+
object_store_config=config.ObjectStoreConfig(
|
| 65 |
+
enabled=True,
|
| 66 |
+
credentials=(
|
| 67 |
+
"credentials/s3_training.secret" if object_store == "s3" else "credentials/gcs_training.secret"
|
| 68 |
+
),
|
| 69 |
+
bucket=bucket,
|
| 70 |
+
),
|
| 71 |
+
wdinfo=wdinfos,
|
| 72 |
+
per_dataset_keys=dataset_keys,
|
| 73 |
+
source=source_name,
|
| 74 |
+
opts={
|
| 75 |
+
"aspect_ratio": "16,9",
|
| 76 |
+
},
|
| 77 |
+
)
|
| 78 |
+
)
|
| 79 |
+
return dataset_infos
|
cosmos_predict2/_src/predict2_multiview/models/multiview_pose_model_rectified_flow.py
ADDED
|
@@ -0,0 +1,261 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
"""Multiview rectified-flow model with pose + warped-frame conditioning for 2-actor joint generation.
|
| 17 |
+
|
| 18 |
+
This subclass only adds *preprocessing* of the three conditioning inputs and leaves the training loop, FSDP,
|
| 19 |
+
EMA and checkpointing of the base ``MultiviewVid2VidModelRectifiedFlow`` untouched:
|
| 20 |
+
* ``control_input_warped`` (pixel RGB) -> normalized to [-1,1] -> VAE-encoded (frozen) -> ``warped_latent``
|
| 21 |
+
* ``control_input_visibility`` (pixel) -> 1ch [0,1] -> down-sampled to latent (T_lat,H_lat,W_lat) -> ``visibility_mask``
|
| 22 |
+
* ``control_input_pose`` (pixel RGB) -> passed through at pixel res -> ``pose_map`` (the net's pose
|
| 23 |
+
encoder normalizes /255 and down-samples /16 spatial, /4 temporal)
|
| 24 |
+
|
| 25 |
+
The preprocessed tensors are written into the data_batch under keys read by the pose conditioner's ReMapkey
|
| 26 |
+
embedders (``pose_map`` / ``warped_latent`` / ``visibility_mask``), so they flow into the MultiViewCondition
|
| 27 |
+
and ultimately into ``MultiViewPoseDiT.forward`` with classifier-free-guidance dropout handled by the
|
| 28 |
+
conditioner. No checkpoint weight surgery is needed: the new pose/cond modules are additive + zero-initialized
|
| 29 |
+
and load as missing keys from the base 2B checkpoint.
|
| 30 |
+
"""
|
| 31 |
+
|
| 32 |
+
from typing import Callable, Dict, Tuple
|
| 33 |
+
|
| 34 |
+
import attrs
|
| 35 |
+
import torch
|
| 36 |
+
import torch.nn.functional as F
|
| 37 |
+
|
| 38 |
+
from cosmos_predict2._src.imaginaire.utils import log
|
| 39 |
+
from einops import rearrange
|
| 40 |
+
from torch import Tensor
|
| 41 |
+
|
| 42 |
+
from cosmos_predict2._src.predict2_multiview.configs.vid2vid.defaults.conditioner import MultiViewCondition
|
| 43 |
+
from cosmos_predict2._src.predict2_multiview.models.multiview_vid2vid_model_rectified_flow import (
|
| 44 |
+
MultiviewVid2VidModelRectifiedFlow,
|
| 45 |
+
MultiviewVid2VidModelRectifiedFlowConfig,
|
| 46 |
+
)
|
| 47 |
+
|
| 48 |
+
_POSE_PREPROCESSED_KEY = "_pose_conditioning_preprocessed"
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
@attrs.define(slots=False)
|
| 52 |
+
class MultiviewVid2VidPoseModelRectifiedFlowConfig(MultiviewVid2VidModelRectifiedFlowConfig):
|
| 53 |
+
# data_batch keys carrying the raw pixel conditioning videos (one per actor, laid out (B,C,V*T,H,W))
|
| 54 |
+
pose_input_key: str = "control_input_pose"
|
| 55 |
+
warped_input_key: str = "control_input_warped"
|
| 56 |
+
visibility_input_key: str = "control_input_visibility"
|
| 57 |
+
# when True, pass the pose RGB through the frozen VAE -> `pose_latent` (16ch latent res) instead of
|
| 58 |
+
# (in addition to) the pixel `pose_map`. The net's pose_mode ("vae_concat"/"vae_mlp_add") consumes it.
|
| 59 |
+
pose_via_vae: bool = False
|
| 60 |
+
# when True, VAE-encode the composite-depth RGB (`control_input_depth`) into `depth_latent` (16ch latent res)
|
| 61 |
+
# for the net's zero-init depth_embedder (net.enable_depth). Depth = warped scene depth + human mesh depth.
|
| 62 |
+
depth_via_vae: bool = False
|
| 63 |
+
depth_input_key: str = "control_input_depth"
|
| 64 |
+
# >0: VAE-encode `reference_frames` (R clean source frames/view) into `reference_latent` for the net's
|
| 65 |
+
# in-context reference appearance conditioning. Must match the net's num_reference_frames.
|
| 66 |
+
num_reference_frames: int = 0
|
| 67 |
+
reference_input_key: str = "reference_frames"
|
| 68 |
+
# compute per-pixel Plücker ray maps (dir+moment, canonical view0-frame0 frame) from camera_w2c/K into
|
| 69 |
+
# `plucker_map` for the net's cross-view shared-space conditioning.
|
| 70 |
+
enable_plucker: bool = False
|
| 71 |
+
# also compute posed-reference Plücker (reference frames' past poses in the same canonical frame) into
|
| 72 |
+
# `reference_plucker_map`, so the in-context reference frames are geometrically grounded too.
|
| 73 |
+
enable_reference_plucker: bool = False
|
| 74 |
+
# VAE-encode the per-person skeleton render of each reference frame (refs_shared `reference_pose`) into
|
| 75 |
+
# `reference_pose_latent` for the net's zero-init reference_pose_embedder (net.enable_reference_pose).
|
| 76 |
+
enable_reference_pose: bool = False
|
| 77 |
+
reference_pose_input_key: str = "reference_pose"
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
class MultiviewVid2VidPoseModelRectifiedFlow(MultiviewVid2VidModelRectifiedFlow):
|
| 81 |
+
def __init__(self, config: MultiviewVid2VidPoseModelRectifiedFlowConfig):
|
| 82 |
+
super().__init__(config)
|
| 83 |
+
|
| 84 |
+
def set_up_model(self):
|
| 85 |
+
# super() builds the net + (if config.use_lora) runs add_lora, which FREEZES the whole base network
|
| 86 |
+
# and leaves only LoRA adapters trainable. Our NEW zero-init conditioning modules are NOT LoRA targets
|
| 87 |
+
# and must stay FULLY trainable, so un-freeze them here (backbone LoRA + full new-modules hybrid).
|
| 88 |
+
super().set_up_model()
|
| 89 |
+
if not getattr(self.config, "use_lora", False):
|
| 90 |
+
return
|
| 91 |
+
base = self.net.base_model.model if hasattr(self.net, "base_model") else self.net
|
| 92 |
+
trainable = ["cond_embedder", "plucker_embedder", "depth_embedder", "reference_pose_embedder", "pose_encoder", "pose_latent_embedder",
|
| 93 |
+
"pose_mlp", "ref_gate"]
|
| 94 |
+
if not getattr(base, "freeze_view_embedding", False):
|
| 95 |
+
trainable.append("view_embeddings")
|
| 96 |
+
n_new = 0
|
| 97 |
+
for name, param in self.net.named_parameters():
|
| 98 |
+
n = name.replace("base_model.model.", "")
|
| 99 |
+
if any(n == p or n.startswith(p + ".") for p in trainable):
|
| 100 |
+
param.requires_grad_(True)
|
| 101 |
+
param.data = param.data.float() # match add_lora's fp32 upcast for trainable params
|
| 102 |
+
n_new += param.numel()
|
| 103 |
+
tot = sum(p.numel() for p in self.net.parameters())
|
| 104 |
+
tr = sum(p.numel() for p in self.net.parameters() if p.requires_grad)
|
| 105 |
+
log.info(f"[LoRA+new-modules] unfroze {n_new:,} new-module params -> trainable {tr:,}/{tot:,} "
|
| 106 |
+
f"({100 * tr / max(1, tot):.3f}%); targets={trainable}")
|
| 107 |
+
|
| 108 |
+
@torch.no_grad()
|
| 109 |
+
def _preprocess_pose_conditioning(self, data_batch: Dict[str, torch.Tensor]) -> None:
|
| 110 |
+
"""Encode/down-sample the raw pixel conditioning videos into the keys read by the conditioner.
|
| 111 |
+
|
| 112 |
+
Idempotent: guarded so it runs once per data_batch even though both training (get_data_and_condition)
|
| 113 |
+
and inference (get_velocity_fn_from_batch) call it.
|
| 114 |
+
"""
|
| 115 |
+
if data_batch.get(_POSE_PREPROCESSED_KEY, False):
|
| 116 |
+
return
|
| 117 |
+
pose_key = self.config.pose_input_key
|
| 118 |
+
warped_key = self.config.warped_input_key
|
| 119 |
+
vis_key = self.config.visibility_input_key
|
| 120 |
+
if pose_key not in data_batch or warped_key not in data_batch or vis_key not in data_batch:
|
| 121 |
+
# nothing to do (e.g. a batch that does not carry pose conditioning)
|
| 122 |
+
return
|
| 123 |
+
|
| 124 |
+
state_t = self.config.state_t
|
| 125 |
+
num_pixel_frames_per_view = self.tokenizer.get_pixel_num_frames(state_t)
|
| 126 |
+
|
| 127 |
+
# --- warped past-frame: pixel RGB -> [-1,1] -> frozen VAE encode (multiview-aware reshape inside) ---
|
| 128 |
+
warped = data_batch[warped_key].to(**self.tensor_kwargs) / 127.5 - 1.0 # (B,C,V*Tpix,H,W)
|
| 129 |
+
warped_latent = self.encode(warped).to(**self.tensor_kwargs) # (B,16,V*state_t,H_lat,W_lat)
|
| 130 |
+
n_views = warped_latent.shape[2] // state_t
|
| 131 |
+
_, _, _, h_lat, w_lat = warped_latent.shape
|
| 132 |
+
data_batch["warped_latent"] = warped_latent
|
| 133 |
+
|
| 134 |
+
# --- visibility mask: pixel -> 1ch [0,1] -> down-sample to (state_t, H_lat, W_lat) per view ---
|
| 135 |
+
vis = data_batch[vis_key].to(**self.tensor_kwargs)
|
| 136 |
+
if vis.shape[1] > 1:
|
| 137 |
+
vis = vis.mean(dim=1, keepdim=True)
|
| 138 |
+
vis = vis / 255.0
|
| 139 |
+
vis = rearrange(vis, "B C (V T) H W -> (B V) C T H W", V=n_views)
|
| 140 |
+
vis = F.interpolate(vis.float(), size=(state_t, h_lat, w_lat), mode="trilinear", align_corners=False)
|
| 141 |
+
vis = rearrange(vis, "(B V) C T H W -> B C (V T) H W", V=n_views).to(**self.tensor_kwargs)
|
| 142 |
+
data_batch["visibility_mask"] = vis
|
| 143 |
+
|
| 144 |
+
# --- pose map: keep pixel resolution & full temporal; the net's pose encoder handles the rest ---
|
| 145 |
+
data_batch["pose_map"] = data_batch[pose_key].float() # (B,3,V*Tpix,H,W) in [0,255]
|
| 146 |
+
|
| 147 |
+
# --- pose latent (optional): pixel RGB pose -> [-1,1] -> frozen VAE encode (same path as warped) ---
|
| 148 |
+
# gives a 16ch latent-res pose representation for the "vae_concat"/"vae_mlp_add" net pose modes.
|
| 149 |
+
if getattr(self.config, "pose_via_vae", False):
|
| 150 |
+
pose = data_batch[pose_key].to(**self.tensor_kwargs) / 127.5 - 1.0 # (B,3,V*Tpix,H,W)
|
| 151 |
+
data_batch["pose_latent"] = self.encode(pose).to(**self.tensor_kwargs) # (B,16,V*state_t,H_lat,W_lat)
|
| 152 |
+
|
| 153 |
+
# --- composite DEPTH latent (optional): RGB-encoded depth (warped scene + human mesh) -> [-1,1] -> frozen
|
| 154 |
+
# VAE encode (same path as warped/pose) -> depth_latent (B,16,V*state_t,H_lat,W_lat). ---
|
| 155 |
+
if getattr(self.config, "depth_via_vae", False):
|
| 156 |
+
dkey = getattr(self.config, "depth_input_key", "control_input_depth")
|
| 157 |
+
if dkey in data_batch:
|
| 158 |
+
depth = data_batch[dkey].to(**self.tensor_kwargs) / 127.5 - 1.0 # (B,3,V*Tpix,H,W)
|
| 159 |
+
data_batch["depth_latent"] = self.encode(depth).to(**self.tensor_kwargs)
|
| 160 |
+
|
| 161 |
+
# --- reference frames (optional): R clean source frames/view, VAE-encoded PER FRAME (1-frame clips) so
|
| 162 |
+
# each of the R frames yields exactly one latent frame -> (B,16,V*R,H_lat,W_lat) reference_latent. ---
|
| 163 |
+
ref_key = getattr(self.config, "reference_input_key", "reference_frames")
|
| 164 |
+
if getattr(self.config, "num_reference_frames", 0) > 0 and ref_key in data_batch:
|
| 165 |
+
refs = data_batch[ref_key].to(**self.tensor_kwargs) / 127.5 - 1.0 # (B,3,V*R,H,W)
|
| 166 |
+
data_batch["reference_latent"] = self._encode_reference_frames(refs).to(**self.tensor_kwargs)
|
| 167 |
+
|
| 168 |
+
# --- reference POSE (optional): per-person skeleton render of each ref frame, VAE-encoded PER FRAME
|
| 169 |
+
# (same path as reference_latent) -> reference_pose_latent (B,16,V*R,H_lat,W_lat). ---
|
| 170 |
+
if getattr(self.config, "enable_reference_pose", False):
|
| 171 |
+
rpkey = getattr(self.config, "reference_pose_input_key", "reference_pose")
|
| 172 |
+
if rpkey in data_batch:
|
| 173 |
+
rpose = data_batch[rpkey].to(**self.tensor_kwargs) / 127.5 - 1.0 # (B,3,V*R,H,W)
|
| 174 |
+
data_batch["reference_pose_latent"] = self._encode_reference_frames(rpose).to(**self.tensor_kwargs)
|
| 175 |
+
|
| 176 |
+
# --- Plücker ray map (optional): per-pixel camera rays in a per-pair canonical frame (view0 frame0),
|
| 177 |
+
# at latent res, for cross-view shared-space grounding. -> plucker_map (B,6,V*state_t,H_lat,W_lat) ---
|
| 178 |
+
if getattr(self.config, "enable_plucker", False) and "camera_w2c" in data_batch:
|
| 179 |
+
src = int(data_batch["camera_src_res"][0, 0]) if "camera_src_res" in data_batch else 504
|
| 180 |
+
# force fp32 (disable autocast): torch.linalg.inv rejects bf16, and validation calls this inside an
|
| 181 |
+
# autocast(bf16) block where einsum/matmul would otherwise downcast the poses to bf16.
|
| 182 |
+
with torch.autocast("cuda", enabled=False):
|
| 183 |
+
w2c_t = data_batch["camera_w2c"].to(torch.float32) # (B,V,Tpix,4,4)
|
| 184 |
+
K = data_batch["camera_K"]
|
| 185 |
+
# per-pair canonical frame = TARGET view0 frame0 (same transform to all views AND references)
|
| 186 |
+
c2w_ref = torch.linalg.inv(w2c_t[:, 0, 0]) # (B,4,4)
|
| 187 |
+
Tp = w2c_t.shape[2]
|
| 188 |
+
idx = torch.linspace(0, Tp - 1, state_t, device=w2c_t.device).round().long()
|
| 189 |
+
data_batch["plucker_map"] = self._plucker_from_w2c(w2c_t[:, :, idx], K, c2w_ref, src, h_lat, w_lat)
|
| 190 |
+
# posed-reference Plücker (optional): reference frames' PAST poses in the SAME canonical frame
|
| 191 |
+
if getattr(self.config, "enable_reference_plucker", False) and "reference_cam_w2c" in data_batch:
|
| 192 |
+
data_batch["reference_plucker_map"] = self._plucker_from_w2c(
|
| 193 |
+
data_batch["reference_cam_w2c"].to(torch.float32), K, c2w_ref, src, h_lat, w_lat
|
| 194 |
+
)
|
| 195 |
+
|
| 196 |
+
data_batch[_POSE_PREPROCESSED_KEY] = True
|
| 197 |
+
|
| 198 |
+
@torch.no_grad()
|
| 199 |
+
def _plucker_from_w2c(self, w2c_BVN, camera_K, c2w_ref, src_res, h_lat, w_lat):
|
| 200 |
+
"""Per-pixel Plücker rays (dir 3 + moment 3) at latent res for a set of poses `w2c_BVN` (B,V,N,4,4),
|
| 201 |
+
expressed in a GIVEN canonical frame `c2w_ref` (B,4,4) = inv(view0-frame0). The SAME c2w_ref is used for
|
| 202 |
+
the real frames AND the references, so both live in one shared frame. camera_K (B,V,3,3) ->
|
| 203 |
+
plucker (B, 6, V*N, H_lat, W_lat)."""
|
| 204 |
+
dt = torch.float32
|
| 205 |
+
dev = w2c_BVN.device
|
| 206 |
+
B, V, N = w2c_BVN.shape[:3]
|
| 207 |
+
w2c = w2c_BVN.to(dt)
|
| 208 |
+
w2c_canon = torch.einsum("bvnij,bjk->bvnik", w2c, c2w_ref.to(dt)) # (B,V,N,4,4)
|
| 209 |
+
c2w = torch.linalg.inv(w2c_canon.reshape(-1, 4, 4)).reshape(B, V, N, 4, 4)
|
| 210 |
+
Rc2w = c2w[..., :3, :3] # (B,V,N,3,3)
|
| 211 |
+
C = c2w[..., :3, 3] # (B,V,N,3) camera centers in canonical frame
|
| 212 |
+
# intrinsics scaled from src_res to the latent grid
|
| 213 |
+
K = camera_K.to(dt).clone() # (B,V,3,3)
|
| 214 |
+
K[..., 0, :] *= w_lat / float(src_res)
|
| 215 |
+
K[..., 1, :] *= h_lat / float(src_res)
|
| 216 |
+
Kinv = torch.linalg.inv(K) # (B,V,3,3)
|
| 217 |
+
ys, xs = torch.meshgrid(
|
| 218 |
+
torch.arange(h_lat, device=dev, dtype=dt), torch.arange(w_lat, device=dev, dtype=dt), indexing="ij"
|
| 219 |
+
)
|
| 220 |
+
pix = torch.stack(
|
| 221 |
+
[xs.reshape(-1) + 0.5, ys.reshape(-1) + 0.5, torch.ones(h_lat * w_lat, device=dev, dtype=dt)], 0
|
| 222 |
+
) # (3, HW)
|
| 223 |
+
d_cam = torch.einsum("bvij,jp->bvip", Kinv, pix) # (B,V,3,HW)
|
| 224 |
+
d_cam = d_cam / d_cam.norm(dim=2, keepdim=True)
|
| 225 |
+
d_world = torch.einsum("bvnij,bvjp->bvnip", Rc2w, d_cam) # (B,V,N,3,HW)
|
| 226 |
+
d_world = d_world / d_world.norm(dim=3, keepdim=True)
|
| 227 |
+
C_exp = C.unsqueeze(-1).expand(-1, -1, -1, -1, h_lat * w_lat) # (B,V,N,3,HW)
|
| 228 |
+
m = torch.cross(C_exp, d_world, dim=3) # moment C x d
|
| 229 |
+
plk = torch.cat([d_world, m], dim=3) # (B,V,N,6,HW)
|
| 230 |
+
plk = plk.reshape(B, V, N, 6, h_lat, w_lat).permute(0, 3, 1, 2, 4, 5).reshape(
|
| 231 |
+
B, 6, V * N, h_lat, w_lat
|
| 232 |
+
)
|
| 233 |
+
return plk.to(**self.tensor_kwargs)
|
| 234 |
+
|
| 235 |
+
@torch.no_grad()
|
| 236 |
+
def _encode_reference_frames(self, refs_B_C_VR_H_W: torch.Tensor) -> torch.Tensor:
|
| 237 |
+
"""Encode each reference frame as an INDEPENDENT 1-frame clip through the frozen VAE.
|
| 238 |
+
A plain multiview encode maps V*R pixel frames via get_pixel_num_frames(state_t) and would mis-infer
|
| 239 |
+
n_views; per-frame encoding gives exactly R latent frames/view. (B,3,V*R,H,W) -> (B,16,V*R,H_lat,W_lat).
|
| 240 |
+
"""
|
| 241 |
+
B = refs_B_C_VR_H_W.shape[0]
|
| 242 |
+
x = rearrange(refs_B_C_VR_H_W, "B C VR H W -> (B VR) C H W").unsqueeze(2) # (B*VR, C, 1, H, W)
|
| 243 |
+
z = self.tokenizer.encode(x) # (B*VR, 16, 1, H_lat, W_lat)
|
| 244 |
+
return rearrange(z.squeeze(2), "(B VR) C H W -> B C VR H W", B=B)
|
| 245 |
+
|
| 246 |
+
def get_data_and_condition(
|
| 247 |
+
self, data_batch: dict[str, torch.Tensor]
|
| 248 |
+
) -> Tuple[Tensor, Tensor, MultiViewCondition]:
|
| 249 |
+
# training path: conditioner runs inside super().get_data_and_condition, so preprocess first
|
| 250 |
+
self._preprocess_pose_conditioning(data_batch)
|
| 251 |
+
return super().get_data_and_condition(data_batch)
|
| 252 |
+
|
| 253 |
+
def get_velocity_fn_from_batch(
|
| 254 |
+
self,
|
| 255 |
+
data_batch: Dict,
|
| 256 |
+
guidance: float = 1.5,
|
| 257 |
+
is_negative_prompt: bool = False,
|
| 258 |
+
) -> Callable:
|
| 259 |
+
# inference path: conditioner is invoked before get_data_and_condition, so preprocess up front
|
| 260 |
+
self._preprocess_pose_conditioning(data_batch)
|
| 261 |
+
return super().get_velocity_fn_from_batch(data_batch, guidance, is_negative_prompt)
|
cosmos_predict2/_src/predict2_multiview/models/multiview_vid2vid_model_rectified_flow.py
ADDED
|
@@ -0,0 +1,655 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
import random
|
| 17 |
+
from dataclasses import field
|
| 18 |
+
from typing import Any, Callable, Dict, Optional, Tuple, cast
|
| 19 |
+
|
| 20 |
+
import attrs
|
| 21 |
+
import torch
|
| 22 |
+
import torch.distributed as dist
|
| 23 |
+
from einops import rearrange
|
| 24 |
+
from megatron.core import parallel_state
|
| 25 |
+
from torch import Tensor
|
| 26 |
+
from torch.distributed import get_process_group_ranks
|
| 27 |
+
|
| 28 |
+
from cosmos_predict2._src.imaginaire.flags import SMOKE
|
| 29 |
+
from cosmos_predict2._src.imaginaire.utils import log
|
| 30 |
+
from cosmos_predict2._src.imaginaire.utils.context_parallel import broadcast, broadcast_split_tensor
|
| 31 |
+
from cosmos_predict2._src.predict2.conditioner import DataType
|
| 32 |
+
from cosmos_predict2._src.predict2.models.text2world_model_rectified_flow import IS_PREPROCESSED_KEY
|
| 33 |
+
from cosmos_predict2._src.predict2.models.video2world_model_rectified_flow import (
|
| 34 |
+
NUM_CONDITIONAL_FRAMES_KEY,
|
| 35 |
+
Video2WorldModelRectifiedFlow,
|
| 36 |
+
Video2WorldModelRectifiedFlowConfig,
|
| 37 |
+
)
|
| 38 |
+
from cosmos_predict2._src.predict2.utils.dtensor_helper import broadcast_dtensor_model_states
|
| 39 |
+
from cosmos_predict2._src.predict2_multiview.configs.vid2vid.defaults.conditioner import (
|
| 40 |
+
ConditionLocationList,
|
| 41 |
+
MultiViewCondition,
|
| 42 |
+
)
|
| 43 |
+
from cosmos_predict2._src.predict2_multiview.models.view_sampling import sample_n_views_from_data_batch
|
| 44 |
+
|
| 45 |
+
TRAIN_SAMPLE_N_VIEWS_KEY = "train_sample_n_views"
|
| 46 |
+
TRAIN_SAMPLING_APPLIED_KEY = "train_sampling_applied"
|
| 47 |
+
_DEFAULT_NEGATIVE_PROMPT = "The video captures a series of frames showing ugly scenes, static with no motion, motion blur, over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. Overall, the video is of poor quality."
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
@attrs.define(slots=False)
|
| 51 |
+
class MultiviewVid2VidModelRectifiedFlowConfig(Video2WorldModelRectifiedFlowConfig):
|
| 52 |
+
min_num_conditional_frames_per_view: int = 1
|
| 53 |
+
max_num_conditional_frames_per_view: int = 2
|
| 54 |
+
train_sample_views_range: Tuple[int, int] | None = None
|
| 55 |
+
condition_locations: ConditionLocationList = field(default_factory=lambda: ConditionLocationList([]))
|
| 56 |
+
state_t: int = 0
|
| 57 |
+
view_condition_dropout_max: int = 0
|
| 58 |
+
online_text_embeddings_as_dict: bool = True # For backward compatibility with old experiments
|
| 59 |
+
conditional_frames_probs: Optional[Dict[int, float]] = None # Probability distribution for conditional frames
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
class MultiviewVid2VidModelRectifiedFlow(Video2WorldModelRectifiedFlow):
|
| 63 |
+
def __init__(self, config: MultiviewVid2VidModelRectifiedFlowConfig):
|
| 64 |
+
super().__init__(config)
|
| 65 |
+
self.state_t = config.state_t
|
| 66 |
+
self.empty_string_text_embeddings = None
|
| 67 |
+
self.neg_text_embeddings = None
|
| 68 |
+
if self.config.text_encoder_config is not None and self.config.text_encoder_config.compute_online:
|
| 69 |
+
compute_empty_and_negative_text_embeddings(self)
|
| 70 |
+
|
| 71 |
+
@torch.no_grad()
|
| 72 |
+
def encode(self, state: torch.Tensor) -> torch.Tensor:
|
| 73 |
+
n_views = state.shape[2] // self.tokenizer.get_pixel_num_frames(self.state_t)
|
| 74 |
+
cp_size = len(get_process_group_ranks(parallel_state.get_context_parallel_group()))
|
| 75 |
+
# let n_views 2 also cp-encoded
|
| 76 |
+
if n_views > 1 and n_views <= cp_size:
|
| 77 |
+
return self.encode_cp(state)
|
| 78 |
+
state = rearrange(state, "B C (V T) H W -> (B V) C T H W", V=n_views)
|
| 79 |
+
encoded_state = super().encode(state)
|
| 80 |
+
encoded_state = rearrange(encoded_state, "(B V) C T H W -> B C (V T) H W", V=n_views)
|
| 81 |
+
return encoded_state
|
| 82 |
+
|
| 83 |
+
@torch.no_grad()
|
| 84 |
+
def decode(self, latent: torch.Tensor) -> torch.Tensor:
|
| 85 |
+
n_views = latent.shape[2] // self.state_t
|
| 86 |
+
cp_size = len(get_process_group_ranks(parallel_state.get_context_parallel_group()))
|
| 87 |
+
# let n_views 2 also cp-decoded
|
| 88 |
+
if n_views > 1 and n_views <= cp_size:
|
| 89 |
+
return self.decode_cp(latent)
|
| 90 |
+
latent = rearrange(latent, "B C (V T) H W -> (B V) C T H W", V=n_views)
|
| 91 |
+
decoded_state = super().decode(latent)
|
| 92 |
+
decoded_state = rearrange(decoded_state, "(B V) C T H W -> B C (V T) H W", V=n_views)
|
| 93 |
+
return decoded_state
|
| 94 |
+
|
| 95 |
+
@torch.no_grad()
|
| 96 |
+
def encode_cp(self, state: torch.Tensor) -> torch.Tensor:
|
| 97 |
+
cp_size = len(get_process_group_ranks(parallel_state.get_context_parallel_group()))
|
| 98 |
+
cp_group = parallel_state.get_context_parallel_group()
|
| 99 |
+
n_views = state.shape[2] // self.tokenizer.get_pixel_num_frames(self.state_t)
|
| 100 |
+
assert n_views <= cp_size, f"n_views must be less than cp_size, got n_views={n_views} and cp_size={cp_size}"
|
| 101 |
+
state_V_B_C_T_H_W = rearrange(state, "B C (V T) H W -> V B C T H W", V=n_views)
|
| 102 |
+
state_input = torch.zeros((cp_size, *state_V_B_C_T_H_W.shape[1:]), **self.tensor_kwargs)
|
| 103 |
+
state_input[0:n_views] = state_V_B_C_T_H_W
|
| 104 |
+
local_state_V_B_C_T_H_W = broadcast_split_tensor(state_input, seq_dim=0, process_group=cp_group)
|
| 105 |
+
local_state = rearrange(local_state_V_B_C_T_H_W, "V B C T H W -> (B V) C T H W")
|
| 106 |
+
encoded_state = super().encode(local_state)
|
| 107 |
+
encoded_state_list = [torch.empty_like(encoded_state) for _ in range(cp_size)]
|
| 108 |
+
dist.all_gather(encoded_state_list, encoded_state, group=cp_group)
|
| 109 |
+
encoded_state = torch.cat(encoded_state_list[0:n_views], dim=2) # [B, C, V * T, H, W]
|
| 110 |
+
return encoded_state
|
| 111 |
+
|
| 112 |
+
@torch.no_grad()
|
| 113 |
+
def decode_cp(self, latent: torch.Tensor) -> torch.Tensor:
|
| 114 |
+
cp_size = len(get_process_group_ranks(parallel_state.get_context_parallel_group()))
|
| 115 |
+
cp_group = parallel_state.get_context_parallel_group()
|
| 116 |
+
n_views = latent.shape[2] // self.state_t
|
| 117 |
+
assert n_views <= cp_size, f"n_views must be less than cp_size, got n_views={n_views} and cp_size={cp_size}"
|
| 118 |
+
latent_V_B_C_T_H_W = rearrange(latent, "B C (V T) H W -> V B C T H W", V=n_views)
|
| 119 |
+
latent_input = torch.zeros((cp_size, *latent_V_B_C_T_H_W.shape[1:]), **self.tensor_kwargs)
|
| 120 |
+
latent_input[0:n_views] = latent_V_B_C_T_H_W
|
| 121 |
+
local_latent_V_B_C_T_H_W = broadcast_split_tensor(latent_input, seq_dim=0, process_group=cp_group)
|
| 122 |
+
local_latent = rearrange(local_latent_V_B_C_T_H_W, "V B C T H W -> (B V) C T H W")
|
| 123 |
+
decoded_state = super().decode(local_latent)
|
| 124 |
+
decoded_state_list = [torch.empty_like(decoded_state) for _ in range(cp_size)]
|
| 125 |
+
dist.all_gather(decoded_state_list, decoded_state, group=cp_group)
|
| 126 |
+
decoded_state = torch.cat(decoded_state_list[0:n_views], dim=2) # [B, C, V * T, H, W]
|
| 127 |
+
return decoded_state
|
| 128 |
+
|
| 129 |
+
def training_step(
|
| 130 |
+
self, data_batch: dict[str, torch.Tensor], iteration: int
|
| 131 |
+
) -> tuple[dict[str, torch.Tensor], torch.Tensor]:
|
| 132 |
+
return training_step_multiview(self, data_batch, iteration)
|
| 133 |
+
|
| 134 |
+
def inplace_compute_text_embeddings_online(self, data_batch: dict[str, torch.Tensor]) -> None:
|
| 135 |
+
inplace_compute_text_embeddings_online_multiview(self, data_batch)
|
| 136 |
+
|
| 137 |
+
def broadcast_split_for_model_parallelsim(
|
| 138 |
+
self,
|
| 139 |
+
x0_B_C_T_H_W: torch.Tensor,
|
| 140 |
+
condition: MultiViewCondition,
|
| 141 |
+
epsilon_B_C_T_H_W: torch.Tensor,
|
| 142 |
+
sigma_B_T: torch.Tensor,
|
| 143 |
+
):
|
| 144 |
+
n_views = x0_B_C_T_H_W.shape[2] // self.state_t
|
| 145 |
+
x0_B_C_T_H_W = rearrange(x0_B_C_T_H_W, "B C (V T) H W -> (B V) C T H W", V=n_views).contiguous()
|
| 146 |
+
if epsilon_B_C_T_H_W is not None:
|
| 147 |
+
epsilon_B_C_T_H_W = rearrange(epsilon_B_C_T_H_W, "B C (V T) H W -> (B V) C T H W", V=n_views).contiguous()
|
| 148 |
+
reshape_sigma_B_T = False
|
| 149 |
+
if sigma_B_T is not None:
|
| 150 |
+
assert sigma_B_T.ndim == 2, "sigma_B_T should be 2D tensor"
|
| 151 |
+
if sigma_B_T.shape[-1] != 1:
|
| 152 |
+
assert sigma_B_T.shape[-1] % n_views == 0, (
|
| 153 |
+
f"sigma_B_T temporal dimension T must either be 1 or a multiple of sample_n_views. Got T={sigma_B_T.shape[-1]} and sample_n_views={n_views}"
|
| 154 |
+
)
|
| 155 |
+
sigma_B_T = rearrange(sigma_B_T, "B (V T) -> (B V) T", V=n_views).contiguous()
|
| 156 |
+
reshape_sigma_B_T = True
|
| 157 |
+
x0_B_C_T_H_W, condition, epsilon_B_C_T_H_W, sigma_B_T = super().broadcast_split_for_model_parallelsim(
|
| 158 |
+
x0_B_C_T_H_W, condition, epsilon_B_C_T_H_W, sigma_B_T
|
| 159 |
+
)
|
| 160 |
+
x0_B_C_T_H_W = rearrange(x0_B_C_T_H_W, "(B V) C T H W -> B C (V T) H W", V=n_views)
|
| 161 |
+
if epsilon_B_C_T_H_W is not None:
|
| 162 |
+
epsilon_B_C_T_H_W = rearrange(epsilon_B_C_T_H_W, "(B V) C T H W -> B C (V T) H W", V=n_views)
|
| 163 |
+
if reshape_sigma_B_T:
|
| 164 |
+
sigma_B_T = rearrange(sigma_B_T, "(B V) T -> B (V T)", V=n_views)
|
| 165 |
+
return x0_B_C_T_H_W, condition, epsilon_B_C_T_H_W, sigma_B_T
|
| 166 |
+
|
| 167 |
+
def get_data_batch_with_latent_view_indices(self, data_batch: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
|
| 168 |
+
num_video_frames_per_view = int(data_batch["num_video_frames_per_view"].cpu().item())
|
| 169 |
+
n_views = data_batch["view_indices"].shape[1] // num_video_frames_per_view
|
| 170 |
+
view_indices_B_V_T = rearrange(data_batch["view_indices"], "B (V T) -> B V T", V=n_views)
|
| 171 |
+
|
| 172 |
+
latent_view_indices_B_V_T = view_indices_B_V_T[:, :, 0 : self.config.state_t]
|
| 173 |
+
latent_view_indices_B_T = rearrange(latent_view_indices_B_V_T, "B V T -> B (V T)")
|
| 174 |
+
data_batch_with_latent_view_indices = data_batch.copy()
|
| 175 |
+
data_batch_with_latent_view_indices["latent_view_indices_B_T"] = latent_view_indices_B_T
|
| 176 |
+
return data_batch_with_latent_view_indices
|
| 177 |
+
|
| 178 |
+
def _normalize_video_databatch_inplace(self, data_batch: dict[str, Tensor], input_key: str = None) -> None:
|
| 179 |
+
input_key = self.input_data_key if input_key is None else input_key
|
| 180 |
+
is_preprocessed = IS_PREPROCESSED_KEY in data_batch and data_batch[IS_PREPROCESSED_KEY] is True
|
| 181 |
+
|
| 182 |
+
num_video_frames_per_view = (
|
| 183 |
+
self.tokenizer.get_pixel_num_frames(self.state_t)
|
| 184 |
+
if is_preprocessed
|
| 185 |
+
else data_batch["num_video_frames_per_view"]
|
| 186 |
+
)
|
| 187 |
+
if isinstance(num_video_frames_per_view, torch.Tensor):
|
| 188 |
+
num_video_frames_per_view = int(num_video_frames_per_view.cpu().item())
|
| 189 |
+
n_views = data_batch[input_key].shape[2] // num_video_frames_per_view
|
| 190 |
+
if input_key in data_batch:
|
| 191 |
+
data_batch[input_key] = rearrange(data_batch[input_key], "B C (V T) H W -> (B V) C T H W", V=n_views)
|
| 192 |
+
super()._normalize_video_databatch_inplace(data_batch, input_key)
|
| 193 |
+
data_batch[input_key] = rearrange(data_batch[input_key], "(B V) C T H W -> B C (V T) H W", V=n_views)
|
| 194 |
+
|
| 195 |
+
def get_data_and_condition(self, data_batch: dict[str, torch.Tensor]) -> Tuple[Tensor, Tensor, MultiViewCondition]:
|
| 196 |
+
data_batch_with_latent_view_indices = self.get_data_batch_with_latent_view_indices(data_batch)
|
| 197 |
+
raw_state, latent_state, condition = super(Video2WorldModelRectifiedFlow, self).get_data_and_condition(
|
| 198 |
+
data_batch_with_latent_view_indices
|
| 199 |
+
)
|
| 200 |
+
condition = cast(MultiViewCondition, condition)
|
| 201 |
+
condition = condition.set_video_condition(
|
| 202 |
+
state_t=self.config.state_t,
|
| 203 |
+
gt_frames=latent_state.to(**self.tensor_kwargs),
|
| 204 |
+
condition_locations=self.config.condition_locations,
|
| 205 |
+
random_min_num_conditional_frames_per_view=self.config.min_num_conditional_frames_per_view,
|
| 206 |
+
random_max_num_conditional_frames_per_view=self.config.max_num_conditional_frames_per_view,
|
| 207 |
+
num_conditional_frames_per_view=None,
|
| 208 |
+
view_condition_dropout_max=self.config.view_condition_dropout_max,
|
| 209 |
+
conditional_frames_probs=self.config.conditional_frames_probs,
|
| 210 |
+
)
|
| 211 |
+
return raw_state, latent_state, condition
|
| 212 |
+
|
| 213 |
+
def get_velocity_fn_from_batch(
|
| 214 |
+
self,
|
| 215 |
+
data_batch: Dict,
|
| 216 |
+
guidance: float = 1.5,
|
| 217 |
+
is_negative_prompt: bool = False,
|
| 218 |
+
) -> Callable:
|
| 219 |
+
"""
|
| 220 |
+
Generates a callable function `x0_fn` based on the provided data batch and guidance factor.
|
| 221 |
+
|
| 222 |
+
This function first processes the input data batch through a conditioning workflow (`conditioner`) to obtain conditioned and unconditioned states. It then defines a nested function `x0_fn` which applies a denoising operation on an input `noise_x` at a given noise level `sigma` using both the conditioned and unconditioned states.
|
| 223 |
+
|
| 224 |
+
Args:
|
| 225 |
+
- data_batch (Dict): A batch of data used for conditioning. The format and content of this dictionary should align with the expectations of the `self.conditioner`
|
| 226 |
+
- guidance (float, optional): A scalar value that modulates the influence of the conditioned state relative to the unconditioned state in the output. Defaults to 1.5.
|
| 227 |
+
- is_negative_prompt (bool): use negative prompt t5 in uncondition if true
|
| 228 |
+
|
| 229 |
+
Returns:
|
| 230 |
+
- Callable: A function `x0_fn(noise_x, sigma)` that takes two arguments, `noise_x` and `sigma`, and return velocity predictoin
|
| 231 |
+
|
| 232 |
+
The returned function is suitable for use in scenarios where a denoised state is required based on both conditioned and unconditioned inputs, with an adjustable level of guidance influence.
|
| 233 |
+
"""
|
| 234 |
+
|
| 235 |
+
data_batch_with_latent_view_indices = self.get_data_batch_with_latent_view_indices(data_batch)
|
| 236 |
+
if NUM_CONDITIONAL_FRAMES_KEY in data_batch_with_latent_view_indices:
|
| 237 |
+
num_conditional_frames = data_batch_with_latent_view_indices[NUM_CONDITIONAL_FRAMES_KEY]
|
| 238 |
+
log.debug(f"Using {num_conditional_frames=} from data batch")
|
| 239 |
+
else:
|
| 240 |
+
num_conditional_frames = 1
|
| 241 |
+
|
| 242 |
+
if is_negative_prompt:
|
| 243 |
+
condition, uncondition = self.conditioner.get_condition_with_negative_prompt(
|
| 244 |
+
data_batch_with_latent_view_indices
|
| 245 |
+
)
|
| 246 |
+
else:
|
| 247 |
+
condition, uncondition = self.conditioner.get_condition_uncondition(data_batch_with_latent_view_indices)
|
| 248 |
+
|
| 249 |
+
is_image_batch = self.is_image_batch(data_batch_with_latent_view_indices)
|
| 250 |
+
condition = condition.edit_data_type(DataType.IMAGE if is_image_batch else DataType.VIDEO)
|
| 251 |
+
uncondition = uncondition.edit_data_type(DataType.IMAGE if is_image_batch else DataType.VIDEO)
|
| 252 |
+
_, x0, _ = self.get_data_and_condition(data_batch_with_latent_view_indices)
|
| 253 |
+
# override condition with inference mode; num_conditional_frames used Here!
|
| 254 |
+
condition = condition.set_video_condition(
|
| 255 |
+
state_t=self.config.state_t,
|
| 256 |
+
gt_frames=x0,
|
| 257 |
+
condition_locations=self.config.condition_locations,
|
| 258 |
+
random_min_num_conditional_frames_per_view=self.config.min_num_conditional_frames_per_view,
|
| 259 |
+
random_max_num_conditional_frames_per_view=self.config.max_num_conditional_frames_per_view,
|
| 260 |
+
num_conditional_frames_per_view=num_conditional_frames,
|
| 261 |
+
view_condition_dropout_max=0,
|
| 262 |
+
conditional_frames_probs=self.config.conditional_frames_probs,
|
| 263 |
+
)
|
| 264 |
+
uncondition = uncondition.set_video_condition(
|
| 265 |
+
state_t=self.config.state_t,
|
| 266 |
+
gt_frames=x0,
|
| 267 |
+
condition_locations=self.config.condition_locations,
|
| 268 |
+
random_min_num_conditional_frames_per_view=self.config.min_num_conditional_frames_per_view,
|
| 269 |
+
random_max_num_conditional_frames_per_view=self.config.max_num_conditional_frames_per_view,
|
| 270 |
+
num_conditional_frames_per_view=num_conditional_frames,
|
| 271 |
+
view_condition_dropout_max=0,
|
| 272 |
+
conditional_frames_probs=self.config.conditional_frames_probs,
|
| 273 |
+
)
|
| 274 |
+
condition = condition.edit_for_inference(
|
| 275 |
+
is_cfg_conditional=True,
|
| 276 |
+
condition_locations=self.config.condition_locations,
|
| 277 |
+
num_conditional_frames_per_view=num_conditional_frames,
|
| 278 |
+
)
|
| 279 |
+
uncondition = uncondition.edit_for_inference(
|
| 280 |
+
is_cfg_conditional=False,
|
| 281 |
+
condition_locations=self.config.condition_locations,
|
| 282 |
+
num_conditional_frames_per_view=num_conditional_frames,
|
| 283 |
+
)
|
| 284 |
+
_, condition, _, _ = self.broadcast_split_for_model_parallelsim(x0, condition, None, None)
|
| 285 |
+
_, uncondition, _, _ = self.broadcast_split_for_model_parallelsim(x0, uncondition, None, None)
|
| 286 |
+
if parallel_state.is_initialized():
|
| 287 |
+
pass
|
| 288 |
+
else:
|
| 289 |
+
assert not self.net.is_context_parallel_enabled, (
|
| 290 |
+
"parallel_state is not initialized, context parallel should be turned off."
|
| 291 |
+
)
|
| 292 |
+
|
| 293 |
+
def velocity_fn(noise: torch.Tensor, noise_x: torch.Tensor, timestep: torch.Tensor) -> torch.Tensor:
|
| 294 |
+
cond_v = self.denoise(noise, noise_x, timestep, condition)
|
| 295 |
+
uncond_v = self.denoise(noise, noise_x, timestep, uncondition)
|
| 296 |
+
velocity_pred = uncond_v + guidance * (cond_v - uncond_v) # align with pred2
|
| 297 |
+
return velocity_pred
|
| 298 |
+
|
| 299 |
+
return velocity_fn
|
| 300 |
+
|
| 301 |
+
@torch.no_grad()
|
| 302 |
+
def generate_samples_from_batch(
|
| 303 |
+
self,
|
| 304 |
+
data_batch: dict[str, torch.Tensor],
|
| 305 |
+
guidance: float = 1.5,
|
| 306 |
+
seed: int = 1,
|
| 307 |
+
state_shape: Tuple | None = None,
|
| 308 |
+
n_sample: int | None = None,
|
| 309 |
+
is_negative_prompt: bool = False,
|
| 310 |
+
num_steps: int = 35,
|
| 311 |
+
shift: float = 5.0,
|
| 312 |
+
**kwargs,
|
| 313 |
+
) -> torch.Tensor:
|
| 314 |
+
data_batch_with_latent_view_indices = self.get_data_batch_with_latent_view_indices(data_batch)
|
| 315 |
+
process_group = parallel_state.get_context_parallel_group()
|
| 316 |
+
cp_size = len(get_process_group_ranks(process_group))
|
| 317 |
+
samples_B_C_T_H_W = super().generate_samples_from_batch(
|
| 318 |
+
data_batch_with_latent_view_indices,
|
| 319 |
+
guidance,
|
| 320 |
+
seed,
|
| 321 |
+
state_shape,
|
| 322 |
+
n_sample,
|
| 323 |
+
is_negative_prompt,
|
| 324 |
+
num_steps,
|
| 325 |
+
shift,
|
| 326 |
+
**kwargs,
|
| 327 |
+
)
|
| 328 |
+
if cp_size > 1:
|
| 329 |
+
samples_B_C_T_H_W = rearrange(
|
| 330 |
+
samples_B_C_T_H_W, "B C (c V T) H W -> B C (V c T) H W", c=cp_size, T=self.state_t // cp_size
|
| 331 |
+
)
|
| 332 |
+
return samples_B_C_T_H_W
|
| 333 |
+
|
| 334 |
+
def set_up_model(self):
|
| 335 |
+
"""Override set_up_model to initialize cross-view attention from base model."""
|
| 336 |
+
# Call parent's set_up_model first
|
| 337 |
+
super().set_up_model()
|
| 338 |
+
|
| 339 |
+
# Load base model and initialize cross-view attention if enabled
|
| 340 |
+
if (
|
| 341 |
+
hasattr(self.net, "enable_cross_view_attn")
|
| 342 |
+
and self.net.enable_cross_view_attn
|
| 343 |
+
and self.net.init_cross_view_attn_weight_from is not None
|
| 344 |
+
):
|
| 345 |
+
log.info("Loading base model for cross-view attention initialization")
|
| 346 |
+
self.net.init_cross_view_attn_with_self_attn_weights(is_ema=False)
|
| 347 |
+
if self.config.ema.enabled:
|
| 348 |
+
self.net_ema.init_cross_view_attn_with_self_attn_weights(is_ema=True)
|
| 349 |
+
|
| 350 |
+
# Broadcast to ensure all ranks have consistent weights
|
| 351 |
+
if self.fsdp_device_mesh is not None:
|
| 352 |
+
log.info("Broadcasting cross-view attention weights for consistency")
|
| 353 |
+
broadcast_dtensor_model_states(self.net, self.fsdp_device_mesh)
|
| 354 |
+
if self.config.ema.enabled:
|
| 355 |
+
broadcast_dtensor_model_states(self.net_ema, self.fsdp_device_mesh)
|
| 356 |
+
|
| 357 |
+
|
| 358 |
+
def compute_text_embeddings_online_multiview_single_caption(
|
| 359 |
+
model, data_batch: dict[str, torch.Tensor]
|
| 360 |
+
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 361 |
+
is_preprocessed = IS_PREPROCESSED_KEY in data_batch and data_batch[IS_PREPROCESSED_KEY] is True
|
| 362 |
+
num_video_frames_per_view = (
|
| 363 |
+
model.tokenizer.get_pixel_num_frames(model.state_t)
|
| 364 |
+
if is_preprocessed
|
| 365 |
+
else data_batch["num_video_frames_per_view"]
|
| 366 |
+
)
|
| 367 |
+
if isinstance(num_video_frames_per_view, torch.Tensor):
|
| 368 |
+
num_video_frames_per_view = int(num_video_frames_per_view.cpu().item())
|
| 369 |
+
n_views = data_batch[model.input_data_key].shape[2] // num_video_frames_per_view
|
| 370 |
+
B, _, _, _, _ = data_batch[model.input_data_key].shape
|
| 371 |
+
|
| 372 |
+
# compute prompt embeddings
|
| 373 |
+
if len(data_batch["ai_caption"]) != 1:
|
| 374 |
+
raise NotImplementedError(f"Expected batch size of 1, got {len(data_batch['ai_caption'])}")
|
| 375 |
+
|
| 376 |
+
if len(data_batch["ai_caption"][0]) != 1:
|
| 377 |
+
raise ValueError(f"Expected a single caption, got {len(data_batch['ai_caption'][0])}")
|
| 378 |
+
|
| 379 |
+
caption = data_batch["ai_caption"][0][0]
|
| 380 |
+
assert isinstance(caption, str)
|
| 381 |
+
view0_text_embeddings_B_L_D = model.text_encoder.compute_text_embeddings_online(
|
| 382 |
+
data_batch={model.input_caption_key: [caption]},
|
| 383 |
+
input_caption_key=model.input_caption_key,
|
| 384 |
+
)
|
| 385 |
+
assert view0_text_embeddings_B_L_D.shape[0] == 1
|
| 386 |
+
assert view0_text_embeddings_B_L_D.shape[1] == 512, (
|
| 387 |
+
f"view0_text_embeddings should be of shape (B, 512, D), got {view0_text_embeddings_B_L_D.shape}"
|
| 388 |
+
)
|
| 389 |
+
output_text_embeddings = model.empty_string_text_embeddings.clone().repeat(B, n_views, 1)
|
| 390 |
+
output_neg_text_embeddings = model.empty_string_text_embeddings.clone().repeat(B, n_views, 1)
|
| 391 |
+
output_text_embeddings = rearrange(output_text_embeddings, "B (V L) D -> V B L D", V=n_views)
|
| 392 |
+
output_neg_text_embeddings = rearrange(output_neg_text_embeddings, "B (V L) D -> V B L D", V=n_views)
|
| 393 |
+
# Assign prompt embeddings to the front camera view
|
| 394 |
+
for i_b in range(B):
|
| 395 |
+
front_cam_view_idx_sample_position = data_batch["front_cam_view_idx_sample_position"][i_b]
|
| 396 |
+
output_text_embeddings[front_cam_view_idx_sample_position, i_b] = view0_text_embeddings_B_L_D[i_b]
|
| 397 |
+
output_neg_text_embeddings[front_cam_view_idx_sample_position, i_b] = model.neg_text_embeddings[0]
|
| 398 |
+
output_text_embeddings = rearrange(output_text_embeddings, "V B L D -> B (V L) D")
|
| 399 |
+
output_neg_text_embeddings = rearrange(output_neg_text_embeddings, "V B L D -> B (V L) D")
|
| 400 |
+
|
| 401 |
+
dropout_text_embeddings = model.empty_string_text_embeddings.clone().repeat(B, n_views, 1)
|
| 402 |
+
|
| 403 |
+
if not model.config.conditioner.text.use_empty_string:
|
| 404 |
+
dropout_text_embeddings *= 0.0
|
| 405 |
+
|
| 406 |
+
return output_text_embeddings, output_neg_text_embeddings, dropout_text_embeddings
|
| 407 |
+
|
| 408 |
+
|
| 409 |
+
def compute_text_embeddings_online_multiview_multiple_captions(
|
| 410 |
+
model, data_batch: dict[str, torch.Tensor]
|
| 411 |
+
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 412 |
+
is_preprocessed = IS_PREPROCESSED_KEY in data_batch and data_batch[IS_PREPROCESSED_KEY] is True
|
| 413 |
+
num_video_frames_per_view = (
|
| 414 |
+
model.tokenizer.get_pixel_num_frames(model.state_t)
|
| 415 |
+
if is_preprocessed
|
| 416 |
+
else data_batch["num_video_frames_per_view"]
|
| 417 |
+
)
|
| 418 |
+
if isinstance(num_video_frames_per_view, torch.Tensor):
|
| 419 |
+
num_video_frames_per_view = int(num_video_frames_per_view.cpu().item())
|
| 420 |
+
n_views = data_batch[model.input_data_key].shape[2] // num_video_frames_per_view
|
| 421 |
+
B, _, _, _, _ = data_batch[model.input_data_key].shape
|
| 422 |
+
|
| 423 |
+
# compute each view's caption separately
|
| 424 |
+
if not len(data_batch["ai_caption"]) == 1:
|
| 425 |
+
raise NotImplementedError(f"Expected batch size of 1, got {len(data_batch['ai_caption'])}")
|
| 426 |
+
|
| 427 |
+
captions = data_batch["ai_caption"][0]
|
| 428 |
+
if len(captions) != n_views:
|
| 429 |
+
raise ValueError(f"Expected {n_views} captions, got {len(captions)}: {captions}")
|
| 430 |
+
view_text_embeddings = []
|
| 431 |
+
for caption in captions:
|
| 432 |
+
data_batch_per_view = {model.input_caption_key: [caption]}
|
| 433 |
+
view_text_embedding = model.text_encoder.compute_text_embeddings_online(
|
| 434 |
+
data_batch_per_view, model.input_caption_key
|
| 435 |
+
)
|
| 436 |
+
view_text_embeddings.append(view_text_embedding)
|
| 437 |
+
|
| 438 |
+
view_text_embeddings_B_V_L_D = torch.stack(view_text_embeddings, dim=1)
|
| 439 |
+
assert view_text_embeddings_B_V_L_D.shape[:3] == (
|
| 440 |
+
B,
|
| 441 |
+
n_views,
|
| 442 |
+
512,
|
| 443 |
+
), f"view_text_embeddings_B_V_L_D should be of shape (B, n_views, 512, D), got {view_text_embeddings_B_V_L_D.shape}"
|
| 444 |
+
output_text_embeddings = rearrange(view_text_embeddings_B_V_L_D, "B V L D -> B (V L) D")
|
| 445 |
+
|
| 446 |
+
# repeat negative embedding for each per view
|
| 447 |
+
output_neg_text_embeddings_L_D = model.neg_text_embeddings[0].clone()
|
| 448 |
+
output_neg_text_embeddings_B_V_L_D = rearrange(output_neg_text_embeddings_L_D, "(B V L) D -> B V L D", B=1, V=1)
|
| 449 |
+
output_neg_text_embeddings_B_V_L_D = output_neg_text_embeddings_B_V_L_D.repeat(B, n_views, 1, 1)
|
| 450 |
+
assert output_neg_text_embeddings_B_V_L_D.shape[:3] == (
|
| 451 |
+
B,
|
| 452 |
+
n_views,
|
| 453 |
+
512,
|
| 454 |
+
), (
|
| 455 |
+
f"output_neg_text_embeddings_B_V_L_D should be of shape (B, n_views, 512, D), got {output_neg_text_embeddings_B_V_L_D.shape}"
|
| 456 |
+
)
|
| 457 |
+
output_neg_text_embeddings = rearrange(output_neg_text_embeddings_B_V_L_D, "B V L D -> B (V L) D")
|
| 458 |
+
|
| 459 |
+
dropout_text_embeddings = model.empty_string_text_embeddings.clone().repeat(B, n_views, 1)
|
| 460 |
+
if not model.config.conditioner.text.use_empty_string:
|
| 461 |
+
dropout_text_embeddings *= 0.0
|
| 462 |
+
|
| 463 |
+
return output_text_embeddings, output_neg_text_embeddings, dropout_text_embeddings
|
| 464 |
+
|
| 465 |
+
|
| 466 |
+
def compute_text_embeddings_online_multiview(
|
| 467 |
+
model, data_batch: dict[str, torch.Tensor]
|
| 468 |
+
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 469 |
+
captions = data_batch["ai_caption"][0]
|
| 470 |
+
is_preprocessed = IS_PREPROCESSED_KEY in data_batch and data_batch[IS_PREPROCESSED_KEY] is True
|
| 471 |
+
num_video_frames_per_view = (
|
| 472 |
+
model.tokenizer.get_pixel_num_frames(model.state_t)
|
| 473 |
+
if is_preprocessed
|
| 474 |
+
else data_batch["num_video_frames_per_view"]
|
| 475 |
+
)
|
| 476 |
+
if isinstance(num_video_frames_per_view, torch.Tensor):
|
| 477 |
+
num_video_frames_per_view = int(num_video_frames_per_view.cpu().item())
|
| 478 |
+
n_views = data_batch[model.input_data_key].shape[2] // num_video_frames_per_view
|
| 479 |
+
assert len(captions) == 1 or len(captions) == n_views, f"Expected 1 or {n_views} captions, got {len(captions)}"
|
| 480 |
+
if len(captions) == 1:
|
| 481 |
+
return compute_text_embeddings_online_multiview_single_caption(model, data_batch)
|
| 482 |
+
else:
|
| 483 |
+
return compute_text_embeddings_online_multiview_multiple_captions(model, data_batch)
|
| 484 |
+
|
| 485 |
+
|
| 486 |
+
def inplace_compute_text_embeddings_online_multiview(model, data_batch: dict[str, torch.Tensor]) -> None:
|
| 487 |
+
output_text_embeddings, output_neg_text_embeddings, dropout_text_embeddings = (
|
| 488 |
+
compute_text_embeddings_online_multiview(model, data_batch)
|
| 489 |
+
)
|
| 490 |
+
t5_text_embeddings = {
|
| 491 |
+
"text_embeddings": output_text_embeddings,
|
| 492 |
+
"dropout_text_embeddings": dropout_text_embeddings,
|
| 493 |
+
}
|
| 494 |
+
neg_t5_text_embeddings = {
|
| 495 |
+
"text_embeddings": output_neg_text_embeddings,
|
| 496 |
+
"dropout_text_embeddings": dropout_text_embeddings,
|
| 497 |
+
}
|
| 498 |
+
data_batch["t5_text_embeddings"] = (
|
| 499 |
+
t5_text_embeddings if model.config.online_text_embeddings_as_dict else t5_text_embeddings["text_embeddings"]
|
| 500 |
+
)
|
| 501 |
+
data_batch["neg_t5_text_embeddings"] = (
|
| 502 |
+
neg_t5_text_embeddings
|
| 503 |
+
if model.config.online_text_embeddings_as_dict
|
| 504 |
+
else neg_t5_text_embeddings["text_embeddings"]
|
| 505 |
+
)
|
| 506 |
+
data_batch["t5_text_mask"] = torch.ones(
|
| 507 |
+
output_text_embeddings.shape[0], output_text_embeddings.shape[1], device="cuda"
|
| 508 |
+
)
|
| 509 |
+
|
| 510 |
+
|
| 511 |
+
def compute_empty_and_negative_text_embeddings(model):
|
| 512 |
+
# Compute empty string embeddings for text embedding dropout
|
| 513 |
+
if model.empty_string_text_embeddings is None:
|
| 514 |
+
empty_string_data_batch = {
|
| 515 |
+
model.input_caption_key: [" "],
|
| 516 |
+
}
|
| 517 |
+
model.empty_string_text_embeddings = model.text_encoder.compute_text_embeddings_online(
|
| 518 |
+
empty_string_data_batch, model.input_caption_key
|
| 519 |
+
)
|
| 520 |
+
|
| 521 |
+
# compute negative prompt embeddings for sampling
|
| 522 |
+
if model.neg_text_embeddings is None:
|
| 523 |
+
neg_promt_data_batch = {
|
| 524 |
+
model.input_caption_key: [_DEFAULT_NEGATIVE_PROMPT],
|
| 525 |
+
}
|
| 526 |
+
model.neg_text_embeddings = model.text_encoder.compute_text_embeddings_online(
|
| 527 |
+
neg_promt_data_batch, model.input_caption_key
|
| 528 |
+
)
|
| 529 |
+
|
| 530 |
+
|
| 531 |
+
def preprocess_databatch(
|
| 532 |
+
data_batch: dict[str, Any],
|
| 533 |
+
train_sample_views_range: Optional[tuple[int, int]],
|
| 534 |
+
) -> dict[str, Any]:
|
| 535 |
+
"""Preprocess data batch with dynamic view sampling."""
|
| 536 |
+
|
| 537 |
+
if TRAIN_SAMPLING_APPLIED_KEY in data_batch and data_batch[TRAIN_SAMPLING_APPLIED_KEY] is True:
|
| 538 |
+
return data_batch
|
| 539 |
+
if train_sample_views_range is not None:
|
| 540 |
+
min_views, max_views = train_sample_views_range
|
| 541 |
+
log.debug(f"Randomly sampling {min_views} to {max_views} views")
|
| 542 |
+
if SMOKE:
|
| 543 |
+
train_sample_n_views = 1
|
| 544 |
+
elif TRAIN_SAMPLE_N_VIEWS_KEY in data_batch:
|
| 545 |
+
train_sample_n_views = data_batch[TRAIN_SAMPLE_N_VIEWS_KEY]
|
| 546 |
+
log.debug(f"Using {TRAIN_SAMPLE_N_VIEWS_KEY} from data batch: {train_sample_n_views}")
|
| 547 |
+
if train_sample_n_views < 1: # No sampling is applied
|
| 548 |
+
return data_batch
|
| 549 |
+
else:
|
| 550 |
+
# Sample n_views to keep
|
| 551 |
+
train_sample_n_views = random.randint(min_views, max_views)
|
| 552 |
+
if dist.is_initialized():
|
| 553 |
+
# Have all ranks globally sample the same number of views such that tensors are same shape for efficiency
|
| 554 |
+
train_sample_n_views_tensor = torch.tensor(
|
| 555 |
+
[train_sample_n_views], device=data_batch["sample_n_views"].device, dtype=torch.int64
|
| 556 |
+
)
|
| 557 |
+
dist.broadcast(train_sample_n_views_tensor, src=0)
|
| 558 |
+
train_sample_n_views = train_sample_n_views_tensor[0].cpu().item()
|
| 559 |
+
|
| 560 |
+
n_views = data_batch["sample_n_views"].cpu().item()
|
| 561 |
+
log.debug(f"Sampling {train_sample_n_views} views out of {n_views}")
|
| 562 |
+
available_view_indices = list(range(n_views))
|
| 563 |
+
if data_batch["front_cam_view_idx_sample_position"][0] is not None and train_sample_n_views < n_views:
|
| 564 |
+
# In cases where we only have a single caption, we shouldn't drop the front camera view
|
| 565 |
+
front_cam_view_idx_sample_position = data_batch["front_cam_view_idx_sample_position"][0].cpu().item()
|
| 566 |
+
available_view_indices.pop(front_cam_view_idx_sample_position)
|
| 567 |
+
keep_view_indices = sorted(random.sample(available_view_indices, train_sample_n_views))
|
| 568 |
+
if parallel_state.get_context_parallel_group() is not None:
|
| 569 |
+
# Broadcast the exact views to keep for each context parallel group
|
| 570 |
+
keep_view_indices = broadcast(keep_view_indices, parallel_state.get_context_parallel_group())
|
| 571 |
+
log.debug(f"Sampled and broadcasted keep_view_indices={keep_view_indices}")
|
| 572 |
+
data_batch = sample_n_views_from_data_batch(data_batch, keep_view_indices)
|
| 573 |
+
data_batch[TRAIN_SAMPLING_APPLIED_KEY] = True
|
| 574 |
+
return data_batch
|
| 575 |
+
|
| 576 |
+
|
| 577 |
+
def training_step_multiview(
|
| 578 |
+
model, data_batch: dict[str, torch.Tensor], iteration: int
|
| 579 |
+
) -> tuple[dict[str, torch.Tensor], torch.Tensor]:
|
| 580 |
+
"""
|
| 581 |
+
Performs a single training step for the diffusion model.
|
| 582 |
+
|
| 583 |
+
This method is responsible for executing one iteration of the model's training. It involves:
|
| 584 |
+
1. Adding noise to the input data using the SDE process.
|
| 585 |
+
2. Passing the noisy data through the network to generate predictions.
|
| 586 |
+
3. Computing the loss based on the difference between the predictions and the original data, \
|
| 587 |
+
considering any configured loss weighting.
|
| 588 |
+
|
| 589 |
+
Args:
|
| 590 |
+
data_batch (dict): raw data batch draw from the training data loader.
|
| 591 |
+
iteration (int): Current iteration number.
|
| 592 |
+
|
| 593 |
+
Returns:
|
| 594 |
+
tuple: A tuple containing two elements:
|
| 595 |
+
- dict: additional data that used to debug / logging / callbacks
|
| 596 |
+
- Tensor: The computed loss for the training step as a PyTorch Tensor.
|
| 597 |
+
|
| 598 |
+
Raises:
|
| 599 |
+
AssertionError: If the class is conditional, \
|
| 600 |
+
but no number of classes is specified in the network configuration.
|
| 601 |
+
|
| 602 |
+
Notes:
|
| 603 |
+
- The method handles different types of conditioning
|
| 604 |
+
- The method also supports Kendall's loss
|
| 605 |
+
"""
|
| 606 |
+
model._update_train_stats(data_batch)
|
| 607 |
+
|
| 608 |
+
# only happens in training
|
| 609 |
+
data_batch = preprocess_databatch(data_batch, model.config.train_sample_views_range)
|
| 610 |
+
|
| 611 |
+
if model.config.text_encoder_config is not None and model.config.text_encoder_config.compute_online:
|
| 612 |
+
model.inplace_compute_text_embeddings_online(data_batch)
|
| 613 |
+
|
| 614 |
+
# Get the input data to noise and denoise~(image, video) and the corresponding conditioner.
|
| 615 |
+
_, x0_B_C_T_H_W, condition = model.get_data_and_condition(data_batch)
|
| 616 |
+
|
| 617 |
+
# Sample pertubation noise levels and N(0, 1) noises
|
| 618 |
+
epsilon_B_C_T_H_W = torch.randn(x0_B_C_T_H_W.size(), **model.tensor_kwargs_fp32)
|
| 619 |
+
batch_size = x0_B_C_T_H_W.size()[0]
|
| 620 |
+
t_B = model.rectified_flow.sample_train_time(batch_size).to(**model.tensor_kwargs_fp32)
|
| 621 |
+
t_B = rearrange(t_B, "b -> b 1") # add a dimension for T, all frames share the same sigma
|
| 622 |
+
x0_B_C_T_H_W, condition, epsilon_B_C_T_H_W, t_B = model.broadcast_split_for_model_parallelsim(
|
| 623 |
+
x0_B_C_T_H_W, condition, epsilon_B_C_T_H_W, t_B
|
| 624 |
+
)
|
| 625 |
+
timesteps = model.rectified_flow.get_discrete_timestamp(t_B, model.tensor_kwargs_fp32)
|
| 626 |
+
sigmas = model.rectified_flow.get_sigmas(
|
| 627 |
+
timesteps,
|
| 628 |
+
model.tensor_kwargs_fp32,
|
| 629 |
+
)
|
| 630 |
+
timesteps = rearrange(timesteps, "b -> b 1")
|
| 631 |
+
sigmas = rearrange(sigmas, "b -> b 1")
|
| 632 |
+
xt_B_C_T_H_W, vt_B_C_T_H_W = model.rectified_flow.get_interpolation(epsilon_B_C_T_H_W, x0_B_C_T_H_W, sigmas)
|
| 633 |
+
|
| 634 |
+
vt_pred_B_C_T_H_W = model.denoise(
|
| 635 |
+
noise=epsilon_B_C_T_H_W,
|
| 636 |
+
xt_B_C_T_H_W=xt_B_C_T_H_W.to(**model.tensor_kwargs),
|
| 637 |
+
timesteps_B_T=timesteps,
|
| 638 |
+
condition=condition,
|
| 639 |
+
)
|
| 640 |
+
|
| 641 |
+
time_weights_B = model.rectified_flow.train_time_weight(timesteps, model.tensor_kwargs_fp32)
|
| 642 |
+
per_instance_loss = torch.mean((vt_pred_B_C_T_H_W - vt_B_C_T_H_W) ** 2, dim=list(range(1, vt_pred_B_C_T_H_W.dim())))
|
| 643 |
+
|
| 644 |
+
loss = torch.mean(time_weights_B * per_instance_loss)
|
| 645 |
+
|
| 646 |
+
output_batch = {
|
| 647 |
+
"x0": x0_B_C_T_H_W,
|
| 648 |
+
"xt": xt_B_C_T_H_W,
|
| 649 |
+
"sigma": sigmas,
|
| 650 |
+
"condition": condition,
|
| 651 |
+
"model_pred": vt_pred_B_C_T_H_W,
|
| 652 |
+
"edm_loss": loss,
|
| 653 |
+
}
|
| 654 |
+
|
| 655 |
+
return output_batch, loss
|
cosmos_predict2/_src/predict2_multiview/models/view_sampling.py
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
"""Functions for sampling views from a multiview control video data batch."""
|
| 17 |
+
|
| 18 |
+
import torch
|
| 19 |
+
|
| 20 |
+
from cosmos_predict2._src.imaginaire.utils import log
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def sample_n_views_from_data_batch(data_batch: dict, keep_view_indices: list[int]) -> dict:
|
| 24 |
+
"""Sample views from data batch, handling both video and control inputs."""
|
| 25 |
+
|
| 26 |
+
keep_view_indices = sorted(keep_view_indices)
|
| 27 |
+
n_keep_views = len(keep_view_indices)
|
| 28 |
+
n_orig_views = data_batch["sample_n_views"].cpu().item()
|
| 29 |
+
if n_keep_views == n_orig_views:
|
| 30 |
+
log.debug("All views are requested to be kept, returning original data batch!")
|
| 31 |
+
return data_batch
|
| 32 |
+
num_video_frames_per_view = data_batch["num_video_frames_per_view"].cpu().item()
|
| 33 |
+
|
| 34 |
+
select_ids = []
|
| 35 |
+
for view_id in keep_view_indices:
|
| 36 |
+
select_ids.extend(list(range(view_id * num_video_frames_per_view, (view_id + 1) * num_video_frames_per_view)))
|
| 37 |
+
select_ids = torch.tensor(select_ids, device=data_batch["sample_n_views"].device, dtype=torch.int64)
|
| 38 |
+
|
| 39 |
+
data_batch["video"] = data_batch["video"][:, :, select_ids] # (B, C, V * T, H, W)
|
| 40 |
+
data_batch["view_indices"] = data_batch["view_indices"][:, select_ids] # (B, V * T)
|
| 41 |
+
data_batch["view_indices_selection"] = data_batch["view_indices_selection"][:, keep_view_indices]
|
| 42 |
+
data_batch["sample_n_views"] = 0 * data_batch["sample_n_views"] + n_keep_views
|
| 43 |
+
|
| 44 |
+
# process control input keys
|
| 45 |
+
for key in data_batch.keys():
|
| 46 |
+
if key.startswith("control_input_"):
|
| 47 |
+
data_batch[key] = data_batch[key][:, :, select_ids] # (B, C, V * T, H, W)
|
| 48 |
+
|
| 49 |
+
# process list-like keys
|
| 50 |
+
for sample_captions in data_batch["ai_caption"]:
|
| 51 |
+
if len(sample_captions) != n_orig_views:
|
| 52 |
+
raise ValueError(
|
| 53 |
+
f"Expected {n_orig_views} captions, got {len(sample_captions)} for key {key}. "
|
| 54 |
+
"If using single caption, view sampling is not currently supported."
|
| 55 |
+
)
|
| 56 |
+
data_batch["ai_caption"] = [
|
| 57 |
+
[caption for i, caption in enumerate(sample_captions) if i in keep_view_indices]
|
| 58 |
+
for sample_captions in data_batch["ai_caption"]
|
| 59 |
+
]
|
| 60 |
+
data_batch["camera_keys_selection"] = [
|
| 61 |
+
[camera_key for i, camera_key in enumerate(sample_camera_keys) if i in keep_view_indices]
|
| 62 |
+
for sample_camera_keys in data_batch["camera_keys_selection"]
|
| 63 |
+
]
|
| 64 |
+
|
| 65 |
+
return data_batch
|
cosmos_predict2/_src/predict2_multiview/networks/multiview_cross_dit.py
ADDED
|
@@ -0,0 +1,1187 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
import os
|
| 17 |
+
from collections import defaultdict, namedtuple
|
| 18 |
+
from dataclasses import dataclass
|
| 19 |
+
from enum import Enum
|
| 20 |
+
from typing import Dict, List, Optional, Tuple
|
| 21 |
+
|
| 22 |
+
import torch
|
| 23 |
+
import torch.amp as amp
|
| 24 |
+
import torch.nn as nn
|
| 25 |
+
from einops import rearrange
|
| 26 |
+
from megatron.core import parallel_state
|
| 27 |
+
from torch.distributed import ProcessGroup, get_process_group_ranks
|
| 28 |
+
from torch.distributed._composable.fsdp import fully_shard
|
| 29 |
+
from torch.utils.checkpoint import CheckpointPolicy, create_selective_checkpoint_contexts
|
| 30 |
+
from torchvision import transforms
|
| 31 |
+
|
| 32 |
+
from cosmos_predict2._src.imaginaire.utils import log
|
| 33 |
+
from cosmos_predict2._src.predict2.conditioner import DataType
|
| 34 |
+
from cosmos_predict2._src.predict2.networks.minimal_v1_lvg_dit import MinimalV1LVGDiT
|
| 35 |
+
from cosmos_predict2._src.predict2.networks.minimal_v4_dit import (
|
| 36 |
+
Attention,
|
| 37 |
+
Block,
|
| 38 |
+
SACConfig,
|
| 39 |
+
)
|
| 40 |
+
from cosmos_predict2._src.predict2_multiview.networks.multiview_dit import (
|
| 41 |
+
MultiCameraSinCosPosEmbAxis,
|
| 42 |
+
MultiCameraVideoRopePosition3DEmb,
|
| 43 |
+
)
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
# implementation of MultiViewCrossAttention changes multiview_dit.py
|
| 47 |
+
class MultiViewCrossAttention(Attention):
|
| 48 |
+
def __init__(self, *args, state_t: int = None, **kwargs) -> None:
|
| 49 |
+
super().__init__(*args, **kwargs)
|
| 50 |
+
assert self.qkv_format == "bshd", "MultiViewCrossAttention only supports qkv_format='bshd'"
|
| 51 |
+
self.state_t = state_t
|
| 52 |
+
|
| 53 |
+
def forward(self, x, context=None, rope_emb=None):
|
| 54 |
+
assert not self.is_selfattn, "MultiViewCrossAttention does not support self-attention"
|
| 55 |
+
B, L, D = x.shape
|
| 56 |
+
|
| 57 |
+
n_cameras = context.shape[1] // 512
|
| 58 |
+
x_B_L_D = rearrange(x, "B (V L) D -> (V B) L D", V=n_cameras)
|
| 59 |
+
context_B_M_D = rearrange(context, "B (V M) D -> (V B) M D", V=n_cameras) if context is not None else None
|
| 60 |
+
x_B_L_D = super().forward(x_B_L_D, context_B_M_D, rope_emb=rope_emb)
|
| 61 |
+
x_B_L_D = rearrange(x_B_L_D, "(V B) L D -> B (V L) D", V=n_cameras)
|
| 62 |
+
return x_B_L_D
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
# ---------------------- Selective Activation Checkpoint Policies -----------------------
|
| 66 |
+
def predict2_2B_crossview_720_context_fn():
|
| 67 |
+
op_count = defaultdict(int)
|
| 68 |
+
|
| 69 |
+
def policy_fn(ctx, func, *args, **kwargs):
|
| 70 |
+
mode = "recompute" if ctx.is_recompute else "forward"
|
| 71 |
+
|
| 72 |
+
if func == torch.ops.aten.mm.default:
|
| 73 |
+
op_count_key = f"{mode}_mm_count"
|
| 74 |
+
if op_count[op_count_key] >= 10:
|
| 75 |
+
result = CheckpointPolicy.MUST_SAVE
|
| 76 |
+
else:
|
| 77 |
+
result = CheckpointPolicy.PREFER_RECOMPUTE
|
| 78 |
+
|
| 79 |
+
# Update count for next operation
|
| 80 |
+
op_count[op_count_key] = (op_count[op_count_key] + 1) % 20
|
| 81 |
+
return result
|
| 82 |
+
|
| 83 |
+
if "flash_attn" in str(func):
|
| 84 |
+
return CheckpointPolicy.MUST_SAVE
|
| 85 |
+
|
| 86 |
+
return CheckpointPolicy.PREFER_RECOMPUTE
|
| 87 |
+
|
| 88 |
+
return create_selective_checkpoint_contexts(policy_fn)
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
class MultiViewCheckpointMode(str, Enum):
|
| 92 |
+
"""Checkpoint modes for MultiViewCrossDiT architecture."""
|
| 93 |
+
|
| 94 |
+
NONE = "none"
|
| 95 |
+
MM_ONLY = "mm_only"
|
| 96 |
+
BLOCK_WISE = "block_wise"
|
| 97 |
+
PREDICT2_2B_CROSSVIEW_720 = "predict2_2b_crossview_720"
|
| 98 |
+
|
| 99 |
+
def __str__(self) -> str:
|
| 100 |
+
return self.value
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
@dataclass
|
| 104 |
+
class MultiViewSACConfig(SACConfig):
|
| 105 |
+
"""Selective Activation Checkpoint Config for MultiViewCrossDiT."""
|
| 106 |
+
|
| 107 |
+
def get_context_fn(self):
|
| 108 |
+
if self.mode == MultiViewCheckpointMode.PREDICT2_2B_CROSSVIEW_720:
|
| 109 |
+
return predict2_2B_crossview_720_context_fn
|
| 110 |
+
else:
|
| 111 |
+
return super().get_context_fn()
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
VideoSize = namedtuple("VideoSize", ["T", "H", "W"])
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
class CrossViewAttention(Attention):
|
| 118 |
+
def __init__(self, *args, cross_view_attn_map: Dict[int, List[int]], **kwargs):
|
| 119 |
+
super().__init__(*args, **kwargs)
|
| 120 |
+
del self.attn_op
|
| 121 |
+
if self.backend == "transformer_engine":
|
| 122 |
+
from transformer_engine.pytorch.attention import DotProductAttention
|
| 123 |
+
|
| 124 |
+
self.attn_op = DotProductAttention(
|
| 125 |
+
self.n_heads,
|
| 126 |
+
self.head_dim,
|
| 127 |
+
num_gqa_groups=self.n_heads,
|
| 128 |
+
attention_dropout=0,
|
| 129 |
+
qkv_format=self.qkv_format,
|
| 130 |
+
attn_mask_type="padding", # important
|
| 131 |
+
attention_type="cross", # important
|
| 132 |
+
)
|
| 133 |
+
else:
|
| 134 |
+
raise NotImplementedError(f"Backend {self.backend} not supported")
|
| 135 |
+
self.cross_view_attn_map = cross_view_attn_map
|
| 136 |
+
self.max_neighbors = max(len(neighbors) for neighbors in cross_view_attn_map.values())
|
| 137 |
+
self.neighbor_indices = None
|
| 138 |
+
self.neighbor_mask = None
|
| 139 |
+
|
| 140 |
+
def forward(self, x, view_indices_B_V, sv_video_size: VideoSize):
|
| 141 |
+
"""
|
| 142 |
+
x: (B, V, L, D)
|
| 143 |
+
view_indices_B_V: (B, V)
|
| 144 |
+
sv_video_size: VideoSize (T, H, W), where T * H * W = L
|
| 145 |
+
"""
|
| 146 |
+
assert not self.is_selfattn, "CrossViewAttention does not support self-attention"
|
| 147 |
+
B, V, L, D = x.shape
|
| 148 |
+
T, H, W = sv_video_size
|
| 149 |
+
assert T * H * W == L, f"T * H * W != L: {T * H * W} != {L}"
|
| 150 |
+
|
| 151 |
+
# move time dimension to batch dimension
|
| 152 |
+
x = rearrange(x, "b v (t h w) d -> (b t) v (h w) d", t=T, h=H, w=W)
|
| 153 |
+
B, V, L, D = x.shape
|
| 154 |
+
|
| 155 |
+
view_indices_B_V = view_indices_B_V.repeat_interleave(T, dim=0).long()
|
| 156 |
+
|
| 157 |
+
# Create neighbor indices and mask on the fly, only once.
|
| 158 |
+
if self.neighbor_indices is None or self.neighbor_indices.device != x.device:
|
| 159 |
+
num_total_views = len(self.cross_view_attn_map)
|
| 160 |
+
neighbor_indices = torch.zeros((num_total_views, self.max_neighbors), dtype=torch.long, device=x.device)
|
| 161 |
+
neighbor_mask = torch.zeros((num_total_views, self.max_neighbors), dtype=torch.bool, device=x.device)
|
| 162 |
+
for i in range(num_total_views):
|
| 163 |
+
neighbors = self.cross_view_attn_map[i]
|
| 164 |
+
for j, neighbor_idx in enumerate(neighbors):
|
| 165 |
+
neighbor_indices[i, j] = neighbor_idx
|
| 166 |
+
neighbor_mask[i, j] = True
|
| 167 |
+
self.neighbor_indices = neighbor_indices
|
| 168 |
+
self.neighbor_mask = neighbor_mask
|
| 169 |
+
|
| 170 |
+
num_total_views = len(self.cross_view_attn_map)
|
| 171 |
+
view_indices_to_tensor_pos = torch.full(
|
| 172 |
+
(B, num_total_views), -1, dtype=torch.long, device=x.device
|
| 173 |
+
) # include out of range view index
|
| 174 |
+
b_indices = torch.arange(B, device=x.device).unsqueeze(1).expand(-1, V).long()
|
| 175 |
+
view_indices_to_tensor_pos[b_indices, view_indices_B_V] = (
|
| 176 |
+
torch.arange(V, device=x.device).unsqueeze(0).expand(B, -1)
|
| 177 |
+
)
|
| 178 |
+
|
| 179 |
+
neighbor_view_indices = self.neighbor_indices[view_indices_B_V] # may include out of range view index
|
| 180 |
+
gather_tensor_pos = view_indices_to_tensor_pos[
|
| 181 |
+
b_indices.unsqueeze(2), neighbor_view_indices
|
| 182 |
+
] # [B, V, max_neighbors], out of range view index will be -1
|
| 183 |
+
|
| 184 |
+
# Sort to move all -1 to the end, which is convenient for creating attention mask.
|
| 185 |
+
gather_tensor_pos, sorted_indices = torch.sort(gather_tensor_pos, dim=-1, descending=True)
|
| 186 |
+
|
| 187 |
+
b_indices_for_gather = torch.arange(B, device=x.device)[:, None, None]
|
| 188 |
+
# Clamp to avoid index error. Masked values will be ignored in attention.
|
| 189 |
+
neighbor_features = x[
|
| 190 |
+
b_indices_for_gather, torch.clamp(gather_tensor_pos, min=0)
|
| 191 |
+
] # [B, V, max_neighbors, L, C]
|
| 192 |
+
|
| 193 |
+
# Prepare for attention
|
| 194 |
+
query = self.q_proj(rearrange(x, "b v l c -> (b v) l c")) # [B*V, L, C]
|
| 195 |
+
context = rearrange(neighbor_features, "b v n l c -> (b v) (n l) c") # [B*V, max_neighbors*L, C]
|
| 196 |
+
key = self.k_proj(context)
|
| 197 |
+
value = self.v_proj(context)
|
| 198 |
+
|
| 199 |
+
q, k, v = map(
|
| 200 |
+
lambda t: rearrange(t, "b ... (h d) -> b ... h d", h=self.n_heads, d=self.head_dim),
|
| 201 |
+
(query, key, value),
|
| 202 |
+
)
|
| 203 |
+
|
| 204 |
+
q = self.q_norm(q)
|
| 205 |
+
k = self.k_norm(k)
|
| 206 |
+
v = self.v_norm(v)
|
| 207 |
+
|
| 208 |
+
# Create attention mask
|
| 209 |
+
is_neighbor_present = gather_tensor_pos != -1 # [B, V, max_neighbors]
|
| 210 |
+
mask_for_input_views = self.neighbor_mask[view_indices_B_V] # [B, V, n]
|
| 211 |
+
|
| 212 |
+
# Reorder mask_for_input_views to match the sorted gather_tensor_pos
|
| 213 |
+
mask_for_input_views = torch.gather(mask_for_input_views, -1, sorted_indices)
|
| 214 |
+
final_mask = is_neighbor_present & mask_for_input_views
|
| 215 |
+
|
| 216 |
+
mask_per_view = rearrange(final_mask, "b v n -> (b v) n") # [BV, n]
|
| 217 |
+
mask_kv = mask_per_view.repeat_interleave(L, dim=1) # [BV, n*L]
|
| 218 |
+
|
| 219 |
+
# Reshape mask to [batch_size, 1, 1, max_seqlen_kv] as per official documentation.
|
| 220 |
+
mask = rearrange(mask_kv, "bv l_kv -> bv 1 1 l_kv") # [BV, 1, 1, n*L]
|
| 221 |
+
atten_mask_kv = ~mask # 0 means keep, 1 means mask
|
| 222 |
+
atten_mask_q = torch.zeros(query.shape[0], 1, 1, query.shape[1]).to(atten_mask_kv)
|
| 223 |
+
|
| 224 |
+
attention_output = self.attn_op(q, k, v, attention_mask=(atten_mask_q, atten_mask_kv))
|
| 225 |
+
attention_output = attention_output.flatten(2) # [B*V, L, H*D]
|
| 226 |
+
output = self.output_dropout(self.output_proj(attention_output))
|
| 227 |
+
output = rearrange(output, "(b v) l d -> b v l d", v=V)
|
| 228 |
+
# recover time dimension from batch to seq
|
| 229 |
+
output = rearrange(output, "(b t) v (h w) d -> b v (t h w) d", t=T, h=H, w=W)
|
| 230 |
+
return output
|
| 231 |
+
|
| 232 |
+
def set_context_parallel_group(self, process_group, ranks, stream):
|
| 233 |
+
raise NotImplementedError("Cross View Attention doesn't need communication")
|
| 234 |
+
|
| 235 |
+
|
| 236 |
+
class MultiViewCrossBlock(Block):
|
| 237 |
+
"""
|
| 238 |
+
A transformer block that takes n_cameras as input.
|
| 239 |
+
Self-Attention (Single View) -> Cross-View Attention -> Cross Attention (text and image)
|
| 240 |
+
"""
|
| 241 |
+
|
| 242 |
+
def __init__(
|
| 243 |
+
self,
|
| 244 |
+
x_dim: int,
|
| 245 |
+
context_dim: int,
|
| 246 |
+
num_heads: int,
|
| 247 |
+
mlp_ratio: float = 4.0,
|
| 248 |
+
use_adaln_lora: bool = False,
|
| 249 |
+
adaln_lora_dim: int = 256,
|
| 250 |
+
cross_view_attn_map: Dict[int, List[int]] = None,
|
| 251 |
+
state_t: int = None,
|
| 252 |
+
backend: str = "transformer_engine",
|
| 253 |
+
image_context_dim: Optional[int] = None,
|
| 254 |
+
use_wan_fp32_strategy: bool = False,
|
| 255 |
+
enable_cross_view_attn: bool = False,
|
| 256 |
+
):
|
| 257 |
+
super().__init__(
|
| 258 |
+
x_dim,
|
| 259 |
+
context_dim,
|
| 260 |
+
num_heads,
|
| 261 |
+
mlp_ratio,
|
| 262 |
+
use_adaln_lora,
|
| 263 |
+
adaln_lora_dim,
|
| 264 |
+
backend,
|
| 265 |
+
image_context_dim,
|
| 266 |
+
use_wan_fp32_strategy,
|
| 267 |
+
)
|
| 268 |
+
self.state_t = state_t
|
| 269 |
+
self.cross_view_attn_map = cross_view_attn_map
|
| 270 |
+
self.enable_cross_view_attn = enable_cross_view_attn
|
| 271 |
+
if image_context_dim is None:
|
| 272 |
+
del self.cross_attn
|
| 273 |
+
# cross attention to text and image condition
|
| 274 |
+
self.cross_attn = MultiViewCrossAttention(
|
| 275 |
+
x_dim,
|
| 276 |
+
context_dim,
|
| 277 |
+
num_heads,
|
| 278 |
+
x_dim // num_heads,
|
| 279 |
+
qkv_format="bshd",
|
| 280 |
+
state_t=state_t,
|
| 281 |
+
use_wan_fp32_strategy=use_wan_fp32_strategy,
|
| 282 |
+
)
|
| 283 |
+
else:
|
| 284 |
+
raise NotImplementedError("image_context_dim is not supported for MultiViewBlock")
|
| 285 |
+
|
| 286 |
+
if enable_cross_view_attn:
|
| 287 |
+
self.cross_view_attn = CrossViewAttention(
|
| 288 |
+
x_dim,
|
| 289 |
+
x_dim, # context_dim, can not set to None
|
| 290 |
+
num_heads,
|
| 291 |
+
x_dim // num_heads,
|
| 292 |
+
qkv_format="bshd",
|
| 293 |
+
use_wan_fp32_strategy=use_wan_fp32_strategy,
|
| 294 |
+
cross_view_attn_map=cross_view_attn_map,
|
| 295 |
+
)
|
| 296 |
+
# no modulation so we set elementwise_affine=True
|
| 297 |
+
self.layer_norm_cross_view_attn = nn.LayerNorm(x_dim, elementwise_affine=True, eps=1e-6)
|
| 298 |
+
|
| 299 |
+
def reset_parameters(self):
|
| 300 |
+
super().reset_parameters()
|
| 301 |
+
if self.enable_cross_view_attn:
|
| 302 |
+
self.layer_norm_cross_view_attn.reset_parameters()
|
| 303 |
+
|
| 304 |
+
def init_weights(self):
|
| 305 |
+
super().init_weights()
|
| 306 |
+
if self.enable_cross_view_attn:
|
| 307 |
+
self.cross_view_attn.init_weights()
|
| 308 |
+
|
| 309 |
+
# Zero-initialize the output projection
|
| 310 |
+
torch.nn.init.zeros_(self.cross_view_attn.output_proj.weight)
|
| 311 |
+
if self.cross_view_attn.output_proj.bias is not None:
|
| 312 |
+
torch.nn.init.zeros_(self.cross_view_attn.output_proj.bias)
|
| 313 |
+
|
| 314 |
+
# most code are copied from parent. insert cross view attn between self attn and cross attn
|
| 315 |
+
def forward(
|
| 316 |
+
self,
|
| 317 |
+
x_B_T_H_W_D: torch.Tensor,
|
| 318 |
+
view_indices_B_T: torch.Tensor,
|
| 319 |
+
emb_B_T_D: torch.Tensor,
|
| 320 |
+
view_embedding_proj_B_V_9D: Optional[torch.Tensor] = None,
|
| 321 |
+
crossattn_emb: Optional[torch.Tensor] = None,
|
| 322 |
+
rope_emb_L_1_1_D: Optional[torch.Tensor] = None,
|
| 323 |
+
adaln_lora_B_T_3D: Optional[torch.Tensor] = None,
|
| 324 |
+
extra_per_block_pos_emb: Optional[torch.Tensor] = None,
|
| 325 |
+
block_idx: Optional[int] = None,
|
| 326 |
+
) -> torch.Tensor:
|
| 327 |
+
"""
|
| 328 |
+
x_B_T_H_W_D: (B, T, H, W, D)
|
| 329 |
+
view_indices_B_T: (B, T)
|
| 330 |
+
"""
|
| 331 |
+
num_cameras = torch.unique(view_indices_B_T[0]).shape[0]
|
| 332 |
+
|
| 333 |
+
if extra_per_block_pos_emb is not None:
|
| 334 |
+
x_B_T_H_W_D = x_B_T_H_W_D + extra_per_block_pos_emb
|
| 335 |
+
|
| 336 |
+
with amp.autocast("cuda", enabled=self.use_wan_fp32_strategy, dtype=torch.float32):
|
| 337 |
+
if self.use_adaln_lora:
|
| 338 |
+
shift_self_attn_B_T_D, scale_self_attn_B_T_D, gate_self_attn_B_T_D = (
|
| 339 |
+
self.adaln_modulation_self_attn(emb_B_T_D) + adaln_lora_B_T_3D
|
| 340 |
+
).chunk(3, dim=-1)
|
| 341 |
+
shift_cross_attn_B_T_D, scale_cross_attn_B_T_D, gate_cross_attn_B_T_D = (
|
| 342 |
+
self.adaln_modulation_cross_attn(emb_B_T_D) + adaln_lora_B_T_3D
|
| 343 |
+
).chunk(3, dim=-1)
|
| 344 |
+
shift_mlp_B_T_D, scale_mlp_B_T_D, gate_mlp_B_T_D = (
|
| 345 |
+
self.adaln_modulation_mlp(emb_B_T_D) + adaln_lora_B_T_3D
|
| 346 |
+
).chunk(3, dim=-1)
|
| 347 |
+
else:
|
| 348 |
+
shift_self_attn_B_T_D, scale_self_attn_B_T_D, gate_self_attn_B_T_D = self.adaln_modulation_self_attn(
|
| 349 |
+
emb_B_T_D
|
| 350 |
+
).chunk(3, dim=-1)
|
| 351 |
+
shift_cross_attn_B_T_D, scale_cross_attn_B_T_D, gate_cross_attn_B_T_D = (
|
| 352 |
+
self.adaln_modulation_cross_attn(emb_B_T_D).chunk(3, dim=-1)
|
| 353 |
+
)
|
| 354 |
+
shift_mlp_B_T_D, scale_mlp_B_T_D, gate_mlp_B_T_D = self.adaln_modulation_mlp(emb_B_T_D).chunk(3, dim=-1)
|
| 355 |
+
|
| 356 |
+
# Reshape tensors from (B, T, D) to (B, T, 1, 1, D) for broadcasting
|
| 357 |
+
shift_self_attn_B_T_1_1_D = rearrange(shift_self_attn_B_T_D, "b t d -> b t 1 1 d").type_as(x_B_T_H_W_D)
|
| 358 |
+
scale_self_attn_B_T_1_1_D = rearrange(scale_self_attn_B_T_D, "b t d -> b t 1 1 d").type_as(x_B_T_H_W_D)
|
| 359 |
+
gate_self_attn_B_T_1_1_D = rearrange(gate_self_attn_B_T_D, "b t d -> b t 1 1 d").type_as(x_B_T_H_W_D)
|
| 360 |
+
|
| 361 |
+
shift_cross_attn_B_T_1_1_D = rearrange(shift_cross_attn_B_T_D, "b t d -> b t 1 1 d").type_as(x_B_T_H_W_D)
|
| 362 |
+
scale_cross_attn_B_T_1_1_D = rearrange(scale_cross_attn_B_T_D, "b t d -> b t 1 1 d").type_as(x_B_T_H_W_D)
|
| 363 |
+
gate_cross_attn_B_T_1_1_D = rearrange(gate_cross_attn_B_T_D, "b t d -> b t 1 1 d").type_as(x_B_T_H_W_D)
|
| 364 |
+
|
| 365 |
+
shift_mlp_B_T_1_1_D = rearrange(shift_mlp_B_T_D, "b t d -> b t 1 1 d").type_as(x_B_T_H_W_D)
|
| 366 |
+
scale_mlp_B_T_1_1_D = rearrange(scale_mlp_B_T_D, "b t d -> b t 1 1 d").type_as(x_B_T_H_W_D)
|
| 367 |
+
gate_mlp_B_T_1_1_D = rearrange(gate_mlp_B_T_D, "b t d -> b t 1 1 d").type_as(x_B_T_H_W_D)
|
| 368 |
+
|
| 369 |
+
if view_embedding_proj_B_V_9D is not None:
|
| 370 |
+
(
|
| 371 |
+
view_shift_self_attn_B_V_D,
|
| 372 |
+
view_scale_self_attn_B_V_D,
|
| 373 |
+
view_gate_self_attn_B_V_D,
|
| 374 |
+
view_shift_cross_attn_B_V_D,
|
| 375 |
+
view_scale_cross_attn_B_V_D,
|
| 376 |
+
view_gate_cross_attn_B_V_D,
|
| 377 |
+
view_shift_mlp_B_V_D,
|
| 378 |
+
view_scale_mlp_B_V_D,
|
| 379 |
+
view_gate_mlp_B_V_D,
|
| 380 |
+
) = view_embedding_proj_B_V_9D.chunk(9, dim=-1)
|
| 381 |
+
|
| 382 |
+
repeat_t = x_B_T_H_W_D.shape[1] // num_cameras
|
| 383 |
+
assert repeat_t * num_cameras == x_B_T_H_W_D.shape[1]
|
| 384 |
+
|
| 385 |
+
def expand_and_rearrange(x_B_V_D):
|
| 386 |
+
expanded_x_B_V_T_D = x_B_V_D.unsqueeze(2).expand(-1, -1, repeat_t, -1)
|
| 387 |
+
return rearrange(expanded_x_B_V_T_D, "b v t d -> b (v t) 1 1 d")
|
| 388 |
+
|
| 389 |
+
shift_self_attn_B_T_1_1_D = shift_self_attn_B_T_1_1_D + expand_and_rearrange(
|
| 390 |
+
view_shift_self_attn_B_V_D
|
| 391 |
+
).type_as(x_B_T_H_W_D)
|
| 392 |
+
scale_self_attn_B_T_1_1_D = scale_self_attn_B_T_1_1_D + expand_and_rearrange(
|
| 393 |
+
view_scale_self_attn_B_V_D
|
| 394 |
+
).type_as(x_B_T_H_W_D)
|
| 395 |
+
gate_self_attn_B_T_1_1_D = gate_self_attn_B_T_1_1_D + expand_and_rearrange(
|
| 396 |
+
view_gate_self_attn_B_V_D
|
| 397 |
+
).type_as(x_B_T_H_W_D)
|
| 398 |
+
shift_cross_attn_B_T_1_1_D = shift_cross_attn_B_T_1_1_D + expand_and_rearrange(
|
| 399 |
+
view_shift_cross_attn_B_V_D
|
| 400 |
+
).type_as(x_B_T_H_W_D)
|
| 401 |
+
scale_cross_attn_B_T_1_1_D = scale_cross_attn_B_T_1_1_D + expand_and_rearrange(
|
| 402 |
+
view_scale_cross_attn_B_V_D
|
| 403 |
+
).type_as(x_B_T_H_W_D)
|
| 404 |
+
gate_cross_attn_B_T_1_1_D = gate_cross_attn_B_T_1_1_D + expand_and_rearrange(
|
| 405 |
+
view_gate_cross_attn_B_V_D
|
| 406 |
+
).type_as(x_B_T_H_W_D)
|
| 407 |
+
shift_mlp_B_T_1_1_D = shift_mlp_B_T_1_1_D + expand_and_rearrange(view_shift_mlp_B_V_D).type_as(x_B_T_H_W_D)
|
| 408 |
+
scale_mlp_B_T_1_1_D = scale_mlp_B_T_1_1_D + expand_and_rearrange(view_scale_mlp_B_V_D).type_as(x_B_T_H_W_D)
|
| 409 |
+
gate_mlp_B_T_1_1_D = gate_mlp_B_T_1_1_D + expand_and_rearrange(view_gate_mlp_B_V_D).type_as(x_B_T_H_W_D)
|
| 410 |
+
|
| 411 |
+
B, T, H, W, D = x_B_T_H_W_D.shape
|
| 412 |
+
|
| 413 |
+
def _fn(_x_B_T_H_W_D, _norm_layer, _scale_B_T_1_1_D, _shift_B_T_1_1_D):
|
| 414 |
+
return _norm_layer(_x_B_T_H_W_D) * (1 + _scale_B_T_1_1_D) + _shift_B_T_1_1_D
|
| 415 |
+
|
| 416 |
+
normalized_x_B_T_H_W_D = _fn(
|
| 417 |
+
x_B_T_H_W_D,
|
| 418 |
+
self.layer_norm_self_attn,
|
| 419 |
+
scale_self_attn_B_T_1_1_D,
|
| 420 |
+
shift_self_attn_B_T_1_1_D,
|
| 421 |
+
)
|
| 422 |
+
|
| 423 |
+
rope_emb_L_1_1_D_sv = rearrange(
|
| 424 |
+
rope_emb_L_1_1_D,
|
| 425 |
+
"(v m) 1 1 d -> v m 1 1 d",
|
| 426 |
+
v=num_cameras,
|
| 427 |
+
)[0]
|
| 428 |
+
|
| 429 |
+
result_B_T_H_W_D = rearrange(
|
| 430 |
+
self.self_attn(
|
| 431 |
+
rearrange(normalized_x_B_T_H_W_D, "b (v t) h w d -> (b v) (t h w) d", v=num_cameras),
|
| 432 |
+
None,
|
| 433 |
+
rope_emb=rope_emb_L_1_1_D_sv,
|
| 434 |
+
),
|
| 435 |
+
"(b v) (t h w) d -> b (v t) h w d",
|
| 436 |
+
v=num_cameras,
|
| 437 |
+
h=H,
|
| 438 |
+
w=W,
|
| 439 |
+
)
|
| 440 |
+
x_B_T_H_W_D = x_B_T_H_W_D + gate_self_attn_B_T_1_1_D * result_B_T_H_W_D
|
| 441 |
+
|
| 442 |
+
# insert cross view attn here. x_B_T_H_W_D
|
| 443 |
+
if self.enable_cross_view_attn:
|
| 444 |
+
num_cameras = torch.unique(view_indices_B_T[0]).shape[0]
|
| 445 |
+
x_B_V_T_H_W_D = rearrange(x_B_T_H_W_D, "b (v t) h w d -> b v t h w d", v=num_cameras)
|
| 446 |
+
sv_video_size = VideoSize(T=x_B_V_T_H_W_D.shape[2], H=x_B_V_T_H_W_D.shape[3], W=x_B_V_T_H_W_D.shape[4])
|
| 447 |
+
x_B_V_L_D = rearrange(x_B_V_T_H_W_D, "b v t h w d -> b v (t h w) d")
|
| 448 |
+
view_indices_B_V = rearrange(view_indices_B_T, "b (v t) -> b v t", v=num_cameras)[..., 0]
|
| 449 |
+
result_cross_view_attn_B_T_H_W_D = rearrange(
|
| 450 |
+
self.cross_view_attn(self.layer_norm_cross_view_attn(x_B_V_L_D), view_indices_B_V, sv_video_size),
|
| 451 |
+
"b v (t h w) d -> b (v t) h w d",
|
| 452 |
+
v=num_cameras,
|
| 453 |
+
t=sv_video_size.T,
|
| 454 |
+
h=sv_video_size.H,
|
| 455 |
+
w=sv_video_size.W,
|
| 456 |
+
)
|
| 457 |
+
x_B_T_H_W_D = x_B_T_H_W_D + result_cross_view_attn_B_T_H_W_D
|
| 458 |
+
|
| 459 |
+
def _x_fn(
|
| 460 |
+
_x_B_T_H_W_D,
|
| 461 |
+
layer_norm_cross_attn,
|
| 462 |
+
_scale_cross_attn_B_T_1_1_D,
|
| 463 |
+
_shift_cross_attn_B_T_1_1_D,
|
| 464 |
+
_gate_cross_attn_B_T_1_1_D,
|
| 465 |
+
):
|
| 466 |
+
_normalized_x_B_T_H_W_D = _fn(
|
| 467 |
+
_x_B_T_H_W_D, layer_norm_cross_attn, _scale_cross_attn_B_T_1_1_D, _shift_cross_attn_B_T_1_1_D
|
| 468 |
+
)
|
| 469 |
+
_result_B_T_H_W_D = rearrange(
|
| 470 |
+
self.cross_attn(
|
| 471 |
+
rearrange(_normalized_x_B_T_H_W_D, "b t h w d -> b (t h w) d"),
|
| 472 |
+
crossattn_emb,
|
| 473 |
+
rope_emb=rope_emb_L_1_1_D,
|
| 474 |
+
),
|
| 475 |
+
"b (t h w) d -> b t h w d",
|
| 476 |
+
t=T,
|
| 477 |
+
h=H,
|
| 478 |
+
w=W,
|
| 479 |
+
)
|
| 480 |
+
# _x_B_T_H_W_D = _x_B_T_H_W_D + _gate_cross_attn_B_T_1_1_D * _result_B_T_H_W_D
|
| 481 |
+
return _result_B_T_H_W_D
|
| 482 |
+
|
| 483 |
+
result_B_T_H_W_D = _x_fn(
|
| 484 |
+
x_B_T_H_W_D,
|
| 485 |
+
self.layer_norm_cross_attn,
|
| 486 |
+
scale_cross_attn_B_T_1_1_D,
|
| 487 |
+
shift_cross_attn_B_T_1_1_D,
|
| 488 |
+
gate_cross_attn_B_T_1_1_D,
|
| 489 |
+
)
|
| 490 |
+
x_B_T_H_W_D = result_B_T_H_W_D * gate_cross_attn_B_T_1_1_D + x_B_T_H_W_D
|
| 491 |
+
|
| 492 |
+
normalized_x_B_T_H_W_D = _fn(
|
| 493 |
+
x_B_T_H_W_D,
|
| 494 |
+
self.layer_norm_mlp,
|
| 495 |
+
scale_mlp_B_T_1_1_D,
|
| 496 |
+
shift_mlp_B_T_1_1_D,
|
| 497 |
+
)
|
| 498 |
+
result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D)
|
| 499 |
+
x_B_T_H_W_D = x_B_T_H_W_D + gate_mlp_B_T_1_1_D * result_B_T_H_W_D
|
| 500 |
+
|
| 501 |
+
return x_B_T_H_W_D
|
| 502 |
+
|
| 503 |
+
|
| 504 |
+
class MultiViewCrossDiT(MinimalV1LVGDiT):
|
| 505 |
+
def __init__(
|
| 506 |
+
self,
|
| 507 |
+
*args,
|
| 508 |
+
timestep_scale: float = 1.0,
|
| 509 |
+
crossattn_emb_channels: int = 1024,
|
| 510 |
+
mlp_ratio: float = 4.0,
|
| 511 |
+
state_t: int,
|
| 512 |
+
n_cameras_emb: int,
|
| 513 |
+
view_condition_dim: int,
|
| 514 |
+
concat_view_embedding: bool,
|
| 515 |
+
adaln_view_embedding: bool,
|
| 516 |
+
layer_mask: Optional[List[bool]] = None,
|
| 517 |
+
sac_config: MultiViewSACConfig = MultiViewSACConfig(),
|
| 518 |
+
enable_cross_view_attn: bool = False,
|
| 519 |
+
cross_view_attn_map_str: Optional[Dict] = None,
|
| 520 |
+
camera_to_view_id: Optional[Dict] = None,
|
| 521 |
+
init_cross_view_attn_weight_from: Optional[str] = None,
|
| 522 |
+
init_cross_view_attn_weight_credentials: Optional[str] = None,
|
| 523 |
+
**kwargs,
|
| 524 |
+
):
|
| 525 |
+
self.crossattn_emb_channels = crossattn_emb_channels
|
| 526 |
+
self.mlp_ratio = mlp_ratio
|
| 527 |
+
self.state_t = state_t
|
| 528 |
+
self.n_cameras_emb = n_cameras_emb
|
| 529 |
+
self.view_condition_dim = view_condition_dim
|
| 530 |
+
self.concat_view_embedding = concat_view_embedding
|
| 531 |
+
self.adaln_view_embedding = adaln_view_embedding
|
| 532 |
+
self.enable_cross_view_attn = enable_cross_view_attn
|
| 533 |
+
self.init_cross_view_attn_weight_from = init_cross_view_attn_weight_from
|
| 534 |
+
self.init_cross_view_attn_weight_credentials = init_cross_view_attn_weight_credentials
|
| 535 |
+
|
| 536 |
+
assert not (self.adaln_view_embedding and self.concat_view_embedding), (
|
| 537 |
+
"adaln_view_embedding and concat_view_embedding cannot be True at the same time"
|
| 538 |
+
)
|
| 539 |
+
assert "in_channels" in kwargs, "in_channels must be provided"
|
| 540 |
+
kwargs["in_channels"] += (
|
| 541 |
+
self.view_condition_dim if self.concat_view_embedding else 0
|
| 542 |
+
) # this avoids overwritting build_patch_embed which still adds padding_mask channel as appropriate
|
| 543 |
+
assert layer_mask is None, "layer_mask is not supported for MultiViewDiT"
|
| 544 |
+
if "n_cameras" in kwargs:
|
| 545 |
+
del kwargs["n_cameras"]
|
| 546 |
+
super().__init__(
|
| 547 |
+
*args,
|
| 548 |
+
mlp_ratio=mlp_ratio,
|
| 549 |
+
timestep_scale=timestep_scale,
|
| 550 |
+
crossattn_emb_channels=crossattn_emb_channels,
|
| 551 |
+
sac_config=sac_config,
|
| 552 |
+
**kwargs,
|
| 553 |
+
)
|
| 554 |
+
|
| 555 |
+
cross_view_attn_map = {}
|
| 556 |
+
for source_view, target_views in cross_view_attn_map_str.items():
|
| 557 |
+
cross_view_attn_map[int(camera_to_view_id[source_view])] = []
|
| 558 |
+
for target_view in target_views:
|
| 559 |
+
cross_view_attn_map[int(camera_to_view_id[source_view])].append(int(camera_to_view_id[target_view]))
|
| 560 |
+
self.cross_view_attn_map = cross_view_attn_map
|
| 561 |
+
|
| 562 |
+
del self.blocks
|
| 563 |
+
self.blocks = nn.ModuleList(
|
| 564 |
+
[
|
| 565 |
+
MultiViewCrossBlock(
|
| 566 |
+
x_dim=self.model_channels,
|
| 567 |
+
context_dim=self.crossattn_emb_channels,
|
| 568 |
+
num_heads=self.num_heads,
|
| 569 |
+
mlp_ratio=self.mlp_ratio,
|
| 570 |
+
use_adaln_lora=self.use_adaln_lora,
|
| 571 |
+
adaln_lora_dim=self.adaln_lora_dim,
|
| 572 |
+
backend=self.atten_backend,
|
| 573 |
+
image_context_dim=None if self.extra_image_context_dim is None else self.model_channels,
|
| 574 |
+
state_t=self.state_t,
|
| 575 |
+
use_wan_fp32_strategy=self.use_wan_fp32_strategy,
|
| 576 |
+
cross_view_attn_map=self.cross_view_attn_map,
|
| 577 |
+
enable_cross_view_attn=self.enable_cross_view_attn,
|
| 578 |
+
)
|
| 579 |
+
for _ in range(self.num_blocks)
|
| 580 |
+
]
|
| 581 |
+
)
|
| 582 |
+
|
| 583 |
+
if self.concat_view_embedding:
|
| 584 |
+
self.view_embeddings = nn.Embedding(self.n_cameras_emb, view_condition_dim)
|
| 585 |
+
|
| 586 |
+
if self.adaln_view_embedding:
|
| 587 |
+
self.adaln_view_embedder = nn.Embedding(self.n_cameras_emb, self.model_channels)
|
| 588 |
+
# cosmos use adaln in self-attn, cross-attn, mlp
|
| 589 |
+
self.adaln_view_proj = nn.Linear(self.model_channels, self.model_channels * 9)
|
| 590 |
+
|
| 591 |
+
self.init_weights()
|
| 592 |
+
self.enable_selective_checkpoint(sac_config, self.blocks)
|
| 593 |
+
|
| 594 |
+
def fully_shard(self, mesh, **fsdp_kwargs):
|
| 595 |
+
for i, block in enumerate(self.blocks):
|
| 596 |
+
reshard_after_forward = i < len(self.blocks) - 1
|
| 597 |
+
fully_shard(block, mesh=mesh, reshard_after_forward=reshard_after_forward, **fsdp_kwargs)
|
| 598 |
+
|
| 599 |
+
fully_shard(self.final_layer, mesh=mesh, reshard_after_forward=True, **fsdp_kwargs)
|
| 600 |
+
if self.extra_per_block_abs_pos_emb:
|
| 601 |
+
for extra_pos_embedder in self.extra_pos_embedders_options.values():
|
| 602 |
+
fully_shard(extra_pos_embedder, mesh=mesh, reshard_after_forward=True, **fsdp_kwargs)
|
| 603 |
+
fully_shard(self.t_embedder, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 604 |
+
if self.extra_image_context_dim is not None:
|
| 605 |
+
fully_shard(self.img_context_proj, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 606 |
+
|
| 607 |
+
if hasattr(self, "view_embeddings"):
|
| 608 |
+
fully_shard(self.view_embeddings, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 609 |
+
|
| 610 |
+
if hasattr(self, "adaln_view_embedder"):
|
| 611 |
+
fully_shard(self.adaln_view_embedder, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 612 |
+
if hasattr(self, "adaln_view_proj"):
|
| 613 |
+
fully_shard(self.adaln_view_proj, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 614 |
+
|
| 615 |
+
def enable_context_parallel(self, process_group: Optional[ProcessGroup] = None):
|
| 616 |
+
# pos_embedder
|
| 617 |
+
for pos_embedder in self.pos_embedder_options.values():
|
| 618 |
+
pos_embedder.enable_context_parallel(process_group=process_group)
|
| 619 |
+
if self.extra_per_block_abs_pos_emb:
|
| 620 |
+
for extra_pos_embedder in self.extra_pos_embedders_options.values():
|
| 621 |
+
extra_pos_embedder.enable_context_parallel(process_group=process_group)
|
| 622 |
+
|
| 623 |
+
# attention
|
| 624 |
+
cp_ranks = get_process_group_ranks(process_group)
|
| 625 |
+
for block in self.blocks:
|
| 626 |
+
block.set_context_parallel_group(
|
| 627 |
+
process_group=process_group,
|
| 628 |
+
ranks=cp_ranks,
|
| 629 |
+
stream=torch.cuda.Stream(),
|
| 630 |
+
)
|
| 631 |
+
|
| 632 |
+
self._is_context_parallel_enabled = True
|
| 633 |
+
|
| 634 |
+
def disable_context_parallel(self):
|
| 635 |
+
# pos_embedder
|
| 636 |
+
for pos_embedder in self.pos_embedder_options.values():
|
| 637 |
+
pos_embedder.disable_context_parallel()
|
| 638 |
+
if self.extra_per_block_abs_pos_emb:
|
| 639 |
+
for extra_pos_embedder in self.extra_pos_embedders_options.values():
|
| 640 |
+
extra_pos_embedder.disable_context_parallel()
|
| 641 |
+
|
| 642 |
+
# attention
|
| 643 |
+
for block in self.blocks:
|
| 644 |
+
block.set_context_parallel_group(
|
| 645 |
+
process_group=None,
|
| 646 |
+
ranks=None,
|
| 647 |
+
stream=torch.cuda.Stream(),
|
| 648 |
+
)
|
| 649 |
+
|
| 650 |
+
self._is_context_parallel_enabled = False
|
| 651 |
+
|
| 652 |
+
def init_weights(self):
|
| 653 |
+
self.x_embedder.init_weights()
|
| 654 |
+
for pos_embedder in self.pos_embedder_options.values():
|
| 655 |
+
pos_embedder.reset_parameters()
|
| 656 |
+
if self.extra_per_block_abs_pos_emb:
|
| 657 |
+
for extra_pos_embedder in self.extra_pos_embedders_options.values():
|
| 658 |
+
extra_pos_embedder.init_weights()
|
| 659 |
+
|
| 660 |
+
self.t_embedder[1].init_weights()
|
| 661 |
+
for block in self.blocks:
|
| 662 |
+
block.init_weights()
|
| 663 |
+
|
| 664 |
+
self.final_layer.init_weights()
|
| 665 |
+
self.t_embedding_norm.reset_parameters()
|
| 666 |
+
|
| 667 |
+
if self.extra_image_context_dim is not None:
|
| 668 |
+
self.img_context_proj[0].reset_parameters()
|
| 669 |
+
|
| 670 |
+
if hasattr(self, "view_embeddings"):
|
| 671 |
+
torch.nn.init.normal_(self.view_embeddings.weight, mean=0.0, std=0.02)
|
| 672 |
+
|
| 673 |
+
if hasattr(self, "adaln_view_embedder"):
|
| 674 |
+
torch.nn.init.normal_(self.adaln_view_embedder.weight, mean=0.0, std=0.05)
|
| 675 |
+
|
| 676 |
+
if hasattr(self, "adaln_view_proj"):
|
| 677 |
+
torch.nn.init.zeros_(self.adaln_view_proj.weight)
|
| 678 |
+
torch.nn.init.zeros_(self.adaln_view_proj.bias)
|
| 679 |
+
|
| 680 |
+
def build_pos_embed(self):
|
| 681 |
+
self.pos_embedder_options = nn.ModuleDict()
|
| 682 |
+
self.extra_pos_embedders_options = nn.ModuleDict()
|
| 683 |
+
for n_cameras in range(1, self.n_cameras_emb + 1):
|
| 684 |
+
pos_embedder, extra_pos_embedder = self.build_pos_embed_for_n_cameras(n_cameras)
|
| 685 |
+
self.pos_embedder_options[f"n_cameras_{n_cameras}"] = pos_embedder
|
| 686 |
+
self.extra_pos_embedders_options[f"n_cameras_{n_cameras}"] = extra_pos_embedder
|
| 687 |
+
|
| 688 |
+
def build_pos_embed_for_n_cameras(self, n_cameras: int):
|
| 689 |
+
if self.pos_emb_cls == "rope3d":
|
| 690 |
+
cls_type = MultiCameraVideoRopePosition3DEmb
|
| 691 |
+
else:
|
| 692 |
+
raise ValueError(f"Unknown pos_emb_cls {self.pos_emb_cls}")
|
| 693 |
+
pos_embedder, extra_pos_embedder = None, None
|
| 694 |
+
kwargs = dict(
|
| 695 |
+
model_channels=self.model_channels,
|
| 696 |
+
len_h=self.max_img_h // self.patch_spatial,
|
| 697 |
+
len_w=self.max_img_w // self.patch_spatial,
|
| 698 |
+
len_t=self.max_frames // self.patch_temporal,
|
| 699 |
+
max_fps=self.max_fps,
|
| 700 |
+
min_fps=self.min_fps,
|
| 701 |
+
is_learnable=self.pos_emb_learnable,
|
| 702 |
+
interpolation=self.pos_emb_interpolation,
|
| 703 |
+
head_dim=self.model_channels // self.num_heads,
|
| 704 |
+
h_extrapolation_ratio=self.rope_h_extrapolation_ratio,
|
| 705 |
+
w_extrapolation_ratio=self.rope_w_extrapolation_ratio,
|
| 706 |
+
t_extrapolation_ratio=self.rope_t_extrapolation_ratio,
|
| 707 |
+
enable_fps_modulation=self.rope_enable_fps_modulation,
|
| 708 |
+
n_cameras=n_cameras,
|
| 709 |
+
)
|
| 710 |
+
pos_embedder = cls_type(
|
| 711 |
+
**kwargs,
|
| 712 |
+
)
|
| 713 |
+
assert pos_embedder.enable_fps_modulation == self.rope_enable_fps_modulation, (
|
| 714 |
+
"enable_fps_modulation must be the same"
|
| 715 |
+
)
|
| 716 |
+
|
| 717 |
+
if self.extra_per_block_abs_pos_emb:
|
| 718 |
+
raise NotImplementedError("extra_per_block_abs_pos_emb is not tested for multi-view DIT")
|
| 719 |
+
kwargs["h_extrapolation_ratio"] = self.extra_h_extrapolation_ratio
|
| 720 |
+
kwargs["w_extrapolation_ratio"] = self.extra_w_extrapolation_ratio
|
| 721 |
+
kwargs["t_extrapolation_ratio"] = self.extra_t_extrapolation_ratio
|
| 722 |
+
extra_pos_embedder = MultiCameraSinCosPosEmbAxis(
|
| 723 |
+
**kwargs,
|
| 724 |
+
)
|
| 725 |
+
return pos_embedder, extra_pos_embedder
|
| 726 |
+
|
| 727 |
+
def prepare_embedded_sequence(
|
| 728 |
+
self,
|
| 729 |
+
x_B_C_T_H_W: torch.Tensor,
|
| 730 |
+
fps: Optional[torch.Tensor] = None,
|
| 731 |
+
padding_mask: Optional[torch.Tensor] = None,
|
| 732 |
+
view_indices_B_T: Optional[torch.Tensor] = None,
|
| 733 |
+
) -> Tuple[torch.Tensor, Optional[torch.Tensor], torch.Tensor]:
|
| 734 |
+
if self.concat_padding_mask:
|
| 735 |
+
padding_mask = transforms.functional.resize(
|
| 736 |
+
padding_mask, list(x_B_C_T_H_W.shape[-2:]), interpolation=transforms.InterpolationMode.NEAREST
|
| 737 |
+
)
|
| 738 |
+
x_B_C_T_H_W = torch.cat(
|
| 739 |
+
[x_B_C_T_H_W, padding_mask.unsqueeze(1).repeat(1, 1, x_B_C_T_H_W.shape[2], 1, 1)], dim=1
|
| 740 |
+
)
|
| 741 |
+
cp_size = (
|
| 742 |
+
len(get_process_group_ranks(parallel_state.get_context_parallel_group()))
|
| 743 |
+
if parallel_state.is_initialized()
|
| 744 |
+
else 1
|
| 745 |
+
)
|
| 746 |
+
n_cameras = (x_B_C_T_H_W.shape[2] * cp_size) // self.state_t
|
| 747 |
+
pos_embedder = self.pos_embedder_options[f"n_cameras_{n_cameras}"] # they are all the same if they are rope
|
| 748 |
+
if self.concat_view_embedding:
|
| 749 |
+
if view_indices_B_T is None:
|
| 750 |
+
view_indices = torch.arange(n_cameras).clamp(
|
| 751 |
+
max=self.n_cameras_emb - 1
|
| 752 |
+
) # View indices [0, 1, ..., V-1]
|
| 753 |
+
view_indices = view_indices.to(x_B_C_T_H_W.device)
|
| 754 |
+
view_embedding = self.view_embeddings(view_indices) # Shape: [V, embedding_dim]
|
| 755 |
+
view_embedding = rearrange(view_embedding, "V D -> D V")
|
| 756 |
+
view_embedding = (
|
| 757 |
+
view_embedding.unsqueeze(0).unsqueeze(3).unsqueeze(4).unsqueeze(5)
|
| 758 |
+
) # Shape: [1, D, V, 1, 1, 1]
|
| 759 |
+
else:
|
| 760 |
+
view_indices_B_T = view_indices_B_T.clamp(max=self.n_cameras_emb - 1)
|
| 761 |
+
view_indices_B_T = view_indices_B_T.to(x_B_C_T_H_W.device).long()
|
| 762 |
+
view_embedding = self.view_embeddings(view_indices_B_T) # B, (V T), D
|
| 763 |
+
view_embedding = rearrange(view_embedding, "B (V T) D -> B D V T", V=n_cameras)
|
| 764 |
+
view_embedding = view_embedding.unsqueeze(-1).unsqueeze(-1) # Shape: [B, D, V, T, 1, 1]
|
| 765 |
+
x_B_C_V_T_H_W = rearrange(x_B_C_T_H_W, "B C (V T) H W -> B C V T H W", V=n_cameras)
|
| 766 |
+
view_embedding = view_embedding.expand(
|
| 767 |
+
x_B_C_V_T_H_W.shape[0],
|
| 768 |
+
view_embedding.shape[1],
|
| 769 |
+
view_embedding.shape[2],
|
| 770 |
+
x_B_C_V_T_H_W.shape[3],
|
| 771 |
+
x_B_C_V_T_H_W.shape[4],
|
| 772 |
+
x_B_C_V_T_H_W.shape[5],
|
| 773 |
+
)
|
| 774 |
+
x_B_C_V_T_H_W = torch.cat([x_B_C_V_T_H_W, view_embedding], dim=1)
|
| 775 |
+
x_B_C_T_H_W = rearrange(x_B_C_V_T_H_W, " B C V T H W -> B C (V T) H W", V=n_cameras)
|
| 776 |
+
|
| 777 |
+
x_B_T_H_W_D = self.x_embedder(x_B_C_T_H_W)
|
| 778 |
+
|
| 779 |
+
if self.extra_per_block_abs_pos_emb:
|
| 780 |
+
extra_pos_embedder = self.extra_pos_embedders_options[str(n_cameras)]
|
| 781 |
+
extra_pos_emb = extra_pos_embedder(x_B_T_H_W_D, fps=fps)
|
| 782 |
+
else:
|
| 783 |
+
extra_pos_emb = None
|
| 784 |
+
|
| 785 |
+
if "rope" in self.pos_emb_cls.lower():
|
| 786 |
+
return x_B_T_H_W_D, pos_embedder(x_B_T_H_W_D, fps=fps), extra_pos_emb
|
| 787 |
+
|
| 788 |
+
if "fps_aware" in self.pos_emb_cls:
|
| 789 |
+
raise NotImplementedError("FPS-aware positional embedding is not supported for multi-view DIT")
|
| 790 |
+
|
| 791 |
+
x_B_T_H_W_D = x_B_T_H_W_D + pos_embedder(x_B_T_H_W_D)
|
| 792 |
+
|
| 793 |
+
return x_B_T_H_W_D, None, extra_pos_emb
|
| 794 |
+
|
| 795 |
+
def forward(
|
| 796 |
+
self,
|
| 797 |
+
x_B_C_T_H_W: torch.Tensor,
|
| 798 |
+
timesteps_B_T: torch.Tensor,
|
| 799 |
+
crossattn_emb: torch.Tensor,
|
| 800 |
+
condition_video_input_mask_B_C_T_H_W: Optional[torch.Tensor] = None,
|
| 801 |
+
fps: Optional[torch.Tensor] = None,
|
| 802 |
+
padding_mask: Optional[torch.Tensor] = None,
|
| 803 |
+
data_type: Optional[DataType] = DataType.VIDEO,
|
| 804 |
+
view_indices_B_T: Optional[torch.Tensor] = None,
|
| 805 |
+
**kwargs,
|
| 806 |
+
) -> torch.Tensor | List[torch.Tensor] | Tuple[torch.Tensor, List[torch.Tensor]]:
|
| 807 |
+
# Deletes elements like condition.use_video_condition that are not used in the forward pass
|
| 808 |
+
del kwargs
|
| 809 |
+
if data_type == DataType.VIDEO:
|
| 810 |
+
x_B_C_T_H_W = torch.cat([x_B_C_T_H_W, condition_video_input_mask_B_C_T_H_W.type_as(x_B_C_T_H_W)], dim=1)
|
| 811 |
+
else:
|
| 812 |
+
B, _, T, H, W = x_B_C_T_H_W.shape
|
| 813 |
+
x_B_C_T_H_W = torch.cat(
|
| 814 |
+
[x_B_C_T_H_W, torch.zeros((B, 1, T, H, W), dtype=x_B_C_T_H_W.dtype, device=x_B_C_T_H_W.device)], dim=1
|
| 815 |
+
)
|
| 816 |
+
|
| 817 |
+
assert isinstance(data_type, DataType), (
|
| 818 |
+
f"Expected DataType, got {type(data_type)}. We need discuss this flag later."
|
| 819 |
+
)
|
| 820 |
+
timesteps_B_T = timesteps_B_T * self.timestep_scale
|
| 821 |
+
x_B_T_H_W_D, rope_emb_L_1_1_D, extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D = self.prepare_embedded_sequence(
|
| 822 |
+
x_B_C_T_H_W,
|
| 823 |
+
fps=fps,
|
| 824 |
+
padding_mask=padding_mask,
|
| 825 |
+
view_indices_B_T=view_indices_B_T,
|
| 826 |
+
)
|
| 827 |
+
if self.use_crossattn_projection:
|
| 828 |
+
crossattn_emb = self.crossattn_proj(crossattn_emb)
|
| 829 |
+
with amp.autocast("cuda", enabled=self.use_wan_fp32_strategy, dtype=torch.float32):
|
| 830 |
+
# (B, 1). input timesteps are (b, 1)
|
| 831 |
+
if timesteps_B_T.ndim == 1:
|
| 832 |
+
timesteps_B_T = timesteps_B_T.unsqueeze(1)
|
| 833 |
+
t_embedding_B_T_D, adaln_lora_B_T_3D = self.t_embedder(timesteps_B_T)
|
| 834 |
+
t_embedding_B_T_D = self.t_embedding_norm(t_embedding_B_T_D)
|
| 835 |
+
|
| 836 |
+
if self.adaln_view_embedding:
|
| 837 |
+
num_cameras = torch.unique(view_indices_B_T[0]).shape[0]
|
| 838 |
+
with amp.autocast("cuda", enabled=self.use_wan_fp32_strategy, dtype=torch.float32):
|
| 839 |
+
view_indices_B_V_T = rearrange(view_indices_B_T, "b (v t) -> b v t", v=num_cameras)
|
| 840 |
+
view_embedding_B_V = self.adaln_view_embedder(view_indices_B_V_T[..., 0]) # B, V, D
|
| 841 |
+
view_embedding_proj_B_V_9D = self.adaln_view_proj(view_embedding_B_V) # B, V, 9D
|
| 842 |
+
else:
|
| 843 |
+
view_embedding_proj_B_V_9D = None
|
| 844 |
+
|
| 845 |
+
# for logging purpose
|
| 846 |
+
affline_scale_log_info = {}
|
| 847 |
+
affline_scale_log_info["t_embedding_B_T_D"] = t_embedding_B_T_D.detach()
|
| 848 |
+
self.affline_scale_log_info = affline_scale_log_info
|
| 849 |
+
self.affline_emb = t_embedding_B_T_D
|
| 850 |
+
self.crossattn_emb = crossattn_emb
|
| 851 |
+
|
| 852 |
+
if extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D is not None:
|
| 853 |
+
assert x_B_T_H_W_D.shape == extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape, (
|
| 854 |
+
f"{x_B_T_H_W_D.shape} != {extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape}"
|
| 855 |
+
)
|
| 856 |
+
|
| 857 |
+
B, T, H, W, D = x_B_T_H_W_D.shape
|
| 858 |
+
|
| 859 |
+
for block_idx, block in enumerate(self.blocks):
|
| 860 |
+
x_B_T_H_W_D = block(
|
| 861 |
+
x_B_T_H_W_D,
|
| 862 |
+
view_indices_B_T,
|
| 863 |
+
t_embedding_B_T_D,
|
| 864 |
+
view_embedding_proj_B_V_9D,
|
| 865 |
+
crossattn_emb,
|
| 866 |
+
rope_emb_L_1_1_D=rope_emb_L_1_1_D,
|
| 867 |
+
adaln_lora_B_T_3D=adaln_lora_B_T_3D,
|
| 868 |
+
extra_per_block_pos_emb=extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D,
|
| 869 |
+
block_idx=block_idx,
|
| 870 |
+
)
|
| 871 |
+
|
| 872 |
+
x_B_T_H_W_O = self.final_layer(x_B_T_H_W_D, t_embedding_B_T_D, adaln_lora_B_T_3D=adaln_lora_B_T_3D)
|
| 873 |
+
x_B_C_Tt_Hp_Wp = self.unpatchify(x_B_T_H_W_O)
|
| 874 |
+
|
| 875 |
+
return x_B_C_Tt_Hp_Wp
|
| 876 |
+
|
| 877 |
+
def init_cross_view_attn_with_self_attn_weights(self, is_ema: bool = False) -> None:
|
| 878 |
+
"""Load self-attention weights from base model checkpoint and initialize cross-view attention."""
|
| 879 |
+
# Check initialization conditions
|
| 880 |
+
if self.init_cross_view_attn_weight_from is None:
|
| 881 |
+
log.info("No checkpoint path provided, skipping cross-view attention initialization")
|
| 882 |
+
return
|
| 883 |
+
|
| 884 |
+
if not self.enable_cross_view_attn:
|
| 885 |
+
log.info("Cross-view attention not enabled, skipping weight loading")
|
| 886 |
+
return
|
| 887 |
+
|
| 888 |
+
log.critical(
|
| 889 |
+
f"Loading base model from {self.init_cross_view_attn_weight_from} for cross-view attention initialization"
|
| 890 |
+
)
|
| 891 |
+
|
| 892 |
+
# Import necessary modules
|
| 893 |
+
import gc
|
| 894 |
+
|
| 895 |
+
import torch.distributed.checkpoint as dcp
|
| 896 |
+
from torch.distributed.checkpoint.default_planner import DefaultLoadPlanner
|
| 897 |
+
|
| 898 |
+
from cosmos_predict2._src.imaginaire.checkpointer.s3_filesystem import S3StorageReader
|
| 899 |
+
|
| 900 |
+
# Prepare checkpoint loading
|
| 901 |
+
checkpoint_path = os.path.join("s3://bucket/" + self.init_cross_view_attn_weight_from, "model")
|
| 902 |
+
storage_reader = S3StorageReader(
|
| 903 |
+
credential_path=self.init_cross_view_attn_weight_credentials,
|
| 904 |
+
path=checkpoint_path,
|
| 905 |
+
)
|
| 906 |
+
|
| 907 |
+
if torch.distributed.is_initialized():
|
| 908 |
+
torch.distributed.barrier()
|
| 909 |
+
|
| 910 |
+
# Build minimal state dict - only load required weights
|
| 911 |
+
log.info("Building minimal state dict (only weights needed for cross-view attention)")
|
| 912 |
+
minimal_state_dict = self._build_minimal_state_dict(is_ema)
|
| 913 |
+
|
| 914 |
+
# Load weights from checkpoint
|
| 915 |
+
log.info(f"Loading {len(minimal_state_dict)} weight tensors from checkpoint")
|
| 916 |
+
dcp.load(minimal_state_dict, storage_reader=storage_reader, planner=DefaultLoadPlanner(allow_partial_load=True))
|
| 917 |
+
|
| 918 |
+
# Verify that loaded weights maintain correct sharding (if FSDP is used)
|
| 919 |
+
if torch.distributed.is_initialized():
|
| 920 |
+
self._verify_loaded_checkpoint_sharding(minimal_state_dict, is_ema)
|
| 921 |
+
|
| 922 |
+
# Initialize cross-view attention weights
|
| 923 |
+
log.info("Starting to initialize cross-view attention from self-attention weights")
|
| 924 |
+
initialized_count, weight_stats = self._copy_weights_to_cross_view_attn(minimal_state_dict, is_ema)
|
| 925 |
+
|
| 926 |
+
log.info(f"Successfully initialized cross-view attention for {initialized_count}/{len(self.blocks)} layers")
|
| 927 |
+
|
| 928 |
+
# Release memory
|
| 929 |
+
log.info("Releasing checkpoint memory")
|
| 930 |
+
del minimal_state_dict
|
| 931 |
+
if torch.cuda.is_available():
|
| 932 |
+
torch.cuda.empty_cache()
|
| 933 |
+
gc.collect()
|
| 934 |
+
|
| 935 |
+
# Verify weight loading
|
| 936 |
+
self.weight_stats_before, self.weight_stats_after = weight_stats
|
| 937 |
+
self._verify_cross_view_attn_weights_loaded()
|
| 938 |
+
|
| 939 |
+
def _build_minimal_state_dict(self, is_ema: bool) -> dict[str, torch.Tensor]:
|
| 940 |
+
"""Build minimal state dict containing only required weights.
|
| 941 |
+
|
| 942 |
+
Note: When using FSDP, torch.empty_like preserves DTensor sharding metadata,
|
| 943 |
+
allowing dcp.load to correctly load only the local shard for each rank.
|
| 944 |
+
"""
|
| 945 |
+
minimal_state_dict = {}
|
| 946 |
+
num_layers = len(self.blocks)
|
| 947 |
+
|
| 948 |
+
# Determine parameter names to load
|
| 949 |
+
param_names = ["q_proj.weight", "k_proj.weight", "v_proj.weight"]
|
| 950 |
+
if hasattr(self.blocks[0], "cross_view_attn"):
|
| 951 |
+
cross_view_attn = self.blocks[0].cross_view_attn
|
| 952 |
+
if hasattr(cross_view_attn, "q_norm") and hasattr(cross_view_attn.q_norm, "weight"):
|
| 953 |
+
param_names.extend(["q_norm.weight", "k_norm.weight"])
|
| 954 |
+
|
| 955 |
+
# Create placeholder for each parameter in each layer
|
| 956 |
+
prefix = f"net{'_ema' if is_ema else ''}"
|
| 957 |
+
for layer_idx in range(num_layers):
|
| 958 |
+
for param_name in param_names:
|
| 959 |
+
key = f"{prefix}.blocks.{layer_idx}.self_attn.{param_name}"
|
| 960 |
+
|
| 961 |
+
# Get target parameter to determine shape and dtype
|
| 962 |
+
# IMPORTANT: This preserves DTensor sharding metadata if model is FSDP-wrapped
|
| 963 |
+
target_param = self._get_nested_attr(self.blocks[layer_idx].cross_view_attn, param_name.split("."))
|
| 964 |
+
|
| 965 |
+
if target_param is not None:
|
| 966 |
+
# empty_like preserves DTensor sharding spec, which tells dcp.load
|
| 967 |
+
# which shard to load for the current rank
|
| 968 |
+
placeholder = torch.empty_like(target_param)
|
| 969 |
+
minimal_state_dict[key] = placeholder
|
| 970 |
+
|
| 971 |
+
# Log sharding info for first layer to verify FSDP setup
|
| 972 |
+
if layer_idx == 0:
|
| 973 |
+
from torch.distributed._tensor.api import DTensor
|
| 974 |
+
|
| 975 |
+
if isinstance(placeholder, DTensor):
|
| 976 |
+
log.info(f"Parameter {param_name} is DTensor with placement: {placeholder.placements}")
|
| 977 |
+
else:
|
| 978 |
+
log.info(f"Parameter {param_name} is regular tensor (not sharded)")
|
| 979 |
+
else:
|
| 980 |
+
log.warning(f"Parameter {param_name} does not exist in layer {layer_idx} cross_view_attn")
|
| 981 |
+
|
| 982 |
+
return minimal_state_dict
|
| 983 |
+
|
| 984 |
+
def _get_nested_attr(self, obj, attr_path: list[str]):
|
| 985 |
+
"""Recursively get nested attribute."""
|
| 986 |
+
try:
|
| 987 |
+
for attr in attr_path:
|
| 988 |
+
obj = getattr(obj, attr)
|
| 989 |
+
return obj
|
| 990 |
+
except AttributeError:
|
| 991 |
+
return None
|
| 992 |
+
|
| 993 |
+
def _verify_loaded_checkpoint_sharding(self, state_dict: dict[str, torch.Tensor], is_ema: bool) -> None:
|
| 994 |
+
"""Verify that checkpoint was loaded with correct FSDP sharding.
|
| 995 |
+
|
| 996 |
+
This checks that:
|
| 997 |
+
1. If local model uses DTensor, loaded weights are also DTensor with matching sharding
|
| 998 |
+
2. The loaded shard size matches what we expect for the current rank
|
| 999 |
+
"""
|
| 1000 |
+
from torch.distributed._tensor.api import DTensor
|
| 1001 |
+
|
| 1002 |
+
rank = torch.distributed.get_rank()
|
| 1003 |
+
world_size = torch.distributed.get_world_size()
|
| 1004 |
+
|
| 1005 |
+
# Check a sample parameter from the first layer
|
| 1006 |
+
prefix = f"net{'_ema' if is_ema else ''}"
|
| 1007 |
+
sample_key = f"{prefix}.blocks.0.self_attn.q_proj.weight"
|
| 1008 |
+
|
| 1009 |
+
if sample_key in state_dict:
|
| 1010 |
+
loaded_param = state_dict[sample_key]
|
| 1011 |
+
target_param = self.blocks[0].cross_view_attn.q_proj.weight
|
| 1012 |
+
|
| 1013 |
+
loaded_is_dtensor = isinstance(loaded_param, DTensor)
|
| 1014 |
+
target_is_dtensor = isinstance(target_param, DTensor)
|
| 1015 |
+
|
| 1016 |
+
if target_is_dtensor and not loaded_is_dtensor:
|
| 1017 |
+
log.warning(
|
| 1018 |
+
f"[Rank {rank}] Mismatch: Local model uses DTensor (FSDP sharded), "
|
| 1019 |
+
f"but checkpoint loaded as regular tensor. This may cause OOM or incorrect behavior."
|
| 1020 |
+
)
|
| 1021 |
+
elif not target_is_dtensor and loaded_is_dtensor:
|
| 1022 |
+
log.warning(
|
| 1023 |
+
f"[Rank {rank}] Mismatch: Local model uses regular tensor, "
|
| 1024 |
+
f"but checkpoint loaded as DTensor. This is unusual."
|
| 1025 |
+
)
|
| 1026 |
+
elif target_is_dtensor and loaded_is_dtensor:
|
| 1027 |
+
# Both are DTensor - verify sharding matches
|
| 1028 |
+
if loaded_param.placements != target_param.placements:
|
| 1029 |
+
log.warning(
|
| 1030 |
+
f"[Rank {rank}] DTensor placement mismatch:\n"
|
| 1031 |
+
f" Loaded: {loaded_param.placements}\n"
|
| 1032 |
+
f" Target: {target_param.placements}\n"
|
| 1033 |
+
f"This may cause incorrect weight copying."
|
| 1034 |
+
)
|
| 1035 |
+
else:
|
| 1036 |
+
log.info(
|
| 1037 |
+
f"[Rank {rank}] ✅ Checkpoint sharding verified: "
|
| 1038 |
+
f"DTensor with placements {loaded_param.placements}"
|
| 1039 |
+
)
|
| 1040 |
+
|
| 1041 |
+
# Check local shard size
|
| 1042 |
+
loaded_local = loaded_param.to_local()
|
| 1043 |
+
target_local = target_param.to_local()
|
| 1044 |
+
log.info(
|
| 1045 |
+
f"[Rank {rank}] Local shard shape - Loaded: {loaded_local.shape}, Target: {target_local.shape}"
|
| 1046 |
+
)
|
| 1047 |
+
else:
|
| 1048 |
+
# Both are regular tensors
|
| 1049 |
+
log.info(f"[Rank {rank}] Both checkpoint and model use regular tensors (no FSDP sharding)")
|
| 1050 |
+
|
| 1051 |
+
def _copy_weights_to_cross_view_attn(
|
| 1052 |
+
self, state_dict: dict[str, torch.Tensor], is_ema: bool
|
| 1053 |
+
) -> tuple[int, tuple[dict, dict]]:
|
| 1054 |
+
"""Copy loaded weights to cross-view attention layers and record statistics."""
|
| 1055 |
+
from torch.distributed._tensor.api import DTensor
|
| 1056 |
+
|
| 1057 |
+
num_layers = len(self.blocks)
|
| 1058 |
+
initialized_count = 0
|
| 1059 |
+
weight_stats_before = {}
|
| 1060 |
+
weight_stats_after = {}
|
| 1061 |
+
prefix = f"net{'_ema' if is_ema else ''}"
|
| 1062 |
+
|
| 1063 |
+
for layer_idx in range(num_layers):
|
| 1064 |
+
block = self.blocks[layer_idx]
|
| 1065 |
+
cross_view_attn = block.cross_view_attn
|
| 1066 |
+
|
| 1067 |
+
# Determine parameters to copy for current layer
|
| 1068 |
+
param_names = ["q_proj.weight", "k_proj.weight", "v_proj.weight"]
|
| 1069 |
+
if hasattr(cross_view_attn, "q_norm") and hasattr(cross_view_attn.q_norm, "weight"):
|
| 1070 |
+
param_names.extend(["q_norm.weight", "k_norm.weight"])
|
| 1071 |
+
|
| 1072 |
+
copied_params = []
|
| 1073 |
+
should_record_stats = layer_idx % 5 == 0 # Sample every 5 layers to record statistics
|
| 1074 |
+
|
| 1075 |
+
for param_name in param_names:
|
| 1076 |
+
key = f"{prefix}.blocks.{layer_idx}.self_attn.{param_name}"
|
| 1077 |
+
|
| 1078 |
+
if key not in state_dict:
|
| 1079 |
+
log.warning(f"Key {key} not found in checkpoint, skipping")
|
| 1080 |
+
continue
|
| 1081 |
+
|
| 1082 |
+
# Get target parameter
|
| 1083 |
+
target_param = self._get_nested_attr(cross_view_attn, param_name.split("."))
|
| 1084 |
+
|
| 1085 |
+
if not isinstance(target_param, torch.nn.Parameter):
|
| 1086 |
+
log.warning(f"Parameter {param_name} in block {layer_idx} is not a Parameter object")
|
| 1087 |
+
continue
|
| 1088 |
+
|
| 1089 |
+
# Record statistics before copy (sampled)
|
| 1090 |
+
if should_record_stats and param_name == "q_proj.weight":
|
| 1091 |
+
before_local = target_param.to_local() if isinstance(target_param, DTensor) else target_param
|
| 1092 |
+
weight_stats_before[layer_idx] = {
|
| 1093 |
+
"mean": before_local.mean().item(),
|
| 1094 |
+
"std": before_local.std().item(),
|
| 1095 |
+
"abs_max": before_local.abs().max().item(),
|
| 1096 |
+
"dtype": str(before_local.dtype),
|
| 1097 |
+
}
|
| 1098 |
+
|
| 1099 |
+
# Copy weights (ensure dtype consistency)
|
| 1100 |
+
source_weight = state_dict[key]
|
| 1101 |
+
if source_weight.dtype != target_param.dtype:
|
| 1102 |
+
log.warning(
|
| 1103 |
+
f"Layer {layer_idx} {param_name}: dtype mismatch ({source_weight.dtype} -> {target_param.dtype}), converting"
|
| 1104 |
+
)
|
| 1105 |
+
source_weight = source_weight.to(dtype=target_param.dtype)
|
| 1106 |
+
|
| 1107 |
+
with torch.no_grad():
|
| 1108 |
+
target_param.copy_(source_weight)
|
| 1109 |
+
|
| 1110 |
+
# Record statistics after copy (sampled)
|
| 1111 |
+
if should_record_stats and param_name == "q_proj.weight":
|
| 1112 |
+
after_local = target_param.to_local() if isinstance(target_param, DTensor) else target_param
|
| 1113 |
+
weight_stats_after[layer_idx] = {
|
| 1114 |
+
"mean": after_local.mean().item(),
|
| 1115 |
+
"std": after_local.std().item(),
|
| 1116 |
+
"abs_max": after_local.abs().max().item(),
|
| 1117 |
+
"dtype": str(after_local.dtype),
|
| 1118 |
+
}
|
| 1119 |
+
|
| 1120 |
+
copied_params.append(param_name)
|
| 1121 |
+
|
| 1122 |
+
if copied_params:
|
| 1123 |
+
initialized_count += 1
|
| 1124 |
+
# Print detailed info only every 10 layers to reduce log noise
|
| 1125 |
+
if layer_idx % 10 == 0:
|
| 1126 |
+
log.info(f"Initialized cross_view_attn for block {layer_idx}, copied parameters: {copied_params}")
|
| 1127 |
+
|
| 1128 |
+
return initialized_count, (weight_stats_before, weight_stats_after)
|
| 1129 |
+
|
| 1130 |
+
def _verify_cross_view_attn_weights_loaded(self) -> None:
|
| 1131 |
+
"""Verify cross-view attention weights are correctly loaded by comparing statistics before and after copy."""
|
| 1132 |
+
rank = torch.distributed.get_rank() if torch.distributed.is_initialized() else 0
|
| 1133 |
+
|
| 1134 |
+
num_updated = 0
|
| 1135 |
+
num_failed = 0
|
| 1136 |
+
dtype_mismatches = []
|
| 1137 |
+
|
| 1138 |
+
# Check weight changes and dtype consistency for sampled layers
|
| 1139 |
+
for layer_idx in sorted(self.weight_stats_before.keys()):
|
| 1140 |
+
if layer_idx not in self.weight_stats_after:
|
| 1141 |
+
num_failed += 1
|
| 1142 |
+
continue
|
| 1143 |
+
|
| 1144 |
+
before = self.weight_stats_before[layer_idx]
|
| 1145 |
+
after = self.weight_stats_after[layer_idx]
|
| 1146 |
+
|
| 1147 |
+
# Verify dtype consistency
|
| 1148 |
+
if before["dtype"] != after["dtype"]:
|
| 1149 |
+
dtype_mismatches.append(f"Layer {layer_idx}: {before['dtype']} -> {after['dtype']}")
|
| 1150 |
+
num_failed += 1
|
| 1151 |
+
continue
|
| 1152 |
+
|
| 1153 |
+
# Calculate weight changes (threshold: 1e-4)
|
| 1154 |
+
mean_change = abs(after["mean"] - before["mean"])
|
| 1155 |
+
std_change = abs(after["std"] - before["std"])
|
| 1156 |
+
max_change = abs(after["abs_max"] - before["abs_max"])
|
| 1157 |
+
|
| 1158 |
+
if mean_change > 1e-4 or std_change > 1e-4 or max_change > 1e-4:
|
| 1159 |
+
num_updated += 1
|
| 1160 |
+
else:
|
| 1161 |
+
num_failed += 1
|
| 1162 |
+
|
| 1163 |
+
# Output verification results
|
| 1164 |
+
total_sampled = len(self.weight_stats_before)
|
| 1165 |
+
if num_failed == 0:
|
| 1166 |
+
log.info(
|
| 1167 |
+
f"[Rank {rank}] ✅ Weight verification passed: {num_updated}/{total_sampled} sampled layers successfully updated",
|
| 1168 |
+
rank0_only=False,
|
| 1169 |
+
)
|
| 1170 |
+
if total_sampled > 0:
|
| 1171 |
+
first_layer_dtype = self.weight_stats_after[list(self.weight_stats_after.keys())[0]]["dtype"]
|
| 1172 |
+
log.info(
|
| 1173 |
+
f"[Rank {rank}] ✅ All weights maintain consistent dtype: {first_layer_dtype}", rank0_only=False
|
| 1174 |
+
)
|
| 1175 |
+
else:
|
| 1176 |
+
log.warning(
|
| 1177 |
+
f"[Rank {rank}] ⚠️ Weight verification: {num_updated}/{total_sampled} layers updated, {num_failed} layers may have failed",
|
| 1178 |
+
rank0_only=False,
|
| 1179 |
+
)
|
| 1180 |
+
if dtype_mismatches:
|
| 1181 |
+
log.error(
|
| 1182 |
+
f"[Rank {rank}] ❌ Detected dtype mismatches:\n" + "\n".join(dtype_mismatches), rank0_only=False
|
| 1183 |
+
)
|
| 1184 |
+
|
| 1185 |
+
# Clean up temporary attributes
|
| 1186 |
+
delattr(self, "weight_stats_before")
|
| 1187 |
+
delattr(self, "weight_stats_after")
|
cosmos_predict2/_src/predict2_multiview/networks/multiview_dit.py
ADDED
|
@@ -0,0 +1,618 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
from typing import List, Optional, Tuple
|
| 17 |
+
|
| 18 |
+
import numpy as np
|
| 19 |
+
import torch
|
| 20 |
+
import torch.amp as amp
|
| 21 |
+
import torch.nn as nn
|
| 22 |
+
from einops import rearrange, repeat
|
| 23 |
+
from megatron.core import parallel_state
|
| 24 |
+
from torch.distributed import ProcessGroup, get_process_group_ranks
|
| 25 |
+
from torch.distributed._composable.fsdp import fully_shard
|
| 26 |
+
from torchvision import transforms
|
| 27 |
+
|
| 28 |
+
from cosmos_predict2._src.imaginaire.utils.context_parallel import split_inputs_cp
|
| 29 |
+
from cosmos_predict2._src.predict2.conditioner import DataType
|
| 30 |
+
from cosmos_predict2._src.predict2.networks.minimal_v1_lvg_dit import MinimalV1LVGDiT
|
| 31 |
+
from cosmos_predict2._src.predict2.networks.minimal_v4_dit import (
|
| 32 |
+
Attention,
|
| 33 |
+
Block,
|
| 34 |
+
SACConfig,
|
| 35 |
+
VideoPositionEmb,
|
| 36 |
+
VideoRopePosition3DEmb,
|
| 37 |
+
)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
class MultiViewCrossAttention(Attention):
|
| 41 |
+
def __init__(self, *args, state_t: int = None, **kwargs) -> None:
|
| 42 |
+
super().__init__(*args, **kwargs)
|
| 43 |
+
assert self.qkv_format == "bshd", "MultiViewCrossAttention only supports qkv_format='bshd'"
|
| 44 |
+
self.state_t = state_t
|
| 45 |
+
|
| 46 |
+
def forward(self, x, context=None, rope_emb=None):
|
| 47 |
+
assert not self.is_selfattn, "MultiViewCrossAttention does not support self-attention"
|
| 48 |
+
B, L, D = x.shape
|
| 49 |
+
|
| 50 |
+
n_cameras = context.shape[1] // 512
|
| 51 |
+
x_B_L_D = rearrange(x, "B (V L) D -> (V B) L D", V=n_cameras)
|
| 52 |
+
context_B_M_D = rearrange(context, "B (V M) D -> (V B) M D", V=n_cameras) if context is not None else None
|
| 53 |
+
x_B_L_D = super().forward(x_B_L_D, context_B_M_D, rope_emb=rope_emb)
|
| 54 |
+
x_B_L_D = rearrange(x_B_L_D, "(V B) L D -> B (V L) D", V=n_cameras)
|
| 55 |
+
return x_B_L_D
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
class MultiViewBlock(Block):
|
| 59 |
+
"""
|
| 60 |
+
A transformer block that takes n_cameras as input. This block
|
| 61 |
+
"""
|
| 62 |
+
|
| 63 |
+
def __init__(
|
| 64 |
+
self,
|
| 65 |
+
x_dim: int,
|
| 66 |
+
context_dim: int,
|
| 67 |
+
num_heads: int,
|
| 68 |
+
mlp_ratio: float = 4.0,
|
| 69 |
+
use_adaln_lora: bool = False,
|
| 70 |
+
adaln_lora_dim: int = 256,
|
| 71 |
+
state_t: int = None,
|
| 72 |
+
backend: str = "transformer_engine",
|
| 73 |
+
image_context_dim: Optional[int] = None,
|
| 74 |
+
use_wan_fp32_strategy: bool = False,
|
| 75 |
+
):
|
| 76 |
+
super().__init__(
|
| 77 |
+
x_dim,
|
| 78 |
+
context_dim,
|
| 79 |
+
num_heads,
|
| 80 |
+
mlp_ratio,
|
| 81 |
+
use_adaln_lora,
|
| 82 |
+
adaln_lora_dim,
|
| 83 |
+
backend,
|
| 84 |
+
image_context_dim,
|
| 85 |
+
use_wan_fp32_strategy,
|
| 86 |
+
)
|
| 87 |
+
self.state_t = state_t
|
| 88 |
+
if image_context_dim is None:
|
| 89 |
+
del self.cross_attn
|
| 90 |
+
self.cross_attn = MultiViewCrossAttention(
|
| 91 |
+
x_dim,
|
| 92 |
+
context_dim,
|
| 93 |
+
num_heads,
|
| 94 |
+
x_dim // num_heads,
|
| 95 |
+
qkv_format="bshd",
|
| 96 |
+
state_t=state_t,
|
| 97 |
+
use_wan_fp32_strategy=use_wan_fp32_strategy,
|
| 98 |
+
)
|
| 99 |
+
else:
|
| 100 |
+
raise NotImplementedError("image_context_dim is not supported for MultiViewBlock")
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
class MultiCameraVideoRopePosition3DEmb(VideoRopePosition3DEmb):
|
| 104 |
+
def __init__(self, *args, n_cameras: int = 1, **kwargs):
|
| 105 |
+
super().__init__(*args, **kwargs)
|
| 106 |
+
self.n_cameras = n_cameras
|
| 107 |
+
|
| 108 |
+
def generate_embeddings(
|
| 109 |
+
self,
|
| 110 |
+
B_T_H_W_C: torch.Size,
|
| 111 |
+
fps: Optional[torch.Tensor] = None,
|
| 112 |
+
h_ntk_factor: Optional[float] = None,
|
| 113 |
+
w_ntk_factor: Optional[float] = None,
|
| 114 |
+
t_ntk_factor: Optional[float] = None,
|
| 115 |
+
):
|
| 116 |
+
B, T, H, W, C = B_T_H_W_C
|
| 117 |
+
single_camera_B_T_H_W_C = (B, T // self.n_cameras, H, W, C)
|
| 118 |
+
em_T_H_W_D = []
|
| 119 |
+
for _ in range(self.n_cameras):
|
| 120 |
+
em_L_1_1_D = super().generate_embeddings(
|
| 121 |
+
single_camera_B_T_H_W_C,
|
| 122 |
+
fps=fps,
|
| 123 |
+
h_ntk_factor=h_ntk_factor,
|
| 124 |
+
w_ntk_factor=w_ntk_factor,
|
| 125 |
+
t_ntk_factor=t_ntk_factor,
|
| 126 |
+
)
|
| 127 |
+
em_T_H_W_D.append(rearrange(em_L_1_1_D, "(t h w) 1 1 d -> t h w d", t=T // self.n_cameras, h=H, w=W))
|
| 128 |
+
em_T_H_W_D = torch.cat(em_T_H_W_D, dim=0)
|
| 129 |
+
return em_T_H_W_D.float()
|
| 130 |
+
|
| 131 |
+
def generate_embeddings_with_refs(
|
| 132 |
+
self,
|
| 133 |
+
B_Te_H_W_C: torch.Size,
|
| 134 |
+
num_ref: int,
|
| 135 |
+
ref_positions: List[int],
|
| 136 |
+
fps: Optional[torch.Tensor] = None,
|
| 137 |
+
h_ntk_factor: Optional[float] = None,
|
| 138 |
+
w_ntk_factor: Optional[float] = None,
|
| 139 |
+
t_ntk_factor: Optional[float] = None,
|
| 140 |
+
):
|
| 141 |
+
"""RoPE for a per-view grid of (state_t real + num_ref reference) frames.
|
| 142 |
+
|
| 143 |
+
Per view the temporal positions are ``[0 .. state_t-1]`` for the real frames followed by the fixed
|
| 144 |
+
``ref_positions`` for the appended reference frames; views are concatenated on the temporal axis
|
| 145 |
+
(view-major), matching the token layout ``[view0 real, view0 ref, view1 real, view1 ref, ...]``.
|
| 146 |
+
Returns ``(V*(state_t+num_ref), H, W, D)`` (un-split; caller applies ``_split_for_context_parallel``).
|
| 147 |
+
"""
|
| 148 |
+
B, Te, H, W, C = B_Te_H_W_C
|
| 149 |
+
frames_per_view = Te // self.n_cameras
|
| 150 |
+
state_t = frames_per_view - num_ref
|
| 151 |
+
assert state_t >= 0 and len(ref_positions) == num_ref, (state_t, num_ref, ref_positions)
|
| 152 |
+
temporal_positions = torch.tensor(
|
| 153 |
+
list(range(state_t)) + list(ref_positions), dtype=torch.float
|
| 154 |
+
) # (frames_per_view,)
|
| 155 |
+
single_camera_B_T_H_W_C = (B, frames_per_view, H, W, C)
|
| 156 |
+
em_T_H_W_D = []
|
| 157 |
+
for _ in range(self.n_cameras):
|
| 158 |
+
em_L_1_1_D = super().generate_embeddings(
|
| 159 |
+
single_camera_B_T_H_W_C,
|
| 160 |
+
fps=fps,
|
| 161 |
+
h_ntk_factor=h_ntk_factor,
|
| 162 |
+
w_ntk_factor=w_ntk_factor,
|
| 163 |
+
t_ntk_factor=t_ntk_factor,
|
| 164 |
+
temporal_positions=temporal_positions,
|
| 165 |
+
)
|
| 166 |
+
em_T_H_W_D.append(
|
| 167 |
+
rearrange(em_L_1_1_D, "(t h w) 1 1 d -> t h w d", t=frames_per_view, h=H, w=W)
|
| 168 |
+
)
|
| 169 |
+
em_T_H_W_D = torch.cat(em_T_H_W_D, dim=0)
|
| 170 |
+
return em_T_H_W_D.float()
|
| 171 |
+
|
| 172 |
+
@property
|
| 173 |
+
def seq_dim(self):
|
| 174 |
+
return 1
|
| 175 |
+
|
| 176 |
+
def _split_for_context_parallel(self, embeddings):
|
| 177 |
+
if self._cp_group is not None:
|
| 178 |
+
embeddings = rearrange(embeddings, "(V T) H W D -> V (T H W) 1 1 D", V=self.n_cameras)
|
| 179 |
+
embeddings = split_inputs_cp(x=embeddings, seq_dim=self.seq_dim, cp_group=self._cp_group)
|
| 180 |
+
embeddings = rearrange(embeddings, "V T 1 1 D -> (V T) 1 1 D", V=self.n_cameras)
|
| 181 |
+
else:
|
| 182 |
+
embeddings = rearrange(embeddings, "t h w d -> (t h w) 1 1 d")
|
| 183 |
+
return embeddings
|
| 184 |
+
|
| 185 |
+
|
| 186 |
+
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
|
| 187 |
+
"""
|
| 188 |
+
embed_dim: output dimension for each position
|
| 189 |
+
pos: a list of positions to be encoded: size (M,)
|
| 190 |
+
out: (M, D)
|
| 191 |
+
"""
|
| 192 |
+
assert embed_dim % 2 == 0
|
| 193 |
+
omega = np.arange(embed_dim // 2, dtype=np.float64)
|
| 194 |
+
omega /= embed_dim / 2.0
|
| 195 |
+
omega = 1.0 / 10000**omega # (D/2,)
|
| 196 |
+
|
| 197 |
+
pos = pos.reshape(-1) # (M,)
|
| 198 |
+
out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product
|
| 199 |
+
|
| 200 |
+
emb_sin = np.sin(out) # (M, D/2)
|
| 201 |
+
emb_cos = np.cos(out) # (M, D/2)
|
| 202 |
+
|
| 203 |
+
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
|
| 204 |
+
return emb
|
| 205 |
+
|
| 206 |
+
|
| 207 |
+
class MultiCameraSinCosPosEmbAxis(VideoPositionEmb):
|
| 208 |
+
def __init__(
|
| 209 |
+
self,
|
| 210 |
+
*, # enforce keyword arguments
|
| 211 |
+
interpolation: str,
|
| 212 |
+
model_channels: int,
|
| 213 |
+
len_h: int,
|
| 214 |
+
len_w: int,
|
| 215 |
+
len_t: int,
|
| 216 |
+
h_extrapolation_ratio: float = 1.0,
|
| 217 |
+
w_extrapolation_ratio: float = 1.0,
|
| 218 |
+
t_extrapolation_ratio: float = 1.0,
|
| 219 |
+
n_cameras: int = 4,
|
| 220 |
+
**kwargs,
|
| 221 |
+
):
|
| 222 |
+
"""
|
| 223 |
+
Args:
|
| 224 |
+
interpolation (str): we curretly only support "crop", ideally when we need extrapolation capacity, we should adjust frequency or other more advanced methods. they are not implemented yet.
|
| 225 |
+
"""
|
| 226 |
+
del kwargs # unused
|
| 227 |
+
self.n_cameras = n_cameras
|
| 228 |
+
super().__init__()
|
| 229 |
+
self.interpolation = interpolation
|
| 230 |
+
assert self.interpolation in ["crop"], f"Unknown interpolation method {self.interpolation}"
|
| 231 |
+
self.model_channels = model_channels
|
| 232 |
+
self.len_h = len_h
|
| 233 |
+
self.len_w = len_w
|
| 234 |
+
self.len_t = len_t
|
| 235 |
+
self.h_extrapolation_ratio = h_extrapolation_ratio
|
| 236 |
+
self.w_extrapolation_ratio = w_extrapolation_ratio
|
| 237 |
+
self.t_extrapolation_ratio = t_extrapolation_ratio
|
| 238 |
+
|
| 239 |
+
emb_h, emb_w, emb_t = self.compute_1d_pos_embeddings()
|
| 240 |
+
|
| 241 |
+
self.register_buffer("pos_emb_h", emb_h, persistent=False)
|
| 242 |
+
self.register_buffer("pos_emb_w", emb_w, persistent=False)
|
| 243 |
+
self.register_buffer("pos_emb_t", emb_t, persistent=False)
|
| 244 |
+
|
| 245 |
+
def compute_1d_pos_embeddings(self):
|
| 246 |
+
"""
|
| 247 |
+
Compute 1d pos embed for each axis
|
| 248 |
+
"""
|
| 249 |
+
dim = self.model_channels
|
| 250 |
+
dim_h = dim // 6 * 2
|
| 251 |
+
dim_w = dim_h
|
| 252 |
+
dim_t = dim - 2 * dim_h
|
| 253 |
+
assert dim == dim_h + dim_w + dim_t, f"bad dim: {dim} != {dim_h} + {dim_w} + {dim_t}"
|
| 254 |
+
|
| 255 |
+
emb_h = get_1d_sincos_pos_embed_from_grid(dim_h, pos=np.arange(self.len_h) * 1.0 / self.h_extrapolation_ratio)
|
| 256 |
+
emb_w = get_1d_sincos_pos_embed_from_grid(dim_w, pos=np.arange(self.len_w) * 1.0 / self.w_extrapolation_ratio)
|
| 257 |
+
emb_t = get_1d_sincos_pos_embed_from_grid(dim_t, pos=np.arange(self.len_t) * 1.0 / self.t_extrapolation_ratio)
|
| 258 |
+
|
| 259 |
+
emb_h = torch.from_numpy(emb_h).float()
|
| 260 |
+
emb_w = torch.from_numpy(emb_w).float()
|
| 261 |
+
emb_t = torch.from_numpy(emb_t).float()
|
| 262 |
+
|
| 263 |
+
return emb_h, emb_w, emb_t
|
| 264 |
+
|
| 265 |
+
def reset_parameters(self):
|
| 266 |
+
emb_h, emb_w, emb_t = self.compute_1d_pos_embeddings()
|
| 267 |
+
self.pos_emb_h = emb_h
|
| 268 |
+
self.pos_emb_w = emb_w
|
| 269 |
+
self.pos_emb_t = emb_t
|
| 270 |
+
|
| 271 |
+
def generate_embeddings(self, B_T_H_W_C: torch.Size, fps=Optional[torch.Tensor]) -> torch.Tensor:
|
| 272 |
+
B, T, H, W, C = B_T_H_W_C
|
| 273 |
+
|
| 274 |
+
single_camera_T = T // self.n_cameras
|
| 275 |
+
|
| 276 |
+
if self.interpolation == "crop":
|
| 277 |
+
emb_h_H = self.pos_emb_h[:H]
|
| 278 |
+
emb_w_W = self.pos_emb_w[:W]
|
| 279 |
+
emb_t_T = self.pos_emb_t[:single_camera_T]
|
| 280 |
+
emb = torch.cat(
|
| 281 |
+
[
|
| 282 |
+
torch.cat(
|
| 283 |
+
[
|
| 284 |
+
repeat(emb_t_T, "t d-> b t h w d", b=B, h=H, w=W),
|
| 285 |
+
repeat(emb_h_H, "h d-> b t h w d", b=B, t=single_camera_T, w=W),
|
| 286 |
+
repeat(emb_w_W, "w d-> b t h w d", b=B, t=single_camera_T, h=H),
|
| 287 |
+
],
|
| 288 |
+
dim=-1,
|
| 289 |
+
)
|
| 290 |
+
for _ in range(self.n_cameras)
|
| 291 |
+
],
|
| 292 |
+
1,
|
| 293 |
+
)
|
| 294 |
+
assert list(emb.shape)[:4] == [B, T, H, W], f"bad shape: {list(emb.shape)[:4]} != {B, T, H, W}"
|
| 295 |
+
return emb
|
| 296 |
+
|
| 297 |
+
@property
|
| 298 |
+
def seq_dim(self):
|
| 299 |
+
return 1
|
| 300 |
+
|
| 301 |
+
def _split_for_context_parallel(self, embeddings):
|
| 302 |
+
if self._cp_group is not None:
|
| 303 |
+
embeddings = rearrange(embeddings, "B (V T) H W C -> (B V) T H W C", V=self.n_cameras)
|
| 304 |
+
embeddings = split_inputs_cp(x=embeddings, seq_dim=self.seq_dim, cp_group=self._cp_group)
|
| 305 |
+
embeddings = rearrange(embeddings, "(B V) T H W C -> B (V T) H W C", V=self.n_cameras)
|
| 306 |
+
return embeddings
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
class MultiViewDiT(MinimalV1LVGDiT):
|
| 310 |
+
def __init__(
|
| 311 |
+
self,
|
| 312 |
+
*args,
|
| 313 |
+
timestep_scale: float = 1.0,
|
| 314 |
+
crossattn_emb_channels: int = 1024,
|
| 315 |
+
mlp_ratio: float = 4.0,
|
| 316 |
+
state_t: int,
|
| 317 |
+
n_cameras_emb: int,
|
| 318 |
+
view_condition_dim: int,
|
| 319 |
+
concat_view_embedding: bool,
|
| 320 |
+
layer_mask: Optional[List[bool]] = None,
|
| 321 |
+
sac_config: SACConfig = SACConfig(),
|
| 322 |
+
**kwargs,
|
| 323 |
+
):
|
| 324 |
+
self.state_t = state_t
|
| 325 |
+
self.n_cameras_emb = n_cameras_emb
|
| 326 |
+
self.view_condition_dim = view_condition_dim
|
| 327 |
+
self.concat_view_embedding = concat_view_embedding
|
| 328 |
+
assert "in_channels" in kwargs, "in_channels must be provided"
|
| 329 |
+
kwargs["in_channels"] += (
|
| 330 |
+
self.view_condition_dim if self.concat_view_embedding else 0
|
| 331 |
+
) # this avoids overwritting build_patch_embed which still adds padding_mask channel as appropriate
|
| 332 |
+
assert layer_mask is None, "layer_mask is not supported for MultiViewDiT"
|
| 333 |
+
if "n_cameras" in kwargs:
|
| 334 |
+
del kwargs["n_cameras"]
|
| 335 |
+
super().__init__(
|
| 336 |
+
*args,
|
| 337 |
+
mlp_ratio=mlp_ratio,
|
| 338 |
+
timestep_scale=timestep_scale,
|
| 339 |
+
crossattn_emb_channels=crossattn_emb_channels,
|
| 340 |
+
sac_config=sac_config,
|
| 341 |
+
**kwargs,
|
| 342 |
+
)
|
| 343 |
+
del self.blocks
|
| 344 |
+
self.blocks = nn.ModuleList(
|
| 345 |
+
[
|
| 346 |
+
MultiViewBlock(
|
| 347 |
+
x_dim=self.model_channels,
|
| 348 |
+
context_dim=crossattn_emb_channels,
|
| 349 |
+
num_heads=self.num_heads,
|
| 350 |
+
mlp_ratio=mlp_ratio,
|
| 351 |
+
use_adaln_lora=self.use_adaln_lora,
|
| 352 |
+
adaln_lora_dim=self.adaln_lora_dim,
|
| 353 |
+
backend=self.atten_backend,
|
| 354 |
+
image_context_dim=None if self.extra_image_context_dim is None else self.model_channels,
|
| 355 |
+
state_t=self.state_t,
|
| 356 |
+
use_wan_fp32_strategy=self.use_wan_fp32_strategy,
|
| 357 |
+
)
|
| 358 |
+
for _ in range(self.num_blocks)
|
| 359 |
+
]
|
| 360 |
+
)
|
| 361 |
+
|
| 362 |
+
if self.concat_view_embedding:
|
| 363 |
+
self.view_embeddings = nn.Embedding(self.n_cameras_emb, view_condition_dim)
|
| 364 |
+
|
| 365 |
+
self.init_weights()
|
| 366 |
+
self.enable_selective_checkpoint(sac_config, self.blocks)
|
| 367 |
+
|
| 368 |
+
def fully_shard(self, mesh, **fsdp_kwargs):
|
| 369 |
+
for i, block in enumerate(self.blocks):
|
| 370 |
+
reshard_after_forward = i < len(self.blocks) - 1
|
| 371 |
+
fully_shard(block, mesh=mesh, reshard_after_forward=reshard_after_forward, **fsdp_kwargs)
|
| 372 |
+
|
| 373 |
+
fully_shard(self.final_layer, mesh=mesh, reshard_after_forward=True, **fsdp_kwargs)
|
| 374 |
+
if self.extra_per_block_abs_pos_emb:
|
| 375 |
+
for extra_pos_embedder in self.extra_pos_embedders_options.values():
|
| 376 |
+
fully_shard(extra_pos_embedder, mesh=mesh, reshard_after_forward=True, **fsdp_kwargs)
|
| 377 |
+
fully_shard(self.t_embedder, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 378 |
+
if self.extra_image_context_dim is not None:
|
| 379 |
+
fully_shard(self.img_context_proj, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 380 |
+
|
| 381 |
+
def enable_context_parallel(self, process_group: Optional[ProcessGroup] = None):
|
| 382 |
+
# pos_embedder
|
| 383 |
+
for pos_embedder in self.pos_embedder_options.values():
|
| 384 |
+
pos_embedder.enable_context_parallel(process_group=process_group)
|
| 385 |
+
if self.extra_per_block_abs_pos_emb:
|
| 386 |
+
for extra_pos_embedder in self.extra_pos_embedders_options.values():
|
| 387 |
+
extra_pos_embedder.enable_context_parallel(process_group=process_group)
|
| 388 |
+
|
| 389 |
+
# attention
|
| 390 |
+
cp_ranks = get_process_group_ranks(process_group)
|
| 391 |
+
for block in self.blocks:
|
| 392 |
+
block.set_context_parallel_group(
|
| 393 |
+
process_group=process_group,
|
| 394 |
+
ranks=cp_ranks,
|
| 395 |
+
stream=torch.cuda.Stream(),
|
| 396 |
+
)
|
| 397 |
+
|
| 398 |
+
self._is_context_parallel_enabled = True
|
| 399 |
+
|
| 400 |
+
def disable_context_parallel(self):
|
| 401 |
+
# pos_embedder
|
| 402 |
+
for pos_embedder in self.pos_embedder_options.values():
|
| 403 |
+
pos_embedder.disable_context_parallel()
|
| 404 |
+
if self.extra_per_block_abs_pos_emb:
|
| 405 |
+
for extra_pos_embedder in self.extra_pos_embedders_options.values():
|
| 406 |
+
extra_pos_embedder.disable_context_parallel()
|
| 407 |
+
|
| 408 |
+
# attention
|
| 409 |
+
for block in self.blocks:
|
| 410 |
+
block.set_context_parallel_group(
|
| 411 |
+
process_group=None,
|
| 412 |
+
ranks=None,
|
| 413 |
+
stream=torch.cuda.Stream(),
|
| 414 |
+
)
|
| 415 |
+
|
| 416 |
+
self._is_context_parallel_enabled = False
|
| 417 |
+
|
| 418 |
+
def init_weights(self):
|
| 419 |
+
self.x_embedder.init_weights()
|
| 420 |
+
for pos_embedder in self.pos_embedder_options.values():
|
| 421 |
+
pos_embedder.reset_parameters()
|
| 422 |
+
if self.extra_per_block_abs_pos_emb:
|
| 423 |
+
for extra_pos_embedder in self.extra_pos_embedders_options.values():
|
| 424 |
+
extra_pos_embedder.init_weights()
|
| 425 |
+
|
| 426 |
+
self.t_embedder[1].init_weights()
|
| 427 |
+
for block in self.blocks:
|
| 428 |
+
block.init_weights()
|
| 429 |
+
|
| 430 |
+
self.final_layer.init_weights()
|
| 431 |
+
self.t_embedding_norm.reset_parameters()
|
| 432 |
+
|
| 433 |
+
if self.extra_image_context_dim is not None:
|
| 434 |
+
self.img_context_proj[0].reset_parameters()
|
| 435 |
+
|
| 436 |
+
def build_pos_embed(self):
|
| 437 |
+
self.pos_embedder_options = nn.ModuleDict()
|
| 438 |
+
self.extra_pos_embedders_options = nn.ModuleDict()
|
| 439 |
+
for n_cameras in range(1, self.n_cameras_emb + 1):
|
| 440 |
+
pos_embedder, extra_pos_embedder = self.build_pos_embed_for_n_cameras(n_cameras)
|
| 441 |
+
self.pos_embedder_options[f"n_cameras_{n_cameras}"] = pos_embedder
|
| 442 |
+
self.extra_pos_embedders_options[f"n_cameras_{n_cameras}"] = extra_pos_embedder
|
| 443 |
+
|
| 444 |
+
def build_pos_embed_for_n_cameras(self, n_cameras: int):
|
| 445 |
+
if self.pos_emb_cls == "rope3d":
|
| 446 |
+
cls_type = MultiCameraVideoRopePosition3DEmb
|
| 447 |
+
else:
|
| 448 |
+
raise ValueError(f"Unknown pos_emb_cls {self.pos_emb_cls}")
|
| 449 |
+
pos_embedder, extra_pos_embedder = None, None
|
| 450 |
+
kwargs = dict(
|
| 451 |
+
model_channels=self.model_channels,
|
| 452 |
+
len_h=self.max_img_h // self.patch_spatial,
|
| 453 |
+
len_w=self.max_img_w // self.patch_spatial,
|
| 454 |
+
len_t=self.max_frames // self.patch_temporal,
|
| 455 |
+
max_fps=self.max_fps,
|
| 456 |
+
min_fps=self.min_fps,
|
| 457 |
+
is_learnable=self.pos_emb_learnable,
|
| 458 |
+
interpolation=self.pos_emb_interpolation,
|
| 459 |
+
head_dim=self.model_channels // self.num_heads,
|
| 460 |
+
h_extrapolation_ratio=self.rope_h_extrapolation_ratio,
|
| 461 |
+
w_extrapolation_ratio=self.rope_w_extrapolation_ratio,
|
| 462 |
+
t_extrapolation_ratio=self.rope_t_extrapolation_ratio,
|
| 463 |
+
enable_fps_modulation=self.rope_enable_fps_modulation,
|
| 464 |
+
n_cameras=n_cameras,
|
| 465 |
+
)
|
| 466 |
+
pos_embedder = cls_type(
|
| 467 |
+
**kwargs,
|
| 468 |
+
)
|
| 469 |
+
assert pos_embedder.enable_fps_modulation == self.rope_enable_fps_modulation, (
|
| 470 |
+
"enable_fps_modulation must be the same"
|
| 471 |
+
)
|
| 472 |
+
|
| 473 |
+
if self.extra_per_block_abs_pos_emb:
|
| 474 |
+
raise NotImplementedError("extra_per_block_abs_pos_emb is not tested for multi-view DIT")
|
| 475 |
+
kwargs["h_extrapolation_ratio"] = self.extra_h_extrapolation_ratio
|
| 476 |
+
kwargs["w_extrapolation_ratio"] = self.extra_w_extrapolation_ratio
|
| 477 |
+
kwargs["t_extrapolation_ratio"] = self.extra_t_extrapolation_ratio
|
| 478 |
+
extra_pos_embedder = MultiCameraSinCosPosEmbAxis(
|
| 479 |
+
**kwargs,
|
| 480 |
+
)
|
| 481 |
+
return pos_embedder, extra_pos_embedder
|
| 482 |
+
|
| 483 |
+
def prepare_embedded_sequence(
|
| 484 |
+
self,
|
| 485 |
+
x_B_C_T_H_W: torch.Tensor,
|
| 486 |
+
fps: Optional[torch.Tensor] = None,
|
| 487 |
+
padding_mask: Optional[torch.Tensor] = None,
|
| 488 |
+
view_indices_B_T: Optional[torch.Tensor] = None,
|
| 489 |
+
) -> Tuple[torch.Tensor, Optional[torch.Tensor], torch.Tensor]:
|
| 490 |
+
if self.concat_padding_mask:
|
| 491 |
+
padding_mask = transforms.functional.resize(
|
| 492 |
+
padding_mask, list(x_B_C_T_H_W.shape[-2:]), interpolation=transforms.InterpolationMode.NEAREST
|
| 493 |
+
)
|
| 494 |
+
x_B_C_T_H_W = torch.cat(
|
| 495 |
+
[x_B_C_T_H_W, padding_mask.unsqueeze(1).repeat(1, 1, x_B_C_T_H_W.shape[2], 1, 1)], dim=1
|
| 496 |
+
)
|
| 497 |
+
cp_size = (
|
| 498 |
+
len(get_process_group_ranks(parallel_state.get_context_parallel_group()))
|
| 499 |
+
if parallel_state.is_initialized()
|
| 500 |
+
else 1
|
| 501 |
+
)
|
| 502 |
+
n_cameras = (x_B_C_T_H_W.shape[2] * cp_size) // self.state_t
|
| 503 |
+
pos_embedder = self.pos_embedder_options[f"n_cameras_{n_cameras}"]
|
| 504 |
+
if self.concat_view_embedding:
|
| 505 |
+
if view_indices_B_T is None:
|
| 506 |
+
view_indices = torch.arange(n_cameras).clamp(
|
| 507 |
+
max=self.n_cameras_emb - 1
|
| 508 |
+
) # View indices [0, 1, ..., V-1]
|
| 509 |
+
view_indices = view_indices.to(x_B_C_T_H_W.device)
|
| 510 |
+
view_embedding = self.view_embeddings(view_indices) # Shape: [V, embedding_dim]
|
| 511 |
+
view_embedding = rearrange(view_embedding, "V D -> D V")
|
| 512 |
+
view_embedding = (
|
| 513 |
+
view_embedding.unsqueeze(0).unsqueeze(3).unsqueeze(4).unsqueeze(5)
|
| 514 |
+
) # Shape: [1, D, V, 1, 1, 1]
|
| 515 |
+
else:
|
| 516 |
+
view_indices_B_T = view_indices_B_T.clamp(max=self.n_cameras_emb - 1)
|
| 517 |
+
view_indices_B_T = view_indices_B_T.to(x_B_C_T_H_W.device).long()
|
| 518 |
+
view_embedding = self.view_embeddings(view_indices_B_T) # B, (V T), D
|
| 519 |
+
view_embedding = rearrange(view_embedding, "B (V T) D -> B D V T", V=n_cameras)
|
| 520 |
+
view_embedding = view_embedding.unsqueeze(-1).unsqueeze(-1) # Shape: [B, D, V, T, 1, 1]
|
| 521 |
+
x_B_C_V_T_H_W = rearrange(x_B_C_T_H_W, "B C (V T) H W -> B C V T H W", V=n_cameras)
|
| 522 |
+
view_embedding = view_embedding.expand(
|
| 523 |
+
x_B_C_V_T_H_W.shape[0],
|
| 524 |
+
view_embedding.shape[1],
|
| 525 |
+
view_embedding.shape[2],
|
| 526 |
+
x_B_C_V_T_H_W.shape[3],
|
| 527 |
+
x_B_C_V_T_H_W.shape[4],
|
| 528 |
+
x_B_C_V_T_H_W.shape[5],
|
| 529 |
+
)
|
| 530 |
+
x_B_C_V_T_H_W = torch.cat([x_B_C_V_T_H_W, view_embedding], dim=1)
|
| 531 |
+
x_B_C_T_H_W = rearrange(x_B_C_V_T_H_W, " B C V T H W -> B C (V T) H W", V=n_cameras)
|
| 532 |
+
|
| 533 |
+
x_B_T_H_W_D = self.x_embedder(x_B_C_T_H_W)
|
| 534 |
+
|
| 535 |
+
if self.extra_per_block_abs_pos_emb:
|
| 536 |
+
extra_pos_embedder = self.extra_pos_embedders_options[str(n_cameras)]
|
| 537 |
+
extra_pos_emb = extra_pos_embedder(x_B_T_H_W_D, fps=fps)
|
| 538 |
+
else:
|
| 539 |
+
extra_pos_emb = None
|
| 540 |
+
|
| 541 |
+
if "rope" in self.pos_emb_cls.lower():
|
| 542 |
+
return x_B_T_H_W_D, pos_embedder(x_B_T_H_W_D, fps=fps), extra_pos_emb
|
| 543 |
+
|
| 544 |
+
if "fps_aware" in self.pos_emb_cls:
|
| 545 |
+
raise NotImplementedError("FPS-aware positional embedding is not supported for multi-view DIT")
|
| 546 |
+
|
| 547 |
+
x_B_T_H_W_D = x_B_T_H_W_D + pos_embedder(x_B_T_H_W_D)
|
| 548 |
+
|
| 549 |
+
return x_B_T_H_W_D, None, extra_pos_emb
|
| 550 |
+
|
| 551 |
+
def forward(
|
| 552 |
+
self,
|
| 553 |
+
x_B_C_T_H_W: torch.Tensor,
|
| 554 |
+
timesteps_B_T: torch.Tensor,
|
| 555 |
+
crossattn_emb: torch.Tensor,
|
| 556 |
+
condition_video_input_mask_B_C_T_H_W: Optional[torch.Tensor] = None,
|
| 557 |
+
fps: Optional[torch.Tensor] = None,
|
| 558 |
+
padding_mask: Optional[torch.Tensor] = None,
|
| 559 |
+
data_type: Optional[DataType] = DataType.VIDEO,
|
| 560 |
+
view_indices_B_T: Optional[torch.Tensor] = None,
|
| 561 |
+
intermediate_feature_ids: Optional[List[int]] = None,
|
| 562 |
+
**kwargs,
|
| 563 |
+
) -> torch.Tensor | List[torch.Tensor] | Tuple[torch.Tensor, List[torch.Tensor]]:
|
| 564 |
+
# Deletes elements like condition.use_video_condition that are not used in the forward pass
|
| 565 |
+
del kwargs
|
| 566 |
+
if data_type == DataType.VIDEO:
|
| 567 |
+
x_B_C_T_H_W = torch.cat([x_B_C_T_H_W, condition_video_input_mask_B_C_T_H_W.type_as(x_B_C_T_H_W)], dim=1)
|
| 568 |
+
else:
|
| 569 |
+
B, _, T, H, W = x_B_C_T_H_W.shape
|
| 570 |
+
x_B_C_T_H_W = torch.cat(
|
| 571 |
+
[x_B_C_T_H_W, torch.zeros((B, 1, T, H, W), dtype=x_B_C_T_H_W.dtype, device=x_B_C_T_H_W.device)], dim=1
|
| 572 |
+
)
|
| 573 |
+
|
| 574 |
+
assert isinstance(data_type, DataType), (
|
| 575 |
+
f"Expected DataType, got {type(data_type)}. We need discuss this flag later."
|
| 576 |
+
)
|
| 577 |
+
timesteps_B_T = timesteps_B_T * self.timestep_scale
|
| 578 |
+
x_B_T_H_W_D, rope_emb_L_1_1_D, extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D = self.prepare_embedded_sequence(
|
| 579 |
+
x_B_C_T_H_W,
|
| 580 |
+
fps=fps,
|
| 581 |
+
padding_mask=padding_mask,
|
| 582 |
+
view_indices_B_T=view_indices_B_T,
|
| 583 |
+
)
|
| 584 |
+
if self.use_crossattn_projection:
|
| 585 |
+
crossattn_emb = self.crossattn_proj(crossattn_emb)
|
| 586 |
+
with amp.autocast("cuda", enabled=self.use_wan_fp32_strategy, dtype=torch.float32):
|
| 587 |
+
if timesteps_B_T.ndim == 1:
|
| 588 |
+
timesteps_B_T = timesteps_B_T.unsqueeze(1)
|
| 589 |
+
t_embedding_B_T_D, adaln_lora_B_T_3D = self.t_embedder(timesteps_B_T)
|
| 590 |
+
t_embedding_B_T_D = self.t_embedding_norm(t_embedding_B_T_D)
|
| 591 |
+
|
| 592 |
+
# for logging purpose
|
| 593 |
+
affline_scale_log_info = {}
|
| 594 |
+
affline_scale_log_info["t_embedding_B_T_D"] = t_embedding_B_T_D.detach()
|
| 595 |
+
self.affline_scale_log_info = affline_scale_log_info
|
| 596 |
+
self.affline_emb = t_embedding_B_T_D
|
| 597 |
+
self.crossattn_emb = crossattn_emb
|
| 598 |
+
|
| 599 |
+
if extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D is not None:
|
| 600 |
+
assert x_B_T_H_W_D.shape == extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape, (
|
| 601 |
+
f"{x_B_T_H_W_D.shape} != {extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape}"
|
| 602 |
+
)
|
| 603 |
+
|
| 604 |
+
B, T, H, W, D = x_B_T_H_W_D.shape
|
| 605 |
+
|
| 606 |
+
for block in self.blocks:
|
| 607 |
+
x_B_T_H_W_D = block(
|
| 608 |
+
x_B_T_H_W_D,
|
| 609 |
+
t_embedding_B_T_D,
|
| 610 |
+
crossattn_emb,
|
| 611 |
+
rope_emb_L_1_1_D=rope_emb_L_1_1_D,
|
| 612 |
+
adaln_lora_B_T_3D=adaln_lora_B_T_3D,
|
| 613 |
+
extra_per_block_pos_emb=extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D,
|
| 614 |
+
)
|
| 615 |
+
|
| 616 |
+
x_B_T_H_W_O = self.final_layer(x_B_T_H_W_D, t_embedding_B_T_D, adaln_lora_B_T_3D=adaln_lora_B_T_3D)
|
| 617 |
+
x_B_C_Tt_Hp_Wp = self.unpatchify(x_B_T_H_W_O)
|
| 618 |
+
return x_B_C_Tt_Hp_Wp
|
cosmos_predict2/_src/predict2_multiview/networks/multiview_pose_dit.py
ADDED
|
@@ -0,0 +1,526 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
"""Multi-view DiT with a pose encoder + warped/visibility conditioning for 2-actor joint generation.
|
| 17 |
+
|
| 18 |
+
Design (additive, zero-initialized -- base checkpoint loads unchanged):
|
| 19 |
+
* ``pose_encoder`` : a 6x Conv3d (UniAnimate-DiT style) stack that turns a per-actor projected 2D pose RGB
|
| 20 |
+
map into a patch-grid feature that is ADDED to the patch tokens.
|
| 21 |
+
* ``cond_embedder`` : a PatchEmbed over ``cat([warped_latent(16), visibility_mask(1)])`` whose output is
|
| 22 |
+
ADDED to the patch tokens. This is mathematically identical to concatenating the
|
| 23 |
+
warped latent + visibility mask onto the input channels of ``x_embedder`` (since
|
| 24 |
+
``Linear([x; w; v]) == W_x x + W_w w + W_v v``), but it leaves ``x_embedder``
|
| 25 |
+
unchanged so the pretrained 2B checkpoint loads with no shape surgery.
|
| 26 |
+
|
| 27 |
+
Both extra paths have a zero-initialized final projection, so at step 0 the network output is identical to
|
| 28 |
+
the warm-started base ``MultiViewDiT`` (useful sanity check + stable fine-tuning start).
|
| 29 |
+
|
| 30 |
+
Actor 1 / Actor 2 are treated as the two "views" of the existing multi-view machinery: tensors are laid out
|
| 31 |
+
``(B, C, V*T, H, W)`` with ``V=2``; self-attention is shared across all ``V*T`` tokens (already the case in
|
| 32 |
+
``MultiViewBlock``), so the two actors are jointly denoised.
|
| 33 |
+
"""
|
| 34 |
+
|
| 35 |
+
from typing import List, Optional, Tuple
|
| 36 |
+
|
| 37 |
+
import torch
|
| 38 |
+
import torch.amp as amp
|
| 39 |
+
import torch.nn as nn
|
| 40 |
+
import torch.nn.functional as F
|
| 41 |
+
from einops import rearrange
|
| 42 |
+
from torch.distributed._composable.fsdp import fully_shard
|
| 43 |
+
|
| 44 |
+
from cosmos_predict2._src.predict2.conditioner import DataType
|
| 45 |
+
from cosmos_predict2._src.predict2.networks.minimal_v4_dit import PatchEmbed
|
| 46 |
+
from cosmos_predict2._src.predict2_multiview.networks.multiview_dit import MultiViewDiT
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
class PoseEncoder(nn.Module):
|
| 50 |
+
"""3D-conv driving-pose encoder, mirroring UniAnimate-DiT (arXiv:2504.11289).
|
| 51 |
+
|
| 52 |
+
Input : ``(N, 3, T_pixel, H_pixel, W_pixel)`` projected-pose RGB in uint8/float [0,255].
|
| 53 |
+
Output: ``(N, out_channels, T_latent, H_pixel/16, W_pixel/16)`` where ``T_latent`` matches the VAE's
|
| 54 |
+
temporal compression (/4). The final projection conv is zero-initialized.
|
| 55 |
+
"""
|
| 56 |
+
|
| 57 |
+
def __init__(self, out_channels: int, hidden_dim: int = 16, in_channels: int = 3, proj_init_std: float = 0.0):
|
| 58 |
+
super().__init__()
|
| 59 |
+
self.proj_init_std = proj_init_std # 0.0 -> zero-init proj (stable); >0 -> small-random (faster pose learning)
|
| 60 |
+
h = hidden_dim
|
| 61 |
+
self.conv = nn.Sequential(
|
| 62 |
+
nn.Conv3d(in_channels, h, (3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1)),
|
| 63 |
+
nn.SiLU(),
|
| 64 |
+
nn.Conv3d(h, h, (3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1)),
|
| 65 |
+
nn.SiLU(),
|
| 66 |
+
nn.Conv3d(h, h, (3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1)),
|
| 67 |
+
nn.SiLU(),
|
| 68 |
+
nn.Conv3d(h, h, (3, 3, 3), stride=(1, 2, 2), padding=(1, 1, 1)), # spatial /2
|
| 69 |
+
nn.SiLU(),
|
| 70 |
+
nn.Conv3d(h, h, (3, 3, 3), stride=(2, 2, 2), padding=(1, 1, 1)), # temporal/2, spatial/2
|
| 71 |
+
nn.SiLU(),
|
| 72 |
+
nn.Conv3d(h, h, (3, 3, 3), stride=(2, 2, 2), padding=(1, 1, 1)), # temporal/2, spatial/2
|
| 73 |
+
nn.SiLU(),
|
| 74 |
+
)
|
| 75 |
+
# final projection to model dim, last spatial /2 -> total spatial /16, temporal /4
|
| 76 |
+
self.proj = nn.Conv3d(h, out_channels, (1, 2, 2), stride=(1, 2, 2), padding=0)
|
| 77 |
+
|
| 78 |
+
def reset_parameters(self) -> None:
|
| 79 |
+
for m in self.conv.modules():
|
| 80 |
+
if isinstance(m, nn.Conv3d):
|
| 81 |
+
nn.init.kaiming_normal_(m.weight, nonlinearity="relu")
|
| 82 |
+
if m.bias is not None:
|
| 83 |
+
nn.init.zeros_(m.bias)
|
| 84 |
+
# final projection: zero-init (proj_init_std=0) keeps day-0 output == base but starves the conv stack
|
| 85 |
+
# of gradient until proj grows; a small-random init lets the conv feature extractor learn from step 0.
|
| 86 |
+
if self.proj_init_std > 0:
|
| 87 |
+
nn.init.normal_(self.proj.weight, mean=0.0, std=self.proj_init_std)
|
| 88 |
+
else:
|
| 89 |
+
nn.init.zeros_(self.proj.weight)
|
| 90 |
+
if self.proj.bias is not None:
|
| 91 |
+
nn.init.zeros_(self.proj.bias)
|
| 92 |
+
|
| 93 |
+
def forward(self, pose_N_C_T_H_W: torch.Tensor) -> torch.Tensor:
|
| 94 |
+
# normalize pose RGB [0,255] -> [0,1] in fp32, then match the conv weight dtype (bf16 during training)
|
| 95 |
+
w_dtype = self.conv[0].weight.dtype
|
| 96 |
+
x = pose_N_C_T_H_W.to(torch.float32) / 255.0
|
| 97 |
+
x = x.to(w_dtype)
|
| 98 |
+
x = self.conv(x)
|
| 99 |
+
x = self.proj(x)
|
| 100 |
+
return x
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
class MultiViewPoseDiT(MultiViewDiT):
|
| 104 |
+
def __init__(
|
| 105 |
+
self,
|
| 106 |
+
*args,
|
| 107 |
+
warped_latent_channels: int = 16,
|
| 108 |
+
visibility_channels: int = 1,
|
| 109 |
+
pose_in_channels: int = 3,
|
| 110 |
+
pose_hidden_dim: int = 16,
|
| 111 |
+
pose_proj_init_std: float = 0.0, # 0 -> zero-init pose proj; >0 -> small-random (faster pose learning, e.g. stage 1)
|
| 112 |
+
freeze_view_embedding: bool = False, # stage-1 single-view: view embedding is meaningless -> freeze it
|
| 113 |
+
# pose routing mode:
|
| 114 |
+
# "encoder" : Conv3d PoseEncoder over the pixel pose RGB, ADDED to tokens (default, original).
|
| 115 |
+
# "vae_concat" : frozen-VAE pose latent (16ch) CHANNEL-CONCATENATED with warped+vis into cond_embedder.
|
| 116 |
+
# "vae_mlp_add" : frozen-VAE pose latent -> trainable PatchEmbed+MLP -> ADDED to tokens.
|
| 117 |
+
pose_mode: str = "encoder",
|
| 118 |
+
pose_latent_channels: int = 16, # VAE latent ch for the pose map (== warped_latent_channels)
|
| 119 |
+
# in-context reference-frame appearance conditioning: append R clean reference frames per view as
|
| 120 |
+
# extra tokens (via the shared x_embedder) that participate in self-attention, at fixed non-contiguous
|
| 121 |
+
# temporal RoPE positions, then strip them before the output. Zero new params except a zero-init gate.
|
| 122 |
+
enable_reference_frames: bool = False,
|
| 123 |
+
num_reference_frames: int = 0, # R reference frames per view-slot (total appended = V*R)
|
| 124 |
+
ref_rope_offset: int = 50, # temporal RoPE position of the first reference frame
|
| 125 |
+
ref_rope_stride: int = 5, # spacing between consecutive reference-frame positions
|
| 126 |
+
# SHARED reference mode: the V*R appended refs are ONE greedy-selected shared set (split across the
|
| 127 |
+
# per-view slots), not per-view pools. With full cross-view self-attention both views attend all of
|
| 128 |
+
# them, so the view-slot placement is arbitrary and the per-view view-embedding on refs is dropped
|
| 129 |
+
# (refs carry no view identity; their only geometric id is the posed-reference Plücker). Real frames
|
| 130 |
+
# are unaffected.
|
| 131 |
+
shared_reference: bool = False,
|
| 132 |
+
# per-pixel Plücker ray conditioning: 6-ch ray map (dir+moment) -> zero-init PatchEmbed -> ADD to tokens
|
| 133 |
+
# (same additive pattern as warped/pose; mathematically == channel-concat + wider x_embedder, but keeps
|
| 134 |
+
# x_embedder unchanged so the base checkpoint warm-starts cleanly). Grounds both views in a shared frame.
|
| 135 |
+
enable_plucker: bool = False,
|
| 136 |
+
plucker_channels: int = 6,
|
| 137 |
+
enable_reference_plucker: bool = False, # also add Plücker to the reference tokens (posed references)
|
| 138 |
+
# also add a POSE latent to the reference tokens: VAE-encoded per-person skeleton render of each ref
|
| 139 |
+
# frame (16ch) -> zero-init PatchEmbed -> ADD to the reference tokens. Gives the in-context references
|
| 140 |
+
# the same pose signal the main frames get from control_input_pose. Same additive pattern as ref-Plücker.
|
| 141 |
+
enable_reference_pose: bool = False,
|
| 142 |
+
# composite DEPTH conditioning (warped scene depth + human mesh depth, RGB-encoded -> frozen VAE ->
|
| 143 |
+
# depth_latent 16ch) -> zero-init PatchEmbed -> ADD to tokens. Same additive pattern as warped/pose/plucker
|
| 144 |
+
# (mathematically == VAE-channel-concat into cond_embedder, but a separate zero-init branch keeps the base
|
| 145 |
+
# checkpoint's cond_embedder unchanged so it warm-starts cleanly).
|
| 146 |
+
enable_depth: bool = False,
|
| 147 |
+
depth_latent_channels: int = 16,
|
| 148 |
+
**kwargs,
|
| 149 |
+
):
|
| 150 |
+
self.enable_depth = enable_depth
|
| 151 |
+
self.depth_latent_channels = depth_latent_channels
|
| 152 |
+
self.warped_latent_channels = warped_latent_channels
|
| 153 |
+
self.visibility_channels = visibility_channels
|
| 154 |
+
self.pose_in_channels = pose_in_channels
|
| 155 |
+
self.pose_hidden_dim = pose_hidden_dim
|
| 156 |
+
self.pose_proj_init_std = pose_proj_init_std
|
| 157 |
+
self.freeze_view_embedding = freeze_view_embedding
|
| 158 |
+
assert pose_mode in ("encoder", "vae_concat", "vae_mlp_add"), pose_mode
|
| 159 |
+
self.pose_mode = pose_mode
|
| 160 |
+
self.pose_latent_channels = pose_latent_channels
|
| 161 |
+
self.enable_reference_frames = enable_reference_frames
|
| 162 |
+
self.shared_reference = shared_reference
|
| 163 |
+
self.num_reference_frames = int(num_reference_frames)
|
| 164 |
+
self.ref_rope_offset = int(ref_rope_offset)
|
| 165 |
+
self.ref_rope_stride = int(ref_rope_stride)
|
| 166 |
+
self.ref_positions = [self.ref_rope_offset + self.ref_rope_stride * i for i in range(self.num_reference_frames)]
|
| 167 |
+
self.enable_plucker = enable_plucker
|
| 168 |
+
self.plucker_channels = plucker_channels
|
| 169 |
+
self.enable_reference_plucker = enable_reference_plucker
|
| 170 |
+
self.enable_reference_pose = enable_reference_pose
|
| 171 |
+
# NOTE: we intentionally do NOT widen x_embedder. The warped+visibility "channel concat" is realized
|
| 172 |
+
# by a separate zero-init PatchEmbed added to the tokens (see module docstring).
|
| 173 |
+
super().__init__(*args, **kwargs)
|
| 174 |
+
if self.enable_reference_frames and self.num_reference_frames > 0:
|
| 175 |
+
# zero-init per-CHANNEL gate on the reference-token contribution -> day-0 output ~= warm-started
|
| 176 |
+
# base. Shape (model_channels,) (not a scalar): FSDP2 shards params per-parameter on dim 0, and a
|
| 177 |
+
# size-1 param yields a degenerate empty shard on some ranks that intermittently breaks the DCP
|
| 178 |
+
# save collective (gather_object NCCL error). A (D,) vector shards evenly and is more expressive.
|
| 179 |
+
self.ref_gate = nn.Parameter(torch.zeros(self.model_channels))
|
| 180 |
+
|
| 181 |
+
# additive conditioning path for warped latent + visibility mask (concatenated on channels). In
|
| 182 |
+
# "vae_concat" mode the VAE pose latent is concatenated into the SAME embedder (channel concat).
|
| 183 |
+
cond_in = self.warped_latent_channels + self.visibility_channels
|
| 184 |
+
if self.pose_mode == "vae_concat":
|
| 185 |
+
cond_in += self.pose_latent_channels
|
| 186 |
+
self.cond_embedder = PatchEmbed(
|
| 187 |
+
spatial_patch_size=self.patch_spatial,
|
| 188 |
+
temporal_patch_size=self.patch_temporal,
|
| 189 |
+
in_channels=cond_in,
|
| 190 |
+
out_channels=self.model_channels,
|
| 191 |
+
)
|
| 192 |
+
# Plücker ray conditioning: PatchEmbed over the 6-ch ray map, ADDED to tokens (zero-init proj)
|
| 193 |
+
if self.enable_plucker:
|
| 194 |
+
self.plucker_embedder = PatchEmbed(
|
| 195 |
+
spatial_patch_size=self.patch_spatial,
|
| 196 |
+
temporal_patch_size=self.patch_temporal,
|
| 197 |
+
in_channels=self.plucker_channels,
|
| 198 |
+
out_channels=self.model_channels,
|
| 199 |
+
)
|
| 200 |
+
# composite-depth conditioning: PatchEmbed over the 16-ch depth VAE latent, ADDED to tokens (zero-init proj)
|
| 201 |
+
if self.enable_depth:
|
| 202 |
+
self.depth_embedder = PatchEmbed(
|
| 203 |
+
spatial_patch_size=self.patch_spatial,
|
| 204 |
+
temporal_patch_size=self.patch_temporal,
|
| 205 |
+
in_channels=self.depth_latent_channels,
|
| 206 |
+
out_channels=self.model_channels,
|
| 207 |
+
)
|
| 208 |
+
# reference-POSE conditioning: PatchEmbed over the 16-ch VAE latent of each reference frame's skeleton
|
| 209 |
+
# render, ADDED to the reference tokens (zero-init). Mirrors the reference-Plücker additive path.
|
| 210 |
+
if self.enable_reference_pose:
|
| 211 |
+
self.reference_pose_embedder = PatchEmbed(
|
| 212 |
+
spatial_patch_size=self.patch_spatial,
|
| 213 |
+
temporal_patch_size=self.patch_temporal,
|
| 214 |
+
in_channels=self.pose_latent_channels,
|
| 215 |
+
out_channels=self.model_channels,
|
| 216 |
+
)
|
| 217 |
+
# pose conditioning path (varies by mode)
|
| 218 |
+
if self.pose_mode == "encoder":
|
| 219 |
+
self.pose_encoder = PoseEncoder(
|
| 220 |
+
out_channels=self.model_channels,
|
| 221 |
+
hidden_dim=self.pose_hidden_dim,
|
| 222 |
+
in_channels=self.pose_in_channels,
|
| 223 |
+
proj_init_std=self.pose_proj_init_std,
|
| 224 |
+
)
|
| 225 |
+
elif self.pose_mode == "vae_mlp_add":
|
| 226 |
+
# frozen-VAE pose latent -> patch-embed to tokens -> trainable MLP -> ADD. MLP last layer
|
| 227 |
+
# zero-initialized so day-0 output == warm-started base.
|
| 228 |
+
self.pose_latent_embedder = PatchEmbed(
|
| 229 |
+
spatial_patch_size=self.patch_spatial,
|
| 230 |
+
temporal_patch_size=self.patch_temporal,
|
| 231 |
+
in_channels=self.pose_latent_channels,
|
| 232 |
+
out_channels=self.model_channels,
|
| 233 |
+
)
|
| 234 |
+
self.pose_mlp = nn.Sequential(
|
| 235 |
+
nn.Linear(self.model_channels, self.model_channels),
|
| 236 |
+
nn.SiLU(),
|
| 237 |
+
nn.Linear(self.model_channels, self.model_channels),
|
| 238 |
+
)
|
| 239 |
+
|
| 240 |
+
# re-run init now that the new submodules exist (zero-inits their final projections)
|
| 241 |
+
self.init_weights()
|
| 242 |
+
|
| 243 |
+
# stage-1 (single-view): the view embedding can't distinguish actors with V=1 -> exclude it from
|
| 244 |
+
# training (architecture/x_embedder unchanged, so weights still transfer cleanly to the joint stage 2).
|
| 245 |
+
if self.freeze_view_embedding and hasattr(self, "view_embeddings"):
|
| 246 |
+
self.view_embeddings.requires_grad_(False)
|
| 247 |
+
|
| 248 |
+
# ------------------------------------------------------------------ weights / sharding
|
| 249 |
+
def init_weights(self):
|
| 250 |
+
super().init_weights()
|
| 251 |
+
if hasattr(self, "cond_embedder"):
|
| 252 |
+
# zero-init the warped/visibility(+pose, in vae_concat) projection -> day-0 contribution is 0 (== base)
|
| 253 |
+
nn.init.zeros_(self.cond_embedder.proj[1].weight)
|
| 254 |
+
if hasattr(self, "plucker_embedder"):
|
| 255 |
+
nn.init.zeros_(self.plucker_embedder.proj[1].weight) # day-0 Plücker contribution is 0 (== base)
|
| 256 |
+
if hasattr(self, "depth_embedder"):
|
| 257 |
+
nn.init.zeros_(self.depth_embedder.proj[1].weight) # day-0 depth contribution is 0 (== base)
|
| 258 |
+
if hasattr(self, "reference_pose_embedder"):
|
| 259 |
+
nn.init.zeros_(self.reference_pose_embedder.proj[1].weight) # day-0 ref-pose contribution is 0 (== base)
|
| 260 |
+
if hasattr(self, "pose_encoder"):
|
| 261 |
+
self.pose_encoder.reset_parameters()
|
| 262 |
+
if hasattr(self, "pose_mlp"):
|
| 263 |
+
# zero-init the MLP's last layer -> the (trainable) VAE-pose path contributes 0 at step 0
|
| 264 |
+
nn.init.zeros_(self.pose_mlp[-1].weight)
|
| 265 |
+
if self.pose_mlp[-1].bias is not None:
|
| 266 |
+
nn.init.zeros_(self.pose_mlp[-1].bias)
|
| 267 |
+
|
| 268 |
+
def fully_shard(self, mesh, **fsdp_kwargs):
|
| 269 |
+
super().fully_shard(mesh, **fsdp_kwargs)
|
| 270 |
+
if hasattr(self, "cond_embedder"):
|
| 271 |
+
fully_shard(self.cond_embedder, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 272 |
+
if hasattr(self, "plucker_embedder"):
|
| 273 |
+
fully_shard(self.plucker_embedder, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 274 |
+
if hasattr(self, "depth_embedder"):
|
| 275 |
+
fully_shard(self.depth_embedder, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 276 |
+
if hasattr(self, "reference_pose_embedder"):
|
| 277 |
+
fully_shard(self.reference_pose_embedder, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 278 |
+
if hasattr(self, "pose_encoder"):
|
| 279 |
+
fully_shard(self.pose_encoder, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 280 |
+
if hasattr(self, "pose_latent_embedder"):
|
| 281 |
+
fully_shard(self.pose_latent_embedder, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 282 |
+
if hasattr(self, "pose_mlp"):
|
| 283 |
+
fully_shard(self.pose_mlp, mesh=mesh, reshard_after_forward=False, **fsdp_kwargs)
|
| 284 |
+
|
| 285 |
+
# ------------------------------------------------------------------ conditioning helpers
|
| 286 |
+
def _embed_warped_visibility(
|
| 287 |
+
self,
|
| 288 |
+
warped_latent_B_C_T_H_W: torch.Tensor,
|
| 289 |
+
visibility_mask_B_C_T_H_W: torch.Tensor,
|
| 290 |
+
pose_latent_B_C_T_H_W: Optional[torch.Tensor] = None,
|
| 291 |
+
) -> torch.Tensor:
|
| 292 |
+
"""PatchEmbed over cat([warped_latent, visibility(, pose_latent)]) -> token grid (B, T, H, W, D).
|
| 293 |
+
|
| 294 |
+
In "vae_concat" mode the frozen-VAE pose latent is concatenated on the channel dim alongside the
|
| 295 |
+
warped latent + visibility mask, so a single linear patch projection mixes all conditioning latents.
|
| 296 |
+
"""
|
| 297 |
+
ref_dtype = self.cond_embedder.proj[1].weight.dtype
|
| 298 |
+
parts = [warped_latent_B_C_T_H_W.to(ref_dtype), visibility_mask_B_C_T_H_W.to(ref_dtype)]
|
| 299 |
+
if pose_latent_B_C_T_H_W is not None:
|
| 300 |
+
parts.append(pose_latent_B_C_T_H_W.to(ref_dtype))
|
| 301 |
+
cond_B_C_T_H_W = torch.cat(parts, dim=1)
|
| 302 |
+
return self.cond_embedder(cond_B_C_T_H_W)
|
| 303 |
+
|
| 304 |
+
def _embed_pose_latent_mlp(self, pose_latent_B_C_T_H_W: torch.Tensor) -> torch.Tensor:
|
| 305 |
+
"""VAE pose latent -> trainable PatchEmbed -> MLP -> token grid (B, T, H, W, D). Last MLP layer
|
| 306 |
+
zero-init so day-0 contribution is 0."""
|
| 307 |
+
ref_dtype = self.pose_latent_embedder.proj[1].weight.dtype
|
| 308 |
+
feat = self.pose_latent_embedder(pose_latent_B_C_T_H_W.to(ref_dtype)) # (B, T, H, W, D)
|
| 309 |
+
return self.pose_mlp(feat)
|
| 310 |
+
|
| 311 |
+
def _embed_pose(
|
| 312 |
+
self, pose_map_B_C_T_H_W: torch.Tensor, n_cameras: int, grid_thw: Tuple[int, int, int]
|
| 313 |
+
) -> torch.Tensor:
|
| 314 |
+
"""Run the pose encoder per-view (no cross-actor temporal leak) -> token grid (B, V*T, H, W, D)."""
|
| 315 |
+
T_tok, H_tok, W_tok = grid_thw
|
| 316 |
+
# split actors so the temporal convs never span the actor boundary
|
| 317 |
+
pose_NV = rearrange(pose_map_B_C_T_H_W, "B C (V T) H W -> (B V) C T H W", V=n_cameras)
|
| 318 |
+
feat = self.pose_encoder(pose_NV) # (B*V, D, t, h, w)
|
| 319 |
+
feat = feat.to(self.cond_embedder.proj[1].weight.dtype)
|
| 320 |
+
# be robust to off-by-one between conv-downsampled grid and the latent patch grid
|
| 321 |
+
if tuple(feat.shape[-3:]) != (T_tok // n_cameras, H_tok, W_tok):
|
| 322 |
+
feat = F.interpolate(
|
| 323 |
+
feat, size=(T_tok // n_cameras, H_tok, W_tok), mode="trilinear", align_corners=False
|
| 324 |
+
)
|
| 325 |
+
feat = rearrange(feat, "(B V) D t h w -> B (V t) h w D", V=n_cameras)
|
| 326 |
+
return feat
|
| 327 |
+
|
| 328 |
+
def _embed_reference_tokens(
|
| 329 |
+
self, reference_latent_B_C_VR_H_W: torch.Tensor, n_cameras: int, reference_plucker_map=None,
|
| 330 |
+
reference_pose_latent=None,
|
| 331 |
+
) -> torch.Tensor:
|
| 332 |
+
"""Assemble the SAME input channels as prepare_embedded_sequence for the reference latents and run the
|
| 333 |
+
SHARED x_embedder -> reference token grid (B, V*R, H_tok, W_tok, D). Reference frames are marked
|
| 334 |
+
'given/clean' via an all-ones condition mask (reusing the existing mask channel; no x_embedder widening).
|
| 335 |
+
Each view's R refs get that view's view-embedding (broadcast), so refs are tagged to their stream."""
|
| 336 |
+
ref_dtype = self.x_embedder.proj[1].weight.dtype
|
| 337 |
+
x = reference_latent_B_C_VR_H_W.to(ref_dtype)
|
| 338 |
+
B, C, VR, H, W = x.shape
|
| 339 |
+
R = VR // n_cameras
|
| 340 |
+
# (1) condition-mask channel = ones (refs are given/clean) -- mirrors forward()'s cond-mask concat
|
| 341 |
+
x = torch.cat([x, torch.ones((B, 1, VR, H, W), dtype=ref_dtype, device=x.device)], dim=1)
|
| 342 |
+
# (2) padding-mask channel = zeros (refs have no padding) -- mirrors prepare_embedded_sequence
|
| 343 |
+
if self.concat_padding_mask:
|
| 344 |
+
x = torch.cat([x, torch.zeros((B, 1, VR, H, W), dtype=ref_dtype, device=x.device)], dim=1)
|
| 345 |
+
# (3) view-embedding channels. The SHARED x_embedder expects these channels (real frames always carry
|
| 346 |
+
# them), so the width must match either way -- "removing the view embedding" for shared refs means
|
| 347 |
+
# feeding a NEUTRAL (zero) view signal, not dropping the channels:
|
| 348 |
+
# - per-view mode: each view's refs get that view's learned embedding (broadcast), tagging them.
|
| 349 |
+
# - shared_reference mode: all refs get an all-ZERO view channel -> no view identity (refs are one
|
| 350 |
+
# shared set with arbitrary view-slot placement); their only geometric id is the posed Plücker.
|
| 351 |
+
if self.concat_view_embedding:
|
| 352 |
+
if self.shared_reference:
|
| 353 |
+
ve = torch.zeros((B, self.view_condition_dim, VR, H, W), dtype=ref_dtype, device=x.device)
|
| 354 |
+
x = torch.cat([x, ve], dim=1)
|
| 355 |
+
else:
|
| 356 |
+
view_idx = torch.arange(n_cameras, device=x.device).clamp(max=self.n_cameras_emb - 1)
|
| 357 |
+
ve = self.view_embeddings(view_idx).to(ref_dtype) # (V, Dv)
|
| 358 |
+
ve = rearrange(ve, "V D -> D V")[None, :, :, None, None, None] # (1, Dv, V, 1, 1, 1)
|
| 359 |
+
xv = rearrange(x, "B C (V R) H W -> B C V R H W", V=n_cameras)
|
| 360 |
+
ve = ve.expand(B, ve.shape[1], n_cameras, R, H, W)
|
| 361 |
+
x = rearrange(torch.cat([xv, ve], dim=1), "B C V R H W -> B C (V R) H W")
|
| 362 |
+
ref_tokens = self.x_embedder(x) # (B, V*R, H_tok, W_tok, D)
|
| 363 |
+
# posed-reference Plücker: ground the reference tokens geometrically (reuse the zero-init plucker embedder)
|
| 364 |
+
if self.enable_reference_plucker and hasattr(self, "plucker_embedder") and reference_plucker_map is not None:
|
| 365 |
+
plk = self.plucker_embedder(reference_plucker_map.to(ref_dtype)) # (B, V*R, H_tok, W_tok, D)
|
| 366 |
+
ref_tokens = ref_tokens + plk.to(ref_tokens.dtype)
|
| 367 |
+
# reference-POSE: add the per-person skeleton latent of each reference frame (zero-init embedder)
|
| 368 |
+
if self.enable_reference_pose and hasattr(self, "reference_pose_embedder") and reference_pose_latent is not None:
|
| 369 |
+
rp = self.reference_pose_embedder(reference_pose_latent.to(ref_dtype)) # (B, V*R, H_tok, W_tok, D)
|
| 370 |
+
ref_tokens = ref_tokens + rp.to(ref_tokens.dtype)
|
| 371 |
+
return ref_tokens
|
| 372 |
+
|
| 373 |
+
# ------------------------------------------------------------------ forward
|
| 374 |
+
def forward(
|
| 375 |
+
self,
|
| 376 |
+
x_B_C_T_H_W: torch.Tensor,
|
| 377 |
+
timesteps_B_T: torch.Tensor,
|
| 378 |
+
crossattn_emb: torch.Tensor,
|
| 379 |
+
condition_video_input_mask_B_C_T_H_W: Optional[torch.Tensor] = None,
|
| 380 |
+
fps: Optional[torch.Tensor] = None,
|
| 381 |
+
padding_mask: Optional[torch.Tensor] = None,
|
| 382 |
+
data_type: Optional[DataType] = DataType.VIDEO,
|
| 383 |
+
view_indices_B_T: Optional[torch.Tensor] = None,
|
| 384 |
+
intermediate_feature_ids: Optional[List[int]] = None,
|
| 385 |
+
# NEW conditioning inputs (declared explicitly so they survive `del kwargs`)
|
| 386 |
+
pose_map_B_C_T_H_W: Optional[torch.Tensor] = None,
|
| 387 |
+
pose_latent_B_C_T_H_W: Optional[torch.Tensor] = None,
|
| 388 |
+
warped_latent_B_C_T_H_W: Optional[torch.Tensor] = None,
|
| 389 |
+
visibility_mask_B_C_T_H_W: Optional[torch.Tensor] = None,
|
| 390 |
+
reference_latent_B_C_R_H_W: Optional[torch.Tensor] = None,
|
| 391 |
+
plucker_map_B_C_T_H_W: Optional[torch.Tensor] = None,
|
| 392 |
+
reference_plucker_map_B_C_R_H_W: Optional[torch.Tensor] = None,
|
| 393 |
+
reference_pose_latent_B_C_R_H_W: Optional[torch.Tensor] = None,
|
| 394 |
+
depth_latent_B_C_T_H_W: Optional[torch.Tensor] = None,
|
| 395 |
+
**kwargs,
|
| 396 |
+
) -> torch.Tensor:
|
| 397 |
+
del kwargs
|
| 398 |
+
if data_type == DataType.VIDEO:
|
| 399 |
+
x_B_C_T_H_W = torch.cat(
|
| 400 |
+
[x_B_C_T_H_W, condition_video_input_mask_B_C_T_H_W.type_as(x_B_C_T_H_W)], dim=1
|
| 401 |
+
)
|
| 402 |
+
else:
|
| 403 |
+
B, _, T, H, W = x_B_C_T_H_W.shape
|
| 404 |
+
x_B_C_T_H_W = torch.cat(
|
| 405 |
+
[x_B_C_T_H_W, torch.zeros((B, 1, T, H, W), dtype=x_B_C_T_H_W.dtype, device=x_B_C_T_H_W.device)],
|
| 406 |
+
dim=1,
|
| 407 |
+
)
|
| 408 |
+
|
| 409 |
+
assert isinstance(data_type, DataType)
|
| 410 |
+
timesteps_B_T = timesteps_B_T * self.timestep_scale
|
| 411 |
+
x_B_T_H_W_D, rope_emb_L_1_1_D, extra_pos_emb = self.prepare_embedded_sequence(
|
| 412 |
+
x_B_C_T_H_W,
|
| 413 |
+
fps=fps,
|
| 414 |
+
padding_mask=padding_mask,
|
| 415 |
+
view_indices_B_T=view_indices_B_T,
|
| 416 |
+
)
|
| 417 |
+
|
| 418 |
+
B, T_tok, H_tok, W_tok, D = x_B_T_H_W_D.shape
|
| 419 |
+
n_cameras = T_tok // self.state_t
|
| 420 |
+
|
| 421 |
+
# additive warped-latent + visibility conditioning (zero-init -> 0 at step 0). In "vae_concat" mode
|
| 422 |
+
# the VAE pose latent is channel-concatenated into the same cond_embedder.
|
| 423 |
+
# cast the additive term back to x's dtype so we never silently upcast the token stream.
|
| 424 |
+
if warped_latent_B_C_T_H_W is not None and visibility_mask_B_C_T_H_W is not None:
|
| 425 |
+
pose_for_concat = pose_latent_B_C_T_H_W if self.pose_mode == "vae_concat" else None
|
| 426 |
+
cond_emb = self._embed_warped_visibility(
|
| 427 |
+
warped_latent_B_C_T_H_W, visibility_mask_B_C_T_H_W, pose_for_concat
|
| 428 |
+
)
|
| 429 |
+
x_B_T_H_W_D = x_B_T_H_W_D + cond_emb.to(x_B_T_H_W_D.dtype)
|
| 430 |
+
|
| 431 |
+
# additive pose conditioning (zero-init -> 0 at step 0)
|
| 432 |
+
if self.pose_mode == "encoder" and pose_map_B_C_T_H_W is not None:
|
| 433 |
+
pose_emb = self._embed_pose(pose_map_B_C_T_H_W, n_cameras, (T_tok, H_tok, W_tok))
|
| 434 |
+
x_B_T_H_W_D = x_B_T_H_W_D + pose_emb.to(x_B_T_H_W_D.dtype)
|
| 435 |
+
elif self.pose_mode == "vae_mlp_add" and pose_latent_B_C_T_H_W is not None:
|
| 436 |
+
pose_emb = self._embed_pose_latent_mlp(pose_latent_B_C_T_H_W)
|
| 437 |
+
x_B_T_H_W_D = x_B_T_H_W_D + pose_emb.to(x_B_T_H_W_D.dtype)
|
| 438 |
+
|
| 439 |
+
# additive Plücker ray conditioning (zero-init -> 0 at step 0). Per-pixel camera rays (canonical frame)
|
| 440 |
+
# patch-embedded and added to the real-frame tokens -> cross-view shared-space grounding.
|
| 441 |
+
if self.enable_plucker and plucker_map_B_C_T_H_W is not None:
|
| 442 |
+
plk_dtype = self.plucker_embedder.proj[1].weight.dtype
|
| 443 |
+
plucker_emb = self.plucker_embedder(plucker_map_B_C_T_H_W.to(plk_dtype)) # (B, V*state_t, h, w, D)
|
| 444 |
+
x_B_T_H_W_D = x_B_T_H_W_D + plucker_emb.to(x_B_T_H_W_D.dtype)
|
| 445 |
+
|
| 446 |
+
# additive composite-DEPTH conditioning (zero-init -> 0 at step 0). Depth VAE latent (warped scene depth +
|
| 447 |
+
# human mesh depth) patch-embedded and added to the real-frame tokens (== VAE-channel-concat, warm-clean).
|
| 448 |
+
if self.enable_depth and depth_latent_B_C_T_H_W is not None:
|
| 449 |
+
d_dtype = self.depth_embedder.proj[1].weight.dtype
|
| 450 |
+
depth_emb = self.depth_embedder(depth_latent_B_C_T_H_W.to(d_dtype))
|
| 451 |
+
x_B_T_H_W_D = x_B_T_H_W_D + depth_emb.to(x_B_T_H_W_D.dtype)
|
| 452 |
+
|
| 453 |
+
# in-context reference-frame appearance conditioning: append R reference frames per view onto the
|
| 454 |
+
# temporal axis (view-major: [view0 real, view0 ref, view1 real, view1 ref, ...]) with fixed
|
| 455 |
+
# non-contiguous temporal RoPE, gated by a zero-init scalar. Stripped before the output layer.
|
| 456 |
+
use_refs = (
|
| 457 |
+
self.enable_reference_frames
|
| 458 |
+
and self.num_reference_frames > 0
|
| 459 |
+
and reference_latent_B_C_R_H_W is not None
|
| 460 |
+
)
|
| 461 |
+
if use_refs:
|
| 462 |
+
R = self.num_reference_frames
|
| 463 |
+
ref_tokens = self._embed_reference_tokens(
|
| 464 |
+
reference_latent_B_C_R_H_W, n_cameras, reference_plucker_map_B_C_R_H_W,
|
| 465 |
+
reference_pose_latent_B_C_R_H_W,
|
| 466 |
+
) # (B, V*R, H, W, D)
|
| 467 |
+
gate = self.ref_gate.to(x_B_T_H_W_D.dtype)
|
| 468 |
+
x_real = rearrange(x_B_T_H_W_D, "B (V t) H W D -> B V t H W D", V=n_cameras)
|
| 469 |
+
x_ref = rearrange(gate * ref_tokens.to(x_B_T_H_W_D.dtype), "B (V r) H W D -> B V r H W D", V=n_cameras)
|
| 470 |
+
x_B_T_H_W_D = rearrange(torch.cat([x_real, x_ref], dim=2), "B V te H W D -> B (V te) H W D")
|
| 471 |
+
# regenerate RoPE for the augmented per-view [real(state_t) ; ref(R)] grid (real positions
|
| 472 |
+
# [0..state_t-1], ref positions [offset, offset+stride, ...]); order matches the token layout above.
|
| 473 |
+
Bx, Te, Hx, Wx, Dx = x_B_T_H_W_D.shape
|
| 474 |
+
pos_embedder = self.pos_embedder_options[f"n_cameras_{n_cameras}"]
|
| 475 |
+
ref_rope = pos_embedder.generate_embeddings_with_refs(
|
| 476 |
+
(Bx, Te, Hx, Wx, Dx), num_ref=R, ref_positions=self.ref_positions, fps=fps
|
| 477 |
+
)
|
| 478 |
+
rope_emb_L_1_1_D = pos_embedder._split_for_context_parallel(ref_rope)
|
| 479 |
+
|
| 480 |
+
if self.use_crossattn_projection:
|
| 481 |
+
crossattn_emb = self.crossattn_proj(crossattn_emb)
|
| 482 |
+
|
| 483 |
+
with amp.autocast("cuda", enabled=self.use_wan_fp32_strategy, dtype=torch.float32):
|
| 484 |
+
if timesteps_B_T.ndim == 1:
|
| 485 |
+
timesteps_B_T = timesteps_B_T.unsqueeze(1)
|
| 486 |
+
t_embedding_B_T_D, adaln_lora_B_T_3D = self.t_embedder(timesteps_B_T)
|
| 487 |
+
t_embedding_B_T_D = self.t_embedding_norm(t_embedding_B_T_D)
|
| 488 |
+
|
| 489 |
+
affline_scale_log_info = {}
|
| 490 |
+
affline_scale_log_info["t_embedding_B_T_D"] = t_embedding_B_T_D.detach()
|
| 491 |
+
self.affline_scale_log_info = affline_scale_log_info
|
| 492 |
+
self.affline_emb = t_embedding_B_T_D
|
| 493 |
+
self.crossattn_emb = crossattn_emb
|
| 494 |
+
|
| 495 |
+
# when reference frames are appended, the per-frame adaLN modulation must span (state_t+R) frames/view.
|
| 496 |
+
# Expand the (per-frame) timestep embedding + adaLN-LoRA by padding each view's R ref slots with its
|
| 497 |
+
# last real-frame value (refs are stripped before the loss, so this padded value never reaches output).
|
| 498 |
+
te_blocks, lora_blocks = t_embedding_B_T_D, adaln_lora_B_T_3D
|
| 499 |
+
if use_refs:
|
| 500 |
+
def _pad_ref_slots(emb):
|
| 501 |
+
if emb is None or emb.shape[1] != n_cameras * self.state_t:
|
| 502 |
+
return emb # T==1 (broadcast) or already augmented -> leave as-is
|
| 503 |
+
e = rearrange(emb, "B (V t) X -> B V t X", V=n_cameras)
|
| 504 |
+
e = torch.cat([e, e[:, :, -1:, :].expand(-1, -1, self.num_reference_frames, -1)], dim=2)
|
| 505 |
+
return rearrange(e, "B V te X -> B (V te) X")
|
| 506 |
+
te_blocks = _pad_ref_slots(t_embedding_B_T_D)
|
| 507 |
+
lora_blocks = _pad_ref_slots(adaln_lora_B_T_3D)
|
| 508 |
+
|
| 509 |
+
for block in self.blocks:
|
| 510 |
+
x_B_T_H_W_D = block(
|
| 511 |
+
x_B_T_H_W_D,
|
| 512 |
+
te_blocks,
|
| 513 |
+
crossattn_emb,
|
| 514 |
+
rope_emb_L_1_1_D=rope_emb_L_1_1_D,
|
| 515 |
+
adaln_lora_B_T_3D=lora_blocks,
|
| 516 |
+
extra_per_block_pos_emb=extra_pos_emb,
|
| 517 |
+
)
|
| 518 |
+
|
| 519 |
+
# strip the appended reference frames per view -> back to (B, V*state_t, H, W, D) for the output layer
|
| 520 |
+
if use_refs:
|
| 521 |
+
xg = rearrange(x_B_T_H_W_D, "B (V te) H W D -> B V te H W D", V=n_cameras)
|
| 522 |
+
x_B_T_H_W_D = rearrange(xg[:, :, : self.state_t], "B V t H W D -> B (V t) H W D")
|
| 523 |
+
|
| 524 |
+
x_B_T_H_W_O = self.final_layer(x_B_T_H_W_D, t_embedding_B_T_D, adaln_lora_B_T_3D=adaln_lora_B_T_3D)
|
| 525 |
+
x_B_C_Tt_Hp_Wp = self.unpatchify(x_B_T_H_W_O)
|
| 526 |
+
return x_B_C_Tt_Hp_Wp
|
cosmos_predict2/_src/predict2_multiview/scripts/inference.py
ADDED
|
@@ -0,0 +1,330 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
"""
|
| 16 |
+
This script is based on projects/cosmos/diffusion/v2/inference/vid2vid.py
|
| 17 |
+
|
| 18 |
+
To run inference on the training data (as visualization/debugging), use:
|
| 19 |
+
```bash
|
| 20 |
+
EXP=buttercup_predict2p5_2b_7views_res720p_fps30_t8_joint_alpamayo1capviewprefix_allcapsviewprefix_29frames_nofps_uniform_dropoutt0
|
| 21 |
+
ckpt_path=s3://bucket/cosmos_predict2_multiview/cosmos2_mv/buttercup_predict2p5_2b_7views_res720p_fps30_t8_joint_alpamayo1capviewprefix_allcapsviewprefix_29frames_nofps_uniform_dropoutt0-0/checkpoints/iter_000012000/
|
| 22 |
+
PYTHONPATH=. torchrun --nproc_per_node=8 --master_port=12341 -m cosmos_predict2._src.predict2_multiview.scripts.inference --experiment ${EXP} --ckpt_path ${ckpt_path} --context_parallel_size 8 --input_is_train_data --max_samples 1 --num_conditional_frames 0 --guidance 3 --save_root results/predict2_multiview/
|
| 23 |
+
|
| 24 |
+
EXP=predict2p5_2b_mv_7train7_res480p_fps15_t24_alpamayo_only_allcaption_uniform_nofps
|
| 25 |
+
ckpt_path=s3://bucket/cosmos_predict2_multiview/cosmos2p5_mv/predict2p5_2b_mv_7train7_res480p_fps15_t24_alpamayo_only_allcaption_uniform_nofps-0/checkpoints/iter_000020000/
|
| 26 |
+
PYTHONPATH=. torchrun --nproc_per_node=1 --master_port=12341 -m cosmos_predict2._src.predict2_multiview.scripts.inference --experiment ${EXP} --ckpt_path ${ckpt_path} --context_parallel_size 1 --input_is_train_data --max_samples 1 --num_conditional_frames 0 --guidance 3 --save_root results/predict2_multiview_480p_20k/
|
| 27 |
+
```
|
| 28 |
+
"""
|
| 29 |
+
|
| 30 |
+
import argparse
|
| 31 |
+
import os
|
| 32 |
+
|
| 33 |
+
import torch as th
|
| 34 |
+
from einops import rearrange
|
| 35 |
+
from megatron.core import parallel_state
|
| 36 |
+
|
| 37 |
+
from cosmos_predict2._src.imaginaire.lazy_config import instantiate
|
| 38 |
+
from cosmos_predict2._src.imaginaire.utils import distributed, log
|
| 39 |
+
from cosmos_predict2._src.imaginaire.visualize.video import save_img_or_video
|
| 40 |
+
from cosmos_predict2._src.predict2.utils.model_loader import load_model_from_checkpoint
|
| 41 |
+
from cosmos_predict2._src.predict2_multiview.scripts.mv_visualize_helper import arrange_video_visualization
|
| 42 |
+
|
| 43 |
+
NUM_CONDITIONAL_FRAMES_KEY = "num_conditional_frames"
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def to_model_input(data_batch, model):
|
| 47 |
+
"""
|
| 48 |
+
Similar to misc.to, but avoid converting uint8 "video" to float
|
| 49 |
+
"""
|
| 50 |
+
for k, v in data_batch.items():
|
| 51 |
+
_v = v
|
| 52 |
+
if isinstance(v, th.Tensor):
|
| 53 |
+
_v = _v.cuda()
|
| 54 |
+
if th.is_floating_point(v):
|
| 55 |
+
_v = _v.to(**model.tensor_kwargs)
|
| 56 |
+
data_batch[k] = _v
|
| 57 |
+
return data_batch
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
class Vid2VidInference:
|
| 61 |
+
"""
|
| 62 |
+
Handles the Vid2Vid inference process, including model loading, data preparation,
|
| 63 |
+
and video generation from an image/video and text prompt. Now supports context parallelism.
|
| 64 |
+
"""
|
| 65 |
+
|
| 66 |
+
def __init__(
|
| 67 |
+
self,
|
| 68 |
+
experiment_name: str,
|
| 69 |
+
ckpt_path: str,
|
| 70 |
+
s3_credential_path: str = "",
|
| 71 |
+
context_parallel_size: int = 1,
|
| 72 |
+
experiment_opts: list[str] = [],
|
| 73 |
+
):
|
| 74 |
+
"""
|
| 75 |
+
Initializes the Vid2VidInference class.
|
| 76 |
+
|
| 77 |
+
Loads the diffusion model and its configuration based on the provided
|
| 78 |
+
experiment name and checkpoint path. Sets up distributed processing if needed.
|
| 79 |
+
|
| 80 |
+
Args:
|
| 81 |
+
experiment_name (str): Name of the experiment configuration.
|
| 82 |
+
ckpt_path (str): Path to the model checkpoint (local or S3).
|
| 83 |
+
s3_credential_path (str): Path to S3 credentials file (if loading from S3).
|
| 84 |
+
context_parallel_size (int): Number of GPUs for context parallelism.
|
| 85 |
+
"""
|
| 86 |
+
self.experiment_name = experiment_name
|
| 87 |
+
self.ckpt_path = ckpt_path
|
| 88 |
+
self.s3_credential_path = s3_credential_path
|
| 89 |
+
self.context_parallel_size = context_parallel_size
|
| 90 |
+
self.process_group = None
|
| 91 |
+
self.experiment_opts = experiment_opts
|
| 92 |
+
|
| 93 |
+
if "RANK" in os.environ:
|
| 94 |
+
self._init_distributed()
|
| 95 |
+
|
| 96 |
+
# Load the model and config
|
| 97 |
+
model, config = load_model_from_checkpoint(
|
| 98 |
+
experiment_name=self.experiment_name,
|
| 99 |
+
s3_checkpoint_dir=self.ckpt_path,
|
| 100 |
+
config_file="cosmos_predict2/_src/predict2_multiview/configs/vid2vid/config.py",
|
| 101 |
+
load_ema_to_reg=True,
|
| 102 |
+
experiment_opts=self.experiment_opts,
|
| 103 |
+
)
|
| 104 |
+
|
| 105 |
+
# Enable context parallel on the model if using context parallelism
|
| 106 |
+
self.rank0 = True
|
| 107 |
+
if self.context_parallel_size > 1:
|
| 108 |
+
model.net.enable_context_parallel(self.process_group)
|
| 109 |
+
self.rank0 = distributed.get_rank() == 0
|
| 110 |
+
|
| 111 |
+
self.model = model
|
| 112 |
+
self.config = config
|
| 113 |
+
self.batch_size = 1
|
| 114 |
+
self.neg_t5_embeddings = None
|
| 115 |
+
|
| 116 |
+
def _init_distributed(self):
|
| 117 |
+
"""Initialize distributed processing for context parallelism."""
|
| 118 |
+
|
| 119 |
+
# Initialize distributed environment
|
| 120 |
+
distributed.init()
|
| 121 |
+
|
| 122 |
+
# Initialize model parallel states
|
| 123 |
+
parallel_state.initialize_model_parallel(
|
| 124 |
+
context_parallel_size=self.context_parallel_size,
|
| 125 |
+
)
|
| 126 |
+
|
| 127 |
+
# Get the process group for context parallel
|
| 128 |
+
self.process_group = parallel_state.get_context_parallel_group()
|
| 129 |
+
|
| 130 |
+
log.info(f"Initialized context parallel with size {self.context_parallel_size}")
|
| 131 |
+
log.info(f"Current rank: {distributed.get_rank()}, World size: {distributed.get_world_size()}")
|
| 132 |
+
|
| 133 |
+
def generate_from_batch(
|
| 134 |
+
self,
|
| 135 |
+
data_batch,
|
| 136 |
+
guidance: int = 7,
|
| 137 |
+
seed: int = 1,
|
| 138 |
+
num_steps: int = 35,
|
| 139 |
+
stack_mode: str = "time",
|
| 140 |
+
use_negative_prompt: bool = True,
|
| 141 |
+
):
|
| 142 |
+
"""Generate video tensor from batch.
|
| 143 |
+
|
| 144 |
+
Returns:
|
| 145 |
+
Tensor with values in the range [0, 1]
|
| 146 |
+
If stack mode is "time", the tensor is of shape (1, 3, v * t, h, w)
|
| 147 |
+
If stack mode is "height", the tensor is of shape (1, 3, t, v * h, w)
|
| 148 |
+
If stack mode is "width", the tensor is of shape (1, 3, t, h, v * w)
|
| 149 |
+
If stack mode is "grid", the tensor is of shape (1, 3, t, 3 * h, 3 * w)
|
| 150 |
+
"""
|
| 151 |
+
data_batch = to_model_input(data_batch, self.model)
|
| 152 |
+
if self.model.config.text_encoder_config is not None and self.model.config.text_encoder_config.compute_online:
|
| 153 |
+
self.model.inplace_compute_text_embeddings_online(data_batch)
|
| 154 |
+
raw_data, x0, condition = self.model.get_data_and_condition(data_batch)
|
| 155 |
+
sample = self.model.generate_samples_from_batch(
|
| 156 |
+
data_batch,
|
| 157 |
+
guidance=guidance,
|
| 158 |
+
# make sure no mismatch and also works for cp
|
| 159 |
+
state_shape=x0.shape[1:],
|
| 160 |
+
n_sample=x0.shape[0],
|
| 161 |
+
seed=seed, # Fixed seed for reproducibility
|
| 162 |
+
num_steps=num_steps,
|
| 163 |
+
is_negative_prompt=use_negative_prompt,
|
| 164 |
+
)
|
| 165 |
+
# (bsz = 1, c = 3, t = n_camera * t, h, w)
|
| 166 |
+
video = ((self.model.decode(sample) + 1.0) / 2.0).clamp(0, 1)
|
| 167 |
+
|
| 168 |
+
# Arrange video according to stack_mode
|
| 169 |
+
video = arrange_video_visualization(video, data_batch, method=stack_mode)
|
| 170 |
+
return video
|
| 171 |
+
|
| 172 |
+
def generate_from_batch_autoregressive(
|
| 173 |
+
self,
|
| 174 |
+
data_batch,
|
| 175 |
+
num_chunks=2,
|
| 176 |
+
chunk_overlap=2,
|
| 177 |
+
guidance: int = 7,
|
| 178 |
+
seed: int = 1,
|
| 179 |
+
num_steps: int = 35,
|
| 180 |
+
stack_mode: str = "time",
|
| 181 |
+
use_negative_prompt: bool = True,
|
| 182 |
+
):
|
| 183 |
+
"""Generate video tensor from batch, with autoregressive mode enabled
|
| 184 |
+
num_chunks: total number of single generation
|
| 185 |
+
chunk_overlap: overlap the
|
| 186 |
+
"""
|
| 187 |
+
data_batch = to_model_input(data_batch, self.model)
|
| 188 |
+
if self.model.config.text_encoder_config is not None and self.model.config.text_encoder_config.compute_online:
|
| 189 |
+
self.model.inplace_compute_text_embeddings_online(data_batch)
|
| 190 |
+
|
| 191 |
+
n_views = len(data_batch["camera_keys_selection"][0])
|
| 192 |
+
num_video_frames_per_view = data_batch["num_video_frames_per_view"][0]
|
| 193 |
+
|
| 194 |
+
generated_chunks = []
|
| 195 |
+
|
| 196 |
+
for i in range(num_chunks):
|
| 197 |
+
log.info(f"start generate chunk {i + 1} / {num_chunks}")
|
| 198 |
+
_, x0, _ = self.model.get_data_and_condition(data_batch)
|
| 199 |
+
sample = self.model.generate_samples_from_batch(
|
| 200 |
+
data_batch,
|
| 201 |
+
guidance=guidance,
|
| 202 |
+
# make sure no mismatch and also works for cp
|
| 203 |
+
state_shape=x0.shape[1:],
|
| 204 |
+
n_sample=x0.shape[0],
|
| 205 |
+
seed=seed, # Fixed seed for reproducibility
|
| 206 |
+
num_steps=num_steps,
|
| 207 |
+
is_negative_prompt=use_negative_prompt,
|
| 208 |
+
)
|
| 209 |
+
# (bsz = 1, c = 3, t = n_camera * t, h, w)
|
| 210 |
+
decoded = self.model.decode(sample)
|
| 211 |
+
chunk_video = ((decoded + 1.0) / 2.0).clamp(0, 1)[0]
|
| 212 |
+
chunk_video = rearrange(chunk_video, "C (V T) H W -> V C T H W", V=n_views)
|
| 213 |
+
if i == 0:
|
| 214 |
+
generated_chunks.append(chunk_video)
|
| 215 |
+
else:
|
| 216 |
+
generated_chunks.append(chunk_video[:, :, chunk_overlap:])
|
| 217 |
+
data_batch["num_conditional_frames"] = chunk_overlap
|
| 218 |
+
data_batch["video"].zero_()
|
| 219 |
+
for v in range(n_views):
|
| 220 |
+
start_idx = num_video_frames_per_view * v
|
| 221 |
+
overlaps = (
|
| 222 |
+
chunk_video[v, :, num_video_frames_per_view - chunk_overlap : num_video_frames_per_view] * 255
|
| 223 |
+
)
|
| 224 |
+
overlaps = overlaps.to(th.uint8).clamp(0, 255)
|
| 225 |
+
data_batch["video"][:, :, start_idx : start_idx + chunk_overlap] = overlaps
|
| 226 |
+
|
| 227 |
+
video = th.cat(generated_chunks, dim=2)
|
| 228 |
+
video = rearrange(video, "V C T H W -> C (V T) H W", V=n_views).unsqueeze(0)
|
| 229 |
+
|
| 230 |
+
# Arrange video according to stack_mode
|
| 231 |
+
video = arrange_video_visualization(video, data_batch, method=stack_mode)
|
| 232 |
+
return video
|
| 233 |
+
|
| 234 |
+
def cleanup(self):
|
| 235 |
+
"""Clean up distributed resources."""
|
| 236 |
+
if "RANK" in os.environ:
|
| 237 |
+
import torch.distributed as dist
|
| 238 |
+
from megatron.core import parallel_state
|
| 239 |
+
|
| 240 |
+
if parallel_state.is_initialized():
|
| 241 |
+
parallel_state.destroy_model_parallel()
|
| 242 |
+
dist.destroy_process_group()
|
| 243 |
+
|
| 244 |
+
|
| 245 |
+
def parse_arguments() -> argparse.Namespace:
|
| 246 |
+
"""Parses command-line arguments for the Vid2Vid inference script."""
|
| 247 |
+
parser = argparse.ArgumentParser(description="Image2World/Video2World inference script")
|
| 248 |
+
parser.add_argument("--experiment", type=str, required=True, help="Experiment config")
|
| 249 |
+
parser.add_argument(
|
| 250 |
+
"--ckpt_path",
|
| 251 |
+
type=str,
|
| 252 |
+
default="",
|
| 253 |
+
help="Path to the checkpoint. If not provided, will use the one specify in the config",
|
| 254 |
+
)
|
| 255 |
+
parser.add_argument("--s3_cred", type=str, default="credentials/s3_checkpoint.secret")
|
| 256 |
+
parser.add_argument(
|
| 257 |
+
"--context_parallel_size",
|
| 258 |
+
type=int,
|
| 259 |
+
default=1,
|
| 260 |
+
help="Context parallel size (number of GPUs to split context over). Set to 8 for 8 GPUs",
|
| 261 |
+
)
|
| 262 |
+
# generation
|
| 263 |
+
parser.add_argument("--guidance", type=int, default=7, help="Guidance value")
|
| 264 |
+
parser.add_argument("--fps", type=int, default=10, help="Frames per second")
|
| 265 |
+
parser.add_argument("--seed", type=int, default=1, help="Guidance value")
|
| 266 |
+
parser.add_argument("--num_conditional_frames", type=int, default=1, help="Number of conditional frames")
|
| 267 |
+
# input
|
| 268 |
+
parser.add_argument(
|
| 269 |
+
"--input_is_train_data",
|
| 270 |
+
action="store_true",
|
| 271 |
+
help="Inference on the training data, the input_root will be ignored if this is set",
|
| 272 |
+
)
|
| 273 |
+
parser.add_argument("--run_mads_verification", action="store_true", help="Run MADS verification")
|
| 274 |
+
parser.add_argument(
|
| 275 |
+
"--mads_verification_prompt",
|
| 276 |
+
type=str,
|
| 277 |
+
default='The video opens with a view from inside a vehicle, positioned at an intersection under a clear blue sky. The camera angle is from the dashboard, offering a first-person perspective of the road ahead. The intersection is marked by multiple traffic lights and street signs, including one that reads "E Garden Blvd." A white van with "TM Stuckateur" branding is seen driving through the intersection, heading in the same direction as the viewer\'s vehicle. Other cars are also present, moving smoothly along the multi-lane road. As the vehicle starts to move forward, the camera pans slightly to the right, revealing more of the surroundings. The road is lined with trees on both sides, providing a natural canopy that filters the sunlight. The trees are lush and green, indicating it might be spring or summer. On the left side of the road, there is a large building with a sign that reads "GROCERY OUTLET," suggesting the presence of a retail store nearby. Further down the road, additional buildings and residential structures can be seen, hinting at a suburban or semi-urban area. The sun is bright and high in the sky, casting long shadows across the road. The light creates a warm, inviting atmosphere, enhancing the clarity of the scene. The road itself is well-maintained, with clear lane markings and directional arrows painted on the asphalt. Overhead, power lines run parallel to the road, supported by poles that also hold traffic lights and street lamps. As the vehicle continues its journey, the camera maintains a steady focus on the road ahead, capturing the smooth flow of traffic and the serene environment. The absence of heavy traffic or congestion adds to the tranquil mood of the scene. The overall ambiance is one of calm and order, with the interplay of natural and man-made elements creating a harmonious urban landscape. The gentle curve of the road and the soft glow of the setting sun add a sense of peacefulness to the drive, making the viewer feel as though they are part of this quiet, picturesque neighborhood.',
|
| 278 |
+
)
|
| 279 |
+
parser.add_argument(
|
| 280 |
+
"--stack_mode",
|
| 281 |
+
type=str,
|
| 282 |
+
default="time",
|
| 283 |
+
choices=["height", "width", "time", "grid"],
|
| 284 |
+
help="Video stacking mode for visualization. grid will create a 3x3 grid of views.",
|
| 285 |
+
)
|
| 286 |
+
parser.add_argument("--input_root", type=str, default="assets/image2world", help="Input root")
|
| 287 |
+
parser.add_argument("--save_root", type=str, default="results/image2world", help="Save root")
|
| 288 |
+
parser.add_argument("--max_samples", type=int, default=20, help="Maximum number of samples to generate")
|
| 289 |
+
return parser.parse_args()
|
| 290 |
+
|
| 291 |
+
|
| 292 |
+
if __name__ == "__main__":
|
| 293 |
+
os.environ["NVTE_FUSED_ATTN"] = "0"
|
| 294 |
+
th.backends.cudnn.benchmark = False
|
| 295 |
+
th.backends.cudnn.deterministic = True
|
| 296 |
+
th.enable_grad(False)
|
| 297 |
+
args = parse_arguments()
|
| 298 |
+
# Initialize the inference handler with context parallel support
|
| 299 |
+
vid2vid_cli = Vid2VidInference(
|
| 300 |
+
args.experiment, args.ckpt_path, args.s3_cred, context_parallel_size=args.context_parallel_size
|
| 301 |
+
)
|
| 302 |
+
mem_bytes = th.cuda.memory_allocated(device=th.device("cuda" if th.cuda.is_available() else "cpu"))
|
| 303 |
+
log.info(f"GPU memory usage after model dcp.load: {mem_bytes / (1024**3):.2f} GB")
|
| 304 |
+
|
| 305 |
+
# Only process files on rank 0 if using distributed processing
|
| 306 |
+
rank0 = True
|
| 307 |
+
if args.context_parallel_size > 1:
|
| 308 |
+
rank0 = distributed.get_rank() == 0
|
| 309 |
+
|
| 310 |
+
os.makedirs(args.save_root, exist_ok=True)
|
| 311 |
+
if args.input_is_train_data:
|
| 312 |
+
dataloader = instantiate(vid2vid_cli.config.dataloader_train)
|
| 313 |
+
for i, batch in enumerate(dataloader):
|
| 314 |
+
if i >= args.max_samples:
|
| 315 |
+
break
|
| 316 |
+
if args.run_mads_verification:
|
| 317 |
+
assert args.num_conditional_frames == 0, "MADS verification only supports 0 conditional frame"
|
| 318 |
+
log.warning(f"Running MADS verification with prompt: {args.mads_verification_prompt[0:100]}...")
|
| 319 |
+
batch["ai_caption"] = [args.mads_verification_prompt]
|
| 320 |
+
batch[NUM_CONDITIONAL_FRAMES_KEY] = args.num_conditional_frames
|
| 321 |
+
video = vid2vid_cli.generate_from_batch(
|
| 322 |
+
batch, guidance=args.guidance, seed=args.seed, stack_mode=args.stack_mode
|
| 323 |
+
)
|
| 324 |
+
if rank0:
|
| 325 |
+
save_name = f"mads_verification_{i}" if args.run_mads_verification else f"infer_from_train_{i}"
|
| 326 |
+
save_img_or_video(video[0], f"{args.save_root}/{save_name}", fps=args.fps)
|
| 327 |
+
if args.run_mads_verification:
|
| 328 |
+
break
|
| 329 |
+
else:
|
| 330 |
+
raise NotImplementedError("Not implemented")
|
cosmos_predict2/_src/predict2_multiview/scripts/inference_cli.py
ADDED
|
@@ -0,0 +1,582 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
"""
|
| 16 |
+
Inference script for constructing data_batch from videos and captions, then running predict2_multiview model.
|
| 17 |
+
|
| 18 |
+
Expected directory structure:
|
| 19 |
+
```
|
| 20 |
+
input_root/
|
| 21 |
+
├── videos/ # Video folder
|
| 22 |
+
│ ├── ftheta_camera_front_wide_120fov/ # or camera_front_wide_120fov/
|
| 23 |
+
│ │ ├── video_id_1.mp4
|
| 24 |
+
│ │ ├── video_id_2.mp4
|
| 25 |
+
│ │ └── ...
|
| 26 |
+
│ ├── ftheta_camera_cross_right_120fov/ # or camera_cross_right_120fov/
|
| 27 |
+
│ │ ├── video_id_1.mp4
|
| 28 |
+
│ │ ├── video_id_2.mp4
|
| 29 |
+
│ │ └── ...
|
| 30 |
+
│ ├── ftheta_camera_rear_right_70fov/ # or camera_rear_right_70fov/
|
| 31 |
+
│ ├── ftheta_camera_rear_tele_30fov/ # or camera_rear_tele_30fov/
|
| 32 |
+
│ ├── ftheta_camera_rear_left_70fov/ # or camera_rear_left_70fov/
|
| 33 |
+
│ ├── ftheta_camera_cross_left_120fov/ # or camera_cross_left_120fov/
|
| 34 |
+
│ └── ftheta_camera_front_tele_30fov/ # or camera_front_tele_30fov/
|
| 35 |
+
│
|
| 36 |
+
└── captions/ # Caption folder (optional, uses default prompt if not present)
|
| 37 |
+
├── ftheta_camera_front_wide_120fov/ # or camera_front_wide_120fov/
|
| 38 |
+
│ ├── video_id_1.txt
|
| 39 |
+
│ ├── video_id_2.txt
|
| 40 |
+
│ └── ...
|
| 41 |
+
├── ftheta_camera_cross_right_120fov/ # or camera_cross_right_120fov/
|
| 42 |
+
│ ├── video_id_1.txt
|
| 43 |
+
│ ├── video_id_2.txt
|
| 44 |
+
│ └── ...
|
| 45 |
+
├── ftheta_camera_rear_right_70fov/ # or camera_rear_right_70fov/
|
| 46 |
+
├── ftheta_camera_rear_tele_30fov/ # or camera_rear_tele_30fov/
|
| 47 |
+
├── ftheta_camera_rear_left_70fov/ # or camera_rear_left_70fov/
|
| 48 |
+
├── ftheta_camera_cross_left_120fov/ # or camera_cross_left_120fov/
|
| 49 |
+
└── ftheta_camera_front_tele_30fov/ # or camera_front_tele_30fov/
|
| 50 |
+
|
| 51 |
+
Notes:
|
| 52 |
+
- The videos/ folder is required (unless num_conditional_frames=0, which uses dummy all-zero videos)
|
| 53 |
+
- The captions/ folder is optional; if not present, a preset default driving scene description is used
|
| 54 |
+
- Each camera's subfolder name supports two formats: "ftheta_{camera_name}" or "{camera_name}"
|
| 55 |
+
- video_id must be consistent across all camera folders
|
| 56 |
+
- All 7 camera views must have corresponding subfolders and files
|
| 57 |
+
|
| 58 |
+
Camera view to View Index mapping:
|
| 59 |
+
- camera_front_wide_120fov: 0
|
| 60 |
+
- camera_cross_right_120fov: 1
|
| 61 |
+
- camera_rear_right_70fov: 2
|
| 62 |
+
- camera_rear_tele_30fov: 3
|
| 63 |
+
- camera_rear_left_70fov: 4
|
| 64 |
+
- camera_cross_left_120fov: 5
|
| 65 |
+
- camera_front_tele_30fov: 6
|
| 66 |
+
```
|
| 67 |
+
|
| 68 |
+
Usage:
|
| 69 |
+
```bash
|
| 70 |
+
EXP=predict2p5_2b_mv_7train7_res720p_fps10_t24_frombase2p5avfinetune_alpamayo_only_allcaption_uniform_nofps
|
| 71 |
+
ckpt_path=s3://bucket/cosmos_predict2_multiview/cosmos2_mv2/predict2p5_2b_mv_7train7_res720p_fps10_t24_frombase2p5avfinetune_alpamayo_only_allcaption_uniform_nofps_resume1-0/checkpoints/iter_000024500/
|
| 72 |
+
|
| 73 |
+
PYTHONPATH=. torchrun --nproc_per_node=8 --master_port=12341 -m cosmos_predict2._src.predict2_multiview.scripts.inference_cli \
|
| 74 |
+
--experiment ${EXP} \
|
| 75 |
+
--ckpt_path ${ckpt_path} \
|
| 76 |
+
--context_parallel_size 8 \
|
| 77 |
+
--input_root /project/cosmos/yiflu/project_official_i4/condition_assets/multiview-inference-assets-1203 \
|
| 78 |
+
--num_conditional_frames 0 \
|
| 79 |
+
--guidance 5 \
|
| 80 |
+
--fps 10 \
|
| 81 |
+
--save_root results/predict2_multiview_av_grid/ \
|
| 82 |
+
--max_samples 5 --stack_mode grid
|
| 83 |
+
```
|
| 84 |
+
|
| 85 |
+
```bash
|
| 86 |
+
EXP=predict2p5_2b_mv_7train7_res480p_fps15_t24_alpamayo_only_allcaption_uniform_nofps
|
| 87 |
+
ckpt_path=s3://bucket/cosmos_predict2_multiview/cosmos2p5_mv/predict2p5_2b_mv_7train7_res480p_fps15_t24_alpamayo_only_allcaption_uniform_nofps-0/checkpoints/iter_000020000/
|
| 88 |
+
|
| 89 |
+
PYTHONPATH=. torchrun --nproc_per_node=8 --master_port=12341 -m cosmos_predict2._src.predict2_multiview.scripts.inference_cli \
|
| 90 |
+
--experiment ${EXP} \
|
| 91 |
+
--ckpt_path ${ckpt_path} \
|
| 92 |
+
--context_parallel_size 8 \
|
| 93 |
+
--input_root /project/cosmos/yiflu/project_official_i4/condition_assets/multiview-inference-assets-1203 \
|
| 94 |
+
--num_conditional_frames 1 \
|
| 95 |
+
--guidance 5 \
|
| 96 |
+
--fps 15 \
|
| 97 |
+
--save_root results/predict2_multiview_av_480p_i2v_grid/ \
|
| 98 |
+
--max_samples 5 --stack_mode grid \
|
| 99 |
+
--target_height 480 --target_width 832 \
|
| 100 |
+
model.config.net.init_cross_view_attn_weight_from=null
|
| 101 |
+
```
|
| 102 |
+
"""
|
| 103 |
+
|
| 104 |
+
import argparse
|
| 105 |
+
import os
|
| 106 |
+
from pathlib import Path
|
| 107 |
+
|
| 108 |
+
import torch as th
|
| 109 |
+
import torchvision
|
| 110 |
+
|
| 111 |
+
from cosmos_predict2._src.imaginaire.utils import distributed, log
|
| 112 |
+
from cosmos_predict2._src.imaginaire.utils.easy_io import easy_io
|
| 113 |
+
from cosmos_predict2._src.imaginaire.visualize.video import save_img_or_video
|
| 114 |
+
from cosmos_predict2._src.predict2_multiview.scripts.inference import Vid2VidInference
|
| 115 |
+
|
| 116 |
+
NUM_CONDITIONAL_FRAMES_KEY = "num_conditional_frames"
|
| 117 |
+
|
| 118 |
+
# Camera name to view index mapping
|
| 119 |
+
CAMERA_TO_VIEW_INDEX = {
|
| 120 |
+
"camera_front_wide_120fov": 0,
|
| 121 |
+
"camera_cross_right_120fov": 1,
|
| 122 |
+
"camera_rear_right_70fov": 2,
|
| 123 |
+
"camera_rear_tele_30fov": 3,
|
| 124 |
+
"camera_rear_left_70fov": 4,
|
| 125 |
+
"camera_cross_left_120fov": 5,
|
| 126 |
+
"camera_front_tele_30fov": 6,
|
| 127 |
+
}
|
| 128 |
+
|
| 129 |
+
DEFAULT_CAMERA_ORDER = list(CAMERA_TO_VIEW_INDEX.keys())
|
| 130 |
+
|
| 131 |
+
# Camera-specific caption prefixes describing camera position and orientation
|
| 132 |
+
CAMERA_TO_CAPTION_PREFIX = {
|
| 133 |
+
"camera_front_wide_120fov": "The video is captured from a camera mounted on a car. The camera is facing forward.",
|
| 134 |
+
"camera_cross_right_120fov": "The video is captured from a camera mounted on a car. The camera is facing to the right.",
|
| 135 |
+
"camera_rear_right_70fov": "The video is captured from a camera mounted on a car. The camera is facing the rear right side.",
|
| 136 |
+
"camera_rear_tele_30fov": "The video is captured from a camera mounted on a car. The camera is facing backwards.",
|
| 137 |
+
"camera_rear_left_70fov": "The video is captured from a camera mounted on a car. The camera is facing the rear left side.",
|
| 138 |
+
"camera_cross_left_120fov": "The video is captured from a camera mounted on a car. The camera is facing to the left.",
|
| 139 |
+
"camera_front_tele_30fov": "The video is captured from a telephoto camera mounted on a car. The camera is facing forward.",
|
| 140 |
+
}
|
| 141 |
+
|
| 142 |
+
DEFAULT_DRIVING_SCENE_PROMPT = """
|
| 143 |
+
A clear daytime driving scene on an open road. The weather is sunny with bright natural lighting and good visibility.
|
| 144 |
+
The sky is partly cloudy with scattered white clouds. The road surface is dry and well-maintained.
|
| 145 |
+
The overall atmosphere is calm and peaceful with moderate traffic conditions. The lighting creates clear
|
| 146 |
+
shadows and provides excellent contrast for safe navigation."""
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
def load_video(video_path: str, target_frames: int = 93, target_size: tuple[int, int] = (720, 1280)) -> th.Tensor:
|
| 150 |
+
"""
|
| 151 |
+
Load video and process it to target size and frame count.
|
| 152 |
+
|
| 153 |
+
Args:
|
| 154 |
+
video_path: Path to video file
|
| 155 |
+
target_frames: Target number of frames
|
| 156 |
+
target_size: Target resolution (H, W)
|
| 157 |
+
|
| 158 |
+
Returns:
|
| 159 |
+
Video tensor with shape (C, T, H, W), dtype uint8
|
| 160 |
+
"""
|
| 161 |
+
try:
|
| 162 |
+
# Load video using easy_io
|
| 163 |
+
video_frames, video_metadata = easy_io.load(video_path) # Returns (T, H, W, C) numpy array
|
| 164 |
+
except Exception as e:
|
| 165 |
+
raise ValueError(f"Failed to load video {video_path}: {e}")
|
| 166 |
+
|
| 167 |
+
# Convert to tensor: (T, H, W, C) -> (C, T, H, W)
|
| 168 |
+
video_tensor = th.from_numpy(video_frames).float() / 255.0
|
| 169 |
+
video_tensor = video_tensor.permute(3, 0, 1, 2) # (T, H, W, C) -> (C, T, H, W)
|
| 170 |
+
|
| 171 |
+
C, T, H, W = video_tensor.shape
|
| 172 |
+
|
| 173 |
+
# Adjust frame count: if video is too long, take first target_frames; if too short, pad with last frame
|
| 174 |
+
if T > target_frames:
|
| 175 |
+
video_tensor = video_tensor[:, :target_frames, :, :]
|
| 176 |
+
elif T < target_frames:
|
| 177 |
+
# Pad with last frame
|
| 178 |
+
last_frame = video_tensor[:, -1:, :, :]
|
| 179 |
+
padding_frames = target_frames - T
|
| 180 |
+
last_frame_repeated = last_frame.repeat(1, padding_frames, 1, 1)
|
| 181 |
+
video_tensor = th.cat([video_tensor, last_frame_repeated], dim=1)
|
| 182 |
+
|
| 183 |
+
# Convert to uint8: (C, T, H, W) -> (T, C, H, W)
|
| 184 |
+
video_tensor = video_tensor.permute(1, 0, 2, 3)
|
| 185 |
+
video_tensor = (video_tensor * 255.0).to(th.uint8)
|
| 186 |
+
|
| 187 |
+
# Adjust resolution
|
| 188 |
+
target_h, target_w = target_size
|
| 189 |
+
if H != target_h or W != target_w:
|
| 190 |
+
# Use resize and center crop
|
| 191 |
+
video_tensor = resize_and_crop(video_tensor, target_size)
|
| 192 |
+
|
| 193 |
+
# Convert back to (C, T, H, W)
|
| 194 |
+
video_tensor = video_tensor.permute(1, 0, 2, 3)
|
| 195 |
+
|
| 196 |
+
return video_tensor
|
| 197 |
+
|
| 198 |
+
|
| 199 |
+
def resize_and_crop(video: th.Tensor, target_size: tuple[int, int]) -> th.Tensor:
|
| 200 |
+
"""
|
| 201 |
+
Resize video and center crop.
|
| 202 |
+
|
| 203 |
+
Args:
|
| 204 |
+
video: Input video with shape (T, C, H, W)
|
| 205 |
+
target_size: Target resolution (H, W)
|
| 206 |
+
|
| 207 |
+
Returns:
|
| 208 |
+
Resized video with shape (T, C, target_H, target_W)
|
| 209 |
+
"""
|
| 210 |
+
orig_h, orig_w = video.shape[2], video.shape[3]
|
| 211 |
+
target_h, target_w = target_size
|
| 212 |
+
|
| 213 |
+
# Calculate scaling ratio to match the smaller dimension to target
|
| 214 |
+
scaling_ratio = max((target_w / orig_w), (target_h / orig_h))
|
| 215 |
+
resizing_shape = (int(scaling_ratio * orig_h), int(scaling_ratio * orig_w))
|
| 216 |
+
|
| 217 |
+
video_resized = torchvision.transforms.functional.resize(video, resizing_shape)
|
| 218 |
+
video_cropped = torchvision.transforms.functional.center_crop(video_resized, target_size)
|
| 219 |
+
|
| 220 |
+
return video_cropped
|
| 221 |
+
|
| 222 |
+
|
| 223 |
+
def load_multiview_videos(
|
| 224 |
+
input_root: Path,
|
| 225 |
+
video_id: str,
|
| 226 |
+
camera_order: list[str],
|
| 227 |
+
target_frames: int = 93,
|
| 228 |
+
target_size: tuple[int, int] = (720, 1280),
|
| 229 |
+
) -> th.Tensor:
|
| 230 |
+
"""
|
| 231 |
+
Load multi-view videos.
|
| 232 |
+
|
| 233 |
+
Args:
|
| 234 |
+
input_root: Input root directory
|
| 235 |
+
video_id: Video ID (filename without extension)
|
| 236 |
+
camera_order: List of camera names in order
|
| 237 |
+
target_frames: Target number of frames per view
|
| 238 |
+
target_size: Target resolution (H, W)
|
| 239 |
+
|
| 240 |
+
Returns:
|
| 241 |
+
Multi-view video tensor with shape (C, V*T, H, W)
|
| 242 |
+
"""
|
| 243 |
+
videos_dir = input_root / "videos"
|
| 244 |
+
video_tensors = []
|
| 245 |
+
|
| 246 |
+
for camera in camera_order:
|
| 247 |
+
if (videos_dir / f"ftheta_{camera}").exists():
|
| 248 |
+
sub_dir = f"ftheta_{camera}"
|
| 249 |
+
elif (videos_dir / camera).exists():
|
| 250 |
+
sub_dir = camera
|
| 251 |
+
else:
|
| 252 |
+
raise FileNotFoundError(f"Folder not found: {videos_dir / f'ftheta_{camera}'} or {videos_dir / camera}")
|
| 253 |
+
|
| 254 |
+
video_path = videos_dir / sub_dir / f"{video_id}.mp4"
|
| 255 |
+
|
| 256 |
+
if not video_path.exists():
|
| 257 |
+
raise FileNotFoundError(f"Video not found: {video_path}")
|
| 258 |
+
|
| 259 |
+
# Load single view video: (C, T, H, W)
|
| 260 |
+
video_tensor = load_video(str(video_path), target_frames, target_size)
|
| 261 |
+
video_tensors.append(video_tensor)
|
| 262 |
+
|
| 263 |
+
# Concatenate all views: (C, V*T, H, W)
|
| 264 |
+
multiview_video = th.cat(video_tensors, dim=1)
|
| 265 |
+
|
| 266 |
+
return multiview_video
|
| 267 |
+
|
| 268 |
+
|
| 269 |
+
def load_multiview_captions(
|
| 270 |
+
input_root: Path, video_id: str, camera_order: list[str], add_camera_prefix: bool = True
|
| 271 |
+
) -> list[str]:
|
| 272 |
+
"""
|
| 273 |
+
Load multi-view captions. Uses default prompt if captions directory does not exist.
|
| 274 |
+
|
| 275 |
+
Args:
|
| 276 |
+
input_root: Input root directory
|
| 277 |
+
video_id: Video ID (filename without extension)
|
| 278 |
+
camera_order: List of camera names in order
|
| 279 |
+
add_camera_prefix: Whether to add camera-specific prefix to captions
|
| 280 |
+
|
| 281 |
+
Returns:
|
| 282 |
+
List of captions, one per view
|
| 283 |
+
"""
|
| 284 |
+
captions_dir = input_root / "captions"
|
| 285 |
+
|
| 286 |
+
# If captions directory does not exist, use default prompt
|
| 287 |
+
if not captions_dir.exists():
|
| 288 |
+
log.warning(
|
| 289 |
+
f"Captions directory not found: {captions_dir}. Using default driving scene prompt for all cameras."
|
| 290 |
+
)
|
| 291 |
+
return [DEFAULT_DRIVING_SCENE_PROMPT] * len(camera_order)
|
| 292 |
+
|
| 293 |
+
captions = []
|
| 294 |
+
|
| 295 |
+
for camera in camera_order:
|
| 296 |
+
if (captions_dir / f"ftheta_{camera}").exists():
|
| 297 |
+
sub_dir = f"ftheta_{camera}"
|
| 298 |
+
elif (captions_dir / camera).exists():
|
| 299 |
+
sub_dir = camera
|
| 300 |
+
else:
|
| 301 |
+
raise FileNotFoundError(f"Folder not found: {captions_dir / f'ftheta_{camera}'} or {captions_dir / camera}")
|
| 302 |
+
|
| 303 |
+
caption_filename = f"{sub_dir}/{video_id}.txt"
|
| 304 |
+
caption_path = captions_dir / caption_filename
|
| 305 |
+
|
| 306 |
+
if not caption_path.exists():
|
| 307 |
+
raise FileNotFoundError(f"Caption file not found: {caption_path}")
|
| 308 |
+
|
| 309 |
+
with open(caption_path, "r", encoding="utf-8") as f:
|
| 310 |
+
caption = f.read().strip()
|
| 311 |
+
|
| 312 |
+
# Add camera-specific prefix if enabled
|
| 313 |
+
if add_camera_prefix and camera in CAMERA_TO_CAPTION_PREFIX:
|
| 314 |
+
caption = f"{CAMERA_TO_CAPTION_PREFIX[camera]} {caption}"
|
| 315 |
+
|
| 316 |
+
captions.append(caption)
|
| 317 |
+
|
| 318 |
+
return captions
|
| 319 |
+
|
| 320 |
+
|
| 321 |
+
def construct_data_batch(
|
| 322 |
+
multiview_video: th.Tensor,
|
| 323 |
+
captions: list[str],
|
| 324 |
+
camera_order: list[str],
|
| 325 |
+
num_conditional_frames: int = 0,
|
| 326 |
+
fps: float = 15.0,
|
| 327 |
+
target_frames_per_view: int = 93,
|
| 328 |
+
) -> dict:
|
| 329 |
+
"""
|
| 330 |
+
Construct data_batch for model inference.
|
| 331 |
+
|
| 332 |
+
Args:
|
| 333 |
+
multiview_video: Multi-view video tensor with shape (C, V*T, H, W)
|
| 334 |
+
captions: List of captions
|
| 335 |
+
camera_order: List of camera names in order
|
| 336 |
+
num_conditional_frames: Number of conditional frames
|
| 337 |
+
fps: Frames per second
|
| 338 |
+
target_frames_per_view: Number of frames per view
|
| 339 |
+
|
| 340 |
+
Returns:
|
| 341 |
+
data_batch dictionary
|
| 342 |
+
"""
|
| 343 |
+
C, VT, H, W = multiview_video.shape
|
| 344 |
+
n_views = len(camera_order)
|
| 345 |
+
T = VT // n_views
|
| 346 |
+
|
| 347 |
+
# Add batch dimension: (C, V*T, H, W) -> (1, C, V*T, H, W)
|
| 348 |
+
multiview_video = multiview_video.unsqueeze(0)
|
| 349 |
+
|
| 350 |
+
# Construct correct view_indices based on camera order
|
| 351 |
+
# Each view's T frames all use that view's corresponding view index
|
| 352 |
+
view_indices_list = []
|
| 353 |
+
for camera in camera_order:
|
| 354 |
+
view_idx = CAMERA_TO_VIEW_INDEX[camera]
|
| 355 |
+
view_indices_list.extend([view_idx] * T)
|
| 356 |
+
view_indices = th.tensor(view_indices_list, dtype=th.int64).unsqueeze(0) # (1, V*T)
|
| 357 |
+
|
| 358 |
+
# Construct view_indices_selection: view indices of cameras in camera_order
|
| 359 |
+
view_indices_selection = th.tensor(
|
| 360 |
+
[CAMERA_TO_VIEW_INDEX[camera] for camera in camera_order], dtype=th.int64
|
| 361 |
+
).unsqueeze(0) # (1, n_views)
|
| 362 |
+
|
| 363 |
+
# Find position of front_wide_120fov in camera_order as ref_cam_view_idx_sample_position
|
| 364 |
+
ref_cam_position = (
|
| 365 |
+
camera_order.index("camera_front_wide_120fov") if "camera_front_wide_120fov" in camera_order else 0
|
| 366 |
+
)
|
| 367 |
+
|
| 368 |
+
# Construct data_batch
|
| 369 |
+
data_batch = {
|
| 370 |
+
"video": multiview_video,
|
| 371 |
+
"ai_caption": [captions],
|
| 372 |
+
"view_indices": view_indices, # (1, V*T), using correct view index
|
| 373 |
+
"fps": th.tensor([fps], dtype=th.float64),
|
| 374 |
+
"chunk_index": th.tensor([0], dtype=th.int64),
|
| 375 |
+
"frame_indices": th.arange(target_frames_per_view).unsqueeze(0), # (1, T)
|
| 376 |
+
"num_video_frames_per_view": th.tensor([target_frames_per_view], dtype=th.int64),
|
| 377 |
+
"view_indices_selection": view_indices_selection, # (1, n_views), using correct view index
|
| 378 |
+
"camera_keys_selection": [camera_order],
|
| 379 |
+
"sample_n_views": th.tensor([n_views], dtype=th.int64),
|
| 380 |
+
"padding_mask": th.zeros(1, 1, H, W, dtype=th.float32),
|
| 381 |
+
"ref_cam_view_idx_sample_position": th.tensor([ref_cam_position], dtype=th.int64),
|
| 382 |
+
"front_cam_view_idx_sample_position": [None],
|
| 383 |
+
"original_hw": th.tensor([[[H, W]] * n_views], dtype=th.int64), # (1, n_views, 2)
|
| 384 |
+
NUM_CONDITIONAL_FRAMES_KEY: num_conditional_frames,
|
| 385 |
+
}
|
| 386 |
+
|
| 387 |
+
return data_batch
|
| 388 |
+
|
| 389 |
+
|
| 390 |
+
def parse_arguments() -> argparse.Namespace:
|
| 391 |
+
"""Parse command line arguments."""
|
| 392 |
+
parser = argparse.ArgumentParser(description="Predict2 Multiview inference from videos and captions")
|
| 393 |
+
parser.add_argument("--experiment", type=str, required=True, help="Experiment config")
|
| 394 |
+
parser.add_argument(
|
| 395 |
+
"--ckpt_path",
|
| 396 |
+
type=str,
|
| 397 |
+
default="",
|
| 398 |
+
help="Path to the checkpoint. If not provided, will use the one specify in the config",
|
| 399 |
+
)
|
| 400 |
+
parser.add_argument("--s3_cred", type=str, default="credentials/s3_checkpoint.secret")
|
| 401 |
+
parser.add_argument(
|
| 402 |
+
"--context_parallel_size",
|
| 403 |
+
type=int,
|
| 404 |
+
default=1,
|
| 405 |
+
help="Context parallel size (number of GPUs to split context over). Set to 8 for 8 GPUs",
|
| 406 |
+
)
|
| 407 |
+
# Generation parameters
|
| 408 |
+
parser.add_argument("--guidance", type=int, default=5, help="Guidance value")
|
| 409 |
+
parser.add_argument("--fps", type=int, default=15, help="Output video FPS")
|
| 410 |
+
parser.add_argument("--seed", type=int, default=1, help="Random seed")
|
| 411 |
+
parser.add_argument("--num_conditional_frames", type=int, default=0, help="Number of conditional frames")
|
| 412 |
+
parser.add_argument("--num_steps", type=int, default=35, help="Number of diffusion steps")
|
| 413 |
+
# Input/output
|
| 414 |
+
parser.add_argument(
|
| 415 |
+
"--input_root",
|
| 416 |
+
type=str,
|
| 417 |
+
required=True,
|
| 418 |
+
help="Input root directory containing videos/ and captions/ subdirectories",
|
| 419 |
+
)
|
| 420 |
+
parser.add_argument("--save_root", type=str, default="results/predict2_multiview_av/", help="Save root")
|
| 421 |
+
parser.add_argument("--max_samples", type=int, default=5, help="Maximum number of samples to generate")
|
| 422 |
+
parser.add_argument(
|
| 423 |
+
"--stack_mode",
|
| 424 |
+
type=str,
|
| 425 |
+
default="time",
|
| 426 |
+
choices=["height", "width", "time", "grid"],
|
| 427 |
+
help="Video stacking mode for visualization. grid will create a 3x3 grid of views.",
|
| 428 |
+
)
|
| 429 |
+
# Video parameters
|
| 430 |
+
parser.add_argument("--target_frames", type=int, default=93, help="Target number of frames per view")
|
| 431 |
+
parser.add_argument("--target_height", type=int, default=720, help="Target video height")
|
| 432 |
+
parser.add_argument("--target_width", type=int, default=1280, help="Target video width")
|
| 433 |
+
# Caption parameters
|
| 434 |
+
parser.add_argument(
|
| 435 |
+
"--add_camera_prefix",
|
| 436 |
+
action="store_true",
|
| 437 |
+
default=True,
|
| 438 |
+
help="Add camera-specific position/orientation prefix to captions",
|
| 439 |
+
)
|
| 440 |
+
parser.add_argument(
|
| 441 |
+
"--no_camera_prefix",
|
| 442 |
+
action="store_false",
|
| 443 |
+
dest="add_camera_prefix",
|
| 444 |
+
help="Do not add camera-specific prefix to captions",
|
| 445 |
+
)
|
| 446 |
+
# Experiment options
|
| 447 |
+
parser.add_argument(
|
| 448 |
+
"opts",
|
| 449 |
+
help="""
|
| 450 |
+
Modify config options at the end of the command. For Yacs configs, use
|
| 451 |
+
space-separated "PATH.KEY VALUE" pairs.
|
| 452 |
+
For python-based LazyConfig, use "path.key=value".
|
| 453 |
+
""".strip(),
|
| 454 |
+
default=None,
|
| 455 |
+
nargs=argparse.REMAINDER,
|
| 456 |
+
)
|
| 457 |
+
|
| 458 |
+
return parser.parse_args()
|
| 459 |
+
|
| 460 |
+
|
| 461 |
+
def main():
|
| 462 |
+
os.environ["NVTE_FUSED_ATTN"] = "0"
|
| 463 |
+
th.backends.cudnn.benchmark = False
|
| 464 |
+
th.backends.cudnn.deterministic = True
|
| 465 |
+
th.enable_grad(False)
|
| 466 |
+
|
| 467 |
+
args = parse_arguments()
|
| 468 |
+
|
| 469 |
+
# Initialize inference handler
|
| 470 |
+
vid2vid_cli = Vid2VidInference(
|
| 471 |
+
args.experiment,
|
| 472 |
+
args.ckpt_path,
|
| 473 |
+
args.s3_cred,
|
| 474 |
+
context_parallel_size=args.context_parallel_size,
|
| 475 |
+
experiment_opts=args.opts,
|
| 476 |
+
)
|
| 477 |
+
mem_bytes = th.cuda.memory_allocated(device=th.device("cuda" if th.cuda.is_available() else "cpu"))
|
| 478 |
+
log.info(f"GPU memory usage after model dcp.load: {mem_bytes / (1024**3):.2f} GB")
|
| 479 |
+
|
| 480 |
+
# Only process files on rank 0
|
| 481 |
+
rank0 = True
|
| 482 |
+
if args.context_parallel_size > 1:
|
| 483 |
+
rank0 = distributed.get_rank() == 0
|
| 484 |
+
|
| 485 |
+
# Create output directory
|
| 486 |
+
os.makedirs(args.save_root, exist_ok=True)
|
| 487 |
+
|
| 488 |
+
input_root = Path(args.input_root)
|
| 489 |
+
videos_dir = input_root / "videos"
|
| 490 |
+
|
| 491 |
+
# Get all video IDs (from first camera directory)
|
| 492 |
+
if (videos_dir / f"ftheta_{DEFAULT_CAMERA_ORDER[0]}").exists():
|
| 493 |
+
first_camera_dir = videos_dir / f"ftheta_{DEFAULT_CAMERA_ORDER[0]}"
|
| 494 |
+
else:
|
| 495 |
+
first_camera_dir = videos_dir / DEFAULT_CAMERA_ORDER[0]
|
| 496 |
+
|
| 497 |
+
video_files = sorted(first_camera_dir.glob("*.mp4"))
|
| 498 |
+
video_ids = [f.stem for f in video_files[: args.max_samples]]
|
| 499 |
+
|
| 500 |
+
log.info(f"Found {len(video_ids)} video IDs, processing {min(len(video_ids), args.max_samples)} samples")
|
| 501 |
+
|
| 502 |
+
for i, video_id in enumerate(video_ids):
|
| 503 |
+
if rank0:
|
| 504 |
+
log.info(f"Processing sample {i + 1}/{len(video_ids)}: {video_id}")
|
| 505 |
+
|
| 506 |
+
try:
|
| 507 |
+
# Load multi-view captions
|
| 508 |
+
captions = load_multiview_captions(
|
| 509 |
+
input_root, video_id, DEFAULT_CAMERA_ORDER, add_camera_prefix=args.add_camera_prefix
|
| 510 |
+
)
|
| 511 |
+
|
| 512 |
+
# Decide whether to load real videos based on num_conditional_frames
|
| 513 |
+
if args.num_conditional_frames == 0:
|
| 514 |
+
# Pure text-to-video generation, use dummy all-zero video
|
| 515 |
+
n_views = len(DEFAULT_CAMERA_ORDER)
|
| 516 |
+
multiview_video = th.zeros(
|
| 517 |
+
3, # C
|
| 518 |
+
n_views * args.target_frames, # V*T
|
| 519 |
+
args.target_height, # H
|
| 520 |
+
args.target_width, # W
|
| 521 |
+
dtype=th.uint8,
|
| 522 |
+
)
|
| 523 |
+
if rank0:
|
| 524 |
+
log.info(f"Using dummy video (all zeros) for text-to-video generation: {multiview_video.shape}")
|
| 525 |
+
else:
|
| 526 |
+
# Need video conditioning, load real videos
|
| 527 |
+
multiview_video = load_multiview_videos(
|
| 528 |
+
input_root,
|
| 529 |
+
video_id,
|
| 530 |
+
DEFAULT_CAMERA_ORDER,
|
| 531 |
+
target_frames=args.target_frames,
|
| 532 |
+
target_size=(args.target_height, args.target_width),
|
| 533 |
+
)
|
| 534 |
+
if rank0:
|
| 535 |
+
log.info(f"Loaded multiview video: {multiview_video.shape}")
|
| 536 |
+
|
| 537 |
+
if rank0:
|
| 538 |
+
log.info(f"Loaded {len(captions)} captions")
|
| 539 |
+
log.info(f"First caption preview: {captions[0][:100]}...")
|
| 540 |
+
|
| 541 |
+
# Construct data_batch
|
| 542 |
+
data_batch = construct_data_batch(
|
| 543 |
+
multiview_video,
|
| 544 |
+
captions,
|
| 545 |
+
DEFAULT_CAMERA_ORDER,
|
| 546 |
+
num_conditional_frames=args.num_conditional_frames,
|
| 547 |
+
fps=args.fps,
|
| 548 |
+
target_frames_per_view=args.target_frames,
|
| 549 |
+
)
|
| 550 |
+
|
| 551 |
+
# Run inference, already arranged the video according to stack_mode
|
| 552 |
+
video = vid2vid_cli.generate_from_batch(
|
| 553 |
+
data_batch,
|
| 554 |
+
guidance=args.guidance,
|
| 555 |
+
seed=args.seed + i,
|
| 556 |
+
num_steps=args.num_steps,
|
| 557 |
+
stack_mode=args.stack_mode,
|
| 558 |
+
)
|
| 559 |
+
|
| 560 |
+
# Save results
|
| 561 |
+
if rank0:
|
| 562 |
+
save_name = f"inference_av_{video_id}"
|
| 563 |
+
save_img_or_video(video[0], f"{args.save_root}/{save_name}", fps=args.fps)
|
| 564 |
+
log.info(f"Saved video to {args.save_root}/{save_name}")
|
| 565 |
+
|
| 566 |
+
except Exception as e:
|
| 567 |
+
log.error(f"Error processing {video_id}: {e}")
|
| 568 |
+
import traceback
|
| 569 |
+
|
| 570 |
+
traceback.print_exc()
|
| 571 |
+
continue
|
| 572 |
+
|
| 573 |
+
# Synchronize all processes
|
| 574 |
+
if args.context_parallel_size > 1:
|
| 575 |
+
th.distributed.barrier()
|
| 576 |
+
|
| 577 |
+
# Cleanup distributed resources
|
| 578 |
+
vid2vid_cli.cleanup()
|
| 579 |
+
|
| 580 |
+
|
| 581 |
+
if __name__ == "__main__":
|
| 582 |
+
main()
|
cosmos_predict2/_src/predict2_multiview/scripts/mv_visualize_helper.py
ADDED
|
@@ -0,0 +1,164 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
import os
|
| 17 |
+
|
| 18 |
+
import torch as th
|
| 19 |
+
|
| 20 |
+
from cosmos_predict2._src.imaginaire.utils import log
|
| 21 |
+
from cosmos_predict2._src.imaginaire.visualize.video import save_img_or_video
|
| 22 |
+
|
| 23 |
+
# Visualization layouts for multi-view video arrangement
|
| 24 |
+
VISUALIZE_LAYOUTS_MADS = {
|
| 25 |
+
"width": [
|
| 26 |
+
[
|
| 27 |
+
"camera_rear_left_70fov",
|
| 28 |
+
"camera_cross_left_120fov",
|
| 29 |
+
"camera_front_wide_120fov",
|
| 30 |
+
"camera_cross_right_120fov",
|
| 31 |
+
"camera_rear_right_70fov",
|
| 32 |
+
"camera_rear_tele_30fov",
|
| 33 |
+
"camera_front_tele_30fov",
|
| 34 |
+
]
|
| 35 |
+
],
|
| 36 |
+
"height": [
|
| 37 |
+
["camera_rear_left_70fov"],
|
| 38 |
+
["camera_cross_left_120fov"],
|
| 39 |
+
["camera_front_wide_120fov"],
|
| 40 |
+
["camera_cross_right_120fov"],
|
| 41 |
+
["camera_rear_right_70fov"],
|
| 42 |
+
["camera_rear_tele_30fov"],
|
| 43 |
+
["camera_front_tele_30fov"],
|
| 44 |
+
],
|
| 45 |
+
"grid": [
|
| 46 |
+
[None, "camera_front_tele_30fov", None],
|
| 47 |
+
["camera_cross_left_120fov", "camera_front_wide_120fov", "camera_cross_right_120fov"],
|
| 48 |
+
["camera_rear_left_70fov", "camera_rear_tele_30fov", "camera_rear_right_70fov"],
|
| 49 |
+
],
|
| 50 |
+
}
|
| 51 |
+
|
| 52 |
+
VISUALIZE_LAYOUTS_AGIBOT = {
|
| 53 |
+
"width": [
|
| 54 |
+
["head_color", "hand_left", "hand_right"],
|
| 55 |
+
],
|
| 56 |
+
}
|
| 57 |
+
|
| 58 |
+
VISUALIZE_LAYOUTS = {
|
| 59 |
+
"mads": VISUALIZE_LAYOUTS_MADS,
|
| 60 |
+
"agibot": VISUALIZE_LAYOUTS_AGIBOT,
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def arrange_video_visualization(mv_video, data_batch, method="width", dataset="mads"):
|
| 65 |
+
"""
|
| 66 |
+
Rearrange multi-view video based on specified layout method.
|
| 67 |
+
|
| 68 |
+
Args:
|
| 69 |
+
mv_video: (B, C, V * T, H, W) - Multi-view video tensor
|
| 70 |
+
data_batch: Batch containing camera order information
|
| 71 |
+
method: Method to arrange video visualization. Can be "width", "height", "grid", or "time".
|
| 72 |
+
- "width": Arrange all 7 views in a single horizontal row
|
| 73 |
+
- "height": Arrange all 7 views in a single vertical column
|
| 74 |
+
- "grid": Arrange views in a 3x3 grid with None values for empty positions
|
| 75 |
+
- "time": Keep original format (V*T in time dimension, no spatial rearrangement)
|
| 76 |
+
Returns:
|
| 77 |
+
Video tensor arranged according to the layout:
|
| 78 |
+
- For "width": (B, C, T, H, V*W) where V=7
|
| 79 |
+
- For "height": (B, C, T, V*H, W) where V=7
|
| 80 |
+
- For "grid": (B, C, T, 3*H, 3*W) with black padding for None positions
|
| 81 |
+
- For "time": (B, C, V*T, H, W) (unchanged)
|
| 82 |
+
"""
|
| 83 |
+
# Handle "time" mode - return video unchanged
|
| 84 |
+
if method == "time":
|
| 85 |
+
return mv_video
|
| 86 |
+
|
| 87 |
+
if method not in VISUALIZE_LAYOUTS[dataset]:
|
| 88 |
+
raise ValueError(
|
| 89 |
+
f"Unsupported visualization method: {method}. Choose from {list(VISUALIZE_LAYOUTS[dataset].keys()) + ['time']}"
|
| 90 |
+
)
|
| 91 |
+
|
| 92 |
+
current_view_order = data_batch["camera_keys_selection"][0]
|
| 93 |
+
n_views = len(current_view_order)
|
| 94 |
+
B, C, VT, H, W = mv_video.shape
|
| 95 |
+
T = VT // n_views
|
| 96 |
+
|
| 97 |
+
# Reshape to separate view and time dimensions: B C (V T) H W -> B C V T H W
|
| 98 |
+
video = mv_video.view(B, C, n_views, T, H, W)
|
| 99 |
+
|
| 100 |
+
# Create mapping from view name to tensor index
|
| 101 |
+
view_name_to_video_tensor_idx = {view_name: idx for idx, view_name in enumerate(current_view_order)}
|
| 102 |
+
|
| 103 |
+
# Create black view for None positions (used in grid layout)
|
| 104 |
+
black_view = th.zeros(B, C, T, H, W, dtype=video.dtype, device=video.device)
|
| 105 |
+
|
| 106 |
+
# Get layout definition
|
| 107 |
+
layout_definition = VISUALIZE_LAYOUTS[dataset][method]
|
| 108 |
+
|
| 109 |
+
# Arrange video according to layout
|
| 110 |
+
grid_rows = []
|
| 111 |
+
for row_of_view_names in layout_definition:
|
| 112 |
+
row_tensors = []
|
| 113 |
+
for view_name in row_of_view_names:
|
| 114 |
+
if view_name is not None and view_name in view_name_to_video_tensor_idx:
|
| 115 |
+
tensor_idx = view_name_to_video_tensor_idx[view_name]
|
| 116 |
+
# video is B C V T H W. Get tensor for view: B C T H W
|
| 117 |
+
row_tensors.append(video[:, :, tensor_idx])
|
| 118 |
+
else:
|
| 119 |
+
# Use black view for None positions or missing views
|
| 120 |
+
row_tensors.append(black_view)
|
| 121 |
+
grid_rows.append(th.cat(row_tensors, dim=-1)) # Concat on W dimension
|
| 122 |
+
|
| 123 |
+
# Concatenate rows on H dimension
|
| 124 |
+
video = th.cat(grid_rows, dim=-2) # Concat on H dimension
|
| 125 |
+
|
| 126 |
+
return video
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
def save_each_view_separately(
|
| 130 |
+
mv_video: th.Tensor,
|
| 131 |
+
data_batch: dict,
|
| 132 |
+
save_dir: str,
|
| 133 |
+
fps: float = 10.0,
|
| 134 |
+
prefix: str = "",
|
| 135 |
+
) -> None:
|
| 136 |
+
"""
|
| 137 |
+
Save each camera view as a separate video file.
|
| 138 |
+
|
| 139 |
+
Args:
|
| 140 |
+
mv_video: Multi-view video tensor with shape (C, V*T, H, W) where V is number of views
|
| 141 |
+
data_batch: Data batch containing camera_keys_selection with actual camera order
|
| 142 |
+
save_dir: Directory to save individual view videos
|
| 143 |
+
fps: Frames per second for saved videos
|
| 144 |
+
prefix: Optional prefix for saved filenames (e.g., "video_" or "control_")
|
| 145 |
+
"""
|
| 146 |
+
os.makedirs(save_dir, exist_ok=True)
|
| 147 |
+
|
| 148 |
+
# Extract actual camera order from data_batch
|
| 149 |
+
camera_order = data_batch["camera_keys_selection"][0]
|
| 150 |
+
|
| 151 |
+
C, VT, H, W = mv_video.shape
|
| 152 |
+
n_views = len(camera_order)
|
| 153 |
+
T = VT // n_views
|
| 154 |
+
|
| 155 |
+
# Reshape to separate views: (C, V*T, H, W) -> (C, V, T, H, W)
|
| 156 |
+
video_views = mv_video.view(C, n_views, T, H, W)
|
| 157 |
+
|
| 158 |
+
# Save each view
|
| 159 |
+
for view_idx, camera_name in enumerate(camera_order):
|
| 160 |
+
view_video = video_views[:, view_idx, :, :, :] # (C, T, H, W)
|
| 161 |
+
filename = f"{prefix}{camera_name}" if prefix else camera_name
|
| 162 |
+
view_save_path = os.path.join(save_dir, filename)
|
| 163 |
+
save_img_or_video(view_video, view_save_path, fps=fps)
|
| 164 |
+
log.info(f"Saved view {camera_name} to {view_save_path}")
|
cosmos_predict2/_src/predict2_multiview/utils/optim_instantiate.py
ADDED
|
@@ -0,0 +1,234 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
#
|
| 4 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 5 |
+
# you may not use this file except in compliance with the License.
|
| 6 |
+
# You may obtain a copy of the License at
|
| 7 |
+
#
|
| 8 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 9 |
+
#
|
| 10 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 11 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 12 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 13 |
+
# See the License for the specific language governing permissions and
|
| 14 |
+
# limitations under the License.
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
import re
|
| 18 |
+
|
| 19 |
+
import hydra
|
| 20 |
+
import torch
|
| 21 |
+
from omegaconf import ListConfig
|
| 22 |
+
from torch import nn
|
| 23 |
+
|
| 24 |
+
from cosmos_predict2._src.imaginaire.utils import log
|
| 25 |
+
from cosmos_predict2._src.predict2.utils.fused_adam_dtensor import FusedAdam
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def get_regular_param_group(net: nn.Module):
|
| 29 |
+
"""
|
| 30 |
+
seperate the parameters of the network into two groups: decay and no_decay.
|
| 31 |
+
based on nano_gpt codebase.
|
| 32 |
+
"""
|
| 33 |
+
param_dict = {pn: p for pn, p in net.named_parameters()}
|
| 34 |
+
param_dict = {pn: p for pn, p in param_dict.items() if p.requires_grad}
|
| 35 |
+
|
| 36 |
+
decay_params = [p for n, p in param_dict.items() if p.dim() >= 2]
|
| 37 |
+
nodecay_params = [p for n, p in param_dict.items() if p.dim() < 2]
|
| 38 |
+
return decay_params, nodecay_params
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def get_base_optimizer(
|
| 42 |
+
model: nn.Module,
|
| 43 |
+
lr: float,
|
| 44 |
+
weight_decay: float,
|
| 45 |
+
optim_type: str = "adamw",
|
| 46 |
+
**kwargs,
|
| 47 |
+
) -> torch.optim.Optimizer:
|
| 48 |
+
net_decay_param, net_nodecay_param = get_regular_param_group(model)
|
| 49 |
+
|
| 50 |
+
num_decay_params = sum(p.numel() for p in net_decay_param)
|
| 51 |
+
num_nodecay_params = sum(p.numel() for p in net_nodecay_param)
|
| 52 |
+
net_param_total = num_decay_params + num_nodecay_params
|
| 53 |
+
log.critical(f"total num parameters : {net_param_total:,}")
|
| 54 |
+
|
| 55 |
+
param_group = [
|
| 56 |
+
{
|
| 57 |
+
"params": net_decay_param + net_nodecay_param,
|
| 58 |
+
"lr": lr,
|
| 59 |
+
"weight_decay": weight_decay,
|
| 60 |
+
},
|
| 61 |
+
]
|
| 62 |
+
|
| 63 |
+
if optim_type == "adamw":
|
| 64 |
+
opt_cls = torch.optim.AdamW
|
| 65 |
+
elif optim_type == "fusedadam":
|
| 66 |
+
opt_cls = FusedAdam
|
| 67 |
+
else:
|
| 68 |
+
raise ValueError(f"Unknown optimizer type: {optim_type}")
|
| 69 |
+
|
| 70 |
+
for k, v in kwargs.items():
|
| 71 |
+
if isinstance(v, ListConfig):
|
| 72 |
+
kwargs[k] = list(v)
|
| 73 |
+
|
| 74 |
+
return opt_cls(param_group, **kwargs)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def get_base_scheduler(
|
| 78 |
+
optimizer: torch.optim.Optimizer,
|
| 79 |
+
model: nn.Module,
|
| 80 |
+
scheduler_config: dict,
|
| 81 |
+
):
|
| 82 |
+
net_scheduler = hydra.utils.instantiate(scheduler_config)
|
| 83 |
+
net_scheduler.model = model
|
| 84 |
+
|
| 85 |
+
return torch.optim.lr_scheduler.LambdaLR(
|
| 86 |
+
optimizer,
|
| 87 |
+
lr_lambda=[
|
| 88 |
+
net_scheduler.schedule,
|
| 89 |
+
],
|
| 90 |
+
)
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
def get_multiple_optimizer(
|
| 94 |
+
model: nn.Module,
|
| 95 |
+
lr: float,
|
| 96 |
+
weight_decay: float,
|
| 97 |
+
optim_type: str = "adamw",
|
| 98 |
+
lr_overrides: list[dict] = None,
|
| 99 |
+
**kwargs,
|
| 100 |
+
) -> torch.optim.Optimizer:
|
| 101 |
+
"""
|
| 102 |
+
Get an optimizer with multiple learning rates for different parts of the model,
|
| 103 |
+
allowing pattern matching for parameter names.
|
| 104 |
+
|
| 105 |
+
The logic is:
|
| 106 |
+
1. All parameters are initially considered for the default learning rate.
|
| 107 |
+
2. We iterate through lr_overrides. If a parameter's name matches a pattern,
|
| 108 |
+
it's moved to a group with the specified learning rate. A parameter is only
|
| 109 |
+
assigned to the *first* pattern it matches.
|
| 110 |
+
|
| 111 |
+
Args:
|
| 112 |
+
model (nn.Module): The model to optimize.
|
| 113 |
+
lr (float): The default learning rate.
|
| 114 |
+
weight_decay (float): The default weight decay.
|
| 115 |
+
optim_type (str): The type of optimizer to use ('adamw' or 'fusedadam').
|
| 116 |
+
lr_overrides (list[dict], optional): A list of dicts with keys:
|
| 117 |
+
- 'pattern' (str): The pattern to match (required)
|
| 118 |
+
- 'lr' (float): The learning rate for matching params (required)
|
| 119 |
+
- 'match_type' (str): 'regex', 'contains', 'startswith', 'endswith' (default: 'contains')
|
| 120 |
+
|
| 121 |
+
Example:
|
| 122 |
+
[
|
| 123 |
+
{'pattern': 'cross_view_attn', 'lr': 2e-4, 'match_type': 'contains'},
|
| 124 |
+
{'pattern': 'text_encoder', 'lr': 1e-5, 'match_type': 'contains'},
|
| 125 |
+
]
|
| 126 |
+
|
| 127 |
+
This is Hydra-friendly and can be overridden from command line like:
|
| 128 |
+
optimizer.lr_overrides.0.lr=1e-4
|
| 129 |
+
optimizer.lr_overrides.1.pattern=vision_encoder
|
| 130 |
+
**kwargs: Additional arguments for the optimizer.
|
| 131 |
+
|
| 132 |
+
Returns:
|
| 133 |
+
torch.optim.Optimizer: The configured optimizer.
|
| 134 |
+
"""
|
| 135 |
+
param_dict = {pn: p for pn, p in model.named_parameters() if p.requires_grad}
|
| 136 |
+
|
| 137 |
+
# Initialize groups for parameters with overridden LR
|
| 138 |
+
override_groups = {} # (lr, has_decay) -> [params]
|
| 139 |
+
override_groups_name = {} # (lr, has_decay) -> [name]
|
| 140 |
+
|
| 141 |
+
# Initialize lists for parameters with default LR
|
| 142 |
+
default_decay_params = []
|
| 143 |
+
default_nodecay_params = []
|
| 144 |
+
|
| 145 |
+
# Temporarily hold all params to check against overrides
|
| 146 |
+
unassigned_params = list(param_dict.items())
|
| 147 |
+
|
| 148 |
+
# First, assign params that match an override pattern
|
| 149 |
+
if lr_overrides:
|
| 150 |
+
override_list = lr_overrides
|
| 151 |
+
|
| 152 |
+
for name, p in list(unassigned_params):
|
| 153 |
+
assigned = False
|
| 154 |
+
for override_item in override_list:
|
| 155 |
+
pattern = override_item["pattern"]
|
| 156 |
+
special_lr = override_item["lr"]
|
| 157 |
+
match_type = override_item.get("match_type", "contains")
|
| 158 |
+
|
| 159 |
+
# Determine if the parameter name matches
|
| 160 |
+
matched = False
|
| 161 |
+
if match_type == "regex":
|
| 162 |
+
matched = re.match(pattern, name) is not None
|
| 163 |
+
elif match_type == "contains":
|
| 164 |
+
matched = pattern in name
|
| 165 |
+
elif match_type == "startswith":
|
| 166 |
+
matched = name.startswith(pattern)
|
| 167 |
+
elif match_type == "endswith":
|
| 168 |
+
matched = name.endswith(pattern)
|
| 169 |
+
else:
|
| 170 |
+
raise ValueError(
|
| 171 |
+
f"Unknown match_type: {match_type}. Must be one of: regex, contains, startswith, endswith"
|
| 172 |
+
)
|
| 173 |
+
|
| 174 |
+
if matched:
|
| 175 |
+
has_decay = p.dim() >= 2
|
| 176 |
+
group_key = (special_lr, has_decay)
|
| 177 |
+
if group_key not in override_groups:
|
| 178 |
+
override_groups[group_key] = []
|
| 179 |
+
if group_key not in override_groups_name:
|
| 180 |
+
override_groups_name[group_key] = []
|
| 181 |
+
override_groups[group_key].append(p)
|
| 182 |
+
override_groups_name[group_key].append(name)
|
| 183 |
+
assigned = True
|
| 184 |
+
break # Assign to first matching pattern
|
| 185 |
+
if assigned:
|
| 186 |
+
# Remove from unassigned list; this is a bit inefficient but clear
|
| 187 |
+
unassigned_params = [(n, param) for n, param in unassigned_params if n != name]
|
| 188 |
+
|
| 189 |
+
# Assign all remaining params to default groups
|
| 190 |
+
for name, p in unassigned_params:
|
| 191 |
+
if p.dim() >= 2:
|
| 192 |
+
default_decay_params.append(p)
|
| 193 |
+
else:
|
| 194 |
+
default_nodecay_params.append(p)
|
| 195 |
+
|
| 196 |
+
# Build final param_groups list for the optimizer
|
| 197 |
+
final_param_groups = []
|
| 198 |
+
if default_decay_params:
|
| 199 |
+
final_param_groups.append({"params": default_decay_params, "lr": lr, "weight_decay": weight_decay})
|
| 200 |
+
if default_nodecay_params:
|
| 201 |
+
final_param_groups.append({"params": default_nodecay_params, "lr": lr, "weight_decay": 0.0})
|
| 202 |
+
|
| 203 |
+
for (special_lr, has_decay), params in override_groups.items():
|
| 204 |
+
final_param_groups.append(
|
| 205 |
+
{"params": params, "lr": special_lr, "weight_decay": weight_decay if has_decay else 0.0}
|
| 206 |
+
)
|
| 207 |
+
|
| 208 |
+
# print the parameter names in each group
|
| 209 |
+
for (special_lr, has_decay), params in override_groups.items():
|
| 210 |
+
log.critical(f"special_lr {special_lr}: {override_groups_name[(special_lr, has_decay)]}")
|
| 211 |
+
|
| 212 |
+
# Log parameter group information
|
| 213 |
+
total_params = 0
|
| 214 |
+
log.critical("Optimizer parameter groups:")
|
| 215 |
+
for i, group in enumerate(final_param_groups):
|
| 216 |
+
group_params = sum(p.numel() for p in group["params"])
|
| 217 |
+
total_params += group_params
|
| 218 |
+
log.critical(
|
| 219 |
+
f" Group {i}: num_params={group_params:,}, lr={group['lr']:.1e}, weight_decay={group['weight_decay']}"
|
| 220 |
+
)
|
| 221 |
+
log.critical(f"Total trainable parameters: {total_params:,}")
|
| 222 |
+
|
| 223 |
+
if optim_type == "adamw":
|
| 224 |
+
opt_cls = torch.optim.AdamW
|
| 225 |
+
elif optim_type == "fusedadam":
|
| 226 |
+
opt_cls = FusedAdam
|
| 227 |
+
else:
|
| 228 |
+
raise ValueError(f"Unknown optimizer type: {optim_type}")
|
| 229 |
+
|
| 230 |
+
for k, v in kwargs.items():
|
| 231 |
+
if isinstance(v, ListConfig):
|
| 232 |
+
kwargs[k] = list(v)
|
| 233 |
+
|
| 234 |
+
return opt_cls(final_param_groups, **kwargs)
|
sh/train_nymeria_longer.sh
ADDED
|
@@ -0,0 +1,38 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env bash
|
| 2 |
+
# Launch 2-actor joint pose-conditioned full fine-tune on the LONGER Nymeria dataset (200+77 frames),
|
| 3 |
+
# 4-GPU FSDP. Warm-starts the 2B multiview backbone from the HF checkpoint; VAE = local Wan2.1.
|
| 4 |
+
#
|
| 5 |
+
# Usage:
|
| 6 |
+
# bash sh/train_nymeria_longer.sh # full run (max_iter from the experiment)
|
| 7 |
+
# NPROC=4 bash sh/train_nymeria_longer.sh # override #GPUs
|
| 8 |
+
# bash sh/train_nymeria_longer.sh trainer.max_iter=5 # pass extra hydra overrides
|
| 9 |
+
#
|
| 10 |
+
# Outputs (checkpoints every save_iter + validation videos) -> /data/model_output/train/<job.name>/
|
| 11 |
+
# W&B: entity=$WANDB_ENTITY, project=$WANDB_PROJECT, run name == <job.name> == the output leaf folder.
|
| 12 |
+
set -euo pipefail
|
| 13 |
+
|
| 14 |
+
REPO="/home/ubuntu/dahyun/cosmos-predict2.5"
|
| 15 |
+
PY="/home/ubuntu/miniforge3/envs/ego_dh/bin"
|
| 16 |
+
EXP="${EXP:-nymeria_pose_2actor_2b_longer}"
|
| 17 |
+
NPROC="${NPROC:-4}"
|
| 18 |
+
PORT="${PORT:-29581}"
|
| 19 |
+
|
| 20 |
+
# secrets (HF_TOKEN, WANDB_API_KEY) — kept outside the repo
|
| 21 |
+
source /home/ubuntu/.ego_dh_secrets.sh
|
| 22 |
+
|
| 23 |
+
export WANDB_ENTITY="${WANDB_ENTITY:-dhyunchung}" # W&B entity
|
| 24 |
+
export WANDB_PROJECT="${WANDB_PROJECT:-multi-ego}" # W&B project (decoupled from the output folder)
|
| 25 |
+
export COSMOS_QWEN_TOKENIZER_DIR="${COSMOS_QWEN_TOKENIZER_DIR:-/data/cosmos_reason1_7b}" # local Qwen tokenizer (offline text encoder)
|
| 26 |
+
export IMAGINAIRE_OUTPUT_ROOT="${IMAGINAIRE_OUTPUT_ROOT:-/data/model_output}" # -> /data/model_output/train/<name>
|
| 27 |
+
export PYTHONPATH="$REPO:${PYTHONPATH:-}"
|
| 28 |
+
|
| 29 |
+
cd "$REPO"
|
| 30 |
+
# stdout/stderr log (in addition to the framework's own console.log under the output dir)
|
| 31 |
+
LOGDIR="${LOGDIR:-/data3/nymeria_train_logs}"
|
| 32 |
+
mkdir -p "$LOGDIR"
|
| 33 |
+
LOGFILE="$LOGDIR/${EXP}_$(date +%Y%m%d_%H%M%S).log"
|
| 34 |
+
echo "logging to $LOGFILE"
|
| 35 |
+
"$PY/torchrun" --nproc_per_node="$NPROC" --master_port="$PORT" \
|
| 36 |
+
-m cosmos_oss.scripts.train \
|
| 37 |
+
--config=cosmos_predict2/_src/predict2_multiview/configs/vid2vid/config.py -- \
|
| 38 |
+
experiment="$EXP" "$@" 2>&1 | tee "$LOGFILE"
|