Skip to content
This repository was archived by the owner on May 20, 2026. It is now read-only.

Commit 7a2715a

Browse files
committed
Adding dataset files for Wan
Signed-off-by: Pranav Prashant Thombre <pthombre@nvidia.com>
1 parent 394eb99 commit 7a2715a

5 files changed

Lines changed: 253 additions & 9 deletions

File tree

dfm/examples/Automodel/finetune/wan2_1_t2v_flow.yaml

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@ model:
1414

1515
data:
1616
dataloader:
17-
_target_: nemo_automodel.components.datasets.diffusion.build_wan21_dataloader
17+
_target_: Automodel.datasets.build_wan21_dataloader
1818
meta_folder: /lustre/fsw/portfolios/coreai/users/linnanw/hdvilla_sample/pika/wan21_codes/1.3B_meta/
1919
batch_size: 1
2020
num_workers: 2
@@ -50,12 +50,12 @@ fsdp:
5050
dp_size: 8
5151

5252
logging:
53-
save_every: 50
53+
save_every: 1000
5454
log_every: 2
5555

5656
checkpoint:
5757
enabled: true
58-
checkpoint_dir: /opt/DFM/wan_t2v_flow_outputs_base_recipe_dfm/
58+
checkpoint_dir: /opt/DFM/wan_t2v_flow_outputs_base_recipe_dfm_test_2/
5959
model_save_format: torch_save
6060
save_consolidated: false
6161
restore_from: null

dfm/src/Automodel/_diffusers/auto_diffusion_pipeline.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -116,7 +116,7 @@ def from_pretrained(
116116
load_for_training: bool = False,
117117
components_to_load: Optional[Iterable[str]] = None,
118118
**kwargs,
119-
) -> DiffusionPipeline:
119+
) -> tuple[DiffusionPipeline, Dict[str, FSDP2Manager]]:
120120
pipe: DiffusionPipeline = DiffusionPipeline.from_pretrained(
121121
pretrained_model_name_or_path,
122122
*model_args,
Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,30 @@
1+
# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
2+
#
3+
# Licensed under the Apache License, Version 2.0 (the "License");
4+
# you may not use this file except in compliance with the License.
5+
# You may obtain a copy of the License at
6+
#
7+
# http://www.apache.org/licenses/LICENSE-2.0
8+
#
9+
# Unless required by applicable law or agreed to in writing, software
10+
# distributed under the License is distributed on an "AS IS" BASIS,
11+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
# See the License for the specific language governing permissions and
13+
# limitations under the License.
14+
15+
from Automodel.datasets.wan21 import (
16+
MetaFilesDataset,
17+
build_node_parallel_sampler,
18+
build_wan21_dataloader,
19+
collate_fn,
20+
create_dataloader,
21+
)
22+
23+
24+
__all__ = [
25+
"MetaFilesDataset",
26+
"build_node_parallel_sampler",
27+
"build_wan21_dataloader",
28+
"collate_fn",
29+
"create_dataloader",
30+
]
Lines changed: 217 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,217 @@
1+
# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
2+
#
3+
# Licensed under the Apache License, Version 2.0 (the "License");
4+
# you may not use this file except in compliance with the License.
5+
# You may obtain a copy of the License at
6+
#
7+
# http://www.apache.org/licenses/LICENSE-2.0
8+
#
9+
# Unless required by applicable law or agreed to in writing, software
10+
# distributed under the License is distributed on an "AS IS" BASIS,
11+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
# See the License for the specific language governing permissions and
13+
# limitations under the License.
14+
15+
from __future__ import annotations
16+
17+
import logging
18+
import os
19+
import pickle
20+
from pathlib import Path
21+
from typing import Callable, Dict, List, Optional, Tuple
22+
23+
import torch
24+
import torch.distributed as dist
25+
from torch.utils.data import DataLoader, Dataset, DistributedSampler
26+
27+
28+
logger = logging.getLogger(__name__)
29+
30+
31+
class MetaFilesDataset(Dataset):
32+
"""PyTorch dataset for WAN2.1 `.meta` files."""
33+
34+
def __init__(
35+
self,
36+
meta_folder: str,
37+
transform_text: Optional[Callable[[torch.Tensor], torch.Tensor]] = None,
38+
transform_video: Optional[Callable[[torch.Tensor], torch.Tensor]] = None,
39+
filter_fn: Optional[Callable[[Dict], bool]] = None,
40+
device: str = "cpu",
41+
max_files: Optional[int] = None,
42+
) -> None:
43+
self.meta_folder = Path(meta_folder)
44+
self.transform_text = transform_text
45+
self.transform_video = transform_video
46+
self.filter_fn = filter_fn
47+
self.device = device
48+
49+
self.meta_files = sorted(self.meta_folder.glob("*.meta"))
50+
if max_files is None:
51+
max_files_env = os.environ.get("MAX_META_FILES")
52+
if max_files_env is not None:
53+
try:
54+
max_files = int(max_files_env)
55+
except ValueError:
56+
logger.warning("Invalid MAX_META_FILES=%s", max_files_env)
57+
58+
if max_files is not None and max_files > 0:
59+
self.meta_files = self.meta_files[:max_files]
60+
logger.info("Limited to first %d meta files", len(self.meta_files))
61+
62+
if not self.meta_files:
63+
raise ValueError(f"No .meta files found in {meta_folder}")
64+
65+
if self.filter_fn:
66+
filtered = []
67+
for path in self.meta_files:
68+
try:
69+
with open(path, "rb") as f:
70+
data = pickle.load(f)
71+
except Exception as exc: # pragma: no cover - best effort logging
72+
logger.warning("Failed to load %s during filtering: %s", path, exc)
73+
continue
74+
if self.filter_fn(data.get("metadata", {})):
75+
filtered.append(path)
76+
self.meta_files = filtered
77+
logger.info("Filtered meta files count: %d", len(self.meta_files))
78+
79+
self._log_dataset_stats()
80+
81+
def _log_dataset_stats(self) -> None:
82+
sample_paths = self.meta_files[: min(5, len(self.meta_files))]
83+
stats: List[Tuple[torch.Size, torch.Size, str]] = []
84+
for path in sample_paths:
85+
try:
86+
with open(path, "rb") as f:
87+
data = pickle.load(f)
88+
stats.append(
89+
(
90+
data["text_embeddings"].shape,
91+
data["video_latents"].shape,
92+
str(data.get("deterministic_latents", "unknown")),
93+
)
94+
)
95+
except Exception as exc: # pragma: no cover - stats only
96+
logger.debug("Failed to sample %s: %s", path, exc)
97+
98+
if stats:
99+
text_shapes, video_shapes, modes = zip(*stats, strict=False)
100+
logger.info("Sample text embeddings: %s", text_shapes)
101+
logger.info("Sample video latents: %s", video_shapes)
102+
logger.info("Sample encoding modes: %s", set(modes))
103+
104+
def __len__(self) -> int:
105+
return len(self.meta_files)
106+
107+
def __getitem__(self, index: int) -> Dict[str, torch.Tensor]: # type: ignore[override]
108+
path = self.meta_files[index]
109+
with open(path, "rb") as f:
110+
data = pickle.load(f)
111+
112+
text_embeddings: torch.Tensor = data["text_embeddings"].to(self.device)
113+
video_latents: torch.Tensor = data["video_latents"].to(self.device)
114+
115+
if self.transform_text is not None:
116+
text_embeddings = self.transform_text(text_embeddings)
117+
if self.transform_video is not None:
118+
video_latents = self.transform_video(video_latents)
119+
120+
file_info = {
121+
"meta_filename": Path(path).name,
122+
"original_filename": data.get("original_filename", "unknown"),
123+
"original_video_path": data.get("original_video_path", "unknown"),
124+
"deterministic_latents": data.get("deterministic_latents", "unknown"),
125+
"memory_optimization": data.get("memory_optimization", "unknown"),
126+
"num_frames": data.get("num_frames", "unknown"),
127+
}
128+
129+
return {
130+
"text_embeddings": text_embeddings,
131+
"video_latents": video_latents,
132+
"metadata": data.get("metadata", {}),
133+
"file_info": file_info,
134+
}
135+
136+
137+
def collate_fn(batch: List[Dict[str, torch.Tensor]]) -> Dict[str, torch.Tensor]:
138+
text_embeddings = torch.stack([item["text_embeddings"] for item in batch])
139+
video_latents = torch.stack([item["video_latents"] for item in batch])
140+
return {
141+
"text_embeddings": text_embeddings,
142+
"video_latents": video_latents,
143+
"metadata": [item["metadata"] for item in batch],
144+
"file_info": [item["file_info"] for item in batch],
145+
}
146+
147+
148+
def build_node_parallel_sampler(
149+
dataset: Dataset,
150+
num_nodes: Optional[int] = None,
151+
shuffle: bool = True,
152+
) -> Optional[DistributedSampler]:
153+
if not dist.is_initialized():
154+
return None
155+
156+
world_size = dist.get_world_size()
157+
local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", world_size))
158+
local_world_size = max(local_world_size, 1)
159+
if num_nodes is None:
160+
num_nodes = max(1, world_size // local_world_size)
161+
162+
node_rank = dist.get_rank() // local_world_size
163+
replicas = num_nodes
164+
165+
return DistributedSampler(
166+
dataset,
167+
num_replicas=replicas,
168+
rank=node_rank,
169+
shuffle=shuffle,
170+
drop_last=False,
171+
)
172+
173+
174+
def build_wan21_dataloader(
175+
*,
176+
meta_folder: str,
177+
batch_size: int,
178+
shuffle: bool = True,
179+
num_workers: int = 2,
180+
device: str = "cpu",
181+
transform_text: Optional[Callable[[torch.Tensor], torch.Tensor]] = None,
182+
transform_video: Optional[Callable[[torch.Tensor], torch.Tensor]] = None,
183+
filter_fn: Optional[Callable[[Dict], bool]] = None,
184+
max_files: Optional[int] = None,
185+
num_nodes: Optional[int] = None,
186+
) -> Tuple[DataLoader, Optional[DistributedSampler]]:
187+
dataset = MetaFilesDataset(
188+
meta_folder=meta_folder,
189+
transform_text=transform_text,
190+
transform_video=transform_video,
191+
filter_fn=filter_fn,
192+
device=device,
193+
max_files=max_files,
194+
)
195+
196+
sampler = build_node_parallel_sampler(dataset, num_nodes, shuffle=shuffle)
197+
198+
use_pin_memory = device == "cpu"
199+
dataloader = DataLoader(
200+
dataset,
201+
batch_size=batch_size,
202+
shuffle=(sampler is None and shuffle),
203+
sampler=sampler,
204+
num_workers=num_workers,
205+
collate_fn=collate_fn,
206+
pin_memory=use_pin_memory,
207+
)
208+
209+
return dataloader, sampler
210+
211+
212+
def create_dataloader(
213+
meta_folder: str,
214+
batch_size: int,
215+
num_nodes: int,
216+
) -> Tuple[DataLoader, Optional[DistributedSampler]]:
217+
return build_wan21_dataloader(meta_folder=meta_folder, batch_size=batch_size, num_nodes=num_nodes)

dfm/src/Automodel/recipes/finetune.py

Lines changed: 2 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,6 @@
2121

2222
import torch
2323
import torch.distributed as dist
24-
import wandb
2524
from Automodel._diffusers.auto_diffusion_pipeline import NeMoAutoDiffusionPipeline
2625
from Automodel.flow_matching.training_step_t2v import (
2726
step_fsdp_transformer_t2v,
@@ -36,6 +35,8 @@
3635
from torch.distributed.fsdp import MixedPrecisionPolicy
3736
from transformers.utils.hub import TRANSFORMERS_CACHE
3837

38+
import wandb
39+
3940

4041
def build_model_and_optimizer(
4142
*,
@@ -462,10 +463,6 @@ def run_train_validation_loop(self):
462463
if is_main_process() and wandb.run is not None:
463464
wandb.log({"epoch/avg_loss": avg_loss, "epoch/num": epoch + 1}, step=global_step)
464465

465-
logging.info("[INFO] Training complete, saving final checkpoint...")
466-
467-
self.save_checkpoint(epoch=self.step_scheduler.epoch, step=global_step)
468-
469466
if is_main_process():
470467
logging.info(f"[INFO] Saved final checkpoint at step {global_step}")
471468
if wandb.run is not None:

0 commit comments

Comments
 (0)