Skip to content

Commit 2b446eb

Browse files
authored
Merge pull request #141 from aryamancodes/aryamang/fix-mv-inference
fix Text2world-MultiView inference by adding missing config options
2 parents b0dff19 + f22e36a commit 2b446eb

5 files changed

Lines changed: 18 additions & 7 deletions

File tree

README.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33
</p>
44

55
<h1 align="center">
6-
6+
77
> 🚨 **Update Notice**
88
>
99
> The latest version of our Cosmos-Predict is now live!

cosmos_predict1/diffusion/config/inference/cosmos-1-diffusion-text2world-multiview.py

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,10 @@
3737
88,
3838
160,
3939
],
40+
net=dict(
41+
extra_per_block_abs_pos_emb=True,
42+
extra_per_block_abs_pos_emb_type="sincos",
43+
),
4044
tokenizer=dict(
4145
video_vae=dict(
4246
pixel_chunk_duration=57,
@@ -59,10 +63,12 @@
5963
name="Cosmos_Predict1_Text2World_7B_Multiview_post_trained",
6064
),
6165
model=dict(
62-
net=dict(
66+
net=dict(
6367
n_views=5,
6468
view_condition_dim=3,
65-
add_repeat_frame_embedding=False,
69+
add_repeat_frame_embedding=False,
70+
extra_per_block_abs_pos_emb=True,
71+
extra_per_block_abs_pos_emb_type="sincos",
6672
),
6773
latent_shape=[
6874
16,

cosmos_predict1/diffusion/config/inference/cosmos-1-diffusion-video2world-multiview.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -84,4 +84,3 @@
8484
Cosmos_Predict1_Video2World_7B_Multiview_post_trained,
8585
]:
8686
cs.store(group="experiment", package="_global_", name=_item["job"]["name"], node=_item)
87-

cosmos_predict1/diffusion/inference/world_generation_pipeline.py

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1012,9 +1012,15 @@ def _run_tokenizer_decoding(self, sample: torch.Tensor) -> np.ndarray:
10121012
video = (1.0 + self.model.decode(sample)).clamp(0, 2) / 2 # [B, 3, T, H, W]
10131013
video_segments = einops.rearrange(video, "b c (v t) h w -> b c v t h w", v=self.n_views)
10141014
video_arrangement = [1, 0, 2, 4, 3, 5]
1015-
# Fill one blank view for 5view
1015+
# Fill one blank view for 5view
10161016
if self.n_views == 5:
1017-
ones_tensor = torch.zeros_like(video_segments[:, :, 0,],).unsqueeze(2)
1017+
ones_tensor = torch.zeros_like(
1018+
video_segments[
1019+
:,
1020+
:,
1021+
0,
1022+
],
1023+
).unsqueeze(2)
10181024
video_segments = torch.cat((video_segments, ones_tensor), dim=2)
10191025
video_arrangement = [1, 0, 2, 3, 5, 4]
10201026
grid_video = torch.stack(

scripts/download_diffusion_checkpoints.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -101,7 +101,7 @@ def convert_pixtral_checkpoint(checkpoint_dir: str, checkpoint_name: str, vit_ty
101101
allow_patterns=["params.json", "consolidated.safetensors"],
102102
local_dir=pixtral_ckpt_dir,
103103
local_dir_use_symlinks=False,
104-
revision="db3e3ed01201248694fcb170c7bd292ecfcad22b"
104+
revision="db3e3ed01201248694fcb170c7bd292ecfcad22b",
105105
)
106106
orig_dtype = torch.get_default_dtype()
107107
dtype = torch.bfloat16

0 commit comments

Comments
 (0)