diff --git a/external/JiT/LICENSE b/external/JiT/LICENSE new file mode 100644 index 000000000..174d5e246 --- /dev/null +++ b/external/JiT/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2025 Tianhong Li + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. \ No newline at end of file diff --git a/external/JiT/README.md b/external/JiT/README.md new file mode 100644 index 000000000..52bde82d4 --- /dev/null +++ b/external/JiT/README.md @@ -0,0 +1,150 @@ +## Just image Transformer (JiT) for Pixel-space Diffusion + +[![arXiv](https://img.shields.io/badge/arXiv%20paper-2511.13720-b31b1b.svg)](https://arxiv.org/abs/2511.13720)  + +

+ +

+ + +This is a PyTorch/GPU re-implementation of the paper [Back to Basics: Let Denoising Generative Models Denoise](https://arxiv.org/abs/2511.13720): + +``` +@article{li2025jit, + title={Back to Basics: Let Denoising Generative Models Denoise}, + author={Li, Tianhong and He, Kaiming}, + journal={arXiv preprint arXiv:2511.13720}, + year={2025} +} +``` + +JiT adopts a minimalist and self-contained design for pixel-level high-resolution image diffusion. +The original implementation was in JAX+TPU. This re-implementation is in PyTorch+GPU. + +

+ +

+ +### Dataset +Download [ImageNet](http://image-net.org/download) dataset, and place it in your `IMAGENET_PATH`. + +### Installation + +Download the code: +``` +git clone https://github.com/LTH14/JiT.git +cd JiT +``` + +A suitable [conda](https://conda.io/) environment named `jit` can be created and activated with: + +``` +conda env create -f environment.yaml +conda activate jit +``` + +If you get ```undefined symbol: iJIT_NotifyEvent``` when importing ```torch```, simply +``` +pip uninstall torch +pip install torch==2.5.1 --index-url https://download.pytorch.org/whl/cu124 +``` +Check this [issue](https://github.com/conda/conda/issues/13812#issuecomment-2071445372) for more details. + +### Training +The below training scripts have been tested on 8 H200 GPUs. + +Example script for training JiT-B/16 on ImageNet 256x256 for 600 epochs: +``` +torchrun --nproc_per_node=8 --nnodes=1 --node_rank=0 \ +main_jit.py \ +--model JiT-B/16 \ +--proj_dropout 0.0 \ +--P_mean -0.8 --P_std 0.8 \ +--img_size 256 --noise_scale 1.0 \ +--batch_size 128 --blr 5e-5 \ +--epochs 600 --warmup_epochs 5 \ +--gen_bsz 128 --num_images 50000 --cfg 2.9 --interval_min 0.1 --interval_max 1.0 \ +--output_dir ${OUTPUT_DIR} --resume ${OUTPUT_DIR} \ +--data_path ${IMAGENET_PATH} --online_eval +``` + +Example script for training JiT-B/32 on ImageNet 512x512 for 600 epochs: +``` +torchrun --nproc_per_node=8 --nnodes=1 --node_rank=0 \ +main_jit.py \ +--model JiT-B/32 \ +--proj_dropout 0.0 \ +--P_mean -0.8 --P_std 0.8 \ +--img_size 512 --noise_scale 2.0 \ +--batch_size 128 --blr 5e-5 \ +--epochs 600 --warmup_epochs 5 \ +--gen_bsz 128 --num_images 50000 --cfg 2.9 --interval_min 0.1 --interval_max 1.0 \ +--output_dir ${OUTPUT_DIR} --resume ${OUTPUT_DIR} \ +--data_path ${IMAGENET_PATH} --online_eval +``` + +Example script for training JiT-H/16 on ImageNet 256x256 for 600 epochs: +``` +torchrun --nproc_per_node=8 --nnodes=1 --node_rank=0 \ +main_jit.py \ +--model JiT-H/16 \ +--proj_dropout 0.2 \ +--P_mean -0.8 --P_std 0.8 \ +--img_size 256 --noise_scale 1.0 \ +--batch_size 128 --blr 5e-5 \ +--epochs 600 --warmup_epochs 5 \ +--gen_bsz 128 --num_images 50000 --cfg 2.2 --interval_min 0.1 --interval_max 1.0 \ +--output_dir ${OUTPUT_DIR} --resume ${OUTPUT_DIR} \ +--data_path ${IMAGENET_PATH} --online_eval +``` + +### Evaluation + +PyTorch pre-trained models are available [here](https://www.dropbox.com/scl/fo/3ken1avtsd81ip67b9qpi/AK218ZNvXKSv74igVvht4PQ?rlkey=14gjrblmljewpl6ygxzlr3njm&st=ffkl77al&dl=0). + +Evaluate pre-trained JiT-B: +``` +torchrun --nproc_per_node=8 --nnodes=1 --node_rank=0 \ +main_jit.py \ +--model JiT-B/16 (or JiT-B/32) \ +--img_size 256 (or 512) --noise_scale 1.0 (or 2.0) \ +--gen_bsz 256 --num_images 50000 --cfg 3.0 --interval_min 0.1 --interval_max 1.0 \ +--output_dir ${CKPT_DIR} --resume ${CKPT_DIR} \ +--data_path ${IMAGENET_PATH} --evaluate_gen +``` + +Evaluate pre-trained JiT-L: +``` +torchrun --nproc_per_node=8 --nnodes=1 --node_rank=0 \ +main_jit.py \ +--model JiT-L/16 (or JiT-L/32) \ +--img_size 256 (or 512) --noise_scale 1.0 (or 2.0) \ +--gen_bsz 256 --num_images 50000 --cfg 2.4 (or 2.5) --interval_min 0.1 --interval_max 1.0 \ +--output_dir ${CKPT_DIR} --resume ${CKPT_DIR} \ +--data_path ${IMAGENET_PATH} --evaluate_gen +``` + +Evaluate pre-trained JiT-H: +``` +torchrun --nproc_per_node=8 --nnodes=1 --node_rank=0 \ +main_jit.py \ +--model JiT-H/16 (or JiT-H/32) \ +--img_size 256 (or 512) --noise_scale 1.0 (or 2.0) \ +--gen_bsz 256 --num_images 50000 --cfg 2.2 (or 2.3) --interval_min 0.1 --interval_max 1.0 \ +--output_dir ${CKPT_DIR} --resume ${CKPT_DIR} \ +--data_path ${IMAGENET_PATH} --evaluate_gen +``` + +We use a customized [```torch-fidelity```](https://github.com/LTH14/torch-fidelity) +to evaluate FID and IS against a reference image folder or statistics. You can use ```prepare_ref.py``` +to prepare the reference image folder, or directly use our pre-computed reference stats +under ```fid_stats```. + +### Acknowledgements + +We thank Google TPU Research Cloud (TRC) for granting us access to TPUs, and the MIT +ORCD Seed Fund Grants for supporting GPU resources. + +### Contact + +If you have any questions, feel free to contact me through email (tianhong@mit.edu). Enjoy! diff --git a/external/JiT/UPSTREAM.md b/external/JiT/UPSTREAM.md new file mode 100644 index 000000000..dd8f30afd --- /dev/null +++ b/external/JiT/UPSTREAM.md @@ -0,0 +1,11 @@ +# Upstream provenance + +This directory vendors the training-relevant files from the MIT-licensed +PyTorch JiT implementation from +https://github.com/LTH14/JiT at commit +`cbc743a2ada5e9762697da2c83f8c4f8379e8c17`. + +The 512px FID statistics and demo images are omitted. The only upstream +compatibility change is device-agnostic rotary buffers in +`util/model_util.py`. The matched experiment lives in `matched_models.py` and +`train_matched.py`; the original JiT model and objective remain intact. diff --git a/external/JiT/denoiser.py b/external/JiT/denoiser.py new file mode 100644 index 000000000..e6c3ff29e --- /dev/null +++ b/external/JiT/denoiser.py @@ -0,0 +1,130 @@ +import torch +import torch.nn as nn +from model_jit import JiT_models + + +class Denoiser(nn.Module): + def __init__( + self, + args + ): + super().__init__() + self.net = JiT_models[args.model]( + input_size=args.img_size, + in_channels=3, + num_classes=args.class_num, + attn_drop=args.attn_dropout, + proj_drop=args.proj_dropout, + ) + self.img_size = args.img_size + self.num_classes = args.class_num + + self.label_drop_prob = args.label_drop_prob + self.P_mean = args.P_mean + self.P_std = args.P_std + self.t_eps = args.t_eps + self.noise_scale = args.noise_scale + + # ema + self.ema_decay1 = args.ema_decay1 + self.ema_decay2 = args.ema_decay2 + self.ema_params1 = None + self.ema_params2 = None + + # generation hyper params + self.method = args.sampling_method + self.steps = args.num_sampling_steps + self.cfg_scale = args.cfg + self.cfg_interval = (args.interval_min, args.interval_max) + + def drop_labels(self, labels): + drop = torch.rand(labels.shape[0], device=labels.device) < self.label_drop_prob + out = torch.where(drop, torch.full_like(labels, self.num_classes), labels) + return out + + def sample_t(self, n: int, device=None): + z = torch.randn(n, device=device) * self.P_std + self.P_mean + return torch.sigmoid(z) + + def forward(self, x, labels): + labels_dropped = self.drop_labels(labels) if self.training else labels + + t = self.sample_t(x.size(0), device=x.device).view(-1, *([1] * (x.ndim - 1))) + e = torch.randn_like(x) * self.noise_scale + + z = t * x + (1 - t) * e + v = (x - z) / (1 - t).clamp_min(self.t_eps) + + x_pred = self.net(z, t.flatten(), labels_dropped) + v_pred = (x_pred - z) / (1 - t).clamp_min(self.t_eps) + + # l2 loss + loss = (v - v_pred) ** 2 + loss = loss.mean(dim=(1, 2, 3)).mean() + + return loss + + @torch.no_grad() + def generate(self, labels): + device = labels.device + bsz = labels.size(0) + z = self.noise_scale * torch.randn(bsz, 3, self.img_size, self.img_size, device=device) + timesteps = torch.linspace(0.0, 1.0, self.steps+1, device=device).view(-1, *([1] * z.ndim)).expand(-1, bsz, -1, -1, -1) + + if self.method == "euler": + stepper = self._euler_step + elif self.method == "heun": + stepper = self._heun_step + else: + raise NotImplementedError + + # ode + for i in range(self.steps - 1): + t = timesteps[i] + t_next = timesteps[i + 1] + z = stepper(z, t, t_next, labels) + # last step euler + z = self._euler_step(z, timesteps[-2], timesteps[-1], labels) + return z + + @torch.no_grad() + def _forward_sample(self, z, t, labels): + # conditional + x_cond = self.net(z, t.flatten(), labels) + v_cond = (x_cond - z) / (1.0 - t).clamp_min(self.t_eps) + + # unconditional + x_uncond = self.net(z, t.flatten(), torch.full_like(labels, self.num_classes)) + v_uncond = (x_uncond - z) / (1.0 - t).clamp_min(self.t_eps) + + # cfg interval + low, high = self.cfg_interval + interval_mask = (t < high) & ((low == 0) | (t > low)) + cfg_scale_interval = torch.where(interval_mask, self.cfg_scale, 1.0) + + return v_uncond + cfg_scale_interval * (v_cond - v_uncond) + + @torch.no_grad() + def _euler_step(self, z, t, t_next, labels): + v_pred = self._forward_sample(z, t, labels) + z_next = z + (t_next - t) * v_pred + return z_next + + @torch.no_grad() + def _heun_step(self, z, t, t_next, labels): + v_pred_t = self._forward_sample(z, t, labels) + + z_next_euler = z + (t_next - t) * v_pred_t + v_pred_t_next = self._forward_sample(z_next_euler, t_next, labels) + + v_pred = 0.5 * (v_pred_t + v_pred_t_next) + z_next = z + (t_next - t) * v_pred + return z_next + + @torch.no_grad() + def update_ema(self): + source_params = list(self.parameters()) + for targ, src in zip(self.ema_params1, source_params): + targ.detach().mul_(self.ema_decay1).add_(src, alpha=1 - self.ema_decay1) + for targ, src in zip(self.ema_params2, source_params): + targ.detach().mul_(self.ema_decay2).add_(src, alpha=1 - self.ema_decay2) diff --git a/external/JiT/engine_jit.py b/external/JiT/engine_jit.py new file mode 100644 index 000000000..8a346ba0a --- /dev/null +++ b/external/JiT/engine_jit.py @@ -0,0 +1,160 @@ +import math +import sys +import os +import shutil + +import torch +import numpy as np +import cv2 + +import util.misc as misc +import util.lr_sched as lr_sched +import torch_fidelity +import copy + + +def train_one_epoch(model, model_without_ddp, data_loader, optimizer, device, epoch, log_writer=None, args=None): + model.train(True) + metric_logger = misc.MetricLogger(delimiter=" ") + metric_logger.add_meter('lr', 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.amp.autocast('cuda', dtype=torch.bfloat16): + loss = model(x, labels) + + 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() + optimizer.step() + + torch.cuda.synchronize() + + model_without_ddp.update_ema() + + metric_logger.update(loss=loss_value) + lr = optimizer.param_groups[0]["lr"] + metric_logger.update(lr=lr) + + loss_value_reduce = misc.all_reduce_mean(loss_value) + + 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) + + +def evaluate(model_without_ddp, args, epoch, batch_size=64, log_writer=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.output_dir, + "{}-steps{}-cfg{}-interval{}-{}-image{}-res{}".format( + model_without_ddp.method, model_without_ddp.steps, model_without_ddp.cfg_scale, + model_without_ddp.cfg_interval[0], model_without_ddp.cfg_interval[1], args.num_images, args.img_size + ) + ) + print("Save to:", save_folder) + if misc.get_rank() == 0 and not os.path.exists(save_folder): + os.makedirs(save_folder) + + # switch to ema params, hard-coded to be the first one + model_state_dict = copy.deepcopy(model_without_ddp.state_dict()) + ema_state_dict = copy.deepcopy(model_without_ddp.state_dict()) + for i, (name, _value) in enumerate(model_without_ddp.named_parameters()): + assert name in ema_state_dict + ema_state_dict[name] = model_without_ddp.ema_params1[i] + print("Switch to ema") + model_without_ddp.load_state_dict(ema_state_dict) + + # 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) + + 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") + model_without_ddp.load_state_dict(model_state_dict) + + # 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, + ) + 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) + print("FID: {:.4f}, Inception Score: {:.4f}".format(fid, inception_score)) + shutil.rmtree(save_folder) + + torch.distributed.barrier() diff --git a/external/JiT/environment.yaml b/external/JiT/environment.yaml new file mode 100644 index 000000000..00716c6fa --- /dev/null +++ b/external/JiT/environment.yaml @@ -0,0 +1,20 @@ +name: jit +channels: + - pytorch + - defaults + - nvidia +dependencies: + - python=3.10 + - pip=22.3 + - pytorch-cuda=12.4 + - pytorch=2.5.1 + - torchvision=0.20.1 + - numpy=1.22 + - pip: + - opencv-python==4.11.0.86 + - timm==0.9.12 + - tensorboard==2.10.0 + - scipy==1.9.1 + - einops==0.8.1 + - gdown==5.2.0 + - -e git+https://github.com/LTH14/torch-fidelity.git@master#egg=torch-fidelity diff --git a/external/JiT/fid_stats/jit_in256_stats.npz b/external/JiT/fid_stats/jit_in256_stats.npz new file mode 100644 index 000000000..3c94ea961 Binary files /dev/null and b/external/JiT/fid_stats/jit_in256_stats.npz differ diff --git a/external/JiT/main_jit.py b/external/JiT/main_jit.py new file mode 100644 index 000000000..cb630d4c9 --- /dev/null +++ b/external/JiT/main_jit.py @@ -0,0 +1,266 @@ +import argparse +import datetime +import numpy as np +import os +import time +from pathlib import Path + +import torch +import torch.backends.cudnn as cudnn +from torch.utils.tensorboard import SummaryWriter +import torchvision.transforms as transforms +import torchvision.datasets as datasets + +from util.crop import center_crop_arr +import util.misc as misc + +import copy +from engine_jit import train_one_epoch, evaluate + +from denoiser import Denoiser + + +def get_args_parser(): + parser = argparse.ArgumentParser('JiT', add_help=False) + + # architecture + parser.add_argument('--model', default='JiT-B/16', type=str, metavar='MODEL', + help='Name of the model to train') + parser.add_argument('--img_size', default=256, type=int, help='Image size') + parser.add_argument('--attn_dropout', type=float, default=0.0, help='Attention dropout rate') + parser.add_argument('--proj_dropout', type=float, default=0.0, help='Projection dropout rate') + + # training + parser.add_argument('--epochs', default=200, type=int) + parser.add_argument('--warmup_epochs', type=int, default=5, metavar='N', + help='Epochs to warm up LR') + parser.add_argument('--batch_size', default=128, type=int, + help='Batch size per GPU (effective batch size = batch_size * # GPUs)') + parser.add_argument('--lr', type=float, default=None, metavar='LR', + help='Learning rate (absolute)') + parser.add_argument('--blr', type=float, default=5e-5, metavar='LR', + help='Base learning rate: absolute_lr = base_lr * total_batch_size / 256') + parser.add_argument('--min_lr', type=float, default=0., metavar='LR', + help='Minimum LR for cyclic schedulers that hit 0') + parser.add_argument('--lr_schedule', type=str, default='constant', + help='Learning rate schedule') + parser.add_argument('--weight_decay', type=float, default=0.0, + help='Weight decay (default: 0.0)') + parser.add_argument('--ema_decay1', type=float, default=0.9999, + help='The first ema to track. Use the first ema for sampling by default.') + parser.add_argument('--ema_decay2', type=float, default=0.9996, + help='The second ema to track') + parser.add_argument('--P_mean', default=-0.8, type=float) + parser.add_argument('--P_std', default=0.8, type=float) + parser.add_argument('--noise_scale', default=1.0, type=float) + parser.add_argument('--t_eps', default=5e-2, type=float) + parser.add_argument('--label_drop_prob', default=0.1, type=float) + + parser.add_argument('--seed', default=0, type=int) + parser.add_argument('--start_epoch', default=0, type=int, metavar='N', + help='Starting epoch') + parser.add_argument('--num_workers', default=12, type=int) + parser.add_argument('--pin_mem', action='store_true', + help='Pin CPU memory in DataLoader for faster GPU transfers') + parser.add_argument('--no_pin_mem', action='store_false', dest='pin_mem') + parser.set_defaults(pin_mem=True) + + # sampling + parser.add_argument('--sampling_method', default='heun', type=str, + help='ODE samping method') + parser.add_argument('--num_sampling_steps', default=50, type=int, + help='Sampling steps') + parser.add_argument('--cfg', default=1.0, type=float, + help='Classifier-free guidance factor') + parser.add_argument('--interval_min', default=0.0, type=float, + help='CFG interval min') + parser.add_argument('--interval_max', default=1.0, type=float, + help='CFG interval max') + parser.add_argument('--num_images', default=50000, type=int, + help='Number of images to generate') + parser.add_argument('--eval_freq', type=int, default=40, + help='Frequency (in epochs) for evaluation') + parser.add_argument('--online_eval', action='store_true') + parser.add_argument('--evaluate_gen', action='store_true') + parser.add_argument('--gen_bsz', type=int, default=256, + help='Generation batch size') + + # dataset + parser.add_argument('--data_path', default='./data/imagenet', type=str, + help='Path to the dataset') + parser.add_argument('--class_num', default=1000, type=int) + + # checkpointing + parser.add_argument('--output_dir', default='./output_dir', + help='Directory to save outputs (empty for no saving)') + parser.add_argument('--resume', default='', + help='Folder that contains checkpoint to resume from') + parser.add_argument('--save_last_freq', type=int, default=5, + help='Frequency (in epochs) to save checkpoints') + parser.add_argument('--log_freq', default=100, type=int) + parser.add_argument('--device', default='cuda', + help='Device to use for training/testing') + + # distributed training + parser.add_argument('--world_size', default=1, type=int, + help='Number of distributed processes') + parser.add_argument('--local_rank', default=-1, type=int) + parser.add_argument('--dist_on_itp', action='store_true') + parser.add_argument('--dist_url', default='env://', + help='URL used to set up distributed training') + + return parser + + +def main(args): + misc.init_distributed_mode(args) + print('Job directory:', os.path.dirname(os.path.realpath(__file__))) + print("Arguments:\n{}".format(args).replace(', ', ',\n')) + + device = torch.device(args.device) + + # Set seeds for reproducibility + seed = args.seed + misc.get_rank() + torch.manual_seed(seed) + np.random.seed(seed) + + cudnn.benchmark = True + + num_tasks = misc.get_world_size() + global_rank = misc.get_rank() + + # Set up TensorBoard logging (only on main process) + if global_rank == 0 and args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + log_writer = SummaryWriter(log_dir=args.output_dir) + else: + log_writer = None + + # Data augmentation transforms + transform_train = transforms.Compose([ + transforms.Lambda(lambda img: center_crop_arr(img, args.img_size)), + transforms.RandomHorizontalFlip(), + transforms.PILToTensor() + ]) + + dataset_train = datasets.ImageFolder(os.path.join(args.data_path, 'train'), transform=transform_train) + print(dataset_train) + + sampler_train = torch.utils.data.DistributedSampler( + dataset_train, num_replicas=num_tasks, rank=global_rank, shuffle=True + ) + print("Sampler_train =", sampler_train) + + data_loader_train = torch.utils.data.DataLoader( + dataset_train, sampler=sampler_train, + batch_size=args.batch_size, + num_workers=args.num_workers, + pin_memory=args.pin_mem, + drop_last=True + ) + + torch._dynamo.config.cache_size_limit = 128 + torch._dynamo.config.optimize_ddp = False + + # Create denoiser + model = Denoiser(args) + + print("Model =", model) + n_params = sum(p.numel() for p in model.parameters() if p.requires_grad) + print("Number of trainable parameters: {:.6f}M".format(n_params / 1e6)) + + model.to(device) + + eff_batch_size = args.batch_size * misc.get_world_size() + if args.lr is None: # only base_lr (blr) is specified + args.lr = args.blr * eff_batch_size / 256 + + print("Base lr: {:.2e}".format(args.lr * 256 / eff_batch_size)) + print("Actual lr: {:.2e}".format(args.lr)) + print("Effective batch size: %d" % eff_batch_size) + + model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[args.gpu]) + model_without_ddp = model.module + + # Set up optimizer with weight decay adjustment for bias and norm layers + param_groups = misc.add_weight_decay(model_without_ddp, args.weight_decay) + optimizer = torch.optim.AdamW(param_groups, lr=args.lr, betas=(0.9, 0.95)) + print(optimizer) + + # Resume from checkpoint if provided + checkpoint_path = os.path.join(args.resume, "checkpoint-last.pth") if args.resume else None + if checkpoint_path and os.path.exists(checkpoint_path): + checkpoint = torch.load(checkpoint_path, map_location='cpu') + model_without_ddp.load_state_dict(checkpoint['model']) + + ema_state_dict1 = checkpoint['model_ema1'] + ema_state_dict2 = checkpoint['model_ema2'] + model_without_ddp.ema_params1 = [ema_state_dict1[name].cuda() for name, _ in model_without_ddp.named_parameters()] + model_without_ddp.ema_params2 = [ema_state_dict2[name].cuda() for name, _ in model_without_ddp.named_parameters()] + print("Resumed checkpoint from", args.resume) + + if 'optimizer' in checkpoint and 'epoch' in checkpoint: + optimizer.load_state_dict(checkpoint['optimizer']) + args.start_epoch = checkpoint['epoch'] + 1 + print("Loaded optimizer & scaler state!") + del checkpoint + else: + model_without_ddp.ema_params1 = copy.deepcopy(list(model_without_ddp.parameters())) + model_without_ddp.ema_params2 = copy.deepcopy(list(model_without_ddp.parameters())) + print("Training from scratch") + + # Evaluate generation + if args.evaluate_gen: + print("Evaluating checkpoint at {} epoch".format(args.start_epoch)) + with torch.random.fork_rng(): + torch.manual_seed(seed) + with torch.no_grad(): + evaluate(model_without_ddp, args, 0, batch_size=args.gen_bsz, log_writer=log_writer) + return + + # Training loop + print(f"Start training for {args.epochs} epochs") + start_time = time.time() + for epoch in range(args.start_epoch, args.epochs): + if args.distributed: + data_loader_train.sampler.set_epoch(epoch) + + train_one_epoch(model, model_without_ddp, data_loader_train, optimizer, device, epoch, log_writer=log_writer, args=args) + + # Save checkpoint periodically + if epoch % args.save_last_freq == 0 or epoch + 1 == args.epochs: + misc.save_model( + args=args, + model_without_ddp=model_without_ddp, + optimizer=optimizer, + epoch=epoch, + epoch_name="last" + ) + + if epoch % 100 == 0 and epoch > 0: + misc.save_model( + args=args, + model_without_ddp=model_without_ddp, + optimizer=optimizer, + epoch=epoch + ) + + # Perform online evaluation at specified intervals + if args.online_eval and (epoch % args.eval_freq == 0 or epoch + 1 == args.epochs): + torch.cuda.empty_cache() + with torch.no_grad(): + evaluate(model_without_ddp, args, epoch, batch_size=args.gen_bsz, log_writer=log_writer) + torch.cuda.empty_cache() + + if misc.is_main_process() and log_writer is not None: + log_writer.flush() + + total_time = time.time() - start_time + total_time_str = str(datetime.timedelta(seconds=int(total_time))) + print('Training time:', total_time_str) + + +if __name__ == '__main__': + args = get_args_parser().parse_args() + Path(args.output_dir).mkdir(parents=True, exist_ok=True) + main(args) diff --git a/external/JiT/matched_models.py b/external/JiT/matched_models.py new file mode 100644 index 000000000..30d8c8dfb --- /dev/null +++ b/external/JiT/matched_models.py @@ -0,0 +1,431 @@ +"""Matched ImageNet objectives for JiT and decoder-only latent denoising. + +The JiT path preserves the upstream clean-image prediction objective. The +endpoint-latent path never reads the target image before terminal decoding. +""" + +from __future__ import annotations + +import contextlib +import math +from dataclasses import dataclass +from typing import Dict, Optional, Tuple + +import torch +from torch import nn +from torch.utils.checkpoint import checkpoint + +from egomimic.models.denoising_nets import CrossBlock, posemb_sincos +from model_jit import JiT_models + + +def _unpatchify(tokens: torch.Tensor, patch_size: int, channels: int = 3) -> torch.Tensor: + batch, count, width = tokens.shape + side = math.isqrt(count) + if side * side != count: + raise ValueError(f"Token count {count} is not a square grid") + expected = patch_size * patch_size * channels + if width != expected: + raise ValueError(f"Patch width {width} does not match expected {expected}") + x = tokens.reshape(batch, side, side, patch_size, patch_size, channels) + x = torch.einsum("nhwpqc->nchpwq", x) + return x.reshape(batch, channels, side * patch_size, side * patch_size) + + +class JiTObjective(nn.Module): + """Official JiT-B/16 network and direct clean-image prediction loss.""" + + architecture = "jit_b16" + + def __init__( + self, + image_size: int = 256, + num_classes: int = 1000, + label_drop_prob: float = 0.1, + p_mean: float = -0.8, + p_std: float = 0.8, + noise_scale: float = 1.0, + t_eps: float = 0.05, + ) -> None: + super().__init__() + self.net = JiT_models["JiT-B/16"]( + input_size=image_size, + in_channels=3, + num_classes=num_classes, + attn_drop=0.0, + proj_drop=0.0, + ) + self.image_size = int(image_size) + self.num_classes = int(num_classes) + self.label_drop_prob = float(label_drop_prob) + self.p_mean = float(p_mean) + self.p_std = float(p_std) + self.noise_scale = float(noise_scale) + self.t_eps = float(t_eps) + + def _drop_labels(self, labels: torch.Tensor) -> torch.Tensor: + if not self.training or self.label_drop_prob <= 0: + return labels + drop = torch.rand(labels.shape[0], device=labels.device) < self.label_drop_prob + return torch.where(drop, torch.full_like(labels, self.num_classes), labels) + + def forward( + self, + images: torch.Tensor, + labels: torch.Tensor, + optimizer_step: int, + force_steps: Optional[int] = None, + ) -> Dict[str, torch.Tensor]: + del optimizer_step, force_steps + labels = self._drop_labels(labels) + logit_t = torch.randn(images.shape[0], device=images.device) * self.p_std + self.p_mean + t = torch.sigmoid(logit_t).reshape(-1, 1, 1, 1) + noise = torch.randn_like(images) * self.noise_scale + state = t * images + (1.0 - t) * noise + target_velocity = (images - state) / (1.0 - t).clamp_min(self.t_eps) + prediction = self.net(state, t.flatten(), labels) + predicted_velocity = (prediction - state) / (1.0 - t).clamp_min(self.t_eps) + loss = (target_velocity - predicted_velocity).square().mean() + return { + "loss": loss, + "prediction_rms": prediction.detach().square().mean().sqrt(), + "noise_rms": noise.detach().square().mean().sqrt(), + "endpoint_rms": prediction.detach().square().mean().sqrt(), + "latent_delta_rms": (prediction.detach() - state.detach()).square().mean().sqrt(), + "unroll_steps": torch.ones((), device=images.device), + } + + @torch.no_grad() + def sample( + self, + labels: torch.Tensor, + num_steps: int = 16, + cfg_scale: float = 1.0, + noise: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + if num_steps <= 0: + raise ValueError("num_steps must be positive") + state = ( + torch.randn( + labels.shape[0], 3, self.image_size, self.image_size, + device=labels.device, + ) + * self.noise_scale + if noise is None + else noise.clone() + ) + for index in range(num_steps): + t_value = index / num_steps + t = torch.full( + (labels.shape[0],), t_value, device=labels.device, dtype=torch.float32 + ) + shaped_t = t.reshape(-1, 1, 1, 1) + conditional = self.net(state, t, labels) + velocity = (conditional - state) / (1.0 - shaped_t).clamp_min(self.t_eps) + if cfg_scale != 1.0: + null_labels = torch.full_like(labels, self.num_classes) + unconditional = self.net(state, t, null_labels) + uncond_velocity = (unconditional - state) / ( + 1.0 - shaped_t + ).clamp_min(self.t_eps) + velocity = uncond_velocity + cfg_scale * (velocity - uncond_velocity) + state = state + velocity / num_steps + return state + + +class ImageCrossTransformer(nn.Module): + """Existing Pipeline cross-transformer with explicit factorized 2-D positions.""" + + def __init__( + self, + grid_size: int, + latent_dim: int, + hidden_dim: int, + depth: int, + num_heads: int, + dropout: float, + mlp_layers: int, + mlp_ratio: float, + ) -> None: + super().__init__() + if hidden_dim < latent_dim: + raise ValueError("hidden_dim must not bottleneck latent_dim") + self.grid_size = int(grid_size) + self.hidden_dim = int(hidden_dim) + self.proj_u = nn.Linear(latent_dim, hidden_dim) + self.proj_d = nn.Linear(hidden_dim, latent_dim) + self.row_position = nn.Parameter(torch.zeros(1, grid_size, 1, hidden_dim)) + self.column_position = nn.Parameter(torch.zeros(1, 1, grid_size, hidden_dim)) + nn.init.normal_(self.row_position, std=0.02) + nn.init.normal_(self.column_position, std=0.02) + self.layers = nn.ModuleList( + [ + CrossBlock( + cond_dim=hidden_dim, + hidden_dim=hidden_dim, + n_heads=num_heads, + dropout=dropout, + mlp_layers=mlp_layers, + mlp_ratio=mlp_ratio, + ) + for _ in range(depth) + ] + ) + + def forward( + self, latent: torch.Tensor, timesteps: torch.Tensor, condition: torch.Tensor + ) -> torch.Tensor: + batch, count, _ = latent.shape + expected = self.grid_size * self.grid_size + if count != expected: + raise ValueError(f"Expected {expected} latent tokens, got {count}") + hidden = self.proj_u(latent).reshape( + batch, self.grid_size, self.grid_size, self.hidden_dim + ) + hidden = hidden + self.row_position + self.column_position + hidden = hidden.reshape(batch, count, self.hidden_dim) + time_embedding = posemb_sincos( + timesteps, self.hidden_dim, min_period=4e-3, max_period=4.0 + ).to(device=hidden.device, dtype=hidden.dtype) + hidden = hidden + time_embedding.unsqueeze(1) + for layer in self.layers: + hidden = layer(hidden, condition) + return self.proj_d(hidden) + + +@dataclass(frozen=True) +class IntegrationResult: + endpoint: torch.Tensor + delta_rms: torch.Tensor + step_sizes: torch.Tensor + + +class EndpointLatentObjective(nn.Module): + """Strict decoder-only endpoint-trained iterative latent image generator.""" + + architecture = "endpoint_latent" + + def __init__( + self, + image_size: int = 256, + patch_size: int = 16, + latent_dim: int = 96, + hidden_dim: int = 352, + depth: int = 16, + num_heads: int = 8, + dropout: float = 0.1, + mlp_layers: int = 4, + mlp_ratio: float = 4.0, + decoder_hidden_dim: int = 512, + num_classes: int = 1000, + label_drop_prob: float = 0.1, + gradient_checkpointing: bool = True, + ) -> None: + super().__init__() + if image_size % patch_size: + raise ValueError("image_size must be divisible by patch_size") + self.image_size = int(image_size) + self.patch_size = int(patch_size) + self.grid_size = image_size // patch_size + self.num_tokens = self.grid_size * self.grid_size + self.latent_dim = int(latent_dim) + self.hidden_dim = int(hidden_dim) + self.num_classes = int(num_classes) + self.label_drop_prob = float(label_drop_prob) + self.gradient_checkpointing = bool(gradient_checkpointing) + self.label_embedding = nn.Embedding(num_classes + 1, hidden_dim) + nn.init.normal_(self.label_embedding.weight, std=0.02) + self.field = ImageCrossTransformer( + grid_size=self.grid_size, + latent_dim=latent_dim, + hidden_dim=hidden_dim, + depth=depth, + num_heads=num_heads, + dropout=dropout, + mlp_layers=mlp_layers, + mlp_ratio=mlp_ratio, + ) + patch_width = patch_size * patch_size * 3 + self.decoder = nn.Sequential( + nn.Linear(latent_dim, decoder_hidden_dim), + nn.SiLU(), + nn.Linear(decoder_hidden_dim, decoder_hidden_dim), + nn.SiLU(), + nn.Linear(decoder_hidden_dim, patch_width), + ) + + @staticmethod + def unroll_steps_at(optimizer_step: int) -> int: + step = max(int(optimizer_step), 1) + if step <= 2000: + return 1 if step % 2 else 2 + cycle = (2,) * 16 + (4,) * 3 + (8,) + return cycle[(step - 2001) % len(cycle)] + + @staticmethod + def sample_step_sizes( + batch_size: int, + num_steps: int, + device: torch.device, + generator: Optional[torch.Generator] = None, + ) -> torch.Tensor: + if num_steps <= 0: + raise ValueError("num_steps must be positive") + if num_steps == 1: + return torch.ones(batch_size, 1, device=device, dtype=torch.float32) + interior = torch.rand( + batch_size, + num_steps - 1, + device=device, + dtype=torch.float64, + generator=generator, + ).sort(dim=-1).values + endpoints = torch.cat( + [ + torch.zeros(batch_size, 1, device=device, dtype=torch.float64), + interior, + torch.ones(batch_size, 1, device=device, dtype=torch.float64), + ], + dim=-1, + ) + steps = endpoints.diff(dim=-1).to(torch.float32) + if not bool(torch.all(steps > 0)): + raise RuntimeError("Integration grid contains a non-positive step") + if not torch.allclose( + steps.sum(dim=-1), torch.ones(batch_size, device=device), atol=1e-6, rtol=1e-6 + ): + raise RuntimeError("Integration grid does not sum to one") + return steps + + def _condition(self, labels: torch.Tensor, drop: bool) -> torch.Tensor: + if drop and self.label_drop_prob > 0: + mask = torch.rand(labels.shape[0], device=labels.device) < self.label_drop_prob + labels = torch.where(mask, torch.full_like(labels, self.num_classes), labels) + return self.label_embedding(labels).unsqueeze(1) + + def _velocity( + self, latent: torch.Tensor, time: torch.Tensor, condition: torch.Tensor + ) -> torch.Tensor: + if self.gradient_checkpointing and self.training and torch.is_grad_enabled(): + return checkpoint(self.field, latent, time, condition, use_reentrant=False) + return self.field(latent, time, condition) + + def integrate( + self, + initial_latent: torch.Tensor, + condition: torch.Tensor, + num_steps: int, + step_sizes: Optional[torch.Tensor] = None, + ) -> IntegrationResult: + batch = initial_latent.shape[0] + if step_sizes is None: + step_sizes = torch.full( + (batch, num_steps), 1.0 / num_steps, + device=initial_latent.device, dtype=torch.float32, + ) + if step_sizes.shape != (batch, num_steps): + raise ValueError( + f"Expected step_sizes {(batch, num_steps)}, got {tuple(step_sizes.shape)}" + ) + if not bool(torch.all(step_sizes > 0)): + raise ValueError("All integration steps must be positive") + if not torch.allclose( + step_sizes.sum(-1), torch.ones(batch, device=step_sizes.device), + atol=1e-6, rtol=1e-6, + ): + raise ValueError("Each integration grid must sum to one") + latent = initial_latent + time = torch.zeros(batch, device=latent.device, dtype=torch.float32) + deltas = [] + for index in range(num_steps): + velocity = self._velocity(latent, time, condition) + dt = step_sizes[:, index].reshape(batch, 1, 1) + delta = dt * velocity + latent = latent + delta + deltas.append(delta.detach().square().mean()) + time = time + step_sizes[:, index] + delta_rms = torch.stack(deltas).mean().sqrt() + return IntegrationResult(latent, delta_rms, step_sizes) + + def decode(self, latent: torch.Tensor) -> torch.Tensor: + patches = self.decoder(latent) + return _unpatchify(patches, self.patch_size, channels=3) + + def predict( + self, + labels: torch.Tensor, + optimizer_step: int, + force_steps: Optional[int] = None, + noise: Optional[torch.Tensor] = None, + step_sizes: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, IntegrationResult]: + batch = labels.shape[0] + if noise is None: + noise = torch.randn( + batch, self.num_tokens, self.latent_dim, device=labels.device + ) + condition = self._condition(labels, drop=self.training) + num_steps = int(force_steps or self.unroll_steps_at(optimizer_step)) + if step_sizes is None: + step_sizes = self.sample_step_sizes(batch, num_steps, labels.device) + result = self.integrate(noise, condition, num_steps, step_sizes) + return self.decode(result.endpoint), result + + def forward( + self, + images: torch.Tensor, + labels: torch.Tensor, + optimizer_step: int, + force_steps: Optional[int] = None, + ) -> Dict[str, torch.Tensor]: + # The target image is deliberately consumed only after latent generation + # and terminal decoding have completed. + prediction, result = self.predict(labels, optimizer_step, force_steps=force_steps) + loss = (prediction - images).square().mean() + return { + "loss": loss, + "prediction_rms": prediction.detach().square().mean().sqrt(), + "noise_rms": result.endpoint.new_tensor(1.0), + "endpoint_rms": result.endpoint.detach().square().mean().sqrt(), + "latent_delta_rms": result.delta_rms.detach(), + "unroll_steps": result.endpoint.new_tensor(result.step_sizes.shape[1]), + } + + @torch.no_grad() + def sample( + self, + labels: torch.Tensor, + num_steps: int = 16, + cfg_scale: float = 1.0, + noise: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + batch = labels.shape[0] + latent = ( + torch.randn(batch, self.num_tokens, self.latent_dim, device=labels.device) + if noise is None else noise.clone() + ) + conditional = self._condition(labels, drop=False) + unconditional = self._condition( + torch.full_like(labels, self.num_classes), drop=False + ) if cfg_scale != 1.0 else None + time = torch.zeros(batch, device=labels.device, dtype=torch.float32) + for _ in range(num_steps): + velocity = self.field(latent, time, conditional) + if unconditional is not None: + uncond_velocity = self.field(latent, time, unconditional) + velocity = uncond_velocity + cfg_scale * (velocity - uncond_velocity) + latent = latent + velocity / num_steps + time = time + 1.0 / num_steps + return self.decode(latent) + + +def build_model(architecture: str, image_size: int = 256, num_classes: int = 1000) -> nn.Module: + if architecture == "jit_b16": + return JiTObjective(image_size=image_size, num_classes=num_classes) + if architecture == "endpoint_latent": + return EndpointLatentObjective(image_size=image_size, num_classes=num_classes) + raise ValueError(f"Unknown architecture: {architecture}") + + +def trainable_parameter_count(module: nn.Module) -> int: + return sum(parameter.numel() for parameter in module.parameters() if parameter.requires_grad) diff --git a/external/JiT/model_jit.py b/external/JiT/model_jit.py new file mode 100644 index 000000000..d2f53abca --- /dev/null +++ b/external/JiT/model_jit.py @@ -0,0 +1,394 @@ +# -------------------------------------------------------- +# References: +# SiT: https://github.com/willisma/SiT +# Lightning-DiT: https://github.com/hustvl/LightningDiT +# -------------------------------------------------------- +import torch +import torch.nn as nn +import math +import torch.nn.functional as F +from util.model_util import VisionRotaryEmbeddingFast, get_2d_sincos_pos_embed, RMSNorm + + +def modulate(x, shift, scale): + return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) + + +class BottleneckPatchEmbed(nn.Module): + """ Image to Patch Embedding + """ + def __init__(self, img_size=224, patch_size=16, in_chans=3, pca_dim=768, embed_dim=768, bias=True): + super().__init__() + img_size = (img_size, img_size) + patch_size = (patch_size, patch_size) + num_patches = (img_size[1] // patch_size[1]) * (img_size[0] // patch_size[0]) + self.img_size = img_size + self.patch_size = patch_size + self.num_patches = num_patches + + self.proj1 = nn.Conv2d(in_chans, pca_dim, kernel_size=patch_size, stride=patch_size, bias=False) + self.proj2 = nn.Conv2d(pca_dim, embed_dim, kernel_size=1, stride=1, bias=bias) + + def forward(self, x): + B, C, H, W = x.shape + assert H == self.img_size[0] and W == self.img_size[1], \ + f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})." + x = self.proj2(self.proj1(x)).flatten(2).transpose(1, 2) + return x + + +class TimestepEmbedder(nn.Module): + """ + Embeds scalar timesteps into vector representations. + """ + def __init__(self, hidden_size, frequency_embedding_size=256): + super().__init__() + self.mlp = nn.Sequential( + nn.Linear(frequency_embedding_size, hidden_size, bias=True), + nn.SiLU(), + nn.Linear(hidden_size, hidden_size, bias=True), + ) + self.frequency_embedding_size = frequency_embedding_size + + @staticmethod + def timestep_embedding(t, dim, max_period=10000): + """ + Create sinusoidal timestep embeddings. + :param t: a 1-D Tensor of N indices, one per batch element. + These may be fractional. + :param dim: the dimension of the output. + :param max_period: controls the minimum frequency of the embeddings. + :return: an (N, D) Tensor of positional embeddings. + """ + # https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py + half = dim // 2 + freqs = torch.exp( + -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half + ).to(device=t.device) + args = t[:, None].float() * freqs[None] + embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) + if dim % 2: + embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) + return embedding + + def forward(self, t): + t_freq = self.timestep_embedding(t, self.frequency_embedding_size) + t_emb = self.mlp(t_freq) + return t_emb + + +class LabelEmbedder(nn.Module): + """ + Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance. + """ + def __init__(self, num_classes, hidden_size): + super().__init__() + self.embedding_table = nn.Embedding(num_classes + 1, hidden_size) + self.num_classes = num_classes + + def forward(self, labels): + embeddings = self.embedding_table(labels) + return embeddings + + +def scaled_dot_product_attention(query, key, value, dropout_p=0.0) -> torch.Tensor: + L, S = query.size(-2), key.size(-2) + scale_factor = 1 / math.sqrt(query.size(-1)) + attn_bias = torch.zeros(query.size(0), 1, L, S, dtype=query.dtype).cuda() + + with torch.cuda.amp.autocast(enabled=False): + attn_weight = query.float() @ key.float().transpose(-2, -1) * scale_factor + attn_weight += attn_bias + attn_weight = torch.softmax(attn_weight, dim=-1) + attn_weight = torch.dropout(attn_weight, dropout_p, train=True) + return attn_weight @ value + + +class Attention(nn.Module): + def __init__(self, dim, num_heads=8, qkv_bias=True, qk_norm=True, attn_drop=0., proj_drop=0.): + super().__init__() + self.num_heads = num_heads + head_dim = dim // num_heads + + self.q_norm = RMSNorm(head_dim) if qk_norm else nn.Identity() + self.k_norm = RMSNorm(head_dim) if qk_norm else nn.Identity() + + self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) + self.attn_drop = nn.Dropout(attn_drop) + self.proj = nn.Linear(dim, dim) + self.proj_drop = nn.Dropout(proj_drop) + + def forward(self, x, rope): + B, N, C = x.shape + qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) + q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple) + + q = self.q_norm(q) + k = self.k_norm(k) + + q = rope(q) + k = rope(k) + + x = scaled_dot_product_attention(q, k, v, dropout_p=self.attn_drop.p if self.training else 0.) + + x = x.transpose(1, 2).reshape(B, N, C) + + x = self.proj(x) + x = self.proj_drop(x) + return x + + +class SwiGLUFFN(nn.Module): + def __init__( + self, + dim: int, + hidden_dim: int, + drop=0.0, + bias=True + ) -> None: + super().__init__() + hidden_dim = int(hidden_dim * 2 / 3) + self.w12 = nn.Linear(dim, 2 * hidden_dim, bias=bias) + self.w3 = nn.Linear(hidden_dim, dim, bias=bias) + self.ffn_dropout = nn.Dropout(drop) + + def forward(self, x): + x12 = self.w12(x) + x1, x2 = x12.chunk(2, dim=-1) + hidden = F.silu(x1) * x2 + return self.w3(self.ffn_dropout(hidden)) + + +class FinalLayer(nn.Module): + """ + The final layer of JiT. + """ + def __init__(self, hidden_size, patch_size, out_channels): + super().__init__() + self.norm_final = RMSNorm(hidden_size) + self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), + nn.Linear(hidden_size, 2 * hidden_size, bias=True) + ) + + @torch.compile + def forward(self, x, c): + shift, scale = self.adaLN_modulation(c).chunk(2, dim=1) + x = modulate(self.norm_final(x), shift, scale) + x = self.linear(x) + return x + + +class JiTBlock(nn.Module): + def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, attn_drop=0.0, proj_drop=0.0): + super().__init__() + self.norm1 = RMSNorm(hidden_size, eps=1e-6) + self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True, qk_norm=True, + attn_drop=attn_drop, proj_drop=proj_drop) + self.norm2 = RMSNorm(hidden_size, eps=1e-6) + mlp_hidden_dim = int(hidden_size * mlp_ratio) + self.mlp = SwiGLUFFN(hidden_size, mlp_hidden_dim, drop=proj_drop) + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), + nn.Linear(hidden_size, 6 * hidden_size, bias=True) + ) + + @torch.compile + def forward(self, x, c, feat_rope=None): + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=-1) + x = x + gate_msa.unsqueeze(1) * self.attn(modulate(self.norm1(x), shift_msa, scale_msa), rope=feat_rope) + x = x + gate_mlp.unsqueeze(1) * self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp)) + return x + + +class JiT(nn.Module): + """ + Just image Transformer. + """ + def __init__( + self, + input_size=256, + patch_size=16, + in_channels=3, + hidden_size=1024, + depth=24, + num_heads=16, + mlp_ratio=4.0, + attn_drop=0.0, + proj_drop=0.0, + num_classes=1000, + bottleneck_dim=128, + in_context_len=32, + in_context_start=8 + ): + super().__init__() + self.in_channels = in_channels + self.out_channels = in_channels + self.patch_size = patch_size + self.num_heads = num_heads + self.hidden_size = hidden_size + self.input_size = input_size + self.in_context_len = in_context_len + self.in_context_start = in_context_start + self.num_classes = num_classes + + # time and class embed + self.t_embedder = TimestepEmbedder(hidden_size) + self.y_embedder = LabelEmbedder(num_classes, hidden_size) + + # linear embed + self.x_embedder = BottleneckPatchEmbed(input_size, patch_size, in_channels, bottleneck_dim, hidden_size, bias=True) + + # use fixed sin-cos embedding + num_patches = self.x_embedder.num_patches + self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, hidden_size), requires_grad=False) + + # in-context cls token + if self.in_context_len > 0: + self.in_context_posemb = nn.Parameter(torch.zeros(1, self.in_context_len, hidden_size), requires_grad=True) + torch.nn.init.normal_(self.in_context_posemb, std=.02) + + # rope + half_head_dim = hidden_size // num_heads // 2 + hw_seq_len = input_size // patch_size + self.feat_rope = VisionRotaryEmbeddingFast( + dim=half_head_dim, + pt_seq_len=hw_seq_len, + num_cls_token=0 + ) + self.feat_rope_incontext = VisionRotaryEmbeddingFast( + dim=half_head_dim, + pt_seq_len=hw_seq_len, + num_cls_token=self.in_context_len + ) + + # transformer + self.blocks = nn.ModuleList([ + JiTBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio, + attn_drop=attn_drop if (depth // 4 * 3 > i >= depth // 4) else 0.0, + proj_drop=proj_drop if (depth // 4 * 3 > i >= depth // 4) else 0.0) + for i in range(depth) + ]) + + # linear predict + self.final_layer = FinalLayer(hidden_size, patch_size, self.out_channels) + + self.initialize_weights() + + def initialize_weights(self): + # Initialize transformer layers: + def _basic_init(module): + if isinstance(module, nn.Linear): + torch.nn.init.xavier_uniform_(module.weight) + if module.bias is not None: + nn.init.constant_(module.bias, 0) + self.apply(_basic_init) + + # Initialize (and freeze) pos_embed by sin-cos embedding: + pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], int(self.x_embedder.num_patches ** 0.5)) + self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0)) + + # Initialize patch_embed like nn.Linear (instead of nn.Conv2d): + w1 = self.x_embedder.proj1.weight.data + nn.init.xavier_uniform_(w1.view([w1.shape[0], -1])) + w2 = self.x_embedder.proj2.weight.data + nn.init.xavier_uniform_(w2.view([w2.shape[0], -1])) + nn.init.constant_(self.x_embedder.proj2.bias, 0) + + # Initialize label embedding table: + nn.init.normal_(self.y_embedder.embedding_table.weight, std=0.02) + + nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02) + nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02) + + # Zero-out adaLN modulation layers: + for block in self.blocks: + nn.init.constant_(block.adaLN_modulation[-1].weight, 0) + nn.init.constant_(block.adaLN_modulation[-1].bias, 0) + + # Zero-out output layers: + nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0) + nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0) + + nn.init.constant_(self.final_layer.linear.weight, 0) + nn.init.constant_(self.final_layer.linear.bias, 0) + + def unpatchify(self, x, p): + """ + x: (N, T, patch_size**2 * C) + imgs: (N, H, W, C) + """ + c = self.out_channels + h = w = int(x.shape[1] ** 0.5) + assert h * w == x.shape[1] + + x = x.reshape(shape=(x.shape[0], h, w, p, p, c)) + x = torch.einsum('nhwpqc->nchpwq', x) + imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p)) + return imgs + + def forward(self, x, t, y): + """ + x: (N, C, H, W) + t: (N,) + y: (N,) + """ + # class and time embeddings + t_emb = self.t_embedder(t) + y_emb = self.y_embedder(y) + c = t_emb + y_emb + + # forward JiT + x = self.x_embedder(x) + x += self.pos_embed + + for i, block in enumerate(self.blocks): + # in-context + if self.in_context_len > 0 and i == self.in_context_start: + in_context_tokens = y_emb.unsqueeze(1).repeat(1, self.in_context_len, 1) + in_context_tokens += self.in_context_posemb + x = torch.cat([in_context_tokens, x], dim=1) + x = block(x, c, self.feat_rope if i < self.in_context_start else self.feat_rope_incontext) + + x = x[:, self.in_context_len:] + + x = self.final_layer(x, c) + output = self.unpatchify(x, self.patch_size) + + return output + + +def JiT_B_16(**kwargs): + return JiT(depth=12, hidden_size=768, num_heads=12, + bottleneck_dim=128, in_context_len=32, in_context_start=4, patch_size=16, **kwargs) + +def JiT_B_32(**kwargs): + return JiT(depth=12, hidden_size=768, num_heads=12, + bottleneck_dim=128, in_context_len=32, in_context_start=4, patch_size=32, **kwargs) + +def JiT_L_16(**kwargs): + return JiT(depth=24, hidden_size=1024, num_heads=16, + bottleneck_dim=128, in_context_len=32, in_context_start=8, patch_size=16, **kwargs) + +def JiT_L_32(**kwargs): + return JiT(depth=24, hidden_size=1024, num_heads=16, + bottleneck_dim=128, in_context_len=32, in_context_start=8, patch_size=32, **kwargs) + +def JiT_H_16(**kwargs): + return JiT(depth=32, hidden_size=1280, num_heads=16, + bottleneck_dim=256, in_context_len=32, in_context_start=10, patch_size=16, **kwargs) + +def JiT_H_32(**kwargs): + return JiT(depth=32, hidden_size=1280, num_heads=16, + bottleneck_dim=256, in_context_len=32, in_context_start=10, patch_size=32, **kwargs) + + +JiT_models = { + 'JiT-B/16': JiT_B_16, + 'JiT-B/32': JiT_B_32, + 'JiT-L/16': JiT_L_16, + 'JiT-L/32': JiT_L_32, + 'JiT-H/16': JiT_H_16, + 'JiT-H/32': JiT_H_32, +} diff --git a/external/JiT/prepare_ref.py b/external/JiT/prepare_ref.py new file mode 100644 index 000000000..776406950 --- /dev/null +++ b/external/JiT/prepare_ref.py @@ -0,0 +1,61 @@ +import os +import argparse +from torchvision import transforms, datasets +from torch.utils.data import DataLoader +from util.crop import center_crop_arr + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument('--data_path', type=str, required=True, + help='Path to ImageNet root directory') + parser.add_argument('--output_path', type=str, default='imagenet-train-256', + help='Folder where transformed images will be saved') + parser.add_argument('--img_size', type=int, default=256, + help='Resolution to center-crop and resize') + args = parser.parse_args() + + transform_train = transforms.Compose([ + transforms.Lambda(lambda pil_image: center_crop_arr(pil_image, args.img_size)), + transforms.CenterCrop(args.img_size), + transforms.ToTensor(), + ]) + + dataset_train = datasets.ImageFolder( + os.path.join(args.data_path, 'train'), + transform=transform_train + ) + + data_loader = DataLoader( + dataset_train, + batch_size=256, + num_workers=32, + shuffle=False, + pin_memory=False + ) + + os.makedirs(args.output_path, exist_ok=True) + + to_pil = transforms.ToPILImage() + global_idx = 0 + + from tqdm import tqdm + for batch_images, batch_labels in tqdm(data_loader): + for i in range(batch_images.size(0)): + img_tensor = batch_images[i] + + pil_img = to_pil(img_tensor) + out_path = os.path.join( + args.output_path, + f"transformed_{global_idx:08d}.png" + ) + pil_img.save(out_path, format='PNG', compress_level=0) + global_idx += 1 + + print(f"Saved batch up to index={global_idx} ...") + + print("Finished saving all images.") + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/external/JiT/test_matched_models.py b/external/JiT/test_matched_models.py new file mode 100644 index 000000000..624387fa0 --- /dev/null +++ b/external/JiT/test_matched_models.py @@ -0,0 +1,86 @@ +import torch + +from matched_models import EndpointLatentObjective + + +def tiny_model() -> EndpointLatentObjective: + return EndpointLatentObjective( + image_size=8, + patch_size=4, + latent_dim=4, + hidden_dim=8, + depth=2, + num_heads=2, + dropout=0.0, + mlp_layers=1, + mlp_ratio=2.0, + decoder_hidden_dim=8, + num_classes=3, + gradient_checkpointing=False, + ) + + +def test_integration_grids_are_positive_fp32_and_sum_to_one(): + for steps in (1, 2, 4, 8, 16): + grid = EndpointLatentObjective.sample_step_sizes(7, steps, torch.device("cpu")) + assert grid.dtype == torch.float32 + assert torch.all(grid > 0) + torch.testing.assert_close(grid.sum(-1), torch.ones(7)) + + +def test_optimizer_step_curriculum_matches_action_reference(): + assert [EndpointLatentObjective.unroll_steps_at(step) for step in range(1, 7)] == [ + 1, + 2, + 1, + 2, + 1, + 2, + ] + cycle = [EndpointLatentObjective.unroll_steps_at(step) for step in range(2001, 2021)] + assert cycle == [2] * 16 + [4] * 3 + [8] + + +def test_target_is_not_an_input_to_latent_generation(): + model = tiny_model().eval() + labels = torch.tensor([0, 1]) + noise = torch.randn(2, model.num_tokens, model.latent_dim) + steps = torch.full((2, 2), 0.5) + first, first_result = model.predict( + labels, optimizer_step=2, force_steps=2, noise=noise, step_sizes=steps + ) + second, second_result = model.predict( + labels, optimizer_step=2, force_steps=2, noise=noise, step_sizes=steps + ) + torch.testing.assert_close(first, second) + torch.testing.assert_close(first_result.endpoint, second_result.endpoint) + + +def test_terminal_image_loss_reaches_decoder_and_field(): + model = tiny_model().train() + images = torch.randn(2, 3, 8, 8) + labels = torch.tensor([0, 1]) + metrics = model(images, labels, optimizer_step=2, force_steps=2) + assert torch.isfinite(metrics["loss"]) + metrics["loss"].backward() + decoder_grad = sum( + parameter.grad.abs().sum() + for parameter in model.decoder.parameters() + if parameter.grad is not None + ) + field_grad = sum( + parameter.grad.abs().sum() + for parameter in model.field.parameters() + if parameter.grad is not None + ) + assert decoder_grad > 0 + assert field_grad > 0 + + +def test_sampling_shape_and_noise_sensitivity(): + model = tiny_model().eval() + labels = torch.tensor([0, 1]) + first = model.sample(labels, num_steps=2) + second = model.sample(labels, num_steps=2) + assert first.shape == (2, 3, 8, 8) + assert not torch.equal(first, second) diff --git a/external/JiT/train_matched.py b/external/JiT/train_matched.py new file mode 100644 index 000000000..8c5d0542b --- /dev/null +++ b/external/JiT/train_matched.py @@ -0,0 +1,623 @@ +"""Distributed matched training for JiT-B/16 and endpoint latent denoising.""" + +from __future__ import annotations + +import argparse +import contextlib +import copy +import json +import math +import os +import random +import signal +import sys +import time +from pathlib import Path +from typing import Dict, Iterable, Optional, Tuple + +import numpy as np +import torch +import torch.distributed as dist +import torchvision.datasets as datasets +import torchvision.transforms as transforms +from torch import nn +from torch.nn.parallel import DistributedDataParallel +from torch.utils.data import DataLoader, DistributedSampler +from torchvision.utils import save_image + +from matched_models import build_model, trainable_parameter_count +from util.crop import center_crop_arr + + +STOP_REQUESTED = False + + +def _request_stop(signum, _frame) -> None: + global STOP_REQUESTED + STOP_REQUESTED = True + print(f"CHECKPOINT_STOP_REQUESTED signal={signum}", flush=True) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--architecture", choices=("jit_b16", "endpoint_latent"), required=True) + parser.add_argument("--data-path", required=True) + parser.add_argument("--output-dir", required=True) + parser.add_argument("--image-size", type=int, default=256) + parser.add_argument("--num-classes", type=int, default=1000) + parser.add_argument("--epochs", type=int, default=600) + parser.add_argument("--batch-size", type=int, default=8, help="Per GPU microbatch") + parser.add_argument("--grad-accum", type=int, default=1) + parser.add_argument("--base-lr", type=float, default=5e-5) + parser.add_argument("--warmup-epochs", type=float, default=5.0) + parser.add_argument("--weight-decay", type=float, default=0.0) + parser.add_argument("--ema-decay", type=float, default=0.9999) + parser.add_argument("--max-optimizer-steps", type=int, default=0) + parser.add_argument("--save-every-steps", type=int, default=1000) + parser.add_argument("--val-every-steps", type=int, default=5000) + parser.add_argument("--val-batches", type=int, default=8) + parser.add_argument("--sample-batch", type=int, default=4) + parser.add_argument("--sample-steps", type=int, default=16) + parser.add_argument("--cfg-scale", type=float, default=1.0) + parser.add_argument("--log-every-steps", type=int, default=10) + parser.add_argument("--num-workers", type=int, default=12) + parser.add_argument("--seed", type=int, default=42) + parser.add_argument("--resume", action="store_true") + parser.add_argument("--overfit-one-batch", action="store_true") + parser.add_argument("--overfit-force-steps", type=int, default=2) + parser.add_argument("--smoke-require-validation", action="store_true") + parser.add_argument("--wandb-project", default="") + parser.add_argument("--wandb-name", default="") + parser.add_argument("--wandb-mode", default="offline", choices=("online", "offline", "disabled")) + parser.add_argument("--expected-min-params", type=int, default=0) + parser.add_argument("--expected-max-params", type=int, default=0) + parser.add_argument("--stop-file", default="") + return parser.parse_args() + + +def init_distributed() -> Tuple[int, int, int, torch.device]: + if not torch.cuda.is_available(): + raise RuntimeError("Matched image training requires CUDA") + rank = int(os.environ.get("RANK", "0")) + world_size = int(os.environ.get("WORLD_SIZE", "1")) + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + torch.cuda.set_device(local_rank) + if world_size > 1: + dist.init_process_group("nccl", init_method="env://") + return rank, world_size, local_rank, torch.device("cuda", local_rank) + + +def is_main(rank: int) -> bool: + return rank == 0 + + +def barrier(world_size: int) -> None: + if world_size > 1: + dist.barrier() + + +def reduce_mean(value: torch.Tensor, world_size: int) -> torch.Tensor: + result = value.detach().float().clone() + if world_size > 1: + dist.all_reduce(result, op=dist.ReduceOp.SUM) + result /= world_size + return result + + +def transform(image_size: int): + return transforms.Compose( + [ + transforms.Lambda(lambda image: center_crop_arr(image, image_size)), + transforms.RandomHorizontalFlip(), + transforms.PILToTensor(), + ] + ) + + +def validation_transform(image_size: int): + return transforms.Compose( + [ + transforms.Lambda(lambda image: center_crop_arr(image, image_size)), + transforms.PILToTensor(), + ] + ) + + +def normalize(images: torch.Tensor, device: torch.device) -> torch.Tensor: + return images.to(device, non_blocking=True).float().div_(255.0).mul_(2.0).sub_(1.0) + + +@torch.no_grad() +def update_ema(ema: nn.Module, source: nn.Module, decay: float) -> None: + source_parameters = dict(source.named_parameters()) + for name, parameter in ema.named_parameters(): + parameter.mul_(decay).add_(source_parameters[name], alpha=1.0 - decay) + source_buffers = dict(source.named_buffers()) + for name, buffer in ema.named_buffers(): + if name in source_buffers and buffer.shape == source_buffers[name].shape: + buffer.copy_(source_buffers[name]) + + +def append_jsonl(path: Path, payload: Dict) -> None: + with path.open("a", encoding="utf-8") as handle: + handle.write(json.dumps(payload, sort_keys=True) + "\n") + + +def atomic_checkpoint(path: Path, payload: Dict) -> None: + temporary = path.with_suffix(path.suffix + ".tmp") + torch.save(payload, temporary) + os.replace(temporary, path) + + +def save_checkpoint( + output_dir: Path, + model: nn.Module, + ema: nn.Module, + optimizer: torch.optim.Optimizer, + args: argparse.Namespace, + epoch: int, + next_batch: int, + optimizer_step: int, + rank: int, +) -> None: + if not is_main(rank): + return + payload = { + "model": model.state_dict(), + "ema": ema.state_dict(), + "optimizer": optimizer.state_dict(), + "args": vars(args), + "epoch": int(epoch), + "next_batch": int(next_batch), + "optimizer_step": int(optimizer_step), + } + atomic_checkpoint(output_dir / "checkpoint-last.pth", payload) + print(f"CHECKPOINT_SAVED step={optimizer_step} epoch={epoch} next_batch={next_batch}", flush=True) + + +def load_checkpoint( + path: Path, + model: nn.Module, + ema: nn.Module, + optimizer: torch.optim.Optimizer, + device: torch.device, +) -> Tuple[int, int, int]: + checkpoint = torch.load(path, map_location=device, weights_only=False) + model.load_state_dict(checkpoint["model"], strict=True) + ema.load_state_dict(checkpoint["ema"], strict=True) + optimizer.load_state_dict(checkpoint["optimizer"]) + return ( + int(checkpoint["epoch"]), + int(checkpoint.get("next_batch", 0)), + int(checkpoint["optimizer_step"]), + ) + + +def current_lr( + base_lr: float, + effective_batch: int, + optimizer_step: int, + updates_per_epoch: int, + warmup_epochs: float, +) -> float: + absolute = base_lr * effective_batch / 256.0 + warmup_updates = max(int(warmup_epochs * updates_per_epoch), 1) + return absolute * min((optimizer_step + 1) / warmup_updates, 1.0) + + +def set_lr(optimizer: torch.optim.Optimizer, value: float) -> None: + for group in optimizer.param_groups: + group["lr"] = value + + +@torch.no_grad() +def validate( + ema: nn.Module, + loader: DataLoader, + sampler: DistributedSampler, + args: argparse.Namespace, + optimizer_step: int, + rank: int, + world_size: int, + device: torch.device, + output_dir: Path, +) -> Dict[str, float]: + ema.eval() + sampler.set_epoch(optimizer_step) + losses = [] + first_images = first_labels = None + with torch.random.fork_rng(devices=[device.index]): + torch.manual_seed(args.seed + 100_000 + rank) + for index, (images, labels) in enumerate(loader): + if index >= args.val_batches: + break + images = normalize(images, device) + labels = labels.to(device, non_blocking=True) + with torch.autocast("cuda", dtype=torch.bfloat16): + metrics = ema(images, labels, optimizer_step) + losses.append(metrics["loss"].detach().float()) + if first_images is None: + first_images, first_labels = images, labels + if not losses: + raise RuntimeError("Validation loader produced no batches") + local_loss = torch.stack(losses).mean() + val_loss = reduce_mean(local_loss, world_size) + result = {"val_loss": float(val_loss.item())} + + if is_main(rank): + count = min(args.sample_batch, first_labels.shape[0]) + labels = first_labels[:count] + targets = first_images[:count] + with torch.random.fork_rng(devices=[device.index]): + torch.manual_seed(args.seed + 200_000) + with torch.autocast("cuda", dtype=torch.bfloat16): + sample_a = ema.sample( + labels, num_steps=args.sample_steps, cfg_scale=args.cfg_scale + ).float() + torch.manual_seed(args.seed + 300_000) + with torch.autocast("cuda", dtype=torch.bfloat16): + sample_b = ema.sample( + labels, num_steps=args.sample_steps, cfg_scale=args.cfg_scale + ).float() + result.update( + { + "sample_mean": float(sample_a.mean().item()), + "sample_std": float(sample_a.std().item()), + "sample_pairwise_mse": float((sample_a - sample_b).square().mean().item()), + "sample_saturation": float((sample_a.abs() > 1.0).float().mean().item()), + } + ) + if args.architecture == "endpoint_latent": + same_seed_rows = [] + for steps in (1, 2, 4, 8, 16): + torch.manual_seed(args.seed + 400_000) + with torch.autocast("cuda", dtype=torch.bfloat16): + variant = ema.sample(labels, num_steps=steps, cfg_scale=args.cfg_scale) + same_seed_rows.append(variant.float()) + grid = torch.cat([targets, *same_seed_rows, sample_b], dim=0) + else: + grid = torch.cat([targets, sample_a, sample_b], dim=0) + sample_path = output_dir / f"samples-step{optimizer_step:08d}.png" + save_image(((grid.clamp(-1, 1) + 1) / 2).cpu(), sample_path, nrow=count) + result["sample_grid"] = str(sample_path) + numeric = [value for value in result.values() if isinstance(value, float)] + if not all(math.isfinite(value) for value in numeric): + raise RuntimeError(f"Validation produced non-finite metrics: {result}") + # JiT intentionally zero-initializes its output layer. Under the exact + # multi-epoch warmup, its first smoke-step sample can therefore be a + # finite constant image; record that fact rather than rejecting a valid + # train-plus-sampling path before the optimizer has moved appreciably. + barrier(world_size) + if is_main(rank): + print("VALIDATION_METRICS " + json.dumps(result, sort_keys=True), flush=True) + return result + + +def maybe_wandb(args: argparse.Namespace, rank: int, config: Dict): + if not is_main(rank) or args.wandb_mode == "disabled" or not args.wandb_project: + return None + import wandb + + return wandb.init( + project=args.wandb_project, + name=args.wandb_name or None, + dir=args.output_dir, + config=config, + mode=args.wandb_mode, + resume="allow", + ) + + +def main() -> int: + args = parse_args() + if args.batch_size <= 0 or args.grad_accum <= 0: + raise ValueError("batch-size and grad-accum must be positive") + rank, world_size, local_rank, device = init_distributed() + signal.signal(signal.SIGUSR1, _request_stop) + output_dir = Path(args.output_dir) + if is_main(rank): + output_dir.mkdir(parents=True, exist_ok=True) + barrier(world_size) + + seed = args.seed + rank + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.backends.cudnn.benchmark = True + torch._dynamo.config.cache_size_limit = 256 + torch._dynamo.config.optimize_ddp = False + + train_dataset = datasets.ImageFolder( + os.path.join(args.data_path, "train"), transform=transform(args.image_size) + ) + val_dataset = datasets.ImageFolder( + os.path.join(args.data_path, "val"), transform=validation_transform(args.image_size) + ) + if train_dataset.classes != val_dataset.classes: + raise RuntimeError("ImageNet train/val class mappings differ") + if len(train_dataset.classes) != args.num_classes: + raise RuntimeError( + f"Expected {args.num_classes} classes, found {len(train_dataset.classes)}" + ) + train_sampler = DistributedSampler( + train_dataset, num_replicas=world_size, rank=rank, shuffle=True, seed=args.seed + ) + val_sampler = DistributedSampler( + val_dataset, num_replicas=world_size, rank=rank, shuffle=False + ) + train_loader = DataLoader( + train_dataset, + sampler=train_sampler, + batch_size=args.batch_size, + num_workers=args.num_workers, + pin_memory=True, + drop_last=True, + persistent_workers=args.num_workers > 0, + ) + val_loader = DataLoader( + val_dataset, + sampler=val_sampler, + batch_size=max(args.sample_batch, 1), + num_workers=min(args.num_workers, 4), + pin_memory=True, + drop_last=False, + persistent_workers=args.num_workers > 0, + ) + + model = build_model(args.architecture, args.image_size, args.num_classes).to(device) + parameter_count = trainable_parameter_count(model) + if args.expected_min_params and parameter_count < args.expected_min_params: + raise RuntimeError(f"Parameter count {parameter_count} below required minimum") + if args.expected_max_params and parameter_count > args.expected_max_params: + raise RuntimeError(f"Parameter count {parameter_count} above required maximum") + ema = copy.deepcopy(model).eval() + for parameter in ema.parameters(): + parameter.requires_grad_(False) + distributed_model = DistributedDataParallel( + model, device_ids=[local_rank], broadcast_buffers=False, find_unused_parameters=False + ) + parameter_groups = [ + { + "params": [ + parameter + for name, parameter in model.named_parameters() + if parameter.requires_grad and parameter.ndim > 1 and not name.endswith("bias") + ], + "weight_decay": args.weight_decay, + }, + { + "params": [ + parameter + for name, parameter in model.named_parameters() + if parameter.requires_grad and (parameter.ndim <= 1 or name.endswith("bias")) + ], + "weight_decay": 0.0, + }, + ] + optimizer = torch.optim.AdamW(parameter_groups, lr=0.0, betas=(0.9, 0.95)) + + effective_batch = args.batch_size * world_size * args.grad_accum + updates_per_epoch = len(train_loader) // args.grad_accum + if updates_per_epoch <= 0: + raise RuntimeError("Gradient accumulation exceeds the epoch's microbatch count") + config = { + **vars(args), + "world_size": world_size, + "effective_batch": effective_batch, + "updates_per_epoch": updates_per_epoch, + "train_examples": len(train_dataset), + "val_examples": len(val_dataset), + "trainable_parameters": parameter_count, + "torch_version": torch.__version__, + } + if is_main(rank): + (output_dir / "resolved_config.json").write_text( + json.dumps(config, indent=2, sort_keys=True) + "\n", encoding="utf-8" + ) + print("RESOLVED_CONFIG " + json.dumps(config, sort_keys=True), flush=True) + run = maybe_wandb(args, rank, config) + + epoch = 0 + next_batch = 0 + optimizer_step = 0 + checkpoint_path = output_dir / "checkpoint-last.pth" + if args.resume and checkpoint_path.exists(): + epoch, next_batch, optimizer_step = load_checkpoint( + checkpoint_path, model, ema, optimizer, device + ) + print( + f"CHECKPOINT_LOADED step={optimizer_step} epoch={epoch} next_batch={next_batch}", + flush=True, + ) + + first_overfit_loss: Optional[float] = None + last_overfit_loss: Optional[float] = None + cached_batch = None + validations_run = 0 + log_path = output_dir / "metrics.jsonl" + optimizer.zero_grad(set_to_none=True) + training_start = time.time() + done = False + + while epoch < args.epochs and not done: + train_sampler.set_epoch(epoch) + micro_in_update = 0 + for batch_index, batch in enumerate(train_loader): + if batch_index < next_batch: + continue + if args.overfit_one_batch: + if cached_batch is None: + cached_batch = batch + batch = cached_batch + images, labels = batch + images = normalize(images, device) + labels = labels.to(device, non_blocking=True) + micro_in_update += 1 + sync_now = micro_in_update == args.grad_accum + context = contextlib.nullcontext() if sync_now else distributed_model.no_sync() + force_steps = args.overfit_force_steps if args.overfit_one_batch else None + rng_context = ( + torch.random.fork_rng(devices=[device.index]) + if args.overfit_one_batch else contextlib.nullcontext() + ) + with context, rng_context: + if args.overfit_one_batch: + torch.manual_seed(args.seed + rank) + with torch.autocast("cuda", dtype=torch.bfloat16): + metrics = distributed_model( + images, labels, optimizer_step + 1, force_steps=force_steps + ) + scaled_loss = metrics["loss"] / args.grad_accum + scaled_loss.backward() + if not sync_now: + continue + + lr = current_lr( + args.base_lr, + effective_batch, + optimizer_step, + updates_per_epoch, + 0.0 if args.overfit_one_batch else args.warmup_epochs, + ) + set_lr(optimizer, lr) + optimizer.step() + optimizer.zero_grad(set_to_none=True) + optimizer_step += 1 + micro_in_update = 0 + update_ema(ema, model, args.ema_decay) + next_batch = batch_index + 1 + + reduced = { + name: float(reduce_mean(value, world_size).item()) + for name, value in metrics.items() + } + reduced.update( + { + "optimizer_step": optimizer_step, + "epoch": epoch, + "batch_index": batch_index, + "lr": lr, + "elapsed_seconds": time.time() - training_start, + } + ) + if args.overfit_one_batch: + if first_overfit_loss is None: + first_overfit_loss = reduced["loss"] + last_overfit_loss = reduced["loss"] + if is_main(rank) and ( + optimizer_step == 1 or optimizer_step % args.log_every_steps == 0 + ): + print("TRAIN_METRICS " + json.dumps(reduced, sort_keys=True), flush=True) + append_jsonl(log_path, {"split": "train", **reduced}) + if run is not None: + run.log({f"train/{key}": value for key, value in reduced.items()}, step=optimizer_step) + + should_validate = ( + args.val_every_steps > 0 and optimizer_step % args.val_every_steps == 0 + ) + if should_validate: + validation = validate( + ema, + val_loader, + val_sampler, + args, + optimizer_step, + rank, + world_size, + device, + output_dir, + ) + validations_run += 1 + if is_main(rank): + append_jsonl( + log_path, + {"split": "validation", "optimizer_step": optimizer_step, **validation}, + ) + if run is not None: + run.log( + { + f"validation/{key}": value + for key, value in validation.items() + if isinstance(value, (int, float)) + }, + step=optimizer_step, + ) + model.train() + + if optimizer_step % args.save_every_steps == 0 or STOP_REQUESTED: + save_checkpoint( + output_dir, + model, + ema, + optimizer, + args, + epoch, + next_batch, + optimizer_step, + rank, + ) + barrier(world_size) + + stop_file_requested = bool(args.stop_file) and Path(args.stop_file).exists() + if stop_file_requested: + save_checkpoint( + output_dir, + model, + ema, + optimizer, + args, + epoch, + next_batch, + optimizer_step, + rank, + ) + barrier(world_size) + if STOP_REQUESTED or stop_file_requested: + done = True + break + if args.max_optimizer_steps and optimizer_step >= args.max_optimizer_steps: + done = True + break + if args.overfit_one_batch: + next_batch = 0 + if not done: + continue + + if not done: + if micro_in_update: + optimizer.zero_grad(set_to_none=True) + epoch += 1 + next_batch = 0 + + save_checkpoint( + output_dir, model, ema, optimizer, args, epoch, next_batch, optimizer_step, rank + ) + barrier(world_size) + if args.overfit_one_batch: + if first_overfit_loss is None or last_overfit_loss is None: + raise RuntimeError("Overfit run produced no optimizer steps") + if not last_overfit_loss < first_overfit_loss: + raise RuntimeError( + f"Overfit loss did not improve: first={first_overfit_loss} last={last_overfit_loss}" + ) + if is_main(rank): + print( + f"OVERFIT_GATE_PASSED first={first_overfit_loss:.8f} " + f"last={last_overfit_loss:.8f}", + flush=True, + ) + if args.smoke_require_validation: + if validations_run < 1: + raise RuntimeError("Smoke ended without scheduled post-training validation") + if is_main(rank): + print(f"SMOKE_GATE_PASSED validations={validations_run}", flush=True) + if run is not None: + run.finish() + barrier(world_size) + if world_size > 1: + dist.destroy_process_group() + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/external/JiT/util/crop.py b/external/JiT/util/crop.py new file mode 100644 index 000000000..7582690e6 --- /dev/null +++ b/external/JiT/util/crop.py @@ -0,0 +1,23 @@ +import numpy as np +from PIL import Image + + +def center_crop_arr(pil_image, image_size): + """ + Center cropping implementation from ADM. + https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126 + """ + while min(*pil_image.size) >= 2 * image_size: + pil_image = pil_image.resize( + tuple(x // 2 for x in pil_image.size), resample=Image.BOX + ) + + scale = image_size / min(*pil_image.size) + pil_image = pil_image.resize( + tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC + ) + + arr = np.array(pil_image) + crop_y = (arr.shape[0] - image_size) // 2 + crop_x = (arr.shape[1] - image_size) // 2 + return Image.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size]) diff --git a/external/JiT/util/lr_sched.py b/external/JiT/util/lr_sched.py new file mode 100644 index 000000000..1ed515188 --- /dev/null +++ b/external/JiT/util/lr_sched.py @@ -0,0 +1,21 @@ +import math + + +def adjust_learning_rate(optimizer, epoch, args): + """Decay the learning rate with half-cycle cosine after warmup""" + if epoch < args.warmup_epochs: + lr = args.lr * epoch / args.warmup_epochs + else: + if args.lr_schedule == "constant": + lr = args.lr + elif args.lr_schedule == "cosine": + lr = args.min_lr + (args.lr - args.min_lr) * 0.5 * \ + (1. + math.cos(math.pi * (epoch - args.warmup_epochs) / (args.epochs - args.warmup_epochs))) + else: + raise NotImplementedError + for param_group in optimizer.param_groups: + if "lr_scale" in param_group: + param_group["lr"] = lr * param_group["lr_scale"] + else: + param_group["lr"] = lr + return lr diff --git a/external/JiT/util/misc.py b/external/JiT/util/misc.py new file mode 100644 index 000000000..b46861959 --- /dev/null +++ b/external/JiT/util/misc.py @@ -0,0 +1,289 @@ +import builtins +import datetime +import os +import time +from collections import defaultdict, deque +from pathlib import Path +import copy + +import torch +import torch.distributed as dist + + +class SmoothedValue(object): + """Track a series of values and provide access to smoothed values over a + window or the global series average. + """ + + def __init__(self, window_size=20, fmt=None): + if fmt is None: + fmt = "{median:.4f} ({global_avg:.4f})" + self.deque = deque(maxlen=window_size) + self.total = 0.0 + self.count = 0 + self.fmt = fmt + + def update(self, value, n=1): + self.deque.append(value) + self.count += n + self.total += value * n + + def synchronize_between_processes(self): + """ + Warning: does not synchronize the deque! + """ + if not is_dist_avail_and_initialized(): + return + t = torch.tensor([self.count, self.total], dtype=torch.float64, device='cuda') + dist.barrier() + dist.all_reduce(t) + t = t.tolist() + self.count = int(t[0]) + self.total = t[1] + + @property + def median(self): + d = torch.tensor(list(self.deque)) + return d.median().item() + + @property + def avg(self): + d = torch.tensor(list(self.deque), dtype=torch.float32) + return d.mean().item() + + @property + def global_avg(self): + return self.total / self.count + + @property + def max(self): + return max(self.deque) + + @property + def value(self): + return self.deque[-1] + + def __str__(self): + return self.fmt.format( + median=self.median, + avg=self.avg, + global_avg=self.global_avg, + max=self.max, + value=self.value) + + +class MetricLogger(object): + def __init__(self, delimiter="\t"): + self.meters = defaultdict(SmoothedValue) + self.delimiter = delimiter + + def update(self, **kwargs): + for k, v in kwargs.items(): + if v is None: + continue + if isinstance(v, torch.Tensor): + v = v.item() + assert isinstance(v, (float, int)) + self.meters[k].update(v) + + def __getattr__(self, attr): + if attr in self.meters: + return self.meters[attr] + if attr in self.__dict__: + return self.__dict__[attr] + raise AttributeError("'{}' object has no attribute '{}'".format( + type(self).__name__, attr)) + + def __str__(self): + loss_str = [] + for name, meter in self.meters.items(): + loss_str.append( + "{}: {}".format(name, str(meter)) + ) + return self.delimiter.join(loss_str) + + def synchronize_between_processes(self): + for meter in self.meters.values(): + meter.synchronize_between_processes() + + def add_meter(self, name, meter): + self.meters[name] = meter + + def log_every(self, iterable, print_freq, header=None): + i = 0 + if not header: + header = '' + start_time = time.time() + end = time.time() + iter_time = SmoothedValue(fmt='{avg:.4f}') + data_time = SmoothedValue(fmt='{avg:.4f}') + space_fmt = ':' + str(len(str(len(iterable)))) + 'd' + log_msg = [ + header, + '[{0' + space_fmt + '}/{1}]', + 'eta: {eta}', + '{meters}', + 'time: {time}', + 'data: {data}' + ] + if torch.cuda.is_available(): + log_msg.append('max mem: {memory:.0f}') + log_msg = self.delimiter.join(log_msg) + MB = 1024.0 * 1024.0 + for obj in iterable: + data_time.update(time.time() - end) + yield obj + iter_time.update(time.time() - end) + if i % print_freq == 0 or i == len(iterable) - 1: + eta_seconds = iter_time.global_avg * (len(iterable) - i) + eta_string = str(datetime.timedelta(seconds=int(eta_seconds))) + if torch.cuda.is_available(): + print(log_msg.format( + i, len(iterable), eta=eta_string, + meters=str(self), + time=str(iter_time), data=str(data_time), + memory=torch.cuda.max_memory_allocated() / MB)) + else: + print(log_msg.format( + i, len(iterable), eta=eta_string, + meters=str(self), + time=str(iter_time), data=str(data_time))) + i += 1 + end = time.time() + total_time = time.time() - start_time + total_time_str = str(datetime.timedelta(seconds=int(total_time))) + print('{} Total time: {} ({:.4f} s / it)'.format( + header, total_time_str, total_time / len(iterable))) + + +def setup_for_distributed(is_master): + """ + This function disables printing when not in master process + """ + builtin_print = builtins.print + + def print(*args, **kwargs): + force = kwargs.pop('force', False) + force = force or (get_world_size() > 8) + if is_master or force: + now = datetime.datetime.now().time() + builtin_print('[{}] '.format(now), end='') # print with time stamp + builtin_print(*args, **kwargs) + + builtins.print = print + + +def is_dist_avail_and_initialized(): + if not dist.is_available(): + return False + if not dist.is_initialized(): + return False + return True + + +def get_world_size(): + if not is_dist_avail_and_initialized(): + return 1 + return dist.get_world_size() + + +def get_rank(): + if not is_dist_avail_and_initialized(): + return 0 + return dist.get_rank() + + +def is_main_process(): + return get_rank() == 0 + + +def save_on_master(*args, **kwargs): + if is_main_process(): + torch.save(*args, **kwargs) + + +def init_distributed_mode(args): + if args.dist_on_itp: + args.rank = int(os.environ['OMPI_COMM_WORLD_RANK']) + args.world_size = int(os.environ['OMPI_COMM_WORLD_SIZE']) + args.gpu = int(os.environ['OMPI_COMM_WORLD_LOCAL_RANK']) + args.dist_url = "tcp://%s:%s" % (os.environ['MASTER_ADDR'], os.environ['MASTER_PORT']) + os.environ['LOCAL_RANK'] = str(args.gpu) + os.environ['RANK'] = str(args.rank) + os.environ['WORLD_SIZE'] = str(args.world_size) + # ["RANK", "WORLD_SIZE", "MASTER_ADDR", "MASTER_PORT", "LOCAL_RANK"] + elif 'RANK' in os.environ and 'WORLD_SIZE' in os.environ: + args.rank = int(os.environ["RANK"]) + args.world_size = int(os.environ['WORLD_SIZE']) + args.gpu = int(os.environ['LOCAL_RANK']) + elif 'SLURM_PROCID' in os.environ: + args.rank = int(os.environ['SLURM_PROCID']) + args.gpu = args.rank % torch.cuda.device_count() + else: + print('Not using distributed mode') + setup_for_distributed(is_master=True) # hack + args.distributed = False + return + + args.distributed = True + + torch.cuda.set_device(args.gpu) + args.dist_backend = 'nccl' + print('| distributed init (rank {}): {}, gpu {}'.format( + args.rank, args.dist_url, args.gpu), flush=True) + torch.distributed.init_process_group(backend=args.dist_backend, init_method=args.dist_url, + world_size=args.world_size, rank=args.rank) + torch.distributed.barrier() + setup_for_distributed(args.rank == 0) + + +def add_weight_decay(model, weight_decay=0, skip_list=()): + decay = [] + no_decay = [] + for name, param in model.named_parameters(): + if not param.requires_grad: + continue # frozen weights + if len(param.shape) == 1 or name.endswith(".bias") or name in skip_list or 'diffloss' in name: + no_decay.append(param) # no weight decay on bias, norm and diffloss + else: + decay.append(param) + return [ + {'params': no_decay, 'weight_decay': 0.}, + {'params': decay, 'weight_decay': weight_decay}] + + +def save_model(args, model_without_ddp, optimizer, epoch, epoch_name=None): + if epoch_name is None: + epoch_name = str(epoch) + output_dir = Path(args.output_dir) + checkpoint_path = output_dir / ('checkpoint-%s.pth' % epoch_name) + + to_save = { + 'model': model_without_ddp.state_dict(), + 'optimizer': optimizer.state_dict(), + 'epoch': epoch, + 'args': args, + } + + # ema + ema_state_dict1 = copy.deepcopy(model_without_ddp.state_dict()) + ema_state_dict2 = copy.deepcopy(model_without_ddp.state_dict()) + for i, (name, _value) in enumerate(model_without_ddp.named_parameters()): + assert name in ema_state_dict1 and name in ema_state_dict2 + ema_state_dict1[name] = model_without_ddp.ema_params1[i] + ema_state_dict2[name] = model_without_ddp.ema_params2[i] + to_save['model_ema1'] = ema_state_dict1 + to_save['model_ema2'] = ema_state_dict2 + + save_on_master(to_save, checkpoint_path) + + +def all_reduce_mean(x): + world_size = get_world_size() + if world_size > 1: + x_reduce = torch.tensor(x).cuda() + dist.all_reduce(x_reduce) + x_reduce /= world_size + return x_reduce.item() + else: + return x \ No newline at end of file diff --git a/external/JiT/util/model_util.py b/external/JiT/util/model_util.py new file mode 100644 index 000000000..9a3bad719 --- /dev/null +++ b/external/JiT/util/model_util.py @@ -0,0 +1,209 @@ +# -------------------------------------------------------- +# References: +# Lightning-DiT: https://github.com/hustvl/LightningDiT +# -------------------------------------------------------- + +from math import pi + +import torch +from torch import nn +import numpy as np + +from einops import rearrange, repeat + + +def broadcat(tensors, dim = -1): + num_tensors = len(tensors) + shape_lens = set(list(map(lambda t: len(t.shape), tensors))) + assert len(shape_lens) == 1, 'tensors must all have the same number of dimensions' + shape_len = list(shape_lens)[0] + dim = (dim + shape_len) if dim < 0 else dim + dims = list(zip(*map(lambda t: list(t.shape), tensors))) + expandable_dims = [(i, val) for i, val in enumerate(dims) if i != dim] + assert all([*map(lambda t: len(set(t[1])) <= 2, expandable_dims)]), 'invalid dimensions for broadcastable concatentation' + max_dims = list(map(lambda t: (t[0], max(t[1])), expandable_dims)) + expanded_dims = list(map(lambda t: (t[0], (t[1],) * num_tensors), max_dims)) + expanded_dims.insert(dim, (dim, dims[dim])) + expandable_shapes = list(zip(*map(lambda t: t[1], expanded_dims))) + tensors = list(map(lambda t: t[0].expand(*t[1]), zip(tensors, expandable_shapes))) + return torch.cat(tensors, dim = dim) + + +def rotate_half(x): + x = rearrange(x, '... (d r) -> ... d r', r = 2) + x1, x2 = x.unbind(dim = -1) + x = torch.stack((-x2, x1), dim = -1) + return rearrange(x, '... d r -> ... (d r)') + + +class VisionRotaryEmbedding(nn.Module): + def __init__( + self, + dim, + pt_seq_len, + ft_seq_len=None, + custom_freqs = None, + freqs_for = 'lang', + theta = 10000, + max_freq = 10, + num_freqs = 1, + ): + super().__init__() + if custom_freqs: + freqs = custom_freqs + elif freqs_for == 'lang': + freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)) + elif freqs_for == 'pixel': + freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi + elif freqs_for == 'constant': + freqs = torch.ones(num_freqs).float() + else: + raise ValueError(f'unknown modality {freqs_for}') + + if ft_seq_len is None: ft_seq_len = pt_seq_len + t = torch.arange(ft_seq_len) / ft_seq_len * pt_seq_len + + freqs_h = torch.einsum('..., f -> ... f', t, freqs) + freqs_h = repeat(freqs_h, '... n -> ... (n r)', r = 2) + + freqs_w = torch.einsum('..., f -> ... f', t, freqs) + freqs_w = repeat(freqs_w, '... n -> ... (n r)', r = 2) + + freqs = broadcat((freqs_h[:, None, :], freqs_w[None, :, :]), dim = -1) + + self.register_buffer("freqs_cos", freqs.cos()) + self.register_buffer("freqs_sin", freqs.sin()) + + def forward(self, t, start_index = 0): + rot_dim = self.freqs_cos.shape[-1] + end_index = start_index + rot_dim + assert rot_dim <= t.shape[-1], f'feature dimension {t.shape[-1]} is not of sufficient size to rotate in all the positions {rot_dim}' + t_left, t, t_right = t[..., :start_index], t[..., start_index:end_index], t[..., end_index:] + t = (t * self.freqs_cos) + (rotate_half(t) * self.freqs_sin) + return torch.cat((t_left, t, t_right), dim = -1) + + +class VisionRotaryEmbeddingFast(nn.Module): + def __init__( + self, + dim, + pt_seq_len=16, + ft_seq_len=None, + custom_freqs = None, + freqs_for = 'lang', + theta = 10000, + max_freq = 10, + num_freqs = 1, + num_cls_token = 0 + ): + super().__init__() + if custom_freqs: + freqs = custom_freqs + elif freqs_for == 'lang': + freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)) + elif freqs_for == 'pixel': + freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi + elif freqs_for == 'constant': + freqs = torch.ones(num_freqs).float() + else: + raise ValueError(f'unknown modality {freqs_for}') + + if ft_seq_len is None: ft_seq_len = pt_seq_len + t = torch.arange(ft_seq_len) / ft_seq_len * pt_seq_len + + freqs = torch.einsum('..., f -> ... f', t, freqs) + freqs = repeat(freqs, '... n -> ... (n r)', r = 2) + freqs = broadcat((freqs[:, None, :], freqs[None, :, :]), dim = -1) + + if num_cls_token > 0: + freqs_flat = freqs.view(-1, freqs.shape[-1]) # [N_img, D] + cos_img = freqs_flat.cos() + sin_img = freqs_flat.sin() + + # prepend in-context cls token + N_img, D = cos_img.shape + cos_pad = torch.ones(num_cls_token, D, dtype=cos_img.dtype, device=cos_img.device) + sin_pad = torch.zeros(num_cls_token, D, dtype=sin_img.dtype, device=sin_img.device) + + self.register_buffer( + "freqs_cos", torch.cat([cos_pad, cos_img], dim=0), persistent=False + ) # [N_cls+N_img, D] + self.register_buffer( + "freqs_sin", torch.cat([sin_pad, sin_img], dim=0), persistent=False + ) + else: + self.register_buffer( + "freqs_cos", freqs.cos().view(-1, freqs.shape[-1]), persistent=False + ) + self.register_buffer( + "freqs_sin", freqs.sin().view(-1, freqs.shape[-1]), persistent=False + ) + + def forward(self, t): return t * self.freqs_cos + rotate_half(t) * self.freqs_sin + + +class RMSNorm(nn.Module): + def __init__(self, hidden_size, eps=1e-6): + """ + LlamaRMSNorm is equivalent to T5LayerNorm + """ + super().__init__() + self.weight = nn.Parameter(torch.ones(hidden_size)) + self.variance_epsilon = eps + + def forward(self, hidden_states): + input_dtype = hidden_states.dtype + hidden_states = hidden_states.to(torch.float32) + variance = hidden_states.pow(2).mean(-1, keepdim=True) + hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) + return (self.weight * hidden_states).to(input_dtype) + + +def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0): + """ + grid_size: int of the grid height and width + return: + pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token) + """ + grid_h = np.arange(grid_size, dtype=np.float32) + grid_w = np.arange(grid_size, dtype=np.float32) + grid = np.meshgrid(grid_w, grid_h) # here w goes first + grid = np.stack(grid, axis=0) + + grid = grid.reshape([2, 1, grid_size, grid_size]) + pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid) + if cls_token and extra_tokens > 0: + pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0) + return pos_embed + + +def get_2d_sincos_pos_embed_from_grid(embed_dim, grid): + assert embed_dim % 2 == 0 + + # use half of dimensions to encode grid_h + emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2) + emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2) + + emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D) + return emb + + +def get_1d_sincos_pos_embed_from_grid(embed_dim, pos): + """ + embed_dim: output dimension for each position + pos: a list of positions to be encoded: size (M,) + out: (M, D) + """ + assert embed_dim % 2 == 0 + omega = np.arange(embed_dim // 2, dtype=np.float64) + omega /= embed_dim / 2. + omega = 1. / 10000**omega # (D/2,) + + pos = pos.reshape(-1) # (M,) + out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product + + emb_sin = np.sin(out) # (M, D/2) + emb_cos = np.cos(out) # (M, D/2) + + emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D) + return emb diff --git a/scripts/train/image_generation_matched.sbatch b/scripts/train/image_generation_matched.sbatch new file mode 100755 index 000000000..3f9d7b8b0 --- /dev/null +++ b/scripts/train/image_generation_matched.sbatch @@ -0,0 +1,190 @@ +#!/bin/bash +#SBATCH --partition=hoffman-lab +#SBATCH --account=hoffman-lab +#SBATCH --qos=long +#SBATCH --nodes=1 +#SBATCH --ntasks=1 +#SBATCH --gres=gpu:a40:4 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=256G +#SBATCH --time=2-00:00:00 +#SBATCH --requeue +#SBATCH --signal=B:USR1@300 + +set -euo pipefail + +ARCHITECTURE=${1:?usage: image_generation_matched.sbatch ARCHITECTURE smoke|full} +MODE=${2:?usage: image_generation_matched.sbatch ARCHITECTURE smoke|full} +case "$ARCHITECTURE" in + jit_b16|endpoint_latent) ;; + *) echo "unsupported architecture: $ARCHITECTURE" >&2; exit 2 ;; +esac +case "$MODE" in + smoke|full) ;; + *) echo "unsupported mode: $MODE" >&2; exit 2 ;; +esac + +SOURCE_ROOT=${SOURCE_ROOT:?submitter must export immutable SOURCE_ROOT} +EXPERIMENT_ROOT=${EXPERIMENT_ROOT:?submitter must export EXPERIMENT_ROOT} +EXPECTED_SHA=${EXPECTED_SHA:?submitter must export EXPECTED_SHA} +DATA_ROOT=${DATA_ROOT:-/coc/dataset/ImageNet/imagenet} +PY_ENV=/coc/flash7/paphiwetsa3/projects/EgoVerse7/.venv +SLURM_BIN=/opt/slurm/Ubuntu-20.04/current/bin +JIT_ROOT="$SOURCE_ROOT/external/JiT" + +test -x "$PY_ENV/bin/python" +test -x "$PY_ENV/bin/torchrun" +test -d "$DATA_ROOT/train" +test -d "$DATA_ROOT/val" +test "$(git -C "$SOURCE_ROOT" rev-parse HEAD)" = "$EXPECTED_SHA" +test -z "$(git -C "$SOURCE_ROOT" status --porcelain=v1 --untracked-files=all)" +test "${SLURM_JOB_PARTITION:?}" = hoffman-lab +test "${SLURM_JOB_ACCOUNT:?}" = hoffman-lab +test "${SLURM_GPUS_ON_NODE:?}" = 4 + +mkdir -p "$EXPERIMENT_ROOT/slurm" "$EXPERIMENT_ROOT/provenance" +PROVENANCE_DIR="$EXPERIMENT_ROOT/provenance/$MODE-$ARCHITECTURE-job${SLURM_JOB_ID}" +mkdir -p "$PROVENANCE_DIR" +cp "$0" "$PROVENANCE_DIR/launcher.sbatch" +cp "$SOURCE_ROOT/scripts/train/verify_image_generation_smoke.py" "$PROVENANCE_DIR/" +git -C "$SOURCE_ROOT" log -1 --format=fuller > "$PROVENANCE_DIR/git_commit.txt" +git -C "$SOURCE_ROOT" status --porcelain=v1 --untracked-files=all > "$PROVENANCE_DIR/git_status.txt" +"$SLURM_BIN/scontrol" show job -dd "$SLURM_JOB_ID" > "$PROVENANCE_DIR/slurm_job.txt" +nvidia-smi -L > "$PROVENANCE_DIR/gpus.txt" +cp "$EXPERIMENT_ROOT/provenance/dataset_inventory.txt" "$PROVENANCE_DIR/" +cp "$EXPERIMENT_ROOT/provenance/dataset_inventory.sha256" "$PROVENANCE_DIR/" +sha256sum "$PROVENANCE_DIR"/* > "$PROVENANCE_DIR/artifacts.sha256" + +source "$PY_ENV/bin/activate" +export PYTHONPATH="$SOURCE_ROOT:$JIT_ROOT" +export PYTHONUNBUFFERED=1 +export TORCHINDUCTOR_CACHE_DIR="$EXPERIMENT_ROOT/torchinductor/$ARCHITECTURE" +mkdir -p "$TORCHINDUCTOR_CACHE_DIR" +cd "$JIT_ROOT" + +PARAM_ARGS=(--expected-min-params 120000000 --expected-max-params 140000000) +COMMON_ARGS=( + --architecture "$ARCHITECTURE" + --data-path "$DATA_ROOT" + --image-size 256 + --num-classes 1000 + --seed 42 + --num-workers 8 + "${PARAM_ARGS[@]}" +) + +if [[ "$ARCHITECTURE" == jit_b16 ]]; then + FULL_BATCH_SIZE=64 + FULL_GRAD_ACCUM=4 +else + FULL_BATCH_SIZE=32 + FULL_GRAD_ACCUM=8 +fi + +if [[ "$MODE" == smoke ]]; then + SMOKE_ROOT="$EXPERIMENT_ROOT/smokes/$ARCHITECTURE/job_${SLURM_JOB_ID}" + OVERFIT_DIR="$SMOKE_ROOT/overfit" + REAL_DIR="$SMOKE_ROOT/real" + mkdir -p "$OVERFIT_DIR" "$REAL_DIR" + + "$PY_ENV/bin/torchrun" --standalone --nproc_per_node=4 train_matched.py \ + "${COMMON_ARGS[@]}" \ + --output-dir "$OVERFIT_DIR" \ + --epochs 1 \ + --batch-size 1 \ + --grad-accum 1 \ + --base-lr 1e-4 \ + --warmup-epochs 0 \ + --max-optimizer-steps 40 \ + --save-every-steps 40 \ + --val-every-steps 0 \ + --log-every-steps 5 \ + --overfit-one-batch \ + --overfit-force-steps 2 \ + --wandb-mode disabled \ + 2>&1 | tee "$SMOKE_ROOT/overfit.log" + + "$PY_ENV/bin/torchrun" --standalone --nproc_per_node=4 train_matched.py \ + "${COMMON_ARGS[@]}" \ + --output-dir "$REAL_DIR" \ + --epochs 1 \ + --batch-size "$FULL_BATCH_SIZE" \ + --grad-accum "$FULL_GRAD_ACCUM" \ + --base-lr 5e-5 \ + --warmup-epochs 5 \ + --max-optimizer-steps 2 \ + --save-every-steps 2 \ + --val-every-steps 1 \ + --val-batches 1 \ + --sample-batch 1 \ + --sample-steps 16 \ + --log-every-steps 1 \ + --smoke-require-validation \ + --wandb-project image-generation-latent-denoise \ + --wandb-name "smoke-${ARCHITECTURE}-${EXPECTED_SHA:0:8}" \ + --wandb-mode offline \ + 2>&1 | tee "$SMOKE_ROOT/real.log" + + "$PY_ENV/bin/python" "$SOURCE_ROOT/scripts/train/verify_image_generation_smoke.py" \ + --overfit-log "$SMOKE_ROOT/overfit.log" \ + --smoke-log "$SMOKE_ROOT/real.log" \ + --smoke-dir "$REAL_DIR" + echo "MATCHED_SMOKE_COMPLETE architecture=$ARCHITECTURE job=$SLURM_JOB_ID" + exit 0 +fi + +RUN_DIR="$EXPERIMENT_ROOT/runs/$ARCHITECTURE" +STOP_FILE="$RUN_DIR/stop-requested" +mkdir -p "$RUN_DIR" +rm -f "$STOP_FILE" + +FULL_ARGS=( + "${COMMON_ARGS[@]}" + --output-dir "$RUN_DIR" + --epochs 600 + --batch-size "$FULL_BATCH_SIZE" + --grad-accum "$FULL_GRAD_ACCUM" + --base-lr 5e-5 + --warmup-epochs 5 + --weight-decay 0 + --ema-decay 0.9999 + --save-every-steps 1000 + --val-every-steps 5000 + --val-batches 8 + --sample-batch 4 + --sample-steps 16 + --cfg-scale 1.0 + --log-every-steps 10 + --resume + --wandb-project image-generation-latent-denoise + --wandb-name "imagenet256-${ARCHITECTURE}-${EXPECTED_SHA:0:8}" + --wandb-mode online + --stop-file "$STOP_FILE" +) + +REQUEUE_REQUESTED=0 +handle_usr1() { + REQUEUE_REQUESTED=1 + touch "$STOP_FILE" + echo "SLURM_REQUEUE_CHECKPOINT_REQUESTED job=$SLURM_JOB_ID" +} +trap handle_usr1 USR1 + +set +e +"$PY_ENV/bin/torchrun" --standalone --nproc_per_node=4 train_matched.py \ + "${FULL_ARGS[@]}" & +TRAIN_PID=$! +while kill -0 "$TRAIN_PID" 2>/dev/null; do + wait "$TRAIN_PID" + TRAIN_STATUS=$? +done +wait "$TRAIN_PID" 2>/dev/null +TRAIN_STATUS=${TRAIN_STATUS:-$?} +set -e + +if [[ "$REQUEUE_REQUESTED" == 1 && "$TRAIN_STATUS" == 0 ]]; then + echo "REQUEUEING_FROM_CHECKPOINT job=$SLURM_JOB_ID" + "$SLURM_BIN/scontrol" requeue "$SLURM_JOB_ID" + exit 0 +fi +exit "$TRAIN_STATUS" diff --git a/scripts/train/inventory_imagenet.py b/scripts/train/inventory_imagenet.py new file mode 100755 index 000000000..8f15fac89 --- /dev/null +++ b/scripts/train/inventory_imagenet.py @@ -0,0 +1,50 @@ +#!/usr/bin/env python3 +"""Create a deterministic metadata inventory for an ImageFolder dataset.""" + +import argparse +import hashlib +import json +import os +from pathlib import Path + + +def split_inventory(path: Path): + digest = hashlib.sha256() + count = 0 + total_bytes = 0 + for directory, dirnames, filenames in os.walk(path, followlinks=True): + dirnames.sort() + filenames.sort() + base = Path(directory) + for filename in filenames: + item = base / filename + stat = item.stat() + relative = item.relative_to(path) + record = f"{relative}\t{stat.st_size}\t{stat.st_mtime_ns}\n".encode() + digest.update(record) + count += 1 + total_bytes += stat.st_size + return {"files": count, "bytes": total_bytes, "metadata_sha256": digest.hexdigest()} + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--data-root", type=Path, required=True) + parser.add_argument("--output", type=Path, required=True) + args = parser.parse_args() + result = { + "data_root": str(args.data_root), + "train_resolved": str((args.data_root / "train").resolve()), + "val_resolved": str((args.data_root / "val").resolve()), + "train": split_inventory(args.data_root / "train"), + "val": split_inventory(args.data_root / "val"), + } + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps(result, indent=2, sort_keys=True) + "\n") + digest = hashlib.sha256(args.output.read_bytes()).hexdigest() + args.output.with_suffix(".sha256").write_text(f"{digest} {args.output.name}\n") + print(json.dumps(result, sort_keys=True)) + + +if __name__ == "__main__": + main() diff --git a/scripts/train/submit_image_generation_matched.sh b/scripts/train/submit_image_generation_matched.sh new file mode 100755 index 000000000..2edae7bdc --- /dev/null +++ b/scripts/train/submit_image_generation_matched.sh @@ -0,0 +1,72 @@ +#!/bin/bash +set -euo pipefail + +WORKTREE=${WORKTREE:-/coc/flash7/paphiwetsa3/worktrees/image-latent-denoise-jit-20260828} +EXPERIMENT_ROOT=${EXPERIMENT_ROOT:-/coc/flash7/paphiwetsa3/experiments/imagenet256_jit_vs_endpoint_latent_20260828} +DATA_ROOT=${DATA_ROOT:-/coc/dataset/ImageNet/imagenet} +SLURM_BIN=/opt/slurm/Ubuntu-20.04/current/bin +PY_ENV=/coc/flash7/paphiwetsa3/projects/EgoVerse7/.venv + +test -z "$(git -C "$WORKTREE" status --porcelain=v1 --untracked-files=all)" +SHA=$(git -C "$WORKTREE" rev-parse HEAD) +SOURCE_ROOT="$EXPERIMENT_ROOT/source_${SHA:0:8}" +mkdir -p "$EXPERIMENT_ROOT/slurm" "$EXPERIMENT_ROOT/provenance" +if [[ ! -e "$SOURCE_ROOT" ]]; then + git -C "$WORKTREE" worktree add --detach "$SOURCE_ROOT" "$SHA" + git -C "$WORKTREE" worktree lock "$SOURCE_ROOT" \ + --reason "immutable ImageNet matched JiT/endpoint-latent source $SHA" +fi +test "$(git -C "$SOURCE_ROOT" rev-parse HEAD)" = "$SHA" +test -z "$(git -C "$SOURCE_ROOT" status --porcelain=v1 --untracked-files=all)" + +INVENTORY="$EXPERIMENT_ROOT/provenance/dataset_inventory.txt" +if [[ ! -s "$INVENTORY" || ! -s "${INVENTORY%.txt}.sha256" ]]; then + "$PY_ENV/bin/python" "$SOURCE_ROOT/scripts/train/inventory_imagenet.py" \ + --data-root "$DATA_ROOT" \ + --output "$INVENTORY" +fi + +EXPORTS="ALL,SOURCE_ROOT=$SOURCE_ROOT,EXPERIMENT_ROOT=$EXPERIMENT_ROOT,EXPECTED_SHA=$SHA,DATA_ROOT=$DATA_ROOT" +LAUNCHER="$SOURCE_ROOT/scripts/train/image_generation_matched.sbatch" + +JIT_SMOKE=$( + "$SLURM_BIN/sbatch" --parsable \ + --job-name=img-jit-smoke \ + --qos=short --time=02:00:00 --no-requeue \ + --output="$EXPERIMENT_ROOT/slurm/%x-%j.out" \ + --error="$EXPERIMENT_ROOT/slurm/%x-%j.err" \ + --export="$EXPORTS" \ + "$LAUNCHER" jit_b16 smoke +) +LATENT_SMOKE=$( + "$SLURM_BIN/sbatch" --parsable \ + --job-name=img-lat-smoke \ + --qos=short --time=02:00:00 --no-requeue \ + --output="$EXPERIMENT_ROOT/slurm/%x-%j.out" \ + --error="$EXPERIMENT_ROOT/slurm/%x-%j.err" \ + --export="$EXPORTS" \ + "$LAUNCHER" endpoint_latent smoke +) + +DEPENDENCY="afterok:${JIT_SMOKE}:${LATENT_SMOKE}" +JIT_FULL=$( + "$SLURM_BIN/sbatch" --parsable \ + --job-name=img-jit-full \ + --dependency="$DEPENDENCY" \ + --output="$EXPERIMENT_ROOT/slurm/%x-%j.out" \ + --error="$EXPERIMENT_ROOT/slurm/%x-%j.err" \ + --export="$EXPORTS" \ + "$LAUNCHER" jit_b16 full +) +LATENT_FULL=$( + "$SLURM_BIN/sbatch" --parsable \ + --job-name=img-lat-full \ + --dependency="$DEPENDENCY" \ + --output="$EXPERIMENT_ROOT/slurm/%x-%j.out" \ + --error="$EXPERIMENT_ROOT/slurm/%x-%j.err" \ + --export="$EXPORTS" \ + "$LAUNCHER" endpoint_latent full +) + +printf 'SHA=%s\nSOURCE_ROOT=%s\nJIT_SMOKE=%s\nLATENT_SMOKE=%s\nJIT_FULL=%s\nLATENT_FULL=%s\n' \ + "$SHA" "$SOURCE_ROOT" "$JIT_SMOKE" "$LATENT_SMOKE" "$JIT_FULL" "$LATENT_FULL" diff --git a/scripts/train/verify_image_generation_smoke.py b/scripts/train/verify_image_generation_smoke.py new file mode 100755 index 000000000..c29553a1b --- /dev/null +++ b/scripts/train/verify_image_generation_smoke.py @@ -0,0 +1,57 @@ +#!/usr/bin/env python3 +"""Fail closed unless overfit, optimization, validation, and reload artifacts exist.""" + +import argparse +import json +import math +from pathlib import Path + + +def require_text(path: Path, marker: str) -> None: + text = path.read_text(encoding="utf-8", errors="replace") + if marker not in text: + raise RuntimeError(f"Missing {marker!r} in {path}") + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--overfit-log", type=Path, required=True) + parser.add_argument("--smoke-log", type=Path, required=True) + parser.add_argument("--smoke-dir", type=Path, required=True) + args = parser.parse_args() + + require_text(args.overfit_log, "OVERFIT_GATE_PASSED") + require_text(args.smoke_log, "TRAIN_METRICS") + require_text(args.smoke_log, "VALIDATION_METRICS") + require_text(args.smoke_log, "SMOKE_GATE_PASSED") + + required = [ + args.smoke_dir / "checkpoint-last.pth", + args.smoke_dir / "resolved_config.json", + args.smoke_dir / "metrics.jsonl", + ] + for path in required: + if not path.is_file() or path.stat().st_size == 0: + raise RuntimeError(f"Missing or empty smoke artifact: {path}") + if not list(args.smoke_dir.glob("samples-step*.png")): + raise RuntimeError("Smoke did not produce a fully sampled validation grid") + + validation_rows = [] + for line in (args.smoke_dir / "metrics.jsonl").read_text().splitlines(): + row = json.loads(line) + if row.get("split") == "validation": + validation_rows.append(row) + if not validation_rows: + raise RuntimeError("No post-training validation rows were recorded") + for row in validation_rows: + for key, value in row.items(): + if isinstance(value, (int, float)) and not math.isfinite(value): + raise RuntimeError(f"Non-finite validation metric {key}={value}") + print( + "IMAGE_GENERATION_SMOKE_VERIFIED " + + json.dumps(validation_rows[-1], sort_keys=True) + ) + + +if __name__ == "__main__": + main()