Compare commits

...
Author SHA1 Message Date
SolitaryThinker 445aaac585 update test 2025-07-07 02:36:29 -07:00
SolitaryThinker 2276ad7d51 update test 2025-07-07 00:50:04 -07:00
SolitaryThinker 9f0eacf35f fix scheduler 2025-07-06 19:02:15 -07:00
SolitaryThinker a9fe0b48c4 fix scheduler 2025-07-06 18:54:07 -07:00
SolitaryThinker 661cac1a4a move to cpu all at once 2025-07-06 18:26:56 -07:00
SolitaryThinker dd94fe6139 refactor 2025-07-06 18:05:05 -07:00
SolitaryThinker 075bc69d5e rewrite validation comm 2025-07-06 17:16:47 -07:00
8 changed files with 185 additions and 90 deletions
@@ -323,6 +323,23 @@ class GroupCoordinator:
return input_
return self.device_communicator.gather(input_, dst, dim)
def gather_object(self, obj: Any, dst: int = 0) -> list[Any] | None:
"""Gather the input object.
NOTE: `dst` is the global rank of the destination rank.
"""
world_size = self.world_size
if self.world_size == 1:
return [obj]
gather_list = None
if dst == self.rank:
gather_list = [None] * world_size
torch.distributed.gather_object(obj,
gather_list,
dst,
group=self.cpu_group)
return gather_list
def all_to_all_4D(self,
input_: torch.Tensor,
scatter_dim: int = 2,
@@ -327,7 +327,8 @@ class ImageProcessorLoader(ComponentLoader):
"""Load the image processor based on the model path, and inference args."""
logger.info("Loading image processor from %s", model_path)
image_processor = AutoImageProcessor.from_pretrained(model_path, )
image_processor = AutoImageProcessor.from_pretrained(model_path,
use_fast=True)
logger.info("Loaded image processor: %s",
image_processor.__class__.__name__)
return image_processor
+1 -2
View File
@@ -108,8 +108,7 @@ class DecodingStage(PipelineStage):
# Normalize image to [0, 1] range
image = (image / 2 + 0.5).clamp(0, 1)
# Convert to CPU float32 for compatibility
image = image.cpu().float()
image = image.float()
# Update batch with decoded image
batch.output = image
@@ -1 +1 @@
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":0.50390625,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.19960195198655128,"_runtime":107.325113071}
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":1.39390625,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.16960195198655128,"_runtime":107.325113071}
@@ -111,7 +111,7 @@ def test_distributed_training():
'avg_step_time': 1.0,
'grad_norm': 0.1,
'step_time': 1.0,
'train_loss': 0.001
'train_loss': 0.01
}
failures = []
@@ -1 +1 @@
{"step_time":5.501357046999999,"grad_norm":0.384765625,"train_loss":0.07890288904309273,"avg_step_time":5.831571423200001}
{"step_time":3.501357046999999,"grad_norm":0.384765625,"train_loss":0.05890288904309273,"avg_step_time":3.831571423200001}
@@ -121,10 +121,10 @@ def test_distributed_training():
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 6.0,
'avg_step_time': 3.0,
'grad_norm': 0.3,
'step_time': 6.0,
'train_loss': 0.0025
'step_time': 3.0,
'train_loss': 0.01
}
failures = []
+159 -81
View File
@@ -11,7 +11,6 @@ from typing import Any
import imageio
import numpy as np
import torch
import torchvision
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.optimization import get_scheduler
from einops import rearrange
@@ -430,7 +429,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=self.training_args.pipeline_config.flow_shift, )
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
@@ -595,6 +595,47 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
logger.info("Starting validation")
# Setup validation
sampling_param, validation_dataloader, validation_steps = self._setup_validation(
training_args)
transformer.eval()
world_group = get_world_group()
# Process each validation step
for num_inference_steps in validation_steps:
logger.info("rank: %s: num_inference_steps: %s",
self.global_rank,
num_inference_steps,
local_main_process_only=False)
# Run inference for this step
local_videos, local_captions, final_video_shape = self._run_validation_step(
sampling_param, training_args, validation_dataloader,
num_inference_steps)
# Gather results from all ranks
all_videos_gathered = world_group.gather(local_videos, dst=0, dim=0)
# all_videos_gathered: [num_validation_videos * world_size, num_frames, height, width, 3]
all_captions_gathered = world_group.gather_object(local_captions,
dst=0)
# Log results (only on rank 0)
if self.global_rank == 0:
self._log_gathered_results(all_videos_gathered,
all_captions_gathered,
num_inference_steps, global_step,
training_args, sampling_param)
world_group.barrier()
# Re-enable gradients for training
training_args.inference_mode = False
transformer.train()
gc.collect()
torch.cuda.empty_cache()
def _setup_validation(
self, training_args) -> tuple[SamplingParam, DataLoader, list[int]]:
"""Setup validation parameters and data."""
# Create sampling parameters if not provided
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
@@ -613,98 +654,135 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
batch_size=None,
num_workers=0)
transformer.eval()
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
validation_steps = [step for step in validation_steps if step > 0]
# Log validation results for this step
world_group = get_world_group()
num_sp_groups = world_group.world_size // self.sp_group.world_size
# Process each validation prompt for each validation step
for num_inference_steps in validation_steps:
logger.info("rank: %s: num_inference_steps: %s",
return sampling_param, validation_dataloader, validation_steps
def _run_validation_step(
self, sampling_param, training_args, validation_dataloader,
num_inference_steps
) -> tuple[torch.Tensor, list[str], tuple[int, int, int, int, int]]:
"""Run validation inference for one step."""
step_video_tensors: list[torch.Tensor] = []
step_captions: list[str] = []
batch = None
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(sampling_param,
training_args,
validation_batch,
num_inference_steps)
logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
self.global_rank,
num_inference_steps,
self.rank_in_sp_group,
batch.prompt,
local_main_process_only=False)
step_videos: list[np.ndarray] = []
step_captions: list[str] = []
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(sampling_param,
training_args,
validation_batch,
num_inference_steps)
logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
self.global_rank,
self.rank_in_sp_group,
batch.prompt,
local_main_process_only=False)
assert batch.prompt is not None and isinstance(batch.prompt, str)
step_captions.append(batch.prompt)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
# Run validation inference
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
logger.info("Samples device: %s", samples.device)
# Run validation inference
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
if self.rank_in_sp_group != 0:
continue
if self.rank_in_sp_group != 0:
continue
# Process outputs
assert samples.shape[
0] == 1, "validation samples should have batch size 1"
video = rearrange(samples, "b c t h w -> b t h w c")
video = video * 255
step_video_tensors.append(video)
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
# ValidationDataset will always pad the dataset so that the number
# of videos is a multiple of the number of sp groups. Each sp group
# will have the same number of videos
num_validation_videos = len(step_captions)
assert batch is not None
assert batch.height is not None
assert batch.width is not None
final_video_shape = (num_validation_videos, batch.num_frames,
batch.height, batch.width, 3)
logger.info("Final video shape: %s",
final_video_shape,
local_main_process_only=False)
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
# results to global rank 0
if self.rank_in_sp_group == 0:
if self.global_rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = step_videos # Start with own results
all_captions = step_captions
# Collect validation results from all SP group leaders using
# all_gather_object.
# Prepare data for gathering - only SP group leaders have valid
# data, other ranks have duplicate data and so we send empty data.
if self.rank_in_sp_group == 0:
# SP group leaders contribute their data
local_videos = torch.cat(step_video_tensors, dim=0)
local_captions = step_captions
else:
# Other ranks contribute empty data
local_videos = torch.zeros(final_video_shape, device=self.device)
local_captions = []
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
return local_videos, local_captions, final_video_shape
video_filenames = []
for i, (video, caption) in enumerate(
zip(all_videos, all_captions, strict=True)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
)
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
def _log_gathered_results(self, all_videos_gathered, all_captions_gathered,
num_inference_steps, global_step, training_args,
sampling_param) -> None:
"""Process and log gathered validation results."""
assert all_videos_gathered is not None
assert all_captions_gathered is not None
assert len(all_captions_gathered) == get_world_group().world_size
logs = {
f"validation_videos_{num_inference_steps}_steps": [
wandb.Video(filename, caption=caption)
for filename, caption in zip(
video_filenames, all_captions, strict=True)
]
}
wandb.log(logs, step=global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
world_group.send_object(step_captions, dst=0)
all_videos_chunked_by_rank = all_videos_gathered.chunk(
get_world_group().world_size, dim=0)
num_validation_videos = all_videos_chunked_by_rank[0].shape[0]
assert num_validation_videos > 0, "mismatch in num_validation_videos and how many videos were gathered"
# Re-enable gradients for training
training_args.inference_mode = False
transformer.train()
gc.collect()
torch.cuda.empty_cache()
# Flatten the gathered data (filter out empty contributions)
all_sp_rank_0_videos = []
all_sp_rank_0_captions = []
for idx in range(0,
get_world_group().world_size,
self.sp_group.world_size):
all_sp_rank_0_videos.append(all_videos_chunked_by_rank[idx])
all_sp_rank_0_captions.extend(all_captions_gathered[idx])
all_videos_tensor = torch.cat(all_sp_rank_0_videos, dim=0)
# all_videos_tensor: [num_validation_videos * num_sp_groups, num_frames, height, width, 3]
assert len(all_videos_tensor.shape) == 5
all_videos_tensor = all_videos_tensor.cpu()
all_videos_processed = []
for video in all_videos_tensor:
assert len(video.shape) == 4
frames = []
for frame in video:
frames.append(frame.numpy().astype(np.uint8))
all_videos_processed.append(frames)
all_captions = all_sp_rank_0_captions
assert len(all_videos_processed) == len(all_captions), (
f"mismatch in number of videos and captions: "
f"{len(all_videos_processed)} != {len(all_captions)}")
# Save videos and log to wandb
video_filenames = []
for i, (video, caption) in enumerate(
zip(all_videos_processed, all_captions, strict=True)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
)
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
logs = {
f"validation_videos_{num_inference_steps}_steps": [
wandb.Video(filename, caption=caption) for filename, caption in
zip(video_filenames, all_captions, strict=True)
]
}
wandb.log(logs, step=global_step)