Compare commits

...
Author SHA1 Message Date
SolitaryThinker 19c1d164c3 fix 2025-12-12 22:51:21 +00:00
2 changed files with 6 additions and 3 deletions
+2 -2
View File
@@ -115,7 +115,7 @@ class ForwardBatch:
# Latent tensors
latents: torch.Tensor | None = None
raw_latent_shape: torch.Tensor | None = None
raw_latent_shape: tuple[int, ...] | None = None
noise_pred: torch.Tensor | None = None
image_latent: torch.Tensor | None = None
@@ -206,7 +206,7 @@ class TrainingBatch:
# Dataloader batch outputs
latents: torch.Tensor | None = None
raw_latent_shape: torch.Tensor | None = None
raw_latent_shape: tuple[int, ...] | None = None
noise_latents: torch.Tensor | None = None
encoder_hidden_states: torch.Tensor | None = None
encoder_attention_mask: torch.Tensor | None = None
@@ -82,6 +82,7 @@ class LatentPreparationStage(PipelineStage):
raise ValueError("Height and width must be provided")
# Calculate latent shape
bcthw_shape: tuple[int, ...] | None = None
if self.use_btchw_layout:
shape = (
batch_size,
@@ -92,6 +93,7 @@ class LatentPreparationStage(PipelineStage):
width // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
)
bcthw_shape = tuple(shape[i] for i in [0, 2, 1, 3, 4])
else:
shape = (
batch_size,
@@ -102,6 +104,7 @@ class LatentPreparationStage(PipelineStage):
width // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
)
bcthw_shape = shape
# Validate generator if it's a list
if isinstance(generator, list) and len(generator) != batch_size:
@@ -123,7 +126,7 @@ class LatentPreparationStage(PipelineStage):
latents = latents * self.scheduler.init_noise_sigma
# Update batch with prepared latents
batch.latents = latents
batch.raw_latent_shape = latents.shape
batch.raw_latent_shape = bcthw_shape
return batch