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