dhyun22 commited on
Commit
3b2f112
·
verified ·
1 Parent(s): e53f5f0

CoMind single-ego 2-view training code + README (what to change)

Browse files
Files changed (41) hide show
  1. README.md +133 -0
  2. cosmos_predict2/_src/predict2/models/video2world_model_rectified_flow.py +247 -0
  3. cosmos_predict2/_src/predict2_multiview/__init__.py +15 -0
  4. cosmos_predict2/_src/predict2_multiview/callbacks/every_n_draw_sample_multiviewvideo.py +485 -0
  5. cosmos_predict2/_src/predict2_multiview/callbacks/frame_loss_log.py +43 -0
  6. cosmos_predict2/_src/predict2_multiview/callbacks/log_weight.py +62 -0
  7. cosmos_predict2/_src/predict2_multiview/callbacks/nymeria_validation_viz.py +335 -0
  8. cosmos_predict2/_src/predict2_multiview/callbacks/sigma_loss_analysis_per_frame.py +338 -0
  9. cosmos_predict2/_src/predict2_multiview/conditioner.py +103 -0
  10. cosmos_predict2/_src/predict2_multiview/configs/__init__.py +15 -0
  11. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/__init__.py +15 -0
  12. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/config.py +47 -0
  13. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/__init__.py +15 -0
  14. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/callbacks.py +50 -0
  15. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/conditioner.py +720 -0
  16. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/dataloader.py +101 -0
  17. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/dataloader_local.py +111 -0
  18. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/model.py +57 -0
  19. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/net.py +190 -0
  20. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/defaults/optimizer.py +72 -0
  21. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/__init__.py +15 -0
  22. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/buttercup2p5_rectified_flow.py +229 -0
  23. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/buttercup2p5_rectified_flow_14b.py +243 -0
  24. cosmos_predict2/_src/predict2_multiview/configs/vid2vid/experiment/nymeria_pose_2actor.py +1005 -0
  25. cosmos_predict2/_src/predict2_multiview/datasets/__init__.py +15 -0
  26. cosmos_predict2/_src/predict2_multiview/datasets/comind_pairs.py +302 -0
  27. cosmos_predict2/_src/predict2_multiview/datasets/local.py +173 -0
  28. cosmos_predict2/_src/predict2_multiview/datasets/multiview.py +547 -0
  29. cosmos_predict2/_src/predict2_multiview/datasets/nymeria_pairs.py +1074 -0
  30. cosmos_predict2/_src/predict2_multiview/datasets/wdinfo_utils.py +79 -0
  31. cosmos_predict2/_src/predict2_multiview/models/multiview_pose_model_rectified_flow.py +261 -0
  32. cosmos_predict2/_src/predict2_multiview/models/multiview_vid2vid_model_rectified_flow.py +655 -0
  33. cosmos_predict2/_src/predict2_multiview/models/view_sampling.py +65 -0
  34. cosmos_predict2/_src/predict2_multiview/networks/multiview_cross_dit.py +1187 -0
  35. cosmos_predict2/_src/predict2_multiview/networks/multiview_dit.py +618 -0
  36. cosmos_predict2/_src/predict2_multiview/networks/multiview_pose_dit.py +526 -0
  37. cosmos_predict2/_src/predict2_multiview/scripts/inference.py +330 -0
  38. cosmos_predict2/_src/predict2_multiview/scripts/inference_cli.py +582 -0
  39. cosmos_predict2/_src/predict2_multiview/scripts/mv_visualize_helper.py +164 -0
  40. cosmos_predict2/_src/predict2_multiview/utils/optim_instantiate.py +234 -0
  41. 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"