Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
445aaac585 | ||
|
|
2276ad7d51 | ||
|
|
9f0eacf35f | ||
|
|
a9fe0b48c4 | ||
|
|
661cac1a4a | ||
|
|
dd94fe6139 | ||
|
|
075bc69d5e |
@@ -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
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user