Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / trainer / dit_trainer.py: 0%
164 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-21 04:29 +0800
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-21 04:29 +0800
1# Copyright 2026 Huawei Technologies Co., Ltd
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"""DiTTrainer — Diffusion Transformer training for Qwen-Image.
17Composition pattern: holds a BaseTrainer and overrides data pipeline steps.
18"""
20import glob
21import logging
22import os
23import torch
24from torch.utils.data import Dataset, DataLoader, DistributedSampler
26from hyper_parallel.trainer.base import BaseTrainer
28logger = logging.getLogger(__name__)
31class DiTTrainer:
32 """Trainer for Qwen-Image DiT diffusion training.
34 Supports:
35 - ``data.type = "dummy_dit"``: deterministic synthetic diffusion tensors
36 for quick single-card validation and loss alignment.
38 Composition pattern — delegates training loop to BaseTrainer.
39 """
41 def __init__(self, args):
42 self.base = BaseTrainer(args)
44 # 13-step init — call base methods, override data steps
45 self.base._setup()
46 self.base._build_model()
47 # 注意:不要手动 .to(device),_build_parallelized_model 会处理 meta → npu
48 self.base._freeze_model()
49 self._build_model_assets()
50 self._build_data_transform()
51 self._build_dataset()
52 self._build_collate_fn()
53 self._build_dataloader()
54 self.base._build_parallelized_model()
55 self._build_optimizer()
56 self.base._build_lr_scheduler()
57 self.base._build_training_context()
58 self.base._init_callbacks()
59 self.base.on_init_end()
61 # ------------------------------------------------------------------
62 # Overridden _build_* methods
63 # ------------------------------------------------------------------
65 def _build_model_assets(self):
66 """DiT does not need tokenizer or processor for dummy data."""
67 self.base.tokenizer = None
68 self.base.processor = None
70 def _build_data_transform(self):
71 """No offline data transform for dummy tensors."""
72 self.base.data_transform = None
74 def _build_dataset(self):
75 """Build deterministic dummy DiT dataset with real diffusion targets.
77 Pre-generates a fixed data pool using a fixed seed so that
78 single-card and multi-card training see identical samples.
79 """
80 data_type = getattr(self.base.args.data, "type", "dummy_dit")
82 model_cfg = self.base.args.model
83 cond_dim = getattr(model_cfg, "joint_attention_dim", 3584)
84 seq_len = 77
85 base_seed = int(getattr(self.base.args, "seed", 42))
86 max_steps = getattr(self.base.args.train, "max_steps", 100)
88 if data_type == "coco_parquet":
89 self._build_parquet_dataset(max_steps)
90 elif data_type == "coco_dit":
91 self._build_coco_dataset(cond_dim, seq_len, base_seed, max_steps)
92 elif data_type == "dummy_dit":
93 self._build_dummy_dataset(model_cfg, cond_dim, seq_len, base_seed, max_steps)
94 else:
95 raise NotImplementedError(
96 f"DiTTrainer supports 'dummy_dit'/'coco_dit'/'coco_parquet', got '{data_type}'"
97 )
99 def _build_dummy_dataset(self, model_cfg, cond_dim, seq_len, base_seed, max_steps):
100 """Deterministic dummy DiT dataset with flow-matching targets."""
101 in_ch = getattr(model_cfg, "in_channels", 64)
102 out_ch = getattr(model_cfg, "out_channels", 64)
103 height = getattr(model_cfg, "height", 256)
104 width = getattr(model_cfg, "width", 256)
105 latent_h = height // 8
106 latent_w = width // 8
108 g = torch.Generator().manual_seed(base_seed)
109 samples = []
110 for _ in range(100):
111 clean = torch.randn(in_ch, latent_h, latent_w, generator=g)
112 eps = torch.randn(out_ch, latent_h, latent_w, generator=g)
113 condition = torch.randn(seq_len, cond_dim, generator=g)
114 ts = torch.randint(1, 1000, (1,), generator=g).squeeze(0)
115 t_norm = ts.float() / 1000.0
116 x_t = (1.0 - t_norm) * clean + t_norm * eps
117 velocity = eps - clean
118 samples.append({
119 "latent": x_t, "timestep": ts,
120 "condition": condition, "target_noise": velocity,
121 "labels": torch.tensor(1, dtype=torch.long),
122 })
124 samples = samples * 2
126 class DummyDiTDataset(Dataset):
127 def __init__(self, samples):
128 self.samples = samples
129 def __len__(self):
130 return len(self.samples)
131 def __getitem__(self, idx):
132 return self.samples[idx % len(self.samples)]
134 self.base.train_dataset = DummyDiTDataset(samples)
135 self.base.state.max_steps = max_steps
136 logger.info_rank0(
137 f"DiT dummy dataset: {len(samples)} samples, "
138 f"latent=({in_ch},{latent_h},{latent_w}) cond=({seq_len},{cond_dim})"
139 )
141 def _build_coco_dataset(self, cond_dim, seq_len, base_seed, max_steps):
142 """Real COCO images (packed VAE latents) + dummy text embeddings."""
143 cache_path = getattr(self.base.args.data, "train_path", None)
144 if not cache_path:
145 raise ValueError("data.train_path must be set for data.type=coco_dit")
146 pth_files = sorted(glob.glob(os.path.join(cache_path, "*.pth")))
147 if not pth_files:
148 raise FileNotFoundError(f"No .pth files found in {cache_path}")
150 class CocoDiTDataset(Dataset):
151 """Map-style dataset over packed-VAE .pth files with on-the-fly
152 flow-matching noise targets."""
153 def __init__(self, files, cond_dim, seq_len, seed):
154 self.files = files
155 self.cond_dim = cond_dim
156 self.seq_len = seq_len
157 self.seed = seed
159 def __len__(self):
160 return len(self.files)
162 def __getitem__(self, idx):
163 data = torch.load(self.files[idx])
164 clean = data["latent_clean"]
165 g = torch.Generator().manual_seed(self.seed + idx)
166 eps = torch.randn(*clean.shape, generator=g)
167 ts = torch.randint(1, 1000, (1,), generator=g).squeeze(0)
168 t_norm = ts.float() / 1000.0
169 x_t = (1.0 - t_norm) * clean + t_norm * eps
170 velocity = eps - clean
171 if "text_embed" in data and data["text_embed"] is not None:
172 condition = data["text_embed"]
173 else:
174 condition = torch.randn(self.seq_len, self.cond_dim, generator=g)
175 return {
176 "latent": x_t, "timestep": ts,
177 "condition": condition, "target_noise": velocity,
178 "labels": torch.tensor(1, dtype=torch.long),
179 }
181 self.base.train_dataset = CocoDiTDataset(pth_files, cond_dim, seq_len, base_seed)
182 self.base.state.max_steps = max_steps
183 logger.info_rank0(
184 f"COCO dataset: {len(pth_files)} samples from {cache_path}, "
185 f"cond=({seq_len},{cond_dim})"
186 )
188 def _build_parquet_dataset(self, max_steps):
189 """Read pre-generated data from parquet via HuggingFace Datasets.
191 Both HP and VeOmni load the same parquet through HuggingFace Datasets,
192 ensuring byte-identical training inputs for cross-framework alignment.
193 """
194 try:
195 from datasets import load_dataset # pylint: disable=import-outside-toplevel
196 import pickle as pk # pylint: disable=import-outside-toplevel
197 except ImportError as exc:
198 raise ImportError("datasets package required: pip install datasets") from exc
200 parquet_path = getattr(self.base.args.data, "train_path", None)
201 if not parquet_path:
202 raise ValueError("data.train_path must point to a .parquet file for data.type=coco_parquet")
204 hf_ds = load_dataset("parquet", data_files=parquet_path, split="train")
206 class ParquetDataset(Dataset):
207 """Map-style dataset reading pre-generated rows from a parquet file."""
208 def __init__(self, hf_ds):
209 self.hf_ds = hf_ds
210 def __len__(self):
211 return len(self.hf_ds)
212 def __getitem__(self, idx):
213 row = self.hf_ds[idx]
214 hidden = pk.loads(row["hidden_states"]).squeeze(0)
215 target = pk.loads(row["training_target"]).squeeze(0)
216 c, h, w = 64, 16, 16
217 return {
218 "latent": hidden.reshape(h, w, c).permute(2, 0, 1).float(),
219 "timestep": pk.loads(row["timestep"]).squeeze(0),
220 "condition": pk.loads(row["encoder_hidden_states"]).squeeze(0).float(),
221 "target_noise": target.reshape(h, w, c).permute(2, 0, 1).float(),
222 "labels": torch.tensor(1, dtype=torch.long),
223 }
225 self.base.train_dataset = ParquetDataset(hf_ds)
226 self.base.state.max_steps = max_steps
227 logger.info_rank0(
228 f"Parquet dataset: {len(hf_ds)} samples from {parquet_path}"
229 )
231 def _build_collate_fn(self):
232 """Stack fixed-size tensors (no padding needed for dummy data)."""
234 def _dit_collate(batch):
235 return {
236 "latent": torch.stack([x["latent"] for x in batch]),
237 "timestep": torch.stack([x["timestep"] for x in batch]),
238 "condition": torch.stack([x["condition"] for x in batch]),
239 "target_noise": torch.stack([x["target_noise"] for x in batch]),
240 "labels": torch.stack([x["labels"] for x in batch]),
241 }
243 self.base.collate_fn = _dit_collate
245 def _build_dataloader(self):
246 """Build DataLoader with no distributed sharding.
248 All ranks (single-card or multi-card) iterate over the *same*
249 samples in the *same* order, which is required for step-by-step
250 loss alignment between single-card and FSDP training.
252 Sampler uses ``shuffle=True`` with ``seed`` from the config so that
253 the index sequence matches VeOmni's
254 ``StatefulDistributedSampler`` (which inherits
255 ``torch.utils.data.distributed.DistributedSampler``); both produce
256 the same ``torch.randperm(len(dataset), generator=g.manual_seed(seed))``
257 sequence, ensuring per-step timestep/noise inputs are byte-identical
258 for cross-framework loss alignment.
259 """
260 # 先调用框架默认方法(正确设置 _grad_accum 等属性)
261 self.base._build_dataloader() # pylint: disable=protected-access
263 # 然后替换 sampler 为不分片的
264 old_dl = self.base.train_dataloader
265 dataset = old_dl.dataset
266 batch_size = old_dl.batch_size
268 sampler = DistributedSampler(
269 dataset,
270 num_replicas=1, # 关键:不分片,所有 rank 按同样顺序取
271 rank=0,
272 shuffle=True,
273 seed=int(getattr(self.base.args.train, "seed", 42)),
274 )
276 self.base.train_dataloader = DataLoader(
277 dataset,
278 batch_size=batch_size,
279 sampler=sampler,
280 collate_fn=old_dl.collate_fn,
281 drop_last=old_dl.drop_last,
282 )
284 def _build_optimizer(self):
285 """Override base._build_optimizer to use fused AdamW (matches VeOmni).
287 Replicates BaseTrainer._build_optimizer but constructs ``torch.optim.AdamW``
288 with ``fused=True, foreach=False`` to match VeOmni's
289 ``build_optimizer(fused=True)`` path, ensuring identical optimizer
290 numerics across frameworks for loss alignment.
291 """
292 lr = getattr(self.base.args.train.optimizer, 'lr', 1e-4)
293 weight_decay = getattr(self.base.args.train.optimizer, 'weight_decay', 0.01)
295 decay_keywords = ("bias", "layernorm", "norm", "rmsnorm")
297 def _is_no_decay(name: str) -> bool:
298 lname = name.lower()
299 return any(kw in lname for kw in decay_keywords)
301 decay_params = []
302 no_decay_params = []
303 seen_ids = set()
304 for n, p in self.base.model.named_parameters():
305 if not p.requires_grad:
306 continue
307 if id(p) in seen_ids:
308 continue
309 seen_ids.add(id(p))
310 if _is_no_decay(n):
311 no_decay_params.append(p)
312 else:
313 decay_params.append(p)
315 param_groups = [
316 {"params": decay_params, "weight_decay": weight_decay},
317 {"params": no_decay_params, "weight_decay": 0.0},
318 ]
319 adam_eps = getattr(self.base.args.train.optimizer, 'eps', 1e-8)
320 adam_betas = getattr(self.base.args.train.optimizer, 'betas', (0.9, 0.999))
321 self.base.optimizer = torch.optim.AdamW(
322 param_groups,
323 lr=lr,
324 betas=adam_betas,
325 eps=adam_eps,
326 foreach=False,
327 fused=True,
328 )
329 logger.info(
330 "Optimizer (DiT override): AdamW fused=True lr=%.2e wd=%.3g "
331 "decay_params=%d no_decay_params=%d",
332 lr, weight_decay, len(decay_params), len(no_decay_params),
333 )
335 # ------------------------------------------------------------------
336 # Delegated methods
337 # ------------------------------------------------------------------
339 def train(self):
340 """Delegate to BaseTrainer.train()."""
341 return self.base.train()