Spaces:
Running on Zero
Running on Zero
Download grn/utils_c2i/engine.py from hanjian/GRN: direct link, hf CLI and curl.
- Browser
- Download file 9.86 kB
-
https://huggingface.co/spaces/hanjian/GRN/resolve/main/grn/utils_c2i/engine.py
- Command line
-
hf download hf://spaces/hanjian/GRN/grn/utils_c2i/engine.py
-
curl -L -o engine.py https://huggingface.co/spaces/hanjian/GRN/resolve/main/grn/utils_c2i/engine.py
9.86 kB
| import math | |
| import sys | |
| import os | |
| import shutil | |
| import json | |
| import copy | |
| import torch | |
| import numpy as np | |
| import cv2 | |
| import torch_fidelity | |
| import grn.utils_c2i.misc as misc | |
| import grn.utils_c2i.lr_sched as lr_sched | |
| import grn.utils.wandb_utils as wandb_utils | |
| from grn.utils_c2i.hbq_util_c2i import raw_feature2label, raw_feature2bit_label | |
| def train_one_epoch(model, model_without_ddp, data_loader, optimizer, device, epoch, log_writer=None, args=None, vae=None): | |
| model.train(True) | |
| metric_logger = misc.MetricLogger(delimiter=" ") | |
| metric_logger.add_meter('lr', misc.SmoothedValue(window_size=1, fmt='{value:.6f}')) | |
| metric_logger.add_meter('grad_norm', misc.SmoothedValue(window_size=1, fmt='{value:.6f}')) | |
| header = 'Epoch: [{}]'.format(epoch) | |
| print_freq = 20 | |
| optimizer.zero_grad() | |
| if log_writer is not None: | |
| print('log_dir: {}'.format(log_writer.log_dir)) | |
| for data_iter_step, (x, labels) in enumerate(metric_logger.log_every(data_loader, print_freq, header)): | |
| # per iteration (instead of per epoch) lr scheduler | |
| lr_sched.adjust_learning_rate(optimizer, data_iter_step / len(data_loader) + epoch, args) | |
| # normalize image to [-1, 1] | |
| x = x.to(device, non_blocking=True).to(torch.float32).div_(255) | |
| x = x * 2.0 - 1.0 | |
| labels = labels.to(device, non_blocking=True) | |
| with torch.no_grad(): | |
| if args.method == 'GRN_ind': | |
| raw_features_, _, _ = vae.encode_for_raw_features(x.unsqueeze(2), scale_schedule=None, slice=True) | |
| raw_features = raw_features_[0].squeeze(2) | |
| x = raw_feature2label(raw_features, hbq_round=args.hbq_round) | |
| elif args.method == 'GRN_bit': | |
| raw_features_, _, _ = vae.encode_for_raw_features(x.unsqueeze(2), scale_schedule=None, slice=True) | |
| raw_features = raw_features_[0].squeeze(2) | |
| x = raw_feature2bit_label(raw_features, hbq_round=args.hbq_round) | |
| with torch.amp.autocast('cuda', dtype=torch.bfloat16): | |
| loss, t_bin2acc, t_bin2freq = model(x, labels) | |
| import torch.distributed as dist | |
| dist.all_reduce(t_bin2acc) | |
| dist.all_reduce(t_bin2freq) | |
| t_bin2acc = t_bin2acc / (t_bin2freq + 1e-8) * 100. | |
| loss_value = loss.item() | |
| if not math.isfinite(loss_value): | |
| print("Loss is {}, stopping training".format(loss_value)) | |
| sys.exit(1) | |
| optimizer.zero_grad() | |
| loss.backward() | |
| grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.clip_grad_norm) | |
| optimizer.step() | |
| torch.cuda.synchronize() | |
| model_without_ddp.update_ema() | |
| loss_value_reduce = misc.all_reduce_mean(loss_value) | |
| grad_norm_reduce = misc.all_reduce_mean(grad_norm.item()) | |
| metric_logger.update(loss=loss_value_reduce) | |
| lr = optimizer.param_groups[0]["lr"] | |
| metric_logger.update(lr=lr) | |
| metric_logger.update(grad_norm=grad_norm_reduce) | |
| if log_writer is not None: | |
| # Use epoch_1000x as the x-axis in TensorBoard to calibrate curves. | |
| epoch_1000x = int((data_iter_step / len(data_loader) + epoch) * 1000) | |
| if data_iter_step % args.log_freq == 0: | |
| log_writer.add_scalar('train_loss', loss_value_reduce, epoch_1000x) | |
| log_writer.add_scalar('lr', lr, epoch_1000x) | |
| if args.wandb: | |
| wandb_utils.log( | |
| { "train loss": loss_value_reduce, "lr": lr, "grad_norm_t": grad_norm_reduce}, | |
| step=epoch_1000x | |
| ) | |
| visual_dict = {} | |
| for t_round in range(10): | |
| if t_bin2freq[t_round] > 0: | |
| visual_dict.update({f"accuracy/signal_{t_round*0.1:.1f}_{(t_round+1)*0.1:.1f}": t_bin2acc[t_round].item()}) | |
| visual_dict.update({f"frequency/signal_{t_round*0.1:.1f}_{(t_round+1)*0.1:.1f}": t_bin2freq[t_round].item()}) | |
| wandb_utils.log(visual_dict, step=epoch_1000x) | |
| def evaluate(model_without_ddp, args, epoch, batch_size=64, log_writer=None, vae=None): | |
| model_without_ddp.eval() | |
| world_size = misc.get_world_size() | |
| local_rank = misc.get_rank() | |
| num_steps = args.num_images // (batch_size * world_size) + 1 | |
| # Construct the folder name for saving generated images. | |
| save_folder = os.path.join( | |
| args.generation_dir, | |
| f'Epoch{epoch:04d}', | |
| "{}-steps{}-cfg{}-tau{}-interval{}-{}-image{}-res{}".format( | |
| model_without_ddp.method, model_without_ddp.steps, model_without_ddp.cfg_scale, args.tau, | |
| model_without_ddp.cfg_interval[0], model_without_ddp.cfg_interval[1], args.num_images, args.img_size | |
| ), | |
| 'images', | |
| ) | |
| print("Save to:", save_folder) | |
| if misc.get_rank() == 0 and not os.path.exists(save_folder): | |
| os.makedirs(save_folder) | |
| json_file = os.path.join( | |
| os.path.dirname(os.environ.get('CKPT_FILE', '/tmp/res')), | |
| f'testing/Epoch{epoch:04d}/images_{args.num_images}', | |
| "{}-steps{}-cfg{}-tau{}-interval{}-{}-image{}-res{}".format( | |
| model_without_ddp.method, model_without_ddp.steps, model_without_ddp.cfg_scale, args.tau, | |
| model_without_ddp.cfg_interval[0], model_without_ddp.cfg_interval[1], args.num_images, args.img_size | |
| ), | |
| 'metrics.json', | |
| ) | |
| # switch to ema params, hard-coded to be the first one | |
| print("Switch to ema") | |
| model_params_backup = [p.detach().clone() for p in model_without_ddp.parameters()] | |
| if hasattr(model_without_ddp, 'module') and hasattr(model_without_ddp.module, 'ema_params1'): | |
| ema_params = model_without_ddp.module.ema_params1 | |
| else: | |
| ema_params = model_without_ddp.ema_params1 | |
| for param, ema_param in zip(model_without_ddp.parameters(), ema_params): | |
| param.data.copy_(ema_param.data) | |
| # ensure that the number of images per class is equal. | |
| class_num = args.class_num | |
| assert args.num_images % class_num == 0, "Number of images per class must be the same" | |
| class_label_gen_world = np.arange(0, class_num).repeat(args.num_images // class_num) | |
| class_label_gen_world = np.hstack([class_label_gen_world, np.zeros(50000)]) | |
| for i in range(num_steps): | |
| print("Generation step {}/{}".format(i, num_steps)) | |
| start_idx = world_size * batch_size * i + local_rank * batch_size | |
| end_idx = start_idx + batch_size | |
| labels_gen = class_label_gen_world[start_idx:end_idx] | |
| labels_gen = torch.Tensor(labels_gen).long().cuda() | |
| with torch.amp.autocast('cuda', dtype=torch.bfloat16): | |
| sampled_images = model_without_ddp.generate(labels_gen) | |
| if args.method == 'GRN_ind': | |
| from grn.utils_c2i.hbq_util_c2i import label2quant_features | |
| sampled_images = label2quant_features(sampled_images, hbq_round=args.hbq_round) | |
| sampled_images = vae.decode(sampled_images.unsqueeze(2), slice=True).squeeze(2) | |
| elif args.method == 'GRN_bit': | |
| from grn.utils_c2i.hbq_util_c2i import bit_label2raw_feature | |
| sampled_images = bit_label2raw_feature(sampled_images, hbq_round=args.hbq_round) | |
| sampled_images = vae.decode(sampled_images.unsqueeze(2), slice=True).squeeze(2) | |
| torch.distributed.barrier() | |
| # denormalize images | |
| sampled_images = (sampled_images + 1) / 2 | |
| sampled_images = sampled_images.detach().cpu() | |
| # distributed save images | |
| for b_id in range(sampled_images.size(0)): | |
| img_id = i * sampled_images.size(0) * world_size + local_rank * sampled_images.size(0) + b_id | |
| if img_id >= args.num_images: | |
| break | |
| gen_img = np.round(np.clip(sampled_images[b_id].numpy().transpose([1, 2, 0]) * 255, 0, 255)) | |
| gen_img = gen_img.astype(np.uint8)[:, :, ::-1] | |
| cv2.imwrite(os.path.join(save_folder, '{}.png'.format(str(img_id).zfill(5))), gen_img) | |
| torch.distributed.barrier() | |
| # back to no ema | |
| print("Switch back from ema") | |
| for param, backup in zip(model_without_ddp.parameters(), model_params_backup): | |
| param.data.copy_(backup.data) | |
| del model_params_backup | |
| torch.cuda.empty_cache() | |
| # compute FID and IS | |
| if log_writer is not None: | |
| if args.img_size == 256: | |
| fid_statistics_file = 'fid_stats/jit_in256_stats.npz' | |
| elif args.img_size == 512: | |
| fid_statistics_file = 'fid_stats/jit_in512_stats.npz' | |
| else: | |
| raise NotImplementedError | |
| metrics_dict = torch_fidelity.calculate_metrics( | |
| input1=save_folder, | |
| input2=None, | |
| fid_statistics_file=fid_statistics_file, | |
| cuda=True, | |
| isc=True, | |
| fid=True, | |
| kid=False, | |
| prc=False, | |
| verbose=False, | |
| restrict_data_size=-1, | |
| shuffle=False, | |
| ) | |
| fid = metrics_dict['frechet_inception_distance'] | |
| inception_score = metrics_dict['inception_score_mean'] | |
| postfix = "_cfg{}_res{}".format(model_without_ddp.cfg_scale, args.img_size) | |
| log_writer.add_scalar('fid{}'.format(postfix), fid, epoch) | |
| log_writer.add_scalar('is{}'.format(postfix), inception_score, epoch) | |
| if args.wandb: | |
| wandb_utils.log( | |
| {"fid": fid, "is": inception_score}, | |
| step=epoch | |
| ) | |
| print("FID: {:.4f}, Inception Score: {:.4f}".format(fid, inception_score)) | |
| os.makedirs(os.path.dirname(json_file), exist_ok=True) | |
| with open(json_file, 'w') as f: | |
| json.dump(metrics_dict, f) | |
| if args.delete_images: | |
| shutil.rmtree(save_folder) | |
| torch.distributed.barrier() | |