Compare commits

...
Author SHA1 Message Date
Will Lin e4ceadb5d5 fix denoising stage init 2025-06-06 11:44:36 -07:00
Zhang PeiyuanandWill Lin 8f8ce6d9e1 [bugfix] [training] Add negative prompt to preprocessing and validation (#479)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-06 11:21:21 -07:00
10 changed files with 138 additions and 21 deletions
+93 -1
View File
@@ -33,7 +33,8 @@ class ParquetVideoTextDataset(Dataset):
world_size: int = 1,
cfg_rate: float = 0.0,
num_latent_t: int = 2,
seed: int = 0):
seed: int = 0,
validation: bool = False):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
@@ -47,6 +48,12 @@ class ParquetVideoTextDataset(Dataset):
self.cfg_rate = cfg_rate
self.num_latent_t = num_latent_t
self.local_indices = None
self.validation = validation
# Negative prompt caching
self.neg_metadata = None
self.cached_neg_prompt: Dict[str, Any] | None = None
self.plan_output_dir = os.path.join(
self.path,
f"data_plan_{self.world_size}_{self.sp_world_size}_{self.dp_world_size}.json"
@@ -75,6 +82,12 @@ class ParquetVideoTextDataset(Dataset):
for row_idx in range(num_rows):
metadatas.append((file_path, row_idx))
# the negative prompt is always the first row in the first
# parquet file
if validation:
self.neg_metadata = metadatas[0]
metadatas = metadatas[1:]
# Generate the plan that distribute rows among workers
random.seed(seed)
random.shuffle(metadatas)
@@ -93,9 +106,88 @@ class ParquetVideoTextDataset(Dataset):
for global_rank in group_ranks_list[sp_group_idx]:
plan[global_rank].append(metadata)
if validation:
assert self.neg_metadata is not None
plan["negative_prompt"] = [self.neg_metadata]
with open(self.plan_output_dir, "w") as f:
json.dump(plan, f)
else:
pass
dist.barrier()
if validation:
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.neg_metadata = plan["negative_prompt"][0]
def _load_and_cache_negative_prompt(self) -> None:
"""Load and cache the negative prompt. Only rank 0 in each SP group should call this."""
if not self.validation or self.neg_metadata is None:
return
if self.cached_neg_prompt is not None:
return
# Only rank 0 in each SP group should read the negative prompt
try:
file_path, row_idx = self.neg_metadata
parquet_file = pq.ParquetFile(file_path)
# Since negative prompt is always the first row (row_idx = 0),
# it's always in the first row group
row_group_index = 0
local_index = row_idx # This will be 0 for the negative prompt
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
row_dict = {k: v[local_index] for k, v in row_group.items()}
del row_group
# Process the negative prompt row
self.cached_neg_prompt = self._process_row(row_dict)
except Exception as e:
logger.error("Failed to load negative prompt: %s", e)
self.cached_neg_prompt = None
def get_validation_negative_prompt(
self
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, Dict[str, Any]]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
Returns the processed negative prompt data (latents, embeddings, masks, info).
"""
if not self.validation:
raise ValueError(
"get_validation_negative_prompt() can only be called in validation mode"
)
# Load and cache if needed (only rank 0 in SP group will actually load)
if self.cached_neg_prompt is None:
self._load_and_cache_negative_prompt()
if self.cached_neg_prompt is None:
raise RuntimeError(
f"Rank {self.rank} (SP rank {self.local_rank}): Could not retrieve negative prompt data"
)
# Extract the components
lat, emb, mask, info = (self.cached_neg_prompt["latents"],
self.cached_neg_prompt["embeddings"],
self.cached_neg_prompt["masks"],
self.cached_neg_prompt["info"])
# Apply the same processing as in __getitem__
if lat.numel() == 0: # Validation parquet
return lat, emb, mask, info
else:
lat = lat[:, -self.num_latent_t:]
if self.sp_world_size > 1:
lat = rearrange(lat,
"t (n s) h w -> t n s h w",
n=self.sp_world_size).contiguous()
lat = lat[:, self.local_rank, :, :, :]
return lat, emb, mask, info
def __len__(self):
if self.local_indices is None:
@@ -161,6 +161,7 @@ class ComposedPipelineBase(ABC):
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
fastvideo_args.num_gpus = int(os.environ.get("WORLD_SIZE", 1))
fastvideo_args.use_cpu_offload = False
# make sure we are in training mode
fastvideo_args.inference_mode = False
@@ -7,7 +7,8 @@ This module defines the dataclasses used to pass state between pipeline componen
in a functional manner, reducing the need for explicit parameter passing.
"""
from dataclasses import dataclass, field
import pprint
from dataclasses import asdict, dataclass, field
from typing import Any, Dict, List, Optional, Union
import torch
@@ -126,4 +127,8 @@ class ForwardBatch:
# Set do_classifier_free_guidance based on guidance scale and negative prompt
if self.guidance_scale > 1.0:
self.do_classifier_free_guidance = True
if self.negative_prompt_embeds is None:
self.negative_prompt_embeds = []
def __str__(self):
return pprint.pformat(asdict(self), indent=2, width=120)
@@ -12,6 +12,7 @@ from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import getdataset
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -300,7 +301,10 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# Prepare batch data for Parquet dataset
batch_data = []
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
if sampling_param.negative_prompt:
prompts = [sampling_param.negative_prompt] + prompts
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(prompts),
desc="Processing validation prompts",
@@ -36,6 +36,7 @@ class ConditioningStage(PipelineStage):
Returns:
The batch with applied conditioning.
"""
# TODO!!
if not batch.do_classifier_free_guidance:
return batch
else:
+11 -8
View File
@@ -47,14 +47,17 @@ class DenoisingStage(PipelineStage):
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
# when used for validation, transformer is None as it is taking from the
# training loop
if transformer is not None:
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
def forward(
self,
@@ -30,6 +30,7 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# We use UniPCMScheduler from Wan2.1 official repo, not the one in diffusers.
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
+16 -5
View File
@@ -178,7 +178,11 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
rank=self.rank,
world_size=self.world_size,
cfg_rate=training_args.cfg,
num_latent_t=training_args.num_latent_t)
num_latent_t=training_args.num_latent_t,
validation=True)
if sampling_param.negative_prompt:
_, negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
)
validation_dataloader = StatefulDataLoader(
validation_dataset,
@@ -194,6 +198,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Add the transformer to the validation pipeline
self.validation_pipeline.add_module("transformer", transformer)
# TODO(Peiyuan): those logic should be inside add_module
self.validation_pipeline.latent_preparation_stage.transformer = transformer # type: ignore[attr-defined]
self.validation_pipeline.denoising_stage.transformer = transformer # type: ignore[attr-defined]
@@ -221,24 +226,30 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
data_type="video",
latents=None,
seed=validation_seed, # Use deterministic seed
generator=torch.Generator(
device="cpu").manual_seed(validation_seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
# make sure we use the same height, width, and num_frames as the training pipeline
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
# TODO(will): validation_sampling_steps and
# validation_guidance_scale are actually passed in as a list of
# values, like "10,20,30". The validation should be run for each
# combination of values.
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=sampling_param.num_inference_steps,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=1,
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
do_classifier_free_guidance=False,
eta=0.0,
)
# Run validation inference
with torch.inference_mode(), torch.autocast("cuda",
dtype=torch.bfloat16):
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
@@ -198,8 +198,7 @@ class WanTrainingPipeline(TrainingPipeline):
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
noise_random_generator = torch.Generator(device="cpu")
noise_random_generator.manual_seed(seed)
noise_random_generator = torch.Generator(device="cpu").manual_seed(seed)
logger.info("Initialized random seeds with seed: %s", seed)
@@ -271,7 +270,7 @@ class WanTrainingPipeline(TrainingPipeline):
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage before train_one_step: %s MB",
gpu_memory_usage)
self._log_validation(self.transformer, self.training_args, 1)
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
@@ -2,8 +2,8 @@
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="your/path/to/Mixkit-Src/merge.txt"
OUTPUT_DIR="your/path"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol/latents"
VALIDATION_PATH="assets/prompt.txt"
torchrun --nproc_per_node=$GPU_NUM \