BlueSourceJY commited on
Commit
8ea4401
·
verified ·
1 Parent(s): 8a6d907

Upload inference_rot_head.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. inference_rot_head.py +407 -0
inference_rot_head.py ADDED
@@ -0,0 +1,407 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Sampling Scripts of LightningDiT.
3
+
4
+ by Maple (Jingfeng Yao) from HUST-VL
5
+ """
6
+
7
+ import os, math, json, pickle, logging, argparse, yaml, torch, numpy as np
8
+ from time import time, strftime
9
+ from glob import glob
10
+ from copy import deepcopy
11
+ from collections import OrderedDict
12
+ from PIL import Image
13
+ from tqdm import tqdm
14
+ import torch.distributed as dist
15
+ from accelerate import Accelerator
16
+ from torch.utils.data import DataLoader
17
+ from torch.nn.parallel import DistributedDataParallel as DDP
18
+ from torch.utils.tensorboard import SummaryWriter
19
+ import torchvision
20
+ # local imports
21
+ from tokenizer.vavae import VA_VAE
22
+ from models.lightningdit_rot_head_vis import LightningDiT_models
23
+ from transport import create_transport, Sampler
24
+ from datasets.img_latent_dataset import ImgLatentDataset
25
+ from visualize_attention import visualize_attention_matrix
26
+
27
+ # sample function
28
+ def do_sample(train_config, accelerator, ckpt_path=None, cfg_scale=None, model=None, vae=None, demo_sample_mode=False):
29
+ """
30
+ Run sampling.
31
+ """
32
+
33
+ folder_name = f"{train_config['model']['model_type'].replace('/', '-')}-ckpt-{ckpt_path.split('/')[-1].split('.')[0]}-{train_config['sample']['sampling_method']}-{train_config['sample']['num_sampling_steps']}".lower()
34
+ # folder_name = "test_speed"
35
+ if cfg_scale is None:
36
+ cfg_scale = train_config['sample']['cfg_scale']
37
+ cfg_interval_start = train_config['sample']['cfg_interval_start'] if 'cfg_interval_start' in train_config['sample'] else 0
38
+ timestep_shift = train_config['sample']['timestep_shift'] if 'timestep_shift' in train_config['sample'] else 0
39
+ if cfg_scale > 1.0:
40
+ folder_name += f"-interval{cfg_interval_start:.2f}"+f"-cfg{cfg_scale:.2f}"
41
+ folder_name += f"-shift{timestep_shift:.2f}"
42
+
43
+ if demo_sample_mode:
44
+ cfg_interval_start = 0
45
+ timestep_shift = 0
46
+ cfg_scale = 9.0
47
+
48
+ sample_folder_dir = os.path.join(train_config['train']['output_dir'], train_config['train']['exp_name'], folder_name)
49
+ if accelerator.process_index == 0:
50
+ if not demo_sample_mode:
51
+ print_with_prefix('Sample_folder_dir=', sample_folder_dir)
52
+ print_with_prefix('ckpt_path=', ckpt_path)
53
+ print_with_prefix('cfg_scale=', cfg_scale)
54
+ print_with_prefix('cfg_interval_start=', cfg_interval_start)
55
+ print_with_prefix('timestep_shift=', timestep_shift)
56
+
57
+ if not os.path.exists(sample_folder_dir):
58
+ if accelerator.process_index == 0:
59
+ os.makedirs(sample_folder_dir, exist_ok=True)
60
+ else:
61
+ png_files = [f for f in os.listdir(sample_folder_dir) if f.endswith('.png')]
62
+ png_count = len(png_files)
63
+ if png_count > train_config['sample']['fid_num']:
64
+ if accelerator.process_index == 0:
65
+ print_with_prefix(f"Found {png_count} PNG files in {sample_folder_dir}, skip sampling.")
66
+ return sample_folder_dir
67
+
68
+ torch.backends.cuda.matmul.allow_tf32 = True # True: fast but may lead to some small numerical differences
69
+ assert torch.cuda.is_available(), "Sampling with DDP requires at least one GPU. sample.py supports CPU-only usage"
70
+ torch.set_grad_enabled(False)
71
+
72
+ # Setup accelerator:
73
+ device = accelerator.device
74
+
75
+ # Setup DDP:
76
+ device = accelerator.device
77
+ seed = train_config['train']['global_seed'] * accelerator.num_processes + accelerator.process_index
78
+ torch.manual_seed(seed)
79
+ # torch.cuda.set_device(device)
80
+ print_with_prefix(f"Starting rank={accelerator.local_process_index}, seed={seed}, world_size={accelerator.num_processes}.")
81
+ rank = accelerator.local_process_index
82
+
83
+ # Load model:
84
+ if 'downsample_ratio' in train_config['vae']:
85
+ downsample_ratio = train_config['vae']['downsample_ratio']
86
+ else:
87
+ downsample_ratio = 16
88
+ latent_size = train_config['data']['image_size'] // downsample_ratio
89
+
90
+ checkpoint = torch.load(ckpt_path, map_location=lambda storage, loc: storage)
91
+ if "ema" in checkpoint: # supports checkpoints from train.py
92
+ checkpoint = checkpoint["ema"]
93
+ model.load_state_dict(checkpoint)
94
+ model.eval() # important!
95
+ model.to(device)
96
+
97
+ transport = create_transport(
98
+ train_config['transport']['path_type'],
99
+ train_config['transport']['prediction'],
100
+ train_config['transport']['loss_weight'],
101
+ train_config['transport']['train_eps'],
102
+ train_config['transport']['sample_eps'],
103
+ use_cosine_loss = train_config['transport']['use_cosine_loss'] if 'use_cosine_loss' in train_config['transport'] else False,
104
+ use_lognorm = train_config['transport']['use_lognorm'] if 'use_lognorm' in train_config['transport'] else False,
105
+ ) # default: velocity;
106
+ sampler = Sampler(transport)
107
+ mode = train_config['sample']['mode']
108
+ if mode == "ODE":
109
+ sample_fn = sampler.sample_ode(
110
+ sampling_method=train_config['sample']['sampling_method'],
111
+ num_steps=train_config['sample']['num_sampling_steps'],
112
+ atol=train_config['sample']['atol'],
113
+ rtol=train_config['sample']['rtol'],
114
+ reverse=train_config['sample']['reverse'],
115
+ timestep_shift=timestep_shift,
116
+ )
117
+ else:
118
+ raise NotImplementedError(f"Sampling mode {mode} is not supported.")
119
+
120
+ if vae is None:
121
+ vae = VA_VAE(
122
+ f'tokenizer/configs/{train_config["vae"]["model_name"]}.yaml',
123
+ )
124
+ if accelerator.process_index == 0:
125
+ print_with_prefix('Loaded VAE model')
126
+
127
+ using_cfg = cfg_scale > 1.0
128
+ if using_cfg:
129
+ if accelerator.process_index == 0:
130
+ print_with_prefix('Using cfg:', using_cfg)
131
+
132
+ if rank == 0:
133
+ os.makedirs(sample_folder_dir, exist_ok=True)
134
+ if accelerator.process_index == 0 and not demo_sample_mode:
135
+ print_with_prefix(f"Saving .png samples at {sample_folder_dir}")
136
+ accelerator.wait_for_everyone()
137
+
138
+ # Figure out how many samples we need to generate on each GPU and how many iterations we need to run:
139
+ n = train_config['sample']['per_proc_batch_size']
140
+ global_batch_size = n * accelerator.num_processes
141
+ # To make things evenly-divisible, we'll sample a bit more than we need and then discard the extra samples:
142
+ num_samples = len([name for name in os.listdir(sample_folder_dir) if (os.path.isfile(os.path.join(sample_folder_dir, name)) and ".png" in name)])
143
+ total_samples = int(math.ceil(train_config['sample']['fid_num'] / global_batch_size) * global_batch_size)
144
+ if rank == 0:
145
+ if accelerator.process_index == 0:
146
+ print_with_prefix(f"Total number of images that will be sampled: {total_samples}")
147
+ assert total_samples % accelerator.num_processes == 0, "total_samples must be divisible by world_size"
148
+ samples_needed_this_gpu = int(total_samples // accelerator.num_processes)
149
+ assert samples_needed_this_gpu % n == 0, "samples_needed_this_gpu must be divisible by the per-GPU batch size"
150
+ iterations = int(samples_needed_this_gpu // n)
151
+ done_iterations = int( int(num_samples // accelerator.num_processes) // n)
152
+ pbar = range(iterations)
153
+ if not demo_sample_mode:
154
+ pbar = tqdm(pbar) if rank == 0 else pbar
155
+ total = 0
156
+
157
+ if accelerator.process_index == 0:
158
+ print_with_prefix("Using latent normalization")
159
+ dataset = ImgLatentDataset(
160
+ data_dir=train_config['data']['data_path'],
161
+ latent_norm=train_config['data']['latent_norm'] if 'latent_norm' in train_config['data'] else False,
162
+ latent_multiplier=train_config['data']['latent_multiplier'] if 'latent_multiplier' in train_config['data'] else 0.18215,
163
+ )
164
+ latent_mean, latent_std = dataset.get_latent_stats()
165
+ latent_multiplier = train_config['data']['latent_multiplier'] if 'latent_multiplier' in train_config['data'] else 0.18215
166
+ # move to device
167
+ latent_mean = latent_mean.clone().detach().to(device)
168
+ latent_std = latent_std.clone().detach().to(device)
169
+
170
+
171
+ # if demo_sample_mode:
172
+ # if accelerator.process_index == 0:
173
+ # images = []
174
+ # for label in tqdm([975, 3, 207, 387, 388, 88, 979, 279], desc="Generating Demo Samples"):
175
+ # z = torch.randn(1, model.in_channels, latent_size, latent_size, device=device)
176
+ # y = torch.tensor([label], device=device)
177
+ # z = torch.cat([z, z], 0)
178
+ # y_null = torch.tensor([1000] * 1, device=device)
179
+ # y = torch.cat([y, y_null], 0)
180
+ # model_kwargs = dict(y=y, cfg_scale=cfg_scale, cfg_interval=False, cfg_interval_start=cfg_interval_start)
181
+ # model_fn = model.forward_with_cfg
182
+ # samples = sample_fn(z, model_fn, **model_kwargs)[-1]
183
+ # samples = (samples * latent_std) / latent_multiplier + latent_mean
184
+ # samples = vae.decode_to_images(samples)
185
+ # images.append(samples)
186
+ # # Combine 8 images into a 2x4 grid
187
+ # os.makedirs('demo_images', exist_ok=True)
188
+ # # Stack all images into a large numpy array
189
+ # all_images = np.stack([img[0] for img in images]) # Take first image from each batch
190
+ # # Rearrange into 2x4 grid
191
+ # h, w = all_images.shape[1:3]
192
+ # grid = np.zeros((2 * h, 4 * w, 3), dtype=np.uint8)
193
+ # for idx, image in enumerate(all_images):
194
+ # i, j = divmod(idx, 4) # Calculate position in 2x4 grid
195
+ # grid[i*h:(i+1)*h, j*w:(j+1)*w] = image
196
+
197
+ # # Save the combined image
198
+ # Image.fromarray(grid).save('demo_images/demo_samples.png')
199
+
200
+ # return None
201
+ if demo_sample_mode:
202
+ # Demo mode: sample exactly one image and save it, then return.
203
+ # Choose a demo label (can be changed or made an argument)
204
+ demo_label = 975
205
+ # create latent noise for one sample
206
+ z = torch.randn(1, model.in_channels, latent_size, latent_size, device=device)
207
+ y = torch.tensor([demo_label], device=device)
208
+
209
+ # Setup classifier-free guidance if needed
210
+ if using_cfg:
211
+ z = torch.cat([z, z], 0)
212
+ y_null = torch.tensor([1000], device=device)
213
+ y = torch.cat([y, y_null], 0)
214
+ model_kwargs = dict(y=y, cfg_scale=cfg_scale, cfg_interval=False, cfg_interval_start=cfg_interval_start)
215
+ model_fn = model.forward_with_cfg
216
+ else:
217
+ model_kwargs = dict(y=y)
218
+ model_fn = model.forward
219
+
220
+ # Run sampling (single batch)
221
+ samples = sample_fn(z, model_fn, **model_kwargs)[-1]
222
+
223
+ if using_cfg:
224
+ # samples contains [cond; uncond] stacked, keep the conditional output
225
+ samples, _ = samples.chunk(2, dim=0)
226
+
227
+ # un-normalize and decode
228
+ samples = (samples * latent_std) / latent_multiplier + latent_mean
229
+ images = vae.decode_to_images(samples)
230
+
231
+ # save first image
232
+ if accelerator.process_index == 0:
233
+ os.makedirs(sample_folder_dir, exist_ok=True)
234
+ Image.fromarray(images[0]).save(os.path.join(sample_folder_dir, 'demo_sample.png'))
235
+ print_with_prefix(f"Saved demo sample to {os.path.join(sample_folder_dir, 'demo_sample.png')}")
236
+
237
+ return sample_folder_dir
238
+
239
+ else:
240
+ # 初始化时间统计变量
241
+ total_sampling_time = 0
242
+ total_vae_decode_time = 0
243
+ total_images_generated = 0
244
+ batch_times = []
245
+
246
+ for i in pbar:
247
+ batch_start_time = time()
248
+ # print("starting batch ", i)
249
+
250
+ # Sample inputs:
251
+ z = torch.randn(n, model.in_channels, latent_size, latent_size, device=device)
252
+ y = torch.randint(0, train_config['data']['num_classes'], (n,), device=device)
253
+
254
+ # Setup classifier-free guidance:
255
+ if using_cfg:
256
+ z = torch.cat([z, z], 0)
257
+ y_null = torch.tensor([1000] * n, device=device)
258
+ y = torch.cat([y, y_null], 0)
259
+ model_kwargs = dict(y=y, cfg_scale=cfg_scale, cfg_interval=True, cfg_interval_start=cfg_interval_start)
260
+ model_fn = model.forward_with_cfg
261
+ else:
262
+ model_kwargs = dict(y=y)
263
+ model_fn = model.forward
264
+
265
+ # 记录采样开始时间
266
+ sampling_start_time = time()
267
+ # print("starting sampling batch ", i)
268
+ samples = sample_fn(z, model_fn, **model_kwargs)[-1]
269
+ sampling_end_time = time()
270
+
271
+ if using_cfg:
272
+ samples, _ = samples.chunk(2, dim=0) # Remove null class samples
273
+
274
+ samples = (samples * latent_std) / latent_multiplier + latent_mean
275
+
276
+ # 记录VAE解码开始时间
277
+ vae_decode_start_time = time()
278
+ # print("starting VAE decode batch ", i)
279
+ samples = vae.decode_to_images(samples)
280
+ vae_decode_end_time = time()
281
+
282
+ # Save samples to disk as individual .png files
283
+ # print("start saving batch ", i)
284
+ for j, sample in enumerate(samples):
285
+ index = j * accelerator.num_processes + accelerator.process_index + total
286
+ Image.fromarray(sample).save(f"{sample_folder_dir}/{index:06d}.png")
287
+
288
+ # 统计时间
289
+ batch_end_time = time()
290
+ batch_time = batch_end_time - batch_start_time
291
+ sampling_time = sampling_end_time - sampling_start_time
292
+ vae_decode_time = vae_decode_end_time - vae_decode_start_time
293
+
294
+ batch_times.append(batch_time)
295
+ total_sampling_time += sampling_time
296
+ total_vae_decode_time += vae_decode_time
297
+ total_images_generated += len(samples)
298
+
299
+ # 每10个batch输出一次统计信息
300
+ if accelerator.process_index == 0 and (i + 1) % 10 == 0:
301
+ avg_sampling_time_per_image = total_sampling_time / total_images_generated
302
+ avg_vae_time_per_image = total_vae_decode_time / total_images_generated
303
+ avg_total_time_per_image = sum(batch_times) / total_images_generated
304
+ # print_with_prefix(f"Batch {i+1}/{iterations}: Avg sampling time per image: {avg_sampling_time_per_image:.3f}s, "
305
+ # f"Avg VAE decode time per image: {avg_vae_time_per_image:.3f}s, "
306
+ # f"Avg total time per image: {avg_total_time_per_image:.3f}s")
307
+
308
+ total += global_batch_size
309
+ accelerator.wait_for_everyone()
310
+
311
+ # 输出最终统计结果
312
+ if accelerator.process_index == 0 and total_images_generated > 0:
313
+ avg_sampling_time_per_image = total_sampling_time / total_images_generated
314
+ avg_vae_time_per_image = total_vae_decode_time / total_images_generated
315
+ avg_total_time_per_image = sum(batch_times) / total_images_generated
316
+ print_with_prefix("=" * 60)
317
+ print_with_prefix("FINAL TIMING STATISTICS:")
318
+ print_with_prefix(f"Total images generated: {total_images_generated}")
319
+ print_with_prefix(f"Average sampling time per image: {avg_sampling_time_per_image:.3f} seconds")
320
+ print_with_prefix(f"Average VAE decode time per image: {avg_vae_time_per_image:.3f} seconds")
321
+ print_with_prefix(f"Average total time per image: {avg_total_time_per_image:.3f} seconds")
322
+ print_with_prefix(f"Total sampling throughput: {total_images_generated / sum(batch_times):.2f} images/second")
323
+ print_with_prefix("=" * 60)
324
+
325
+ return sample_folder_dir
326
+
327
+ # some utils
328
+ def print_with_prefix(*messages):
329
+ prefix = f"\033[34m[LightningDiT-Sampling {strftime('%Y-%m-%d %H:%M:%S')}]\033[0m"
330
+ combined_message = ' '.join(map(str, messages))
331
+ print(f"{prefix}: {combined_message}")
332
+
333
+ def load_config(config_path):
334
+ with open(config_path, "r") as file:
335
+ config = yaml.safe_load(file)
336
+ return config
337
+
338
+ if __name__ == "__main__":
339
+
340
+ # read config
341
+ parser = argparse.ArgumentParser()
342
+ parser.add_argument('--config', type=str, default='configs/lightningdit_b_ldmvae_f16d16.yaml')
343
+ parser.add_argument('--demo', action='store_true', default=False)
344
+ args = parser.parse_args()
345
+ accelerator = Accelerator()
346
+ train_config = load_config(args.config)
347
+
348
+ # get ckpt_dir
349
+ assert 'ckpt_path' in train_config, "ckpt_path must be specified in config"
350
+ if accelerator.process_index == 0:
351
+ print_with_prefix('Using ckpt:', train_config['ckpt_path'])
352
+ ckpt_dir = train_config['ckpt_path']
353
+
354
+ if 'downsample_ratio' in train_config['vae']:
355
+ latent_size = train_config['data']['image_size'] // train_config['vae']['downsample_ratio']
356
+ else:
357
+ latent_size = train_config['data']['image_size'] // 16
358
+
359
+ # get model
360
+ model = LightningDiT_models[train_config['model']['model_type']](
361
+ input_size=latent_size,
362
+ num_classes=train_config['data']['num_classes'],
363
+ use_qknorm=train_config['model']['use_qknorm'],
364
+ use_swiglu=train_config['model']['use_swiglu'] if 'use_swiglu' in train_config['model'] else False,
365
+ use_rope=train_config['model']['use_rope'] if 'use_rope' in train_config['model'] else False,
366
+ use_rmsnorm=train_config['model']['use_rmsnorm'] if 'use_rmsnorm' in train_config['model'] else False,
367
+ wo_shift=train_config['model']['wo_shift'] if 'wo_shift' in train_config['model'] else False,
368
+ in_channels=train_config['model']['in_chans'] if 'in_chans' in train_config['model'] else 4,
369
+ learn_sigma=train_config['model']['learn_sigma'] if 'learn_sigma' in train_config['model'] else False,
370
+ num_rot=train_config['model']['num_rot'] if 'num_rot' in train_config['model'] else 4
371
+ )
372
+
373
+ # naive sample
374
+ sample_folder_dir = do_sample(train_config, accelerator, ckpt_path=ckpt_dir, model=model, demo_sample_mode=args.demo)
375
+
376
+ #visualize attention map
377
+ # attn_weights_list = []
378
+ # for block in model.blocks:
379
+ # attn = block.attn
380
+ # attn_weights = attn.attn_weights[:1] #cfg情况下取条件分支的注意力权重
381
+ # # print("attn_weights.shape:", attn_weights.shape)
382
+ # attn_weights_list.append(attn_weights)
383
+ # attn_weights = torch.cat(attn_weights_list, dim=0) # (num_layers, num_heads, N, N)
384
+ # # print("Concatenated attn_weights shape:", attn_weights.shape)
385
+ # print("visualizing attention maps...")
386
+ # visualize_attention_matrix(attn_weights,
387
+ # save_path="/home/jiayou.zhang/hom/personal/jinyuan/LightningDiT/attention_maps_new/xl800",
388
+ # model_name="xl800-cfg-cond",
389
+ # method="heatmap")
390
+
391
+ if not args.demo:
392
+ # calculate FID
393
+ # Important: FID is only for reference, please use ADM evaluation for paper reporting
394
+ if accelerator.process_index == 0:
395
+ from tools.calculate_fid import calculate_fid_given_paths
396
+ print_with_prefix('Calculating FID with {} number of samples'.format(train_config['sample']['fid_num']))
397
+ assert 'fid_reference_file' in train_config['data'], "fid_reference_file must be specified in config"
398
+ fid_reference_file = train_config['data']['fid_reference_file']
399
+ fid = calculate_fid_given_paths(
400
+ [fid_reference_file, sample_folder_dir],
401
+ batch_size=50,
402
+ dims=2048,
403
+ device='cuda',
404
+ num_workers=8,
405
+ sp_len = train_config['sample']['fid_num']
406
+ )
407
+ print_with_prefix('fid=',fid)