From 8ac72ae691f7a00be6502b5bb902a46af371aacc Mon Sep 17 00:00:00 2001 From: Bartosz Stefaniak Date: Tue, 23 Jun 2026 12:57:09 +0000 Subject: [PATCH] Fix timestep passed to time_embedder Signed-off-by: Bartosz Stefaniak --- .../_torch/visual_gen/models/cosmos3/pipeline_cosmos3.py | 3 ++- .../visual_gen/models/cosmos3/transformer_cosmos3.py | 9 ++++++--- 2 files changed, 8 insertions(+), 4 deletions(-) diff --git a/tensorrt_llm/_torch/visual_gen/models/cosmos3/pipeline_cosmos3.py b/tensorrt_llm/_torch/visual_gen/models/cosmos3/pipeline_cosmos3.py index f33f81fcb6a5..2de422dcefb5 100644 --- a/tensorrt_llm/_torch/visual_gen/models/cosmos3/pipeline_cosmos3.py +++ b/tensorrt_llm/_torch/visual_gen/models/cosmos3/pipeline_cosmos3.py @@ -586,7 +586,8 @@ def forward_fn( """ noise_pred = self.transformer( hidden_states=latent_input, - timestep=timestep / self.scheduler.config.num_train_timesteps, + timestep=timestep, + attention_timestep=timestep / self.scheduler.config.num_train_timesteps, text_ids=extra_tensors["text_ids"], text_mask=extra_tensors["text_mask"], video_shape=video_shape, diff --git a/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py b/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py index 24016d761de0..96a103e39d94 100644 --- a/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py +++ b/tensorrt_llm/_torch/visual_gen/models/cosmos3/transformer_cosmos3.py @@ -866,6 +866,7 @@ def forward( self, hidden_states: torch.Tensor, timestep: Optional[torch.Tensor] = None, + attention_timestep: Optional[torch.Tensor] = None, text_ids: Optional[torch.Tensor] = None, text_mask: Optional[torch.Tensor] = None, video_shape: Optional[Tuple[int, int, int]] = None, @@ -878,7 +879,9 @@ def forward( Args: hidden_states: [B, C, T, H, W] noisy latents - timestep: Normalized diffusion timestep in [0, 1], shape [B] + timestep: Raw scheduler diffusion timestep, shape [B] + attention_timestep: Normalized diffusion timestep in [0, 1], shape [B], + for attention backends that use timestep-dependent behavior. text_ids: [B, S_text] tokenized text input text_mask: [B, S_text] attention mask for text (1=real, 0=pad) video_shape: (T, H, W) in latent space @@ -931,7 +934,7 @@ def forward( text_ids, text_mask, freqs_und, - timestep=timestep, + timestep=attention_timestep, ) self.cached_freqs_gen = freqs_gen @@ -971,7 +974,7 @@ def forward( k_und, v_und, freqs_gen, - timestep=timestep, + timestep=attention_timestep, ) hidden_gen = self.sharder.gather(hidden_gen, dim=1, unpad_to=S_gen)