GRN / grn /utils_c2i /engine.py
hanjian.thu123
[update] app.py
17a8581
Raw History Blame Contribute Delete
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()