Skip to content
Draft
12 changes: 12 additions & 0 deletions examples/inference/basic/basic_cosmos_predict.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
# SPDX-License-Identifier: Apache-2.0
"""Inference script for Cosmos Predict video generation.

Example usage:
python examples/inference/basic/basic_cosmos_predict.py \
--model_name nvidia/Cosmos-1.0-Prompt2World-7B-Video \
--prompt "A cute dog walking."
"""
from fastvideo.utils.cli import inference_entry

if __name__ == "__main__":
inference_entry()
1 change: 1 addition & 0 deletions fastvideo/configs/models/dits/cosmos2_5.py
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,7 @@ class Cosmos25ArchConfig(DiTArchConfig):
qk_norm: str = "rms_norm"
eps: float = 1e-6
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
use_condition_mask: bool = True

def __post_init__(self):
super().__post_init__()
Expand Down
15 changes: 15 additions & 0 deletions fastvideo/configs/models/encoders/cosmos_predict_text_encoder.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig, TextEncoderConfig
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig

@dataclass
class CosmosPredictTextEncoderArchConfig(Reason1ArchConfig):
"""Arch config for Cosmos Predict text encoder. It is basically Qwen2.5-VL-7B-Instruct."""
pass

@dataclass
class CosmosPredictTextEncoderConfig(TextEncoderConfig):
"""Cosmos Predict text encoder config."""
arch_config: CosmosPredictTextEncoderArchConfig = field(default_factory=CosmosPredictTextEncoderArchConfig)
tokenizer_type: str = "Qwen/Qwen2.5-VL-7B-Instruct"
107 changes: 107 additions & 0 deletions fastvideo/configs/pipelines/cosmos_predict.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,107 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Callable
from dataclasses import dataclass, field

import torch

from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import Cosmos25VideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import (
Cosmos25ArchConfig,
Cosmos25_14BArchConfig,
Cosmos25_14BVideoConfig,
)
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.cosmos_predict_text_encoder import CosmosPredictTextEncoderConfig, CosmosPredictTextEncoderArchConfig
from fastvideo.configs.models.vaes import Cosmos25VAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig


def _identity_preprocess_text(prompt: str) -> str:
return prompt


def cosmos_predict_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
# Just return hidden states unmodified. The text encoder handles its own logic if needed.
hidden_states = getattr(outputs, "hidden_states", None)
if hidden_states is None:
raise ValueError("Cosmos Predict postprocess requires outputs.hidden_states")
return hidden_states


@dataclass
class CosmosPredictConfig(PipelineConfig):
"""Configuration for Cosmos Predict (Text-to-Video/Video-to-Video) generation pipeline."""

dit_config: DiTConfig = field(default_factory=lambda: Cosmos25VideoConfig(arch_config=Cosmos25ArchConfig(
num_attention_heads=16,
attention_head_dim=128,
in_channels=16,
out_channels=16,
num_layers=28,
patch_size=[1, 2, 2],
max_size=[128, 240, 240],
rope_scale=[1.0, 3.0, 3.0],
text_embed_dim=1024,
mlp_ratio=4.0,
adaln_lora_dim=256,
use_adaln_lora=True,
concat_padding_mask=True,
extra_pos_embed_type=None,
use_crossattn_projection=True,
rope_enable_fps_modulation=False,
qk_norm="rms_norm",
use_condition_mask=False, # Cosmos Predict does not concat condition mask
)))

vae_config: VAEConfig = field(default_factory=Cosmos25VAEConfig)

text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (CosmosPredictTextEncoderConfig(
arch_config=CosmosPredictTextEncoderArchConfig()), ))

preprocess_text_funcs: tuple[Callable[[str], str],
...] = field(default_factory=lambda: (_identity_preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda: (cosmos_predict_postprocess_text, ))

dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))

embedded_cfg_scale: float = 0.0
flow_shift: float = 5.0

vae_tiling: bool = False
vae_sp: bool = False

def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
self._vae_latent_dim = 16


@dataclass
class CosmosPredict14BConfig(CosmosPredictConfig):
"""Configuration for Cosmos Predict 14B pipeline."""

dit_config: DiTConfig = field(default_factory=lambda: Cosmos25_14BVideoConfig(arch_config=Cosmos25_14BArchConfig(
num_attention_heads=40,
attention_head_dim=128,
in_channels=16,
out_channels=16,
num_layers=36,
patch_size=[1, 2, 2],
max_size=[128, 240, 240],
rope_scale=[1.0, 3.0, 3.0],
text_embed_dim=1024,
mlp_ratio=4.0,
adaln_lora_dim=256,
use_adaln_lora=True,
concat_padding_mask=True,
extra_pos_embed_type=None,
use_crossattn_projection=True,
rope_enable_fps_modulation=False,
qk_norm="rms_norm",
use_condition_mask=False,
)))

23 changes: 16 additions & 7 deletions fastvideo/models/dits/cosmos2_5.py
Original file line number Diff line number Diff line change
Expand Up @@ -749,14 +749,15 @@ def __init__(self, config: Cosmos25VideoConfig, hf_config: dict[str, Any]) -> No
self.adaln_lora_dim = getattr(config, "adaln_lora_dim", 256)
self.extra_pos_embed_type = getattr(config, "extra_pos_embed_type", None)
self.use_crossattn_projection = getattr(config, "use_crossattn_projection", False)
self.use_condition_mask = getattr(config, "use_condition_mask", True)

# 1. Patch Embedding
# Account for: VAE channels + condition_mask (1) + padding_mask (1 if concat_padding_mask)
# Account for: VAE channels + condition_mask (1, optional) + padding_mask (1 if concat_padding_mask)
patch_embed_in_channels = config.in_channels # Base VAE channels (16)
patch_embed_in_channels += 1 # Always add 1 for condition_mask
if self.use_condition_mask:
patch_embed_in_channels += 1 # Add 1 for condition_mask
if config.concat_padding_mask:
patch_embed_in_channels += 1 # Add 1 for padding_mask
# Total: 16 + 1 + 1 = 18 (with concat_padding_mask=True)

self.patch_embed = Cosmos25PatchEmbed(patch_embed_in_channels, inner_dim, config.patch_size)

Expand Down Expand Up @@ -845,11 +846,19 @@ def forward(

batch_size, num_channels, num_frames, height, width = hidden_states.shape

# 1. Concatenate condition mask if provided
if condition_mask is not None:
hidden_states = torch.cat([hidden_states, condition_mask], dim=1)
# 1. Concatenate condition mask if provided and expected
if self.use_condition_mask:
if condition_mask is not None:
hidden_states = torch.cat([hidden_states, condition_mask], dim=1)
else:
# If not provided, create a dummy zero mask
dummy_mask = torch.zeros(
batch_size, 1, num_frames, height, width,
dtype=hidden_states.dtype, device=hidden_states.device
)
hidden_states = torch.cat([hidden_states, dummy_mask], dim=1)

# 2. Concatenate padding mask if needed
# 2. Concatenate padding mask if required
if self.concat_padding_mask and padding_mask is not None:
padding_mask = transforms.functional.resize(
padding_mask,
Expand Down
37 changes: 37 additions & 0 deletions fastvideo/models/encoders/cosmos_predict_text_encoder.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
import torch
import torch.nn as nn
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import Qwen2_5_VLForConditionalGeneration
from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import Qwen2_5_VLConfig

class CosmosPredictTextEncoder(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig = None):
super().__init__()
if config is None:
# Default fallback for testing
config = Qwen2_5_VLConfig()

# We instantiate the standard HF model used by Cosmos Predict
self.model = Qwen2_5_VLForConditionalGeneration(config)

def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
"""
Forward pass for text encoding in Cosmos Predict.
Extracts hidden states from all layers (except the embedding layer 0),
normalizes them, and concatenates them to form prompt_embeds.
"""
outputs = self.model(
input_ids=input_ids,
output_hidden_states=True,
return_dict=True,
)
hidden_states = outputs.hidden_states

normalized_hidden_states = []
for layer_idx in range(1, len(hidden_states)):
normalized_state = (hidden_states[layer_idx] - hidden_states[layer_idx].mean(dim=-1, keepdim=True)) / (
hidden_states[layer_idx].std(dim=-1, keepdim=True) + 1e-8
)
normalized_hidden_states.append(normalized_state)

prompt_embeds = torch.cat(normalized_hidden_states, dim=-1)
return prompt_embeds
Loading
Loading