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

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. 

16 

17Composition pattern: holds a BaseTrainer and overrides data pipeline steps. 

18""" 

19 

20import glob 

21import logging 

22import os 

23import torch 

24from torch.utils.data import Dataset, DataLoader, DistributedSampler 

25 

26from hyper_parallel.trainer.base import BaseTrainer 

27 

28logger = logging.getLogger(__name__) 

29 

30 

31class DiTTrainer: 

32 """Trainer for Qwen-Image DiT diffusion training. 

33 

34 Supports: 

35 - ``data.type = "dummy_dit"``: deterministic synthetic diffusion tensors 

36 for quick single-card validation and loss alignment. 

37 

38 Composition pattern — delegates training loop to BaseTrainer. 

39 """ 

40 

41 def __init__(self, args): 

42 self.base = BaseTrainer(args) 

43 

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() 

60 

61 # ------------------------------------------------------------------ 

62 # Overridden _build_* methods 

63 # ------------------------------------------------------------------ 

64 

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 

69 

70 def _build_data_transform(self): 

71 """No offline data transform for dummy tensors.""" 

72 self.base.data_transform = None 

73 

74 def _build_dataset(self): 

75 """Build deterministic dummy DiT dataset with real diffusion targets. 

76 

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") 

81 

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) 

87 

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 ) 

98 

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 

107 

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 }) 

123 

124 samples = samples * 2 

125 

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)] 

133 

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 ) 

140 

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}") 

149 

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 

158 

159 def __len__(self): 

160 return len(self.files) 

161 

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 } 

180 

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 ) 

187 

188 def _build_parquet_dataset(self, max_steps): 

189 """Read pre-generated data from parquet via HuggingFace Datasets. 

190 

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 

199 

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") 

203 

204 hf_ds = load_dataset("parquet", data_files=parquet_path, split="train") 

205 

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 } 

224 

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 ) 

230 

231 def _build_collate_fn(self): 

232 """Stack fixed-size tensors (no padding needed for dummy data).""" 

233 

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 } 

242 

243 self.base.collate_fn = _dit_collate 

244 

245 def _build_dataloader(self): 

246 """Build DataLoader with no distributed sharding. 

247 

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. 

251 

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 

262 

263 # 然后替换 sampler 为不分片的 

264 old_dl = self.base.train_dataloader 

265 dataset = old_dl.dataset 

266 batch_size = old_dl.batch_size 

267 

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 ) 

275 

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 ) 

283 

284 def _build_optimizer(self): 

285 """Override base._build_optimizer to use fused AdamW (matches VeOmni). 

286 

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) 

294 

295 decay_keywords = ("bias", "layernorm", "norm", "rmsnorm") 

296 

297 def _is_no_decay(name: str) -> bool: 

298 lname = name.lower() 

299 return any(kw in lname for kw in decay_keywords) 

300 

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) 

314 

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 ) 

334 

335 # ------------------------------------------------------------------ 

336 # Delegated methods 

337 # ------------------------------------------------------------------ 

338 

339 def train(self): 

340 """Delegate to BaseTrainer.train().""" 

341 return self.base.train()