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
+
+[](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()