diff --git a/README.md b/README.md
index 4948c3933d0d1fa55310b7f318835018bdb424dc..3a5b58640e0a950e4a9be31ba4a39fe4836c9b45 100644
--- a/README.md
+++ b/README.md
@@ -1,27 +1,7 @@
----
-title: CoTyle
-emoji: 🎨
-colorFrom: gray
-colorTo: purple
-sdk: gradio
-sdk_version: 5.49.1
-app_file: app.py
-python_version: 3.10
-# 移除 license 字段,因为它不是官方支持的字段
-gpu: true
-suggested_hardware: a100-large
-models:
- - Kwai-Kolors/CoTyle
-tags:
- - image-generation
- - code-to-style
- - gradio
----
-
# A Style is Worth One Code: Unlocking Code-to-Style Image Generation with Discrete Style Space
-
+
@@ -78,6 +58,8 @@ git clone https://github.com/Kwai-Kolors/CoTyle
cd CoTyle
conda create -n cotyle python=3.10
conda activate cotyle
+pip install torch==2.6.0 torchvision==0.21.0
+pip install -e git+https://github.com/Lakonik/piFlow.git@b1ef16e5e305251bccdfeac2a0e3d0ef339b974a#egg=lakonlab
pip install -r requirements.txt
```
diff --git a/lakonlab/__init__.py b/lakonlab/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..a32ec7cce43929108d9d1d6bf8a030aadce2b6da
--- /dev/null
+++ b/lakonlab/__init__.py
@@ -0,0 +1,20 @@
+import warnings
+
+# suppress warnings from MMCV about optional dependencies
+warnings.filterwarnings(
+ 'ignore',
+ category=UserWarning,
+ message=r'^Fail to import ``MultiScaleDeformableAttention`` from ``mmcv\.ops\.multi_scale_deform_attn``.*',
+ module=r'^mmcv\.cnn\.bricks\.transformer$',
+)
+
+# import all modules for registration
+from .apis import *
+from .datasets import *
+from .models import *
+from .ops import *
+from .runner import *
+from .evaluation import *
+from .utils import *
+
+from .version import __version__
diff --git a/lakonlab/__pycache__/__init__.cpython-310.pyc b/lakonlab/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..bc6bfedba3e8f8f710187a42a67b584c47b99cea
Binary files /dev/null and b/lakonlab/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/__pycache__/__init__.cpython-313.pyc b/lakonlab/__pycache__/__init__.cpython-313.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..9d64e62eff6043439194269c97dbbd3e162c0ff7
Binary files /dev/null and b/lakonlab/__pycache__/__init__.cpython-313.pyc differ
diff --git a/lakonlab/__pycache__/version.cpython-310.pyc b/lakonlab/__pycache__/version.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..847738f853b187f86bf8d1afb72a6f371e9b9a65
Binary files /dev/null and b/lakonlab/__pycache__/version.cpython-310.pyc differ
diff --git a/lakonlab/apis/__init__.py b/lakonlab/apis/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..71730d8668a3ea1d1c60f05e99383d3faee05598
--- /dev/null
+++ b/lakonlab/apis/__init__.py
@@ -0,0 +1,4 @@
+from .train import train_model
+from .inference import init_model
+
+__all__ = ['init_model', 'train_model']
diff --git a/lakonlab/apis/__pycache__/__init__.cpython-310.pyc b/lakonlab/apis/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..7d95e1fca06cc6c258e37b231f1132bad42f1c47
Binary files /dev/null and b/lakonlab/apis/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/apis/__pycache__/__init__.cpython-313.pyc b/lakonlab/apis/__pycache__/__init__.cpython-313.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..eed7d535a5aa87b4cd6ea3766695ae30d99548c5
Binary files /dev/null and b/lakonlab/apis/__pycache__/__init__.cpython-313.pyc differ
diff --git a/lakonlab/apis/__pycache__/inference.cpython-310.pyc b/lakonlab/apis/__pycache__/inference.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..6d66cd027c27826f409a7c765691d6fe79a0e356
Binary files /dev/null and b/lakonlab/apis/__pycache__/inference.cpython-310.pyc differ
diff --git a/lakonlab/apis/__pycache__/train.cpython-310.pyc b/lakonlab/apis/__pycache__/train.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..fa90cbeb52ff93011564cbcd7347ee9e77c8e65c
Binary files /dev/null and b/lakonlab/apis/__pycache__/train.cpython-310.pyc differ
diff --git a/lakonlab/apis/__pycache__/train.cpython-313.pyc b/lakonlab/apis/__pycache__/train.cpython-313.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..9476d2255ecc9743e96329f27a8688cfbd51b482
Binary files /dev/null and b/lakonlab/apis/__pycache__/train.cpython-313.pyc differ
diff --git a/lakonlab/apis/inference.py b/lakonlab/apis/inference.py
new file mode 100644
index 0000000000000000000000000000000000000000..666efff65eca57649c411286e9f7fdb30a386d58
--- /dev/null
+++ b/lakonlab/apis/inference.py
@@ -0,0 +1,56 @@
+import torch
+import mmcv
+from mmcv.runner import load_checkpoint
+from mmgen.models import build_model
+from lakonlab.runner.hooks.ema_hook import get_ori_key
+
+
+def init_model(
+ config, checkpoint=None, device='cuda:0', cfg_options=None,
+ ema_only=True, use_fp16=False, use_bf16=False):
+ if isinstance(config, str):
+ config = mmcv.Config.fromfile(config)
+ elif not isinstance(config, mmcv.Config):
+ raise TypeError('config must be a filename or Config object, '
+ f'but got {type(config)}')
+ if cfg_options is not None:
+ config.merge_from_dict(cfg_options)
+
+ model = build_model(
+ config.model, train_cfg=config.train_cfg, test_cfg=config.test_cfg)
+
+ if ema_only:
+ module_keys = []
+ for hook in config.get('custom_hooks', []):
+ if hook['type'] in ('ExponentialMovingAverageHookMod', 'ExponentialMovingAverageHook'):
+ if isinstance(hook['module_keys'], str):
+ module_keys.append(hook['module_keys'])
+ else:
+ module_keys.extend(hook['module_keys'])
+ for key in module_keys:
+ ori_key = get_ori_key(key)
+ del model._modules[ori_key]
+
+ if checkpoint is not None:
+ load_checkpoint(model, checkpoint, map_location='cpu')
+
+ model._cfg = config # save the config in the model for convenience
+
+ for module in model.modules():
+ if hasattr(module, 'bake_lora_weights'):
+ module.bake_lora_weights()
+
+ if use_fp16 or use_bf16:
+ for m in model.modules():
+ if hasattr(m, 'autocast_dtype'):
+ setattr(m, 'autocast_dtype', None)
+ if use_fp16:
+ assert not use_bf16
+ model.to(dtype=torch.float16)
+ elif use_bf16:
+ model.to(dtype=torch.bfloat16)
+
+ model.to(device)
+ model.eval()
+
+ return model
diff --git a/lakonlab/apis/train.py b/lakonlab/apis/train.py
new file mode 100644
index 0000000000000000000000000000000000000000..5e67b93595c66092652fe20fa0bc8fb64bcd71be
--- /dev/null
+++ b/lakonlab/apis/train.py
@@ -0,0 +1,166 @@
+# Modified from https://github.com/open-mmlab/mmgeneration
+
+import warnings
+import re
+from copy import deepcopy
+
+from mmcv.parallel import MMDataParallel
+from mmcv.runner import HOOKS, IterBasedRunner, OptimizerHook, build_runner
+from mmcv.utils import build_from_cfg
+
+from mmgen.datasets import build_dataset
+from mmgen.utils import get_root_logger
+
+from lakonlab.parallel import apply_module_wrapper
+from lakonlab.runner.optimizer import build_optimizers
+from lakonlab.runner.checkpoint import exists_ckpt
+from lakonlab.datasets import build_dataloader
+
+
+def train_model(model,
+ dataset,
+ cfg,
+ distributed=False,
+ validate=False,
+ timestamp=None,
+ meta=None):
+ logger = get_root_logger(cfg.log_level)
+
+ # prepare data loaders
+ dataset = dataset if isinstance(dataset, (list, tuple)) else [dataset]
+
+ # default loader config
+ loader_cfg = dict(
+ # cfg.gpus will be ignored if distributed
+ num_gpus=len(cfg.gpu_ids),
+ seed=cfg.seed)
+
+ # The overall dataloader settings
+ loader_cfg.update({
+ k: v
+ for k, v in cfg.data.items()
+ if k not in [
+ 'train', 'train_dataloader', 'val_dataloader', 'test_dataloader'
+ ] and not re.fullmatch(r'(val|test)\d*', k)
+ })
+
+ # The specific datalaoder settings
+ train_loader_cfg = {**loader_cfg, **cfg.data.get('train_dataloader', {})}
+
+ data_loaders = [build_dataloader(ds, **train_loader_cfg) for ds in dataset]
+
+ if cfg.get('apex_amp', None):
+ raise NotImplementedError('Apex AMP is no longer supported.')
+
+ # put model on gpus
+ if distributed:
+ module_wrapper = cfg.get('module_wrapper', None)
+ model = apply_module_wrapper(model, module_wrapper, cfg)
+ else:
+ model = MMDataParallel(model, device_ids=cfg.gpu_ids)
+
+ # build optimizer
+ if cfg.optimizer:
+ optimizer = build_optimizers(model, cfg.optimizer)
+ # In GANs, we allow building optimizer in GAN model.
+ else:
+ optimizer = None
+
+ # allow users to define the runner
+ if cfg.get('runner', None):
+ runner = build_runner(
+ cfg.runner,
+ dict(
+ model=model,
+ optimizer=optimizer,
+ work_dir=cfg.work_dir,
+ logger=logger,
+ use_apex_amp=False,
+ meta=meta))
+ else:
+ runner = IterBasedRunner(
+ model,
+ optimizer=optimizer,
+ work_dir=cfg.work_dir,
+ logger=logger,
+ meta=meta)
+ # set if use dynamic ddp in training
+ # is_dynamic_ddp=cfg.get('is_dynamic_ddp', False))
+ # an ugly walkaround to make the .log and .log.json filenames the same
+ runner.timestamp = timestamp
+
+ # fp16 setting
+ fp16_cfg = cfg.get('fp16', None)
+
+ # In GANs, we can directly optimize parameter in `train_step` function.
+ if cfg.get('optimizer_cfg', None) is None:
+ optimizer_config = None
+ elif fp16_cfg is not None:
+ raise NotImplementedError('Fp16 has not been supported.')
+ # optimizer_config = Fp16OptimizerHook(
+ # **cfg.optimizer_config, **fp16_cfg, distributed=distributed)
+ # default to use OptimizerHook
+ elif distributed and 'type' not in cfg.optimizer_config:
+ optimizer_config = OptimizerHook(**cfg.optimizer_config)
+ else:
+ optimizer_config = cfg.optimizer_config
+
+ # # update `out_dir` in ckpt hook
+ # if cfg.checkpoint_config is not None:
+ # cfg.checkpoint_config['out_dir'] = os.path.join(
+ # cfg.work_dir, cfg.checkpoint_config.get('out_dir', 'ckpt'))
+
+ # register hooks
+ runner.register_training_hooks(cfg.lr_config, optimizer_config,
+ cfg.checkpoint_config, cfg.log_config,
+ cfg.get('momentum_config', None))
+
+ # # DistSamplerSeedHook should be used with EpochBasedRunner
+ # if distributed:
+ # runner.register_hook(DistSamplerSeedHook())
+
+ # In general, we do NOT adopt standard evaluation hook in GAN training.
+ # Thus, if you want a eval hook, you need further define the key of
+ # 'evaluation' in the config.
+ # register eval hooks
+ if validate and cfg.get('evaluation', None) is not None:
+ assert isinstance(cfg.evaluation, list)
+ for eval_cfg_ in cfg.evaluation:
+ val_dataset = build_dataset(cfg.data[eval_cfg_.data])
+ val_loader_cfg = {
+ **loader_cfg, 'shuffle': False,
+ **cfg.data.get('val_dataloader', {})
+ }
+ val_dataloader = build_dataloader(val_dataset, **val_loader_cfg)
+ eval_cfg = deepcopy(eval_cfg_)
+ priority = eval_cfg.pop('priority', 'LOW')
+ eval_cfg.update(dict(dist=distributed, dataloader=val_dataloader))
+ eval_hook = build_from_cfg(eval_cfg, HOOKS)
+ runner.register_hook(eval_hook, priority=priority)
+
+ # user-defined hooks
+ if cfg.get('custom_hooks', None):
+ custom_hooks = cfg.custom_hooks
+ assert isinstance(custom_hooks, list), \
+ f'custom_hooks expect list type, but got {type(custom_hooks)}'
+ for hook_cfg in cfg.custom_hooks:
+ assert isinstance(hook_cfg, dict), \
+ 'Each item in custom_hooks expects dict type, but got ' \
+ f'{type(hook_cfg)}'
+ hook_cfg = hook_cfg.copy()
+ priority = hook_cfg.pop('priority', 'NORMAL')
+ hook = build_from_cfg(hook_cfg, HOOKS)
+ runner.register_hook(hook, priority=priority)
+
+ ckpt_kwargs = dict()
+ if distributed and module_wrapper.lower() in ['fsdp', 'fsdp2']:
+ ckpt_kwargs.update(map_location='cpu')
+ if exists_ckpt(cfg.resume_from):
+ runner.resume(cfg.resume_from, **ckpt_kwargs)
+ for data_loader in data_loaders:
+ data_loader.sampler.set_epoch(runner.epoch)
+ data_loader.sampler.set_iter(runner.iter)
+ elif exists_ckpt(cfg.load_from):
+ runner.load_checkpoint(cfg.load_from, **ckpt_kwargs)
+
+ runner.run(data_loaders, cfg.workflow, cfg.total_iters)
diff --git a/lakonlab/datasets/__init__.py b/lakonlab/datasets/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..eea27615a10b25d2ff093a9d306c8516a19e6125
--- /dev/null
+++ b/lakonlab/datasets/__init__.py
@@ -0,0 +1,8 @@
+from .builder import build_dataloader
+from .imagenet import ImageNet
+from .checkerboard import CheckerboardData
+from .image_prompts import ImagePrompt
+
+__all__ = [
+ 'build_dataloader', 'ImageNet', 'CheckerboardData', 'ImagePrompt'
+]
diff --git a/lakonlab/datasets/__pycache__/__init__.cpython-310.pyc b/lakonlab/datasets/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..dc9d4c4d89e59d2921cce30816e8288674dbc390
Binary files /dev/null and b/lakonlab/datasets/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/datasets/__pycache__/builder.cpython-310.pyc b/lakonlab/datasets/__pycache__/builder.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..3c872f7d0e094b7718b180cf2981f96947ee01ba
Binary files /dev/null and b/lakonlab/datasets/__pycache__/builder.cpython-310.pyc differ
diff --git a/lakonlab/datasets/__pycache__/checkerboard.cpython-310.pyc b/lakonlab/datasets/__pycache__/checkerboard.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..2052c08eaeedd9d05a7e576c8ed688a4aef9f941
Binary files /dev/null and b/lakonlab/datasets/__pycache__/checkerboard.cpython-310.pyc differ
diff --git a/lakonlab/datasets/__pycache__/image_prompts.cpython-310.pyc b/lakonlab/datasets/__pycache__/image_prompts.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..10aa81eff558b4adeb9106f73cf287e3878052ee
Binary files /dev/null and b/lakonlab/datasets/__pycache__/image_prompts.cpython-310.pyc differ
diff --git a/lakonlab/datasets/__pycache__/imagenet.cpython-310.pyc b/lakonlab/datasets/__pycache__/imagenet.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..86fa66014a42bf88886a03ea0b0cdb8916185c0e
Binary files /dev/null and b/lakonlab/datasets/__pycache__/imagenet.cpython-310.pyc differ
diff --git a/lakonlab/datasets/builder.py b/lakonlab/datasets/builder.py
new file mode 100644
index 0000000000000000000000000000000000000000..68bf2a37fcbb42635c841db5971410143207a731
--- /dev/null
+++ b/lakonlab/datasets/builder.py
@@ -0,0 +1,61 @@
+import warnings
+from functools import partial
+
+from mmcv.parallel import collate
+from mmcv.runner import get_dist_info
+from mmcv.utils import TORCH_VERSION, digit_version
+from torch.utils.data import DataLoader
+
+from mmgen.datasets.builder import worker_init_fn
+from .samplers import DistributedSampler
+
+
+def build_dataloader(dataset,
+ samples_per_gpu,
+ workers_per_gpu,
+ num_gpus=1,
+ dist=True,
+ shuffle=True,
+ seed=None,
+ persistent_workers=False,
+ sampler=None,
+ **kwargs):
+ rank, world_size = get_dist_info()
+ if dist:
+ assert sampler is None, 'sampler is not supported in distributed mode'
+ sampler = DistributedSampler(
+ dataset,
+ world_size,
+ rank,
+ shuffle=shuffle,
+ samples_per_gpu=samples_per_gpu,
+ seed=seed)
+ shuffle = False
+ batch_size = samples_per_gpu
+ num_workers = workers_per_gpu
+ else:
+ batch_size = num_gpus * samples_per_gpu
+ num_workers = num_gpus * workers_per_gpu
+
+ init_fn = partial(
+ worker_init_fn, num_workers=num_workers, rank=rank,
+ seed=seed) if seed is not None else None
+
+ if (digit_version(TORCH_VERSION) >= digit_version('1.7.0')
+ and TORCH_VERSION != 'parrots'):
+ kwargs['persistent_workers'] = persistent_workers
+ elif persistent_workers is True:
+ warnings.warn('persistent_workers is invalid because your pytorch '
+ 'version is lower than 1.7.0')
+
+ data_loader = DataLoader(
+ dataset,
+ batch_size=batch_size,
+ sampler=sampler,
+ num_workers=num_workers,
+ collate_fn=partial(collate, samples_per_gpu=samples_per_gpu),
+ shuffle=shuffle,
+ worker_init_fn=init_fn,
+ **kwargs)
+
+ return data_loader
diff --git a/lakonlab/datasets/checkerboard.py b/lakonlab/datasets/checkerboard.py
new file mode 100644
index 0000000000000000000000000000000000000000..425005c8dd7c8d5264ad535ade8f50edd8fd13b6
--- /dev/null
+++ b/lakonlab/datasets/checkerboard.py
@@ -0,0 +1,59 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import torch
+
+from torch.utils.data import Dataset
+from mmgen.datasets.builder import DATASETS
+
+
+@DATASETS.register_module()
+class CheckerboardData(Dataset):
+ def __init__(
+ self,
+ n_rc=4,
+ n_samples=1e8,
+ thickness=1.0,
+ scale=1,
+ shift=[0.0, 0.0],
+ rotation=0.0,
+ test_mode=False):
+ super().__init__()
+ self.n_rc = n_rc
+ self.n_samples = int(n_samples)
+ self.thickness = thickness
+ self.scale = scale
+ self.shift = torch.tensor(shift, dtype=torch.float32)
+ self.rotation = rotation
+ white_squares = [(i, j) for i in range(n_rc) for j in range(n_rc) if (i + j) % 2 == 0]
+ self.white_squares = torch.tensor(white_squares, dtype=torch.float32)
+ self.n_squares = len(white_squares)
+ self.samples = self.draw_samples(self.n_samples)
+
+ def draw_samples(self, n_samples):
+ chosen_indices = torch.randint(0, self.n_squares, size=(n_samples, ))
+ chosen_squares = self.white_squares[chosen_indices]
+ square_samples = torch.rand(n_samples, 2, dtype=torch.float32)
+ if self.thickness < 1:
+ square_samples = square_samples - 0.5
+ square_samples_r = square_samples.square().sum(dim=-1, keepdims=True)
+ square_samples_angle = torch.atan2(square_samples[:, 1], square_samples[:, 0]).unsqueeze(-1)
+ max_r = torch.minimum(
+ 0.5 / square_samples_angle.cos().abs().clamp(min=1e-6),
+ 0.5 / square_samples_angle.sin().abs().clamp(min=1e-6)).square()
+ square_samples_r_scaled = max_r - (max_r - square_samples_r) * self.thickness ** 0.5
+ square_samples *= (square_samples_r_scaled / square_samples_r).sqrt()
+ square_samples = square_samples + 0.5
+ samples = (chosen_squares + square_samples) * (2 / self.n_rc) - 1
+ if self.rotation != 0.0:
+ angle = torch.tensor(self.rotation, dtype=torch.float32) * torch.pi / 180
+ rotation_matrix = torch.tensor([[torch.cos(angle), -torch.sin(angle)],
+ [torch.sin(angle), torch.cos(angle)]])
+ samples = samples @ rotation_matrix
+ return samples * self.scale + self.shift
+
+ def __len__(self):
+ return self.n_samples
+
+ def __getitem__(self, idx):
+ data = dict(x=self.samples[idx])
+ return data
diff --git a/lakonlab/datasets/image_prompts.py b/lakonlab/datasets/image_prompts.py
new file mode 100644
index 0000000000000000000000000000000000000000..0d5b0173217d7f2b653977ffea015c0d5607112a
--- /dev/null
+++ b/lakonlab/datasets/image_prompts.py
@@ -0,0 +1,432 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import logging
+import os
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+import zstandard as zstd
+import pickle
+import gzip
+import orjson
+import mmcv
+import torch.storage
+torch.storage.UntypedStorage.dtype = torch.uint8 # hot patch for torch 2.6 deserialization
+
+from io import BytesIO
+from typing import Optional, Tuple, Union
+from torch.utils.data import Dataset
+from datasets import load_dataset, DatasetDict, Dataset as HFDataset
+from mmcv.fileio import FileClient
+from mmcv.parallel import DataContainer as DC
+from mmgen.utils import get_root_logger
+from mmgen.datasets.builder import DATASETS
+from lakonlab.utils.io_utils import load_image
+
+
+@DATASETS.register_module()
+class ImagePrompt(Dataset):
+ """Initialize an image/prompt dataset that reads either cached pickled records
+ (zstd-compressed) or a HuggingFace prompt dataset (optionally paired with images).
+
+ Args:
+ data_root (str): Root path for IO, resolved via `mmcv.FileClient`.
+ cache_dir (Optional[str]): Subdirectory of `data_root` containing `.zst`
+ cache shards. Enables cache mode when provided and exists. Caches must
+ contain pickled dicts with keys `"prompt"` and `"prompt_embed_kwargs"`,
+ and optionally `"latents"` or `"latent_size"`.
+ cache_datalist_path (Optional[str]): Optional datalist path for `cache_dir`.
+ Supports `.jsonl`, `.jsonl.gz`, or `.json`. If not exists, files are
+ discovered by listing the directory.
+ ignore_cached_latents (bool): If True, ignores any cached latents and
+ prioritizes loading images from `image_dir`. Defaults to False.
+ prompt_dataset_kwargs (Optional[dict]): Keyword arguments forwarded to
+ `datasets.load_dataset(...)`. Enables prompt-dataset mode when provided.
+ If a `DatasetDict` is returned, a split (e.g., "train") is selected
+ internally.
+ image_dir (Optional[str]): Subdirectory of `data_root` with images to pair
+ with prompts (used only in prompt-dataset mode).
+ image_datalist_path (Optional[str]): Optional datalist for `image_dir`
+ (same formats as above). When `bucketize=True`, JSONL entries must
+ include `"size_idx"`.
+ image_extension (Optional[str]): Image file extension used to compose
+ paths when `image_dir` is set. Defaults to ".png".
+ image_scale_factor (float): Scale factor applied to image spatial dimensions
+ after loading. Defaults to 1.0 (no scaling).
+ negative_prompt_embeds_path (Optional[str]): Path to a `torch.load`-able
+ file containing keyward arguments forwarded to the diffusion model
+ for negative prompt embeddings. Added to each sample when provided.
+ negative_prompt_kwargs (Optional[dict]): Keyword arguments forwarded to
+ the text encoder for negative prompts. Added to each sample when provided.
+ pad_seq_len (int): If set, pads/truncates `"encoder_hidden_states"` and
+ `"encoder_hidden_states_mask"` along the sequence dimension to this
+ length. Defaults to None.
+ latent_size (Optional[Tuple[int]]): Default latent shape `(C, H, W)` used
+ when no cached latents exist and no image size is provided. Defaults to
+ `(16, 128, 128)`.
+ vae_scale_factor (Optional[Union[int, Tuple[int]]]): Downscale factor(s)
+ applied to image dimensions when deriving latent sizes from image size.
+ If an `int`, applies to each spatial dim; if a `tuple`, its length must
+ match the provided image spatial size (e.g., `(H, W)` or `(T, H, W)` for
+ video VAEs).
+ repeat (int): Virtual repetition factor for each underlying sample.
+ Affects `__len__` and index mapping.
+ start_ind (Optional[int]): Start index (inclusive) into the underlying
+ dataset. Defaults to 0.
+ end_ind (int): End index (exclusive) into the underlying dataset. Defaults
+ to dataset length.
+ bucketize (bool): If True, enables bucketing in `DistributedSampler` so that
+ each rank receives samples of the same size. Expects `"size_idx"` in JSONL
+ datalists and collects bucket ids. Defaults to False.
+ test_mode (bool): If True, return deterministic noise per sample instead of
+ reading/allocating real latents or images.
+ """
+
+ PROMPT_KEY_MAPS = {
+ 'prompt_embeds': 'encoder_hidden_states',
+ 'prompt_embeds_scale': 'encoder_hidden_states_scale',
+ 'pooled_prompt_embeds': 'pooled_projections',
+ 'prompt_embeds_mask': 'encoder_hidden_states_mask'
+ }
+
+ def __init__(self,
+ data_root: str,
+ cache_dir: Optional[str] = None,
+ cache_datalist_path: Optional[str] = None,
+ ignore_cached_latents: bool = False,
+ prompt_dataset_kwargs: Optional[dict] = None,
+ image_dir: Optional[str] = None,
+ image_datalist_path: Optional[str] = None,
+ image_extension: Optional[str] = '.png',
+ image_scale_factor: float = 1.0,
+ negative_prompt_embeds_path: Optional[str] = None,
+ negative_prompt_kwargs: Optional[dict] = None,
+ pad_seq_len: int = None,
+ latent_size: Optional[Tuple[int]] = (16, 128, 128),
+ vae_scale_factor: Optional[Union[int, Tuple[int]]] = 8,
+ repeat: int = 1,
+ start_ind: Optional[int] = None,
+ end_ind: int = None,
+ bucketize: bool = False,
+ test_mode: bool = False):
+ super().__init__()
+ self.data_root = data_root
+ self.file_client = FileClient.infer_client(uri=self.data_root)
+
+ self.pad_seq_len = pad_seq_len
+
+ self.cache_dir_path = self.cache_datalist_path = None
+ self.prompt_dataset = self.image_dir_path = self.image_datalist_path = None
+ self.image_extension = image_extension
+ self.image_scale_factor = image_scale_factor
+ self.ignore_cached_latents = ignore_cached_latents
+ self.bucketize = bucketize
+ bucket_ids = None
+
+ if (cache_dir is not None
+ and self.file_client.isdir(self.file_client.join_path(data_root, cache_dir))
+ and cache_datalist_path is not None
+ and FileClient.infer_client(uri=cache_datalist_path).isfile(cache_datalist_path)):
+ self.cache_dir_path = self.file_client.join_path(data_root, cache_dir)
+ self.cache_datalist, bucket_ids = self.parse_datalist(
+ self.cache_dir_path, cache_datalist_path)
+ dataset_len = len(self.cache_datalist)
+
+ elif prompt_dataset_kwargs is not None:
+ self.prompt_dataset = load_dataset(**prompt_dataset_kwargs)
+ if isinstance(self.prompt_dataset, DatasetDict):
+ split = 'train' if 'train' in self.prompt_dataset else list(self.prompt_dataset.keys())[0]
+ self.prompt_dataset = self.prompt_dataset[split]
+ assert isinstance(self.prompt_dataset, HFDataset), \
+ f"Expected HF Dataset/DatasetDict, got {type(self.prompt_dataset)}."
+ dataset_len = len(self.prompt_dataset)
+
+ else:
+ raise ValueError('Either `cache_dir` or `prompt_dataset_kwargs` must be provided.')
+
+ if image_dir is not None and self.file_client.isdir(self.file_client.join_path(data_root, image_dir)):
+ self.image_dir_path = self.file_client.join_path(data_root, image_dir)
+ self.image_datalist, bucket_ids = self.parse_datalist(
+ self.image_dir_path, image_datalist_path, datalist_must_exist=True)
+ assert dataset_len == len(self.image_datalist)
+
+ if bucket_ids is None and self.bucketize:
+ assert self.prompt_dataset is not None
+ bucket_ids = self.get_bucket_ids_from_prompt_dataset()
+
+ self.negative_prompt_embed_kwargs = None
+ if negative_prompt_embeds_path is not None:
+ negative_prompt_embeds_bytesio = BytesIO(
+ FileClient.infer_client(uri=negative_prompt_embeds_path).get(negative_prompt_embeds_path))
+ self.negative_prompt_embed_kwargs = self.parse_prompt_embeds(
+ torch.load(negative_prompt_embeds_bytesio, map_location='cpu'))
+ self.negative_prompt_kwargs = negative_prompt_kwargs
+
+ self.latent_size = latent_size
+ self.vae_scale_factor = vae_scale_factor
+
+ self.repeat = repeat
+ if start_ind is not None:
+ start_ind = max(min(start_ind, dataset_len - 1), -dataset_len) % dataset_len
+ else:
+ start_ind = 0
+ if end_ind is not None:
+ end_ind = max(min(end_ind - 1, dataset_len - 1), -dataset_len) % dataset_len + 1
+ else:
+ end_ind = dataset_len
+ assert start_ind < end_ind, f'Invalid start_ind and end_ind.'
+ self.start_ind = start_ind
+ self.end_ind = end_ind
+
+ if self.bucketize:
+ assert bucket_ids is not None and len(bucket_ids) == dataset_len
+ self.bucket_ids = [bucket_ids[self._map_idx(i)] for i in range(len(self))]
+
+ self.test_mode = test_mode
+
+ def get_bucket_ids_from_prompt_dataset(self):
+ ds = self.prompt_dataset
+ assert 'height' in ds.column_names and 'width' in ds.column_names, \
+ 'When bucketize=True and no datalist is provided, the prompt dataset ' \
+ 'must contain `height` and `width` columns.'
+ cols = ['height', 'width']
+ if 'frames' in ds.column_names:
+ cols = ['frames'] + cols
+ ds_arrow = ds.with_format('arrow', columns=cols)
+ batch = ds_arrow[:]
+
+ arrs = [batch[c].combine_chunks().to_numpy(zero_copy_only=False) for c in cols]
+ arrs = np.stack(arrs, axis=1)
+
+ _, inv = np.unique(arrs, axis=0, return_inverse=True)
+ return inv.tolist()
+
+ def parse_datalist(self, dir_path, datalist_path=None, datalist_must_exist=False):
+ logger = get_root_logger()
+
+ if datalist_path is not None and FileClient.infer_client(uri=datalist_path).isfile(datalist_path):
+ filenames = []
+ bucket_ids = []
+
+ datalist_bytesio = BytesIO(FileClient.infer_client(uri=datalist_path).get(datalist_path))
+ if datalist_path.endswith('.jsonl.gz') or datalist_path.endswith('.jsonl'):
+ if datalist_path.endswith('.jsonl.gz'):
+ with gzip.open(datalist_bytesio, 'rt', encoding='utf-8') as f:
+ datalist = f.readlines()
+ else:
+ datalist = datalist_bytesio.read().decode('utf-8').splitlines()
+ for line in datalist:
+ data_item = orjson.loads(line)
+ if 'filename' in data_item:
+ filenames.append(data_item['filename'])
+ elif 'image_hash' in data_item:
+ filenames.append(data_item['image_hash'])
+ else:
+ raise ValueError('No valid key to identify data item.')
+ if self.bucketize:
+ assert 'size_idx' in data_item, 'size_idx must be provided for bucketize.'
+ bucket_ids.append(data_item['size_idx'])
+ elif datalist_path.endswith('.json'):
+ assert not self.bucketize, 'Bucketize not supported for json datalist.'
+ datalist = orjson.loads(datalist_bytesio.read())
+ for data_item in datalist:
+ filenames.append(os.path.splitext(os.path.basename(data_item))[0])
+ else:
+ raise ValueError('Datalist file must be .jsonl, .jsonl.gz or .json')
+
+ else:
+ assert not datalist_must_exist, f'Datalist file {datalist_path} does not exist.'
+ assert not self.bucketize, 'Bucketize not supported when datalist is not provided.'
+ mmcv.print_log(
+ f'Datalist file {datalist_path} does not exist, directly list all files in the directory.',
+ logger=logger,
+ level=logging.WARNING)
+ # list all files in the directory
+ filenames = [os.path.splitext(p)[0] for p in self.file_client.list_dir_or_file(dir_path)]
+ filenames.sort()
+ bucket_ids = None
+ # save the datalist if datalist_path is provided
+ if datalist_path is not None:
+ if datalist_path.endswith('.jsonl.gz') or datalist_path.endswith('.jsonl'):
+ datalist = []
+ for filename in filenames:
+ datalist.append(orjson.dumps({'filename': filename}).decode('utf-8'))
+ datalist_str = '\n'.join(datalist)
+ if datalist_path.endswith('.jsonl.gz'):
+ datalist_bytesio = BytesIO()
+ with gzip.open(datalist_bytesio, 'wt', encoding='utf-8') as f:
+ f.write(datalist_str)
+ FileClient.infer_client(uri=datalist_path).put(datalist_bytesio.getvalue(), datalist_path)
+ else:
+ FileClient.infer_client(uri=datalist_path).put_text(datalist_str, datalist_path)
+ elif datalist_path.endswith('.json'):
+ datalist = filenames
+ FileClient.infer_client(uri=datalist_path).put_text(
+ orjson.dumps(datalist).decode('utf-8'), datalist_path)
+
+ mmcv.print_log(f'Loaded {len(filenames)} samples.', logger=logger)
+
+ return filenames, bucket_ids
+
+ def pad_prompt_embeds(self, prompt_embeds):
+ if self.pad_seq_len is not None:
+ if prompt_embeds.size(0) > self.pad_seq_len:
+ prompt_embeds = prompt_embeds[:self.pad_seq_len]
+ else:
+ zeros_size = (self.pad_seq_len - prompt_embeds.size(0),) + prompt_embeds.shape[1:]
+ prompt_embeds = torch.cat([prompt_embeds, prompt_embeds.new_zeros(zeros_size)], dim=0)
+ return prompt_embeds
+
+ def parse_prompt_embeds(self, data):
+ prompt_embed_kwargs = data.get('prompt_embed_kwargs', {}).copy()
+
+ # Map legacy keys to new ones if not already present
+ for legacy_key, new_key in self.PROMPT_KEY_MAPS.items():
+ if legacy_key in data and new_key not in prompt_embed_kwargs:
+ prompt_embed_kwargs[new_key] = data[legacy_key]
+
+ # Common post-processing
+ encoder_hidden_states_scale = prompt_embed_kwargs.pop('encoder_hidden_states_scale', None)
+ if 'encoder_hidden_states' in prompt_embed_kwargs:
+ encoder_hidden_states = prompt_embed_kwargs['encoder_hidden_states'].float()
+ if encoder_hidden_states_scale is not None:
+ encoder_hidden_states = encoder_hidden_states * encoder_hidden_states_scale
+ prompt_embed_kwargs['encoder_hidden_states'] = self.pad_prompt_embeds(encoder_hidden_states)
+
+ if 'pooled_projections' in prompt_embed_kwargs:
+ prompt_embed_kwargs['pooled_projections'] = prompt_embed_kwargs['pooled_projections'].float()
+
+ if 'encoder_hidden_states_mask' in prompt_embed_kwargs:
+ prompt_embed_kwargs['encoder_hidden_states_mask'] = self.pad_prompt_embeds(
+ prompt_embed_kwargs['encoder_hidden_states_mask'])
+
+ return prompt_embed_kwargs
+
+ def calculate_latent_size(self, image_spatial_size):
+ if isinstance(self.vae_scale_factor, int):
+ latent_spatial_size = tuple(s // self.vae_scale_factor for s in image_spatial_size)
+ else:
+ assert len(self.vae_scale_factor) == len(image_spatial_size)
+ latent_spatial_size = tuple(
+ s // f for s, f in zip(image_spatial_size, self.vae_scale_factor))
+ latent_size = (self.latent_size[0],) + latent_spatial_size
+ return latent_size
+
+ def calculate_scaled_image_size(self, image_spatial_size):
+ if self.image_scale_factor != 1:
+ if len(image_spatial_size) == 2:
+ new_spatial_size = (int(round(image_spatial_size[0] * self.image_scale_factor)),
+ int(round(image_spatial_size[1] * self.image_scale_factor)))
+ elif len(image_spatial_size) == 3:
+ new_spatial_size = (image_spatial_size[0],
+ int(round(image_spatial_size[1] * self.image_scale_factor)),
+ int(round(image_spatial_size[2] * self.image_scale_factor)))
+ else:
+ raise ValueError(f'Unsupported image spatial size {image_spatial_size}.')
+ else:
+ new_spatial_size = image_spatial_size
+ return new_spatial_size
+
+ def scale_image(self, image):
+ if self.image_scale_factor != 1:
+ new_spatial_size = self.calculate_scaled_image_size(image.shape[1:])
+ if len(new_spatial_size) == 2:
+ image = F.interpolate(
+ image[None], size=new_spatial_size, mode='bicubic', align_corners=False, antialias=True
+ )[0].clamp(min=0, max=1)
+ elif len(new_spatial_size) == 3:
+ image = F.interpolate(
+ image, size=new_spatial_size[1:], mode='bicubic', align_corners=False, antialias=True
+ ).clamp(min=0, max=1)
+ else:
+ raise ValueError(f'Unsupported image spatial size {image.shape[1:]}.')
+ return image
+
+ def _map_idx(self, idx):
+ return self.start_ind + (idx // self.repeat)
+
+ def __len__(self):
+ return self.repeat * (self.end_ind - self.start_ind)
+
+ def __getitem__(self, idx):
+ mapped_idx = self._map_idx(idx)
+
+ prompt_data = None
+
+ if self.cache_dir_path is not None:
+ data_path = self.file_client.join_path(
+ self.cache_dir_path, f'{self.cache_datalist[mapped_idx]}.zst')
+ data_bytesio = BytesIO(self.file_client.get(data_path))
+ with zstd.ZstdDecompressor().stream_reader(data_bytesio) as f:
+ raw_data = pickle.load(f)
+ data = dict(
+ ids=DC(idx, cpu_only=True),
+ name=DC(raw_data['prompt'], cpu_only=True),
+ prompt_embed_kwargs=self.parse_prompt_embeds(raw_data))
+
+ if not self.ignore_cached_latents: # load latents
+ if 'latents' in raw_data:
+ latents = raw_data['latents']
+ if self.test_mode:
+ data['noise'] = torch.randn(
+ latents.size(), dtype=torch.float32, generator=torch.Generator().manual_seed(idx))
+ else:
+ data['latents'] = latents.float()
+ latents_scale = raw_data.get('latents_scale', None)
+ if latents_scale is not None:
+ data['latents'] = data['latents'] * latents_scale
+ else:
+ latent_size = raw_data.get('latent_size', self.latent_size)
+ if self.test_mode:
+ data['noise'] = torch.randn(
+ latent_size, dtype=torch.float32, generator=torch.Generator().manual_seed(idx))
+ else:
+ data['latents'] = torch.empty(latent_size, dtype=torch.float32)
+
+ else:
+ prompt_data = self.prompt_dataset[mapped_idx]
+ if 'prompt_kwargs' in prompt_data:
+ prompt_kwargs = {k: DC(v, cpu_only=True) for k, v in prompt_data['prompt_kwargs'].items()}
+ else:
+ prompt_kwargs = dict(prompt=DC(prompt_data['prompt'], cpu_only=True))
+ data = dict(
+ ids=DC(idx, cpu_only=True),
+ name=DC(prompt_data['prompt'], cpu_only=True),
+ prompt_kwargs=prompt_kwargs)
+
+ if self.image_dir_path is not None:
+ image_path = self.file_client.join_path(
+ self.image_dir_path, self.image_datalist[mapped_idx] + self.image_extension)
+ image = load_image(image_path, self.file_client)
+ image = np.moveaxis(image, -1, 0) # channel first
+ if self.test_mode:
+ data['noise'] = torch.randn(
+ self.calculate_latent_size(self.calculate_scaled_image_size(image.shape[1:])),
+ dtype=torch.float32, generator=torch.Generator().manual_seed(idx))
+ else:
+ images = torch.from_numpy(image)
+ if images.dtype == torch.uint8:
+ images = images.float() / 255.0
+ assert torch.is_floating_point(images), f'Image dtype {images.dtype} not supported.'
+ data['images'] = self.scale_image(images.float())
+ elif 'latents' not in data and 'noise' not in data: # allocate latents if not already loaded
+ if prompt_data is not None and 'height' in prompt_data and 'width' in prompt_data:
+ image_spatial_size = (prompt_data['height'], prompt_data['width'])
+ if 'frames' in prompt_data:
+ image_spatial_size = (prompt_data['frames'],) + image_spatial_size
+ latent_size = self.calculate_latent_size(self.calculate_scaled_image_size(image_spatial_size))
+ else:
+ latent_size = self.latent_size
+ if self.test_mode:
+ data['noise'] = torch.randn(
+ latent_size, dtype=torch.float32, generator=torch.Generator().manual_seed(idx))
+ else:
+ data['latents'] = torch.empty(latent_size, dtype=torch.float32)
+
+ if self.negative_prompt_embed_kwargs is not None:
+ data.update(negative_prompt_embed_kwargs=self.negative_prompt_embed_kwargs)
+ if self.negative_prompt_kwargs is not None:
+ data.update(negative_prompt_kwargs=self.negative_prompt_kwargs)
+
+ return data
diff --git a/lakonlab/datasets/imagenet.py b/lakonlab/datasets/imagenet.py
new file mode 100644
index 0000000000000000000000000000000000000000..f67ad954d2a151e032e638e687161dc8b9d4bae2
--- /dev/null
+++ b/lakonlab/datasets/imagenet.py
@@ -0,0 +1,155 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import os
+
+import numpy as np
+import torch
+import mmcv
+
+from io import BytesIO
+from PIL import Image
+from torch.utils.data import Dataset
+from mmcv.fileio import FileClient
+from mmcv.parallel import DataContainer as DC
+from mmgen.datasets.builder import DATASETS
+from mmgen.utils import get_root_logger
+
+
+def image_preproc(pil_image, image_size, random_flip=False):
+ """
+ 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
+ arr = arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size]
+
+ if random_flip and np.random.rand() < 0.5:
+ arr = np.ascontiguousarray(arr[:, ::-1])
+
+ if arr.ndim == 2:
+ arr = np.stack([arr] * 3, axis=-1)
+ elif arr.ndim == 3:
+ if arr.shape[2] == 1:
+ arr = np.concatenate([arr] * 3, axis=-1)
+ elif arr.shape[2] == 4:
+ arr = arr[:, :, :3]
+ else:
+ assert arr.shape[2] == 3
+ else:
+ raise ValueError(f'Unexpected number of dimensions: {arr.ndim}')
+ return arr
+
+
+@DATASETS.register_module()
+class ImageNet(Dataset):
+ def __init__(
+ self,
+ data_root='data/imagenet/train',
+ datalist_path='data/imagenet/train.txt',
+ label2name_path='data/imagenet/imagenet1000_clsidx_to_labels.txt',
+ random_flip=True,
+ negative_label=1000,
+ image_size=256,
+ latent_size=(4, 32, 32),
+ test_label_repeat=1,
+ test_mode=False,
+ num_test_images=50000):
+ super().__init__()
+ self.data_root = data_root
+ self.file_client = FileClient.infer_client(uri=self.data_root)
+
+ self.datalist_path = datalist_path
+ self.label2name_path = label2name_path
+ self.random_flip = random_flip
+ self.negative_label = negative_label
+ self.image_size = image_size
+ self.latent_size = latent_size
+ self.test_label_repeat = test_label_repeat
+ self.test_mode = test_mode
+ self.num_test_images = num_test_images
+
+ self.label2name = {}
+ label2name_text = FileClient.infer_client(uri=self.label2name_path).get_text(self.label2name_path)
+ for line in label2name_text.split('\n'):
+ line = line.strip()
+ if len(line) == 0:
+ continue
+ idx, name = line.split(':')
+ idx, name = idx.strip(), name.strip()
+ if name[-1] == ',':
+ name = name[:-1]
+ if name[0] == '"' and name[-1] == '"':
+ name = name[1:-1]
+ if name[0] == "'" and name[-1] == "'":
+ name = name[1:-1]
+ self.label2name[int(idx)] = name
+
+ if not test_mode:
+ self.all_paths = []
+ self.all_labels = []
+ datalist_text = FileClient.infer_client(uri=self.datalist_path).get_text(self.datalist_path)
+ for line in datalist_text.split('\n'):
+ line = line.strip()
+ if len(line) == 0:
+ continue
+ path_label = line.split(' ')
+ self.all_paths.append(path_label[0])
+ if len(path_label) > 1:
+ self.all_labels.append(int(path_label[1]))
+
+ logger = get_root_logger()
+ mmcv.print_log(f'Data root: {self.data_root}', logger=logger)
+ mmcv.print_log(f'Data list path: {self.datalist_path}', logger=logger)
+ mmcv.print_log(f'Number of images: {len(self.all_paths)}', logger=logger)
+
+ def __len__(self):
+ return self.num_test_images if self.test_mode else len(self.all_paths)
+
+ def __getitem__(self, idx):
+ data = dict(ids=DC(idx, cpu_only=True))
+
+ if self.test_mode:
+ label_generator = torch.Generator().manual_seed(idx // self.test_label_repeat)
+ label = torch.randint(0, 1000, (), generator=label_generator).long()
+ noise_generator = torch.Generator().manual_seed(idx + 1000)
+ noise = torch.randn(self.latent_size, generator=noise_generator)
+ data.update(noise=noise)
+
+ else:
+ rel_data_path = self.all_paths[idx]
+ data.update(paths=DC(rel_data_path, cpu_only=True))
+ data_path = self.file_client.join_path(self.data_root, rel_data_path)
+ data_bytesio = BytesIO(self.file_client.get(data_path))
+ ext = os.path.splitext(data_path)[-1]
+ if ext.lower() in ('.pth', '.pt'):
+ torch_data = torch.load(data_bytesio, map_location='cpu')
+ label = torch_data['y'].long()
+ data.update(latents=torch_data['x'].float())
+ elif ext.lower() in ('.jpg', '.jpeg', '.png'):
+ label = torch.tensor(self.all_labels[idx], dtype=torch.long)
+ img_data = Image.open(data_bytesio)
+ data.update(
+ images=torch.from_numpy(image_preproc(
+ img_data, self.image_size, random_flip=self.random_flip)).float().permute(2, 0, 1) / 255.0)
+ else:
+ raise ValueError(f'Unsupported file extension: {ext}')
+
+ name = self.label2name[label.item()]
+ data.update(labels=label, name=DC(name, cpu_only=True))
+
+ if self.negative_label is not None:
+ if isinstance(self.negative_label, int):
+ data.update(negative_labels=torch.tensor(self.negative_label, dtype=torch.long))
+ else:
+ raise ValueError(f'Unsupported negative label: {self.negative_label}')
+
+ return data
diff --git a/lakonlab/datasets/samplers/__init__.py b/lakonlab/datasets/samplers/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..7251f1fdb8561c8b730f7e58bf523dfd1e806c6b
--- /dev/null
+++ b/lakonlab/datasets/samplers/__init__.py
@@ -0,0 +1 @@
+from .distributed_sampler import DistributedSampler
diff --git a/lakonlab/datasets/samplers/__pycache__/__init__.cpython-310.pyc b/lakonlab/datasets/samplers/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..ee7ba09a1fffc12430ff0f6686364b18de1b629e
Binary files /dev/null and b/lakonlab/datasets/samplers/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/datasets/samplers/__pycache__/distributed_sampler.cpython-310.pyc b/lakonlab/datasets/samplers/__pycache__/distributed_sampler.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..755d38f7a308fb009b713df89f06e95b37a6cb85
Binary files /dev/null and b/lakonlab/datasets/samplers/__pycache__/distributed_sampler.cpython-310.pyc differ
diff --git a/lakonlab/datasets/samplers/distributed_sampler.py b/lakonlab/datasets/samplers/distributed_sampler.py
new file mode 100644
index 0000000000000000000000000000000000000000..398796bf2fb227714a3c6a4ae18e1b9596c54459
--- /dev/null
+++ b/lakonlab/datasets/samplers/distributed_sampler.py
@@ -0,0 +1,158 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import numpy as np
+import torch
+
+from torch.utils.data import DistributedSampler as _DistributedSampler
+from mmgen.utils import sync_random_seed
+
+
+def reverse_index_map(bucket_ids):
+ bucket_map = dict()
+ for data_id, bucket_id in enumerate(bucket_ids):
+ if bucket_id not in bucket_map:
+ bucket_map[bucket_id] = []
+ bucket_map[bucket_id].append(data_id)
+ return bucket_map
+
+
+class DistributedSampler(_DistributedSampler):
+
+ def __init__(self,
+ dataset,
+ num_replicas=None,
+ rank=None,
+ shuffle=True,
+ samples_per_gpu=1,
+ seed=None):
+ super().__init__(dataset, num_replicas=num_replicas, rank=rank)
+
+ self.shuffle = shuffle
+ self.samples_per_gpu = samples_per_gpu
+
+ self.bucket_map = self.total_size_bucketwise = None
+ if hasattr(dataset, 'bucket_ids'):
+ self._init_bucket_sampler(dataset)
+ else:
+ self._init_sampler(dataset)
+
+ self.seed = sync_random_seed(seed)
+ self.skip_iter = 0
+
+ def _init_sampler(self, dataset):
+ data_len = len(dataset)
+ # to avoid padding bug when meeting too small dataset
+ if data_len < self.num_replicas * self.samples_per_gpu:
+ raise ValueError(
+ 'You may use too small dataset and our distributed '
+ 'sampler cannot pad your dataset correctly. Please '
+ 'use fewer GPUs or smaller batch sizes per GPU.')
+
+ num_batches = int(np.ceil(data_len / self.num_replicas / self.samples_per_gpu))
+ self.num_samples = num_batches * self.samples_per_gpu
+ self.total_size = self.num_samples * self.num_replicas
+
+ def _init_bucket_sampler(self, dataset):
+ self.bucket_map = reverse_index_map(dataset.bucket_ids)
+ self.bucket_map = dict(sorted(self.bucket_map.items())) # sort by bucket_id
+
+ data_len = 0
+ self.total_size_bucketwise = {}
+
+ for bucket_id, data_indices in self.bucket_map.items():
+ _data_len = len(data_indices)
+ if _data_len < self.samples_per_gpu:
+ raise ValueError(
+ 'You may use too small dataset and our distributed '
+ 'sampler cannot pad your dataset correctly. Please '
+ 'use smaller batch sizes per GPU.')
+
+ _total_num_batches = int(np.ceil(_data_len / self.samples_per_gpu))
+ _total_size = _total_num_batches * self.samples_per_gpu
+
+ data_len += _total_size
+ self.total_size_bucketwise[bucket_id] = _total_size
+
+ if data_len < self.num_replicas * self.samples_per_gpu:
+ raise ValueError(
+ 'You may use too small dataset and our distributed '
+ 'sampler cannot pad your dataset correctly. Please '
+ 'use fewer GPUs or smaller batch sizes per GPU.')
+
+ num_batches = int(np.ceil(data_len / self.num_replicas / self.samples_per_gpu))
+ self.num_samples = num_batches * self.samples_per_gpu
+ self.total_size = self.num_samples * self.num_replicas
+
+ def update_sampler(self, dataset, samples_per_gpu=None):
+ self.dataset = dataset
+ if samples_per_gpu is not None:
+ self.samples_per_gpu = samples_per_gpu
+ self.bucket_map = self.total_size_bucketwise = None
+ if hasattr(dataset, 'bucket_ids'):
+ self._init_bucket_sampler(dataset)
+ else:
+ self._init_sampler(dataset)
+
+ def set_iter(self, iteration):
+ num_batches = self.num_samples // self.samples_per_gpu
+ self.skip_iter = iteration % num_batches
+
+ def __iter__(self):
+ if self.bucket_map is None:
+ if self.shuffle:
+ g = torch.Generator()
+ g.manual_seed(self.seed + self.epoch)
+ indices = torch.randperm(len(self.dataset), generator=g).tolist()
+ else:
+ indices = torch.arange(len(self.dataset)).tolist()
+ # add extra samples to make it evenly divisible
+ indices += indices[:(self.total_size - len(indices))]
+ assert len(indices) == self.total_size
+ # subsample
+ indices = indices[self.rank:self.total_size:self.num_replicas]
+
+ else: # guarantees that batch samples are from the same bucket
+ if self.shuffle:
+ g = torch.Generator()
+ g.manual_seed(self.seed + self.epoch)
+ else:
+ g = None
+ indices = []
+ for bucket_id, data_indices in self.bucket_map.items():
+ data_indices = torch.tensor(data_indices)
+ if g is not None:
+ data_indices = data_indices[torch.randperm(len(data_indices), generator=g)]
+ pad = self.total_size_bucketwise[bucket_id] - data_indices.numel()
+ if pad:
+ data_indices = torch.cat([data_indices, data_indices[:pad]], dim=0)
+ assert data_indices.numel() == self.total_size_bucketwise[bucket_id]
+ _total_num_batches = self.total_size_bucketwise[bucket_id] // self.samples_per_gpu
+ _num_batches = _total_num_batches // self.num_replicas
+ _total_leftover_batches = _total_num_batches % self.num_replicas
+ # data_indices_a: evenly split batches for full round-robins across replicas
+ # data_indices_b: the leftover partial round-robin
+ data_indices_a = data_indices[:(_num_batches * self.num_replicas * self.samples_per_gpu)].reshape(
+ _num_batches, self.samples_per_gpu, self.num_replicas
+ ).permute(0, 2, 1).reshape(
+ _num_batches * self.num_replicas, self.samples_per_gpu)
+ data_indices_b = data_indices[(_num_batches * self.num_replicas * self.samples_per_gpu):].reshape(
+ self.samples_per_gpu, _total_leftover_batches
+ ).permute(1, 0)
+ indices.extend([data_indices_a, data_indices_b])
+ indices = torch.cat(indices, dim=0) # (total_num_batches, samples_per_gpu)
+ if g is not None:
+ indices = indices[torch.randperm(indices.size(0), generator=g)]
+ total_num_batches = self.total_size // self.samples_per_gpu
+ pad = total_num_batches - indices.size(0)
+ if pad:
+ indices = torch.cat([indices, indices[:pad]], dim=0)
+ assert indices.numel() == self.total_size
+ indices = indices[self.rank:total_num_batches:self.num_replicas].flatten().tolist()
+
+ assert len(indices) == self.num_samples
+ skip_len = self.skip_iter * self.samples_per_gpu
+ assert skip_len < self.num_samples
+ indices = indices[skip_len:]
+ self.skip_iter = 0
+
+ return iter(indices)
diff --git a/lakonlab/evaluation/__init__.py b/lakonlab/evaluation/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..ffb1d4fc8ac173bd3e4146da900367b10b19db1d
--- /dev/null
+++ b/lakonlab/evaluation/__init__.py
@@ -0,0 +1,8 @@
+from .metrics import FIDKID, PR, InceptionMetrics, ColorStats, HPSv2, CLIPSimilarity
+from .vqa_score import VQAScore
+from .hpsv3 import HPSv3
+from .eval_hooks import GenerativeEvalHook
+
+__all__ = ['GenerativeEvalHook', 'FIDKID', 'PR',
+ 'InceptionMetrics', 'ColorStats', 'HPSv2', 'VQAScore', 'CLIPSimilarity',
+ 'HPSv3']
diff --git a/lakonlab/evaluation/__pycache__/__init__.cpython-310.pyc b/lakonlab/evaluation/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..e665cf2c4a4594f1fc6cc0df40e3b18dd984c90c
Binary files /dev/null and b/lakonlab/evaluation/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/evaluation/__pycache__/eval_hooks.cpython-310.pyc b/lakonlab/evaluation/__pycache__/eval_hooks.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..f8a54c42847c2e45039dfed8088dc5d88e2f7906
Binary files /dev/null and b/lakonlab/evaluation/__pycache__/eval_hooks.cpython-310.pyc differ
diff --git a/lakonlab/evaluation/__pycache__/hpsv3.cpython-310.pyc b/lakonlab/evaluation/__pycache__/hpsv3.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..9cb3146360ccc941fad774bd388996d39c6be039
Binary files /dev/null and b/lakonlab/evaluation/__pycache__/hpsv3.cpython-310.pyc differ
diff --git a/lakonlab/evaluation/__pycache__/metrics.cpython-310.pyc b/lakonlab/evaluation/__pycache__/metrics.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..256dc1331c98471a691ffd6c281bd5ecbff80633
Binary files /dev/null and b/lakonlab/evaluation/__pycache__/metrics.cpython-310.pyc differ
diff --git a/lakonlab/evaluation/__pycache__/vqa_score.cpython-310.pyc b/lakonlab/evaluation/__pycache__/vqa_score.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..b53acec90b1033f7a8ab1bcc2667c65531904d16
Binary files /dev/null and b/lakonlab/evaluation/__pycache__/vqa_score.cpython-310.pyc differ
diff --git a/lakonlab/evaluation/eval_hooks.py b/lakonlab/evaluation/eval_hooks.py
new file mode 100644
index 0000000000000000000000000000000000000000..73a782886dda20972f95b7a7ab6d29ab0426315a
--- /dev/null
+++ b/lakonlab/evaluation/eval_hooks.py
@@ -0,0 +1,318 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import sys
+import os
+import re
+import unicodedata
+import numpy as np
+import torch
+import torch.distributed as dist
+import mmcv
+
+from copy import deepcopy
+from concurrent.futures import ThreadPoolExecutor
+from mmcv.runner import HOOKS, get_dist_info
+from mmcv.fileio import FileClient
+from mmgen.models.architectures.common import get_module_device
+from mmgen.core import GenerativeEvalHook as _GenerativeEvalHook
+from lakonlab.utils.io_utils import save_image, save_video, load_images_parallel
+from lakonlab.runner.timer import default_timers
+from lakonlab.utils import gc_context
+from lakonlab.ui.media_viewer import write_html
+
+default_timers.add_timer('total time')
+
+_reserved = {"CON", "PRN", "AUX", "NUL", *(f"COM{i}" for i in range(1, 10)), *(f"LPT{i}" for i in range(1, 10))}
+_invalid = re.compile(r'[<>:"/\\|?*\:%]+')
+
+
+def flatten_list(lst):
+ for item in lst:
+ if isinstance(item, list):
+ yield from flatten_list(item) # recurse into sub-list
+ else:
+ yield item
+
+
+def _safe_name(name: str, max_len: int = 160) -> str:
+ s = unicodedata.normalize("NFKC", str(name))
+ s = "".join(c for c in s if 32 <= ord(c) != 127) # drop control chars
+ s = _invalid.sub("", s).strip(" ._") # strip invalid chars
+ if s.startswith("."):
+ s = s.lstrip(".") # avoid hidden files
+ if s.upper() in _reserved:
+ s += "_" # avoid reserved names
+ s = (s or "untitled")[:max_len].rstrip(" .") # truncate & clean
+ return s or "untitled"
+
+
+def evaluate(model, dataloader, metrics=None,
+ feed_batch_size=32, viz_dir=None, viz_num=None, sample_kwargs=dict(),
+ fps=16, enable_timers=False, reuse_viz=False):
+ has_metrics = metrics is not None and len(metrics) > 0
+ if has_metrics:
+ for metric in metrics:
+ if hasattr(metric, 'load_to_gpu'):
+ metric.load_to_gpu()
+
+ if enable_timers:
+ default_timers.enable_all()
+
+ batch_size = dataloader.batch_size
+ rank, ws = get_dist_info()
+ total_batch_size = batch_size * ws
+
+ max_num_fakes = len(dataloader.dataset)
+
+ if viz_dir is not None:
+ html_entries = []
+ saved_data_ids = set()
+ file_client = FileClient.infer_client(uri=viz_dir)
+ executor = ThreadPoolExecutor(max_workers=(os.cpu_count() or 4) * 4)
+ else:
+ html_entries = saved_data_ids = file_client = executor = None
+
+ if rank == 0:
+ mmcv.print_log(
+ f'Generate {max_num_fakes} fake samples for evaluation', 'mmgen')
+ pbar = mmcv.ProgressBar(max_num_fakes)
+
+ log_vars = dict()
+ batch_size_list = []
+
+ for i, data in enumerate(dataloader):
+ can_reuse = False
+ loaded_imgs_tensor = None
+ loaded_html_entries = None
+
+ if viz_dir is not None and reuse_viz:
+ batch_names = list(flatten_list(data['name'].data))
+ ids = list(flatten_list(data['ids'].data))
+
+ reuse_candidates = []
+ all_png_exist = True
+
+ for name, data_id in zip(batch_names, ids):
+ safe_name = _safe_name(name)
+ filename = f'{data_id:09d}_{safe_name}.png'
+ filepath = file_client.join_path(viz_dir, filename)
+ if file_client.isfile(filepath):
+ reuse_candidates.append((data_id, name, filename, filepath))
+ else:
+ all_png_exist = False
+ break
+
+ if all_png_exist:
+ filepaths = []
+ loaded_html_entries = []
+ for (data_id, name, filename, filepath) in reuse_candidates:
+ filepaths.append(filepath)
+ rel_filepath = file_client.join_path(os.path.basename(viz_dir), filename)
+ loaded_html_entries.append((data_id, rel_filepath, name))
+ try:
+ loaded_imgs_tensor = torch.from_numpy(
+ np.stack(load_images_parallel(filepaths, file_client), axis=0)
+ ).permute(0, 3, 1, 2).float() / 255.0
+ can_reuse = True
+ except Exception as e:
+ can_reuse = False
+ loaded_imgs_tensor = None
+ loaded_html_entries = None
+
+ if can_reuse:
+ outputs_dict = dict(
+ pred_imgs=loaded_imgs_tensor, # [0,1] float32
+ num_samples=loaded_imgs_tensor.size(0))
+ html_entries.extend(loaded_html_entries)
+
+ else:
+ # fall back to generation
+ sample_kwargs_ = deepcopy(sample_kwargs)
+
+ with default_timers['total time']:
+ outputs_dict = model.val_step(data, show_pbar=rank == 0, **sample_kwargs_)
+
+ if viz_dir is not None:
+ batch_names = list(flatten_list(data['name'].data))
+ for batch_id, data_id in enumerate(flatten_list(data['ids'].data)):
+ if (viz_num is not None and data_id >= viz_num) or (data_id in saved_data_ids):
+ continue
+ name = batch_names[batch_id]
+ safe_name = _safe_name(name)
+
+ image_viz = (outputs_dict['pred_imgs'][batch_id] * 255).round().to(torch.uint8)
+ if image_viz.dim() == 3: # image
+ image_viz = image_viz.permute(1, 2, 0).cpu().numpy()
+ filename = f'{data_id:09d}_{safe_name}.png'
+ executor.submit(
+ save_image,
+ image_viz, file_client.join_path(viz_dir, filename), file_client)
+ elif image_viz.dim() == 4: # video
+ image_viz = image_viz.permute(1, 2, 3, 0).cpu().numpy() # (t, h, w, c)
+ filename = f'{data_id:09d}_{safe_name}.mp4'
+ executor.submit(
+ save_video,
+ image_viz, file_client.join_path(viz_dir, filename), file_client, fps)
+ else:
+ raise ValueError(f'Unsupported image dimension: {image_viz.dim()}')
+ rel_filepath = file_client.join_path(os.path.basename(viz_dir), filename)
+ html_entries.append((data_id, rel_filepath, name))
+ saved_data_ids.add(data_id)
+
+ if 'log_vars' in outputs_dict:
+ for k, v in outputs_dict['log_vars'].items():
+ if k in log_vars:
+ log_vars[k].append(outputs_dict['log_vars'][k])
+ else:
+ log_vars[k] = [outputs_dict['log_vars'][k]]
+ batch_size_list.append(outputs_dict['num_samples'])
+
+ if has_metrics:
+ pred_imgs = outputs_dict['pred_imgs'].split(feed_batch_size, dim=0)
+ real_imgs = None
+ if 'images' in data:
+ real_imgs = data['images']
+ elif 'target_imgs' in outputs_dict:
+ real_imgs = outputs_dict['target_imgs']
+ if real_imgs is not None:
+ real_imgs = real_imgs.split(feed_batch_size, dim=0)
+ requires_prompt = False
+ for metric in metrics:
+ requires_prompt |= getattr(metric, 'requires_prompt', False)
+ if requires_prompt:
+ prompts = list(flatten_list(data['name'].data)) # list of prompts
+ prompts = [prompts[i:i + feed_batch_size] for i in range(0, len(prompts), feed_batch_size)]
+ for metric in metrics:
+ for batch_id, batch_imgs in enumerate(pred_imgs):
+ if getattr(metric, 'requires_prompt', False):
+ metric.feed(
+ dict(imgs=batch_imgs * 2 - 1, prompts=prompts[batch_id]), 'fakes')
+ if real_imgs is not None:
+ metric.feed(
+ dict(imgs=real_imgs[batch_id] * 2 - 1, prompts=prompts[batch_id]), 'reals')
+ else:
+ metric.feed(batch_imgs * 2 - 1, 'fakes')
+ if real_imgs is not None:
+ metric.feed(real_imgs[batch_id] * 2 - 1, 'reals')
+
+ if rank == 0:
+ pbar.update(total_batch_size)
+
+ if ws > 1:
+ device = get_module_device(model)
+ batch_size_list = torch.tensor(batch_size_list, dtype=torch.float, device=device)
+ batch_size_sum = torch.sum(batch_size_list)
+ dist.all_reduce(batch_size_sum, op=dist.ReduceOp.SUM)
+ for k, v in log_vars.items():
+ weigted_values = torch.tensor(log_vars[k], dtype=torch.float, device=device) * batch_size_list
+ weigted_values_sum = torch.sum(weigted_values)
+ dist.all_reduce(weigted_values_sum, op=dist.ReduceOp.SUM)
+ log_vars[k] = float(weigted_values_sum / batch_size_sum)
+ else:
+ for k, v in log_vars.items():
+ log_vars[k] = np.average(log_vars[k], weights=batch_size_list)
+
+ if viz_dir is not None:
+ if ws > 1:
+ gathered = [None for _ in range(ws)]
+ dist.all_gather_object(gathered, html_entries)
+ if rank == 0:
+ html_entries = [e for sub in gathered for e in (sub or [])]
+ if rank == 0:
+ unique_entries = dict()
+ for entry in html_entries:
+ data_id = entry[0]
+ if data_id not in unique_entries:
+ unique_entries[data_id] = entry
+ html_entries = list(unique_entries.values())
+ html_entries.sort(key=lambda item: item[0])
+ html_path = file_client.join_path(os.path.dirname(viz_dir), os.path.basename(viz_dir) + '.html')
+ write_html(html_path, html_entries, file_client)
+ executor.shutdown(wait=True)
+
+ return log_vars
+
+
+@HOOKS.register_module(force=True)
+class GenerativeEvalHook(_GenerativeEvalHook):
+ greater_keys = ['acc', 'top', 'AR@', 'auc', 'precision', 'mAP', 'is', 'test_ssim', 'test_psnr']
+ less_keys = ['loss', 'fid', 'kid', 'test_lpips']
+ _supported_best_metrics = ['fid', 'kid', 'is', 'test_ssim', 'test_psnr', 'test_lpips']
+
+ def __init__(self,
+ *args,
+ data='',
+ viz_dir=None,
+ feed_batch_size=32,
+ viz_num=None,
+ clear_reals=False,
+ prefix='',
+ metric_cpu_offload=False,
+ **kwargs):
+ super(GenerativeEvalHook, self).__init__(*args, **kwargs)
+ self.data = data
+ self.viz_dir = viz_dir
+ self.file_client = FileClient.infer_client(
+ uri=viz_dir) if viz_dir is not None else None
+ self.feed_batch_size = feed_batch_size
+ self.viz_num = viz_num
+ self.clear_reals = clear_reals
+ self.prefix = prefix
+ self.metric_cpu_offload = metric_cpu_offload
+
+ @torch.no_grad()
+ def after_train_iter(self, runner):
+ with gc_context(enable=True):
+ interval = self.get_current_interval(runner)
+ if not self.every_n_iters(runner, interval):
+ return
+
+ runner.model.eval()
+ rank, ws = get_dist_info()
+
+ if self.viz_dir is not None:
+ viz_dir = self.file_client.join_path(self.viz_dir, str(runner.iter + 1))
+ if rank == 0:
+ if self.file_client.exists(viz_dir):
+ for name in self.file_client.list_dir_or_file(viz_dir):
+ self.file_client.remove(self.file_client.join_path(viz_dir, name))
+ if ws > 1:
+ dist.barrier()
+ else:
+ viz_dir = None
+ log_vars = evaluate(
+ runner.model, self.dataloader, self.metrics, self.feed_batch_size,
+ viz_dir, self.viz_num, self.sample_kwargs)
+
+ if len(runner.log_buffer.output) == 0:
+ runner.log_buffer.clear()
+
+ # a dirty walkround to change the line at the end of pbar
+ if rank == 0:
+ sys.stdout.write('\n')
+ for metric in self.metrics:
+ metric.summary()
+ for name, val in metric._result_dict.items():
+ prefix_name = self.prefix + '_' + name if len(self.prefix) > 0 else name
+ runner.log_buffer.output[self.data + '_' + prefix_name] = val
+ # record best metric and save the best ckpt
+ if self.save_best_ckpt and name in self.best_metric:
+ self._save_best_ckpt(runner, val, name)
+ for name, val in log_vars.items():
+ prefix_name = self.prefix + '_' + name if len(self.prefix) > 0 else name
+ # print(self.data + '_' + prefix_name + ' = {}'.format(val))
+ runner.log_buffer.output[self.data + '_' + prefix_name] = val
+ # record best metric and save the best ckpt
+ if self.save_best_ckpt and name in self.best_metric:
+ self._save_best_ckpt(runner, val, name)
+ runner.log_buffer.ready = True
+
+ runner.model.train()
+
+ for metric in self.metrics:
+ metric.clear(clear_reals=self.clear_reals)
+ if self.metric_cpu_offload:
+ if hasattr(metric, 'offload_to_cpu'):
+ metric.offload_to_cpu()
+
+ torch.cuda.empty_cache()
diff --git a/lakonlab/evaluation/hpsv3.py b/lakonlab/evaluation/hpsv3.py
new file mode 100644
index 0000000000000000000000000000000000000000..814aea66441085db5aa6645f3acd0e33ce5db340
--- /dev/null
+++ b/lakonlab/evaluation/hpsv3.py
@@ -0,0 +1,587 @@
+# Modified from https://github.com/MizzenAI/HPSv3
+
+import math
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torch.distributed as dist
+import mmcv
+
+from typing import List, Optional, Union
+from torch.distributed.fsdp import MixedPrecision, ShardingStrategy, FullyShardedDataParallel
+from torch.distributed.fsdp.wrap import ModuleWrapPolicy
+from accelerate import init_empty_weights
+from transformers import Qwen2VLForConditionalGeneration, AutoProcessor
+from transformers.image_processing_utils import BaseImageProcessor, BatchFeature
+from transformers.image_utils import (
+ OPENAI_CLIP_MEAN,
+ OPENAI_CLIP_STD,
+)
+from transformers.utils import TensorType
+from transformers.models.qwen2_vl.modeling_qwen2_vl import Qwen2VLVisionBlock, Qwen2VLDecoderLayer
+from mmcv.runner import get_dist_info
+from mmgen.core.registry import METRICS
+from mmgen.core.evaluation.metrics import Metric
+from lakonlab.runner.checkpoint import _load_checkpoint
+
+
+INSTRUCTION = """
+You are tasked with evaluating a generated image based on Visual Quality and Text Alignment and give a overall score to estimate the human preference. Please provide a rating from 0 to 10, with 0 being the worst and 10 being the best.
+
+**Visual Quality:**
+Evaluate the overall visual quality of the image. The following sub-dimensions should be considered:
+- **Reasonableness:** The image should not contain any significant biological or logical errors, such as abnormal body structures or nonsensical environmental setups.
+- **Clarity:** Evaluate the sharpness and visibility of the image. The image should be clear and easy to interpret, with no blurring or indistinct areas.
+- **Detail Richness:** Consider the level of detail in textures, materials, lighting, and other visual elements (e.g., hair, clothing, shadows).
+- **Aesthetic and Creativity:** Assess the artistic aspects of the image, including the color scheme, composition, atmosphere, depth of field, and the overall creative appeal. The scene should convey a sense of harmony and balance.
+- **Safety:** The image should not contain harmful or inappropriate content, such as political, violent, or adult material. If such content is present, the image quality and satisfaction score should be the lowest possible.
+
+**Text Alignment:**
+Assess how well the image matches the textual prompt across the following sub-dimensions:
+- **Subject Relevance** Evaluate how accurately the subject(s) in the image (e.g., person, animal, object) align with the textual description. The subject should match the description in terms of number, appearance, and behavior.
+- **Style Relevance:** If the prompt specifies a particular artistic or stylistic style, evaluate how well the image adheres to this style.
+- **Contextual Consistency**: Assess whether the background, setting, and surrounding elements in the image logically fit the scenario described in the prompt. The environment should support and enhance the subject without contradictions.
+- **Attribute Fidelity**: Check if specific attributes mentioned in the prompt (e.g., colors, clothing, accessories, expressions, actions) are faithfully represented in the image. Minor deviations may be acceptable, but critical attributes should be preserved.
+- **Semantic Coherence**: Evaluate whether the overall meaning and intent of the prompt are captured in the image. The generated content should not introduce elements that conflict with or distort the original description.
+Textual prompt - {text_prompt}
+
+
+"""
+
+prompt_with_special_token = """
+Please provide the overall ratings of this image: <|Reward|>
+
+END
+"""
+
+prompt_without_special_token = """
+Please provide the overall ratings of this image:
+"""
+
+
+def smart_resize(
+ height: int, width: int, factor: int = 28, min_pixels: int = 56 * 56, max_pixels: int = 14 * 14 * 4 * 1280):
+ """Rescales the image so that the following conditions are met:
+
+ 1. Both dimensions (height and width) are divisible by 'factor'.
+
+ 2. The total number of pixels is within the range ['min_pixels', 'max_pixels'].
+
+ 3. The aspect ratio of the image is maintained as closely as possible.
+
+ """
+ if height < factor or width < factor:
+ raise ValueError(f"height:{height} or width:{width} must be larger than factor:{factor}")
+ elif max(height, width) / min(height, width) > 200:
+ raise ValueError(
+ f"absolute aspect ratio must be smaller than 200, got {max(height, width) / min(height, width)}"
+ )
+ h_bar = round(height / factor) * factor
+ w_bar = round(width / factor) * factor
+ if h_bar * w_bar > max_pixels:
+ beta = math.sqrt((height * width) / max_pixels)
+ h_bar = math.floor(height / beta / factor) * factor
+ w_bar = math.floor(width / beta / factor) * factor
+ elif h_bar * w_bar < min_pixels:
+ beta = math.sqrt(min_pixels / (height * width))
+ h_bar = math.ceil(height * beta / factor) * factor
+ w_bar = math.ceil(width * beta / factor) * factor
+ return h_bar, w_bar
+
+
+class Qwen2VLImageProcessor(BaseImageProcessor):
+ model_input_names = ["pixel_values", "image_grid_thw", "pixel_values_videos", "video_grid_thw"]
+
+ def __init__(
+ self,
+ do_resize: bool = True,
+ do_normalize: bool = True,
+ image_mean: Optional[Union[float, List[float]]] = None,
+ image_std: Optional[Union[float, List[float]]] = None,
+ min_pixels: int = 256 * 28 * 28,
+ max_pixels: int = 256 * 28 * 28,
+ patch_size: int = 14,
+ temporal_patch_size: int = 2,
+ merge_size: int = 2,
+ **kwargs):
+ super().__init__(**kwargs)
+ self.do_resize = do_resize
+ self.do_normalize = do_normalize
+ self.image_mean = image_mean if image_mean is not None else OPENAI_CLIP_MEAN
+ self.image_std = image_std if image_std is not None else OPENAI_CLIP_STD
+ self.min_pixels = min_pixels
+ self.max_pixels = max_pixels
+ self.patch_size = patch_size
+ self.temporal_patch_size = temporal_patch_size
+ self.merge_size = merge_size
+
+ def _preprocess(
+ self,
+ images: torch.Tensor,
+ do_resize: bool = None,
+ do_normalize: bool = None,
+ image_mean: Optional[Union[float, List[float]]] = None,
+ image_std: Optional[Union[float, List[float]]] = None):
+ batch_size, channel, height, width = images.size()
+ if do_resize:
+ resized_height, resized_width = smart_resize(
+ height,
+ width,
+ factor=self.patch_size * self.merge_size,
+ min_pixels=self.min_pixels,
+ max_pixels=self.max_pixels,
+ )
+ images = F.interpolate(
+ images,
+ size=(resized_height, resized_width),
+ mode='bicubic',
+ align_corners=False,
+ antialias=True,
+ ).clamp(min=0, max=1)
+ else:
+ resized_height, resized_width = height, width
+
+ if do_normalize:
+ mean = torch.tensor(image_mean, device=images.device, dtype=images.dtype).view(-1, 1, 1)
+ std = torch.tensor(image_std, device=images.device, dtype=images.dtype).view(-1, 1, 1)
+ images = (images - mean) / std
+
+ patches = images.unsqueeze(1).expand(-1, self.temporal_patch_size, -1, -1, -1)
+
+ grid_t = 1
+ grid_h, grid_w = resized_height // self.patch_size, resized_width // self.patch_size
+
+ patches = patches.reshape(
+ batch_size * grid_t,
+ self.temporal_patch_size,
+ channel,
+ grid_h // self.merge_size,
+ self.merge_size,
+ self.patch_size,
+ grid_w // self.merge_size,
+ self.merge_size,
+ self.patch_size,
+ )
+ patches = patches.permute(0, 3, 6, 4, 7, 2, 1, 5, 8)
+ flatten_patches = patches.reshape(
+ batch_size * grid_t * grid_h * grid_w, channel * self.temporal_patch_size * self.patch_size * self.patch_size
+ )
+
+ return flatten_patches, np.array((grid_t, grid_h, grid_w)).reshape(1, 3).repeat(batch_size, axis=0)
+
+ def preprocess(
+ self,
+ images: torch.Tensor,
+ return_tensors: Optional[Union[str, TensorType]] = None):
+ pixel_values, vision_grid_thws = self._preprocess(
+ images,
+ do_resize=self.do_resize,
+ do_normalize=self.do_normalize,
+ image_mean=self.image_mean,
+ image_std=self.image_std,
+ )
+ data = {"pixel_values": pixel_values, "image_grid_thw": vision_grid_thws}
+ return BatchFeature(data=data, tensor_type=return_tensors)
+
+
+class Qwen2VLRewardModelBT(Qwen2VLForConditionalGeneration):
+
+ def __init__(
+ self,
+ config,
+ output_dim=4,
+ reward_token="last",
+ special_token_ids=None,
+ rm_head_type="default",
+ rm_head_kwargs=None,
+ ):
+ super().__init__(config)
+ # pdb.set_trace()
+ self.output_dim = output_dim
+ if rm_head_type == "default":
+ self.rm_head = nn.Linear(config.hidden_size, output_dim, bias=False)
+ elif rm_head_type == "ranknet":
+ if rm_head_kwargs is not None:
+ for layer in range(rm_head_kwargs.get("num_layers", 3)):
+ if layer == 0:
+ self.rm_head = nn.Sequential(
+ nn.Linear(config.hidden_size, rm_head_kwargs["hidden_size"]),
+ nn.ReLU(),
+ nn.Dropout(rm_head_kwargs.get("dropout", 0.1)),
+ )
+ elif layer < rm_head_kwargs.get("num_layers", 3) - 1:
+ self.rm_head.add_module(
+ f"layer_{layer}",
+ nn.Sequential(
+ nn.Linear(rm_head_kwargs["hidden_size"], rm_head_kwargs["hidden_size"]),
+ nn.ReLU(),
+ nn.Dropout(rm_head_kwargs.get("dropout", 0.1)),
+ ),
+ )
+ else:
+ self.rm_head.add_module(
+ f"output_layer",
+ nn.Linear(rm_head_kwargs["hidden_size"], output_dim, bias=rm_head_kwargs.get("bias", False)),
+ )
+
+ else:
+ self.rm_head = nn.Sequential(
+ nn.Linear(config.hidden_size, 1024),
+ nn.ReLU(),
+ nn.Dropout(0.05),
+ nn.Linear(1024, 16),
+ nn.ReLU(),
+ nn.Linear(16, output_dim),
+ )
+
+ self.rm_head.to(torch.float32)
+ self.reward_token = reward_token
+
+ self.special_token_ids = special_token_ids
+ if self.special_token_ids is not None:
+ self.reward_token = "special"
+
+ def forward(
+ self,
+ input_ids: torch.LongTensor = None,
+ attention_mask: Optional[torch.Tensor] = None,
+ position_ids: Optional[torch.LongTensor] = None,
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
+ inputs_embeds: Optional[torch.FloatTensor] = None,
+ labels: Optional[torch.LongTensor] = None,
+ use_cache: Optional[bool] = None,
+ output_attentions: Optional[bool] = None,
+ output_hidden_states: Optional[bool] = None,
+ return_dict: Optional[bool] = None,
+ pixel_values: Optional[torch.Tensor] = None,
+ pixel_values_videos: Optional[torch.FloatTensor] = None,
+ image_grid_thw: Optional[torch.LongTensor] = None,
+ video_grid_thw: Optional[torch.LongTensor] = None,
+ rope_deltas: Optional[torch.LongTensor] = None,
+ ):
+ # modified from the origin class Qwen2VLForConditionalGeneration
+ output_attentions = (
+ output_attentions
+ if output_attentions is not None
+ else self.config.output_attentions
+ )
+ output_hidden_states = (
+ output_hidden_states
+ if output_hidden_states is not None
+ else self.config.output_hidden_states
+ )
+ return_dict = (
+ return_dict if return_dict is not None else self.config.use_return_dict
+ )
+ # pdb.set_trace()
+ if inputs_embeds is None:
+ inputs_embeds = self.model.language_model.embed_tokens(input_ids)
+ if pixel_values is not None:
+ pixel_values = pixel_values.type(self.model.visual.get_dtype())
+ image_embeds = self.model.visual(pixel_values, grid_thw=image_grid_thw)
+ image_mask = (
+ (input_ids == self.config.image_token_id)
+ .unsqueeze(-1)
+ .expand_as(inputs_embeds)
+ )
+ image_embeds = image_embeds.to(
+ inputs_embeds.device, inputs_embeds.dtype
+ )
+ inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)
+
+ if pixel_values_videos is not None:
+ pixel_values_videos = pixel_values_videos.type(self.model.visual.get_dtype())
+ video_embeds = self.model.visual(pixel_values_videos, grid_thw=video_grid_thw)
+ video_mask = (
+ (input_ids == self.config.video_token_id)
+ .unsqueeze(-1)
+ .expand_as(inputs_embeds)
+ )
+ video_embeds = video_embeds.to(
+ inputs_embeds.device, inputs_embeds.dtype
+ )
+ inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds)
+
+ if attention_mask is not None:
+ attention_mask = attention_mask.to(inputs_embeds.device)
+
+ outputs = self.model.language_model(
+ input_ids=None,
+ position_ids=position_ids,
+ attention_mask=attention_mask,
+ past_key_values=past_key_values,
+ inputs_embeds=inputs_embeds,
+ use_cache=use_cache,
+ output_attentions=output_attentions,
+ output_hidden_states=output_hidden_states,
+ return_dict=return_dict,
+ )
+
+ hidden_states = outputs[0] # [B, L, D]
+ with torch.autocast(device_type='cuda', dtype=torch.float32):
+ logits = self.rm_head(hidden_states) # [B, L, N]
+
+ if input_ids is not None:
+ batch_size = input_ids.shape[0]
+ else:
+ batch_size = inputs_embeds.shape[0]
+
+ # get sequence length
+ if self.config.pad_token_id is None and batch_size != 1:
+ raise ValueError(
+ "Cannot handle batch sizes > 1 if no padding token is defined."
+ )
+ if self.config.pad_token_id is None:
+ sequence_lengths = -1
+ else:
+ if input_ids is not None:
+ # if no pad token found, use modulo instead of reverse indexing for ONNX compatibility
+ sequence_lengths = (
+ torch.eq(input_ids, self.config.pad_token_id).int().argmax(-1) - 1
+ )
+ sequence_lengths = sequence_lengths % input_ids.shape[-1]
+ sequence_lengths = sequence_lengths.to(logits.device)
+ else:
+ sequence_lengths = -1
+
+ # get the last token's logits
+ if self.reward_token == "last":
+ pooled_logits = logits[
+ torch.arange(batch_size, device=logits.device), sequence_lengths
+ ]
+ elif self.reward_token == "mean":
+ # get the mean of all valid tokens' logits
+ valid_lengths = torch.clamp(sequence_lengths, min=0, max=logits.size(1) - 1)
+ pooled_logits = torch.stack(
+ [logits[i, : valid_lengths[i]].mean(dim=0) for i in range(batch_size)]
+ )
+ elif self.reward_token == "special":
+ # special_token_ids = self.tokenizer.convert_tokens_to_ids(self.special_tokens)
+ # create a mask for special tokens
+ special_token_mask = torch.zeros_like(input_ids, dtype=torch.bool)
+ for special_token_id in self.special_token_ids:
+ special_token_mask = special_token_mask | (
+ input_ids == special_token_id
+ )
+ pooled_logits = logits[special_token_mask, ...]
+ pooled_logits = pooled_logits.view(
+ batch_size, 1, -1
+ ) # [B, 3, N] assert 3 attributes
+ pooled_logits = pooled_logits.view(batch_size, -1)
+
+ # pdb.set_trace()
+ else:
+ raise ValueError("Invalid reward_token")
+
+ return {"logits": pooled_logits}
+
+
+_hpsv3_cache = {}
+
+
+def load_hpsv3(device, dtype, use_fsdp=True):
+ # Create cache key from arguments
+ cache_key = f"{device}_{dtype}_{use_fsdp}"
+
+ # Check if model is already cached
+ if cache_key in _hpsv3_cache:
+ return _hpsv3_cache[cache_key]
+
+ processor = AutoProcessor.from_pretrained(
+ 'Qwen/Qwen2-VL-7B-Instruct', padding_side='right',
+ )
+ processor.image_processor = Qwen2VLImageProcessor()
+ special_tokens = ['<|Reward|>']
+ processor.tokenizer.add_special_tokens(
+ {'additional_special_tokens': special_tokens}
+ )
+ special_token_ids = processor.tokenizer.convert_tokens_to_ids(special_tokens)
+
+ with init_empty_weights():
+ config = Qwen2VLRewardModelBT.config_class.from_pretrained(
+ 'Qwen/Qwen2-VL-7B-Instruct',
+ )
+ model = Qwen2VLRewardModelBT(
+ config,
+ output_dim=2,
+ reward_token='special',
+ special_token_ids=special_token_ids,
+ rm_head_type='ranknet',
+ )
+ model.requires_grad_(False)
+
+ model.resize_token_embeddings(len(processor.tokenizer))
+
+ model.config.tokenizer_padding_side = processor.tokenizer.padding_side
+ model.config.pad_token_id = processor.tokenizer.pad_token_id
+
+ state_dict = _load_checkpoint(
+ 'huggingface://MizzenAI/HPSv3/HPSv3.safetensors', map_location='cpu'
+ )
+ new_state_dict = dict()
+ for k, v in state_dict.items(): # fix transformers version mismatch
+ if k.startswith('model.'):
+ new_k = 'model.language_model.' + k[len('model.'):]
+ elif k.startswith('visual.'):
+ new_k = 'model.visual.' + k[len('visual.'):]
+ else:
+ new_k = k
+ new_state_dict[new_k] = v
+ model.load_state_dict(new_state_dict, strict=True, assign=True)
+ model.rm_head.to(torch.float32)
+
+ if use_fsdp:
+ mmcv.print_log('Wrapping HPSv3 model with FSDP.')
+ ignored_states = []
+ for p in model.rm_head.parameters():
+ p.data = p.data.cuda()
+ ignored_states.append(p)
+ model = FullyShardedDataParallel(
+ model,
+ device_id=torch.cuda.current_device(),
+ use_orig_params=False,
+ mixed_precision=MixedPrecision(
+ param_dtype=dtype,
+ reduce_dtype=dtype,
+ buffer_dtype=dtype,
+ cast_root_forward_inputs=False),
+ sharding_strategy=ShardingStrategy.HYBRID_SHARD,
+ auto_wrap_policy=ModuleWrapPolicy([Qwen2VLVisionBlock, Qwen2VLDecoderLayer]),
+ ignored_states=ignored_states
+ )
+ else:
+ model.to(device)
+
+ result = model, processor
+ _hpsv3_cache[cache_key] = result
+ return result
+
+
+@METRICS.register_module()
+class HPSv3(Metric):
+ name = 'HPSv3'
+ requires_prompt = True
+
+ def __init__(self,
+ num_images=None,
+ use_fsdp=True):
+ super().__init__(num_images)
+ use_fsdp = use_fsdp and torch.cuda.is_available() and dist.is_initialized() and dist.get_world_size() > 0
+
+ self.use_fsdp = use_fsdp
+ self.dtype = torch.bfloat16
+ self.device = 'cuda' if use_fsdp else 'cpu'
+
+ self.model, self.processor = load_hpsv3(device=self.device, dtype=self.dtype, use_fsdp=use_fsdp)
+ self.model.eval()
+
+ def prepare(self):
+ self.scores = []
+
+ @torch.no_grad()
+ def feed_op(self, batch, mode):
+ imgs = batch['imgs']
+ prompts = batch['prompts']
+
+ imgs = (imgs.to(device=self.device, dtype=torch.float32) / 2 + 0.5).clamp(0, 1)
+
+ message_list = []
+ for text in prompts:
+ out_message = [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "image",
+ "min_pixels": self.processor.image_processor.min_pixels,
+ "max_pixels": self.processor.image_processor.max_pixels,
+ },
+ {
+ "type": "text",
+ "text": (
+ INSTRUCTION.format(text_prompt=text)
+ + prompt_with_special_token
+ ),
+ },
+ ],
+ }
+ ]
+ message_list.append(out_message)
+
+ batch = self.processor(
+ text=self.processor.apply_chat_template(message_list, tokenize=False, add_generation_prompt=True),
+ images=imgs,
+ padding=True,
+ return_tensors="pt",
+ videos_kwargs={"do_rescale": True})
+ batch = {k: v.to(self.device) for k, v in batch.items()}
+ rewards = self.model(
+ return_dict=True,
+ **batch
+ )["logits"][:, 0]
+
+ if dist.is_initialized():
+ ws = dist.get_world_size()
+ placeholder = [torch.empty_like(rewards) for _ in range(ws)]
+ dist.all_gather(placeholder, rewards)
+ rewards = torch.cat(placeholder, dim=0)
+
+ if (dist.is_initialized() and dist.get_rank() == 0) or not dist.is_initialized():
+ self.scores.append(rewards.float().cpu())
+
+ def feed(self, batch, mode):
+ if mode == 'reals':
+ return 0
+
+ if self.num_images is None:
+ self.feed_op(batch, mode)
+
+ else:
+ _, ws = get_dist_info()
+
+ if self.num_fake_feeded == self.num_fake_need:
+ return 0
+
+ if isinstance(batch, dict):
+ batch_size = len(list(batch.values())[0])
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = {k: v[:end] for k, v in batch.items()}
+ else:
+ batch_size = batch.shape[0]
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = batch[:end]
+
+ global_end = min(batch_size * ws,
+ self.num_fake_need - self.num_fake_feeded)
+ self.feed_op(batch_to_feed, mode)
+ self.num_fake_feeded += global_end
+ return end
+
+ @torch.no_grad()
+ def summary(self):
+ scores = torch.cat(self.scores, dim=0)
+ if self.num_images is not None:
+ assert scores.shape[0] >= self.num_images
+ scores = scores[:self.num_images]
+ mean_score = scores.mean().item()
+ self._result_dict = dict(hpsv3=mean_score)
+ self._result_str = f'HPSv3: {mean_score:.4f}'
+ return mean_score
+
+ def clear_fake_data(self):
+ self.scores = []
+ self.num_fake_feeded = 0
+
+ def clear(self, clear_reals=False):
+ self.clear_fake_data()
+
+ def load_to_gpu(self):
+ if torch.cuda.is_available() and not isinstance(self.model, FullyShardedDataParallel):
+ self.model.cuda()
+ self.device = 'cuda'
+
+ def offload_to_cpu(self):
+ if not isinstance(self.model, FullyShardedDataParallel):
+ self.model.cpu()
+ self.device = 'cpu'
diff --git a/lakonlab/evaluation/metrics.py b/lakonlab/evaluation/metrics.py
new file mode 100644
index 0000000000000000000000000000000000000000..675c68e08e331c62a873d24a4bc909702ccdae6d
--- /dev/null
+++ b/lakonlab/evaluation/metrics.py
@@ -0,0 +1,1329 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import os
+import sys
+import logging
+import pickle
+import warnings
+import numpy as np
+import torch
+import torch.distributed as dist
+import torch.nn.functional as F
+import mmcv
+import hashlib
+
+from copy import deepcopy
+from contextlib import contextmanager, redirect_stdout, nullcontext
+from scipy import linalg
+from scipy.stats import entropy
+from torchvision import models
+from mmcv.runner import get_dist_info, load_checkpoint
+from mmgen.utils import get_root_logger
+from mmgen.core.registry import METRICS
+from mmgen.core.evaluation.metrics import (
+ Metric, TERO_INCEPTION_URL, _load_inception_torch, MMGEN_CACHE_DIR)
+from mmgen.core.evaluation.metrics import FID as _FID
+from mmgen.core.evaluation.metrics import PR as _PR
+from open_clip import get_tokenizer, create_model
+from lakonlab.utils.io_utils import download_from_huggingface, download_from_url
+
+
+# Global caches for model loading
+_inception_cache = {}
+_hpsv2_cache = {}
+_clip_cache = {}
+
+
+def _argv_ctx(argv):
+ class _Argv:
+
+ def __enter__(self):
+ self._old = sys.argv
+ sys.argv = argv
+
+ def __exit__(self, exc_type, exc, tb):
+ sys.argv = self._old
+
+ return _Argv()
+
+
+def _redirect_stdout(to_buf):
+ return redirect_stdout(to_buf) if to_buf is not None else nullcontext()
+
+
+@contextmanager
+def _quarantine_openclip_logging():
+ """
+ Guard against open_clip (and friends) mutating global logging.
+ Snapshots root handlers/level, runs the block, then removes any
+ NEW handlers and restores the level. Also disables propagation
+ for the open_clip logger so logs don’t bubble to root.
+ """
+ root = logging.getLogger()
+ before_handlers = tuple(root.handlers) # snapshot by identity
+ before_ids = {id(h) for h in before_handlers}
+ before_level = root.level
+
+ try:
+ yield
+ finally:
+ # Remove only handlers that were added during the block
+ for h in list(root.handlers):
+ if id(h) not in before_ids:
+ root.removeHandler(h)
+ try:
+ h.close()
+ except Exception:
+ pass
+ root.setLevel(before_level)
+
+ # Clamp open_clip logger so it won’t re-emit to root
+ oc = logging.getLogger("open_clip")
+ oc.propagate = False
+ oc.handlers.clear()
+
+
+def _load_inception_from_path(inception_path, map_location=None):
+ mmcv.print_log(
+ 'Try to load Tero\'s Inception Model from '
+ f'\'{inception_path}\'.', 'mmgen')
+ try:
+ model = torch.jit.load(inception_path, map_location=map_location)
+ mmcv.print_log('Load Tero\'s Inception Model successfully.', 'mmgen')
+ except Exception as e:
+ model = None
+ mmcv.print_log(
+ 'Load Tero\'s Inception Model failed. '
+ f'\'{e}\' occurs.', 'mmgen')
+ return model
+
+
+def _load_inception_from_url(inception_url, map_location=None):
+ """
+ Fix multi-node downloading issue in MMGen.
+ """
+ inception_url = inception_url if inception_url else TERO_INCEPTION_URL
+ mmcv.print_log(f'Try to download Inception Model from {inception_url}...',
+ 'mmgen')
+ try:
+ path = download_from_url(inception_url, dest_dir=MMGEN_CACHE_DIR)
+ mmcv.print_log('Download Finished.')
+ return _load_inception_from_path(path, map_location=map_location)
+ except Exception as e:
+ mmcv.print_log(f'Download Failed. {e} occurs.')
+ return None
+
+
+def load_inception(inception_args, metric, map_location=None):
+ """
+ Fix multi-node downloading issue in MMGen.
+ """
+ if not isinstance(inception_args, dict):
+ raise TypeError('Receive invalid \'inception_args\': '
+ f'\'{inception_args}\'')
+
+ # Create cache key from arguments
+ cache_key = hashlib.md5(str(sorted(inception_args.items())).encode()).hexdigest()
+ cache_key += f"_{metric}"
+
+ # Check if model is already cached
+ if cache_key in _inception_cache:
+ return _inception_cache[cache_key]
+
+ _inception_args = deepcopy(inception_args)
+ inceptoin_type = _inception_args.pop('type', None)
+
+ if torch.__version__ < '1.6.0':
+ mmcv.print_log(
+ 'Current Pytorch Version not support script module, load '
+ 'Inception Model from torch model zoo. If you want to use '
+ 'Tero\' script model, please update your Pytorch higher '
+ f'than \'1.6\' (now is {torch.__version__})', 'mmgen')
+ result = _load_inception_torch(_inception_args, metric), 'pytorch'
+ _inception_cache[cache_key] = result
+ return result
+
+ # load pytorch version is specific
+ if inceptoin_type != 'StyleGAN':
+ result = _load_inception_torch(_inception_args, metric), 'pytorch'
+ _inception_cache[cache_key] = result
+ return result
+
+ # try to load Tero's version
+ path = _inception_args.get('inception_path', TERO_INCEPTION_URL)
+
+ # try to parse `path` as web url and download
+ if 'http' not in path:
+ model = _load_inception_from_path(path, map_location=map_location)
+ if isinstance(model, torch.nn.Module):
+ result = model, 'StyleGAN'
+ _inception_cache[cache_key] = result
+ return result
+
+ # try to parse `path` as path on disk
+ model = _load_inception_from_url(path, map_location=map_location)
+ if isinstance(model, torch.nn.Module):
+ result = model, 'StyleGAN'
+ _inception_cache[cache_key] = result
+ return result
+
+ raise RuntimeError('Cannot Load Inception Model, please check the input '
+ f'`inception_args`: {inception_args}')
+
+
+def load_hpsv2(hps_version, device='cpu', precision='fp16'):
+ assert hps_version in ['v2', 'v2.1']
+
+ # Create cache key from arguments
+ cache_key = f"{hps_version}_{device}_{precision}"
+
+ # Check if model is already cached
+ if cache_key in _hpsv2_cache:
+ return _hpsv2_cache[cache_key]
+
+ with _quarantine_openclip_logging():
+ model = create_model(
+ 'ViT-H-14-quickgelu',
+ precision=precision,
+ device=device,
+ output_dict=True)
+ model.requires_grad_(False)
+ tokenizer = get_tokenizer('ViT-H-14')
+ load_checkpoint(
+ model,
+ f'huggingface://xswu/HPSv2/HPS_{hps_version}_compressed.pt',
+ map_location='cpu', strict=True)
+
+ result = model, tokenizer
+ _hpsv2_cache[cache_key] = result
+ return result
+
+
+def load_openclip(
+ model_name='ViT-L-14-336-quickgelu',
+ pretrained='openai',
+ device='cpu',
+ precision='fp16'):
+ cache_key = f'{model_name}_{pretrained}_{device}_{precision}'
+ if cache_key in _clip_cache:
+ return _clip_cache[cache_key]
+
+ with _quarantine_openclip_logging():
+ model = create_model(
+ model_name,
+ pretrained=pretrained,
+ precision=precision,
+ device=device,
+ output_dict=True)
+ model.requires_grad_(False)
+ tokenizer = get_tokenizer(model_name)
+ _clip_cache[cache_key] = (model, tokenizer)
+ return _clip_cache[cache_key]
+
+
+def compute_pr_distances(row_features,
+ col_features,
+ col_batch_size=10000):
+ dist_batches = []
+ for col_batch in col_features.split(col_batch_size):
+ dist_batch = torch.cdist(
+ row_features.unsqueeze(0), col_batch.unsqueeze(0))[0]
+ dist_batches.append(dist_batch.cpu())
+ return torch.cat(dist_batches, dim=1)
+
+
+@METRICS.register_module(force=True)
+class PR(_PR):
+
+ def __init__(
+ self,
+ num_images=None,
+ image_shape=None,
+ feats_pkl=None,
+ k=3,
+ bgr2rgb=True,
+ vgg16_script=None,
+ inception_args=None,
+ row_batch_size=10000,
+ col_batch_size=10000):
+ super(_PR, self).__init__(num_images, image_shape)
+
+ self.feats_pkl = feats_pkl
+
+ self.vgg16 = self.inception_net = None
+ self.device = 'cpu'
+
+ if vgg16_script is not None:
+ mmcv.print_log('loading vgg16 for improved precision and recall...',
+ 'mmgen')
+ if os.path.isfile(vgg16_script):
+ self.vgg16 = torch.jit.load('work_dirs/cache/vgg16.pt', map_location=self.device).eval()
+ self.use_tero_scirpt = True
+ else:
+ mmcv.print_log(
+ 'Cannot load Tero\'s script module. Use official '
+ 'vgg16 instead', 'mmgen')
+ self.vgg16 = models.vgg16(pretrained=True).eval()
+ self.use_tero_scirpt = False
+ elif inception_args is not None:
+ self.inception_net, self.inception_style = load_inception(
+ inception_args, 'FID')
+ else:
+ raise ValueError('Please provide either vgg16_script or inception_args')
+
+ self.k = k
+ self.bgr2rgb = bgr2rgb
+ self.row_batch_size = row_batch_size
+ self.col_batch_size = col_batch_size
+
+ def prepare(self):
+ self.features_of_reals = []
+ self.features_of_fakes = []
+ if self.feats_pkl is not None:
+ assert mmcv.is_filepath(self.feats_pkl)
+ with open(self.feats_pkl, 'rb') as f:
+ reference = pickle.load(f)
+ self.features_of_reals = [torch.from_numpy(feat) for feat in reference['features_of_reals']]
+ self.num_real_feeded = reference['num_real_feeded']
+ mmcv.print_log(
+ f'Load reference inception pkl from {self.feats_pkl}',
+ 'mmgen')
+
+ def extract_features(self, batch):
+ if self.vgg16 is not None:
+ if self.use_tero_scirpt:
+ batch = (batch * 127.5 + 128).clamp(0, 255).to(torch.uint8)
+ feat = self.vgg16(batch, return_features=True)
+ else:
+ batch = F.interpolate(batch, size=(224, 224))
+ before_fc = self.vgg16.features(batch)
+ before_fc = before_fc.view(-1, 7 * 7 * 512)
+ feat = self.vgg16.classifier[:4](before_fc)
+ else:
+ if self.inception_style == 'StyleGAN':
+ batch = (batch * 127.5 + 128).clamp(0, 255).to(torch.uint8)
+ feat = self.inception_net(batch, return_features=True)
+ else:
+ feat = self.inception_net(batch)[0].view(batch.shape[0], -1)
+ return feat
+
+ @torch.no_grad()
+ def feed_op(self, batch, mode):
+ batch = batch.to(self.device)
+ if self.bgr2rgb:
+ batch = batch[:, [2, 1, 0]]
+
+ feat = self.extract_features(batch)
+
+ if dist.is_initialized():
+ ws = dist.get_world_size()
+ placeholder = [torch.zeros_like(feat) for _ in range(ws)]
+ dist.all_gather(placeholder, feat)
+ feat = torch.stack(placeholder, dim=1).reshape(feat.size(0) * ws, *feat.shape[1:])
+
+ if (dist.is_initialized() and dist.get_rank() == 0) or not dist.is_initialized():
+ if mode == 'reals':
+ self.features_of_reals.append(feat)
+ elif mode == 'fakes':
+ self.features_of_fakes.append(feat)
+ else:
+ raise ValueError(f'{mode} is not a implemented feed mode.')
+
+ def feed(self, batch, mode):
+ if self.num_images is not None:
+ return super().feed(batch, mode)
+ else:
+ self.feed_op(batch, mode)
+
+ @torch.no_grad()
+ def summary(self):
+ gen_features = torch.cat(self.features_of_fakes)
+ real_features = torch.cat(self.features_of_reals).to(device=gen_features.device)
+ if self.num_images is not None:
+ assert gen_features.shape[0] >= self.num_images
+ gen_features = gen_features[:self.num_images]
+ if self.feats_pkl is None: # real feats not pre-calculated
+ assert real_features.shape[0] >= self.num_images
+ real_features = real_features[:self.num_images]
+
+ self._result_dict = {}
+
+ for name, manifold, probes in [
+ ('precision', real_features, gen_features),
+ ('recall', gen_features, real_features)
+ ]:
+ kth = []
+ for manifold_batch in manifold.split(self.row_batch_size):
+ distance = compute_pr_distances(
+ row_features=manifold_batch,
+ col_features=manifold,
+ col_batch_size=self.col_batch_size)
+ kth.append(
+ distance.to(torch.float32).kthvalue(self.k + 1).values.to(torch.float16))
+ kth = torch.cat(kth)
+ pred = []
+ for probes_batch in probes.split(self.row_batch_size):
+ distance = compute_pr_distances(
+ row_features=probes_batch,
+ col_features=manifold,
+ col_batch_size=self.col_batch_size)
+ pred.append((distance <= kth).any(dim=1))
+ self._result_dict[name] = float(torch.cat(pred).to(torch.float32).mean())
+
+ precision = self._result_dict['precision']
+ recall = self._result_dict['recall']
+ self._result_str = f'precision: {precision}, recall:{recall}'
+ return self._result_dict
+
+ def clear_fake_data(self):
+ self.features_of_fakes = []
+ self.num_fake_feeded = 0
+
+ def clear(self, clear_reals=False):
+ self.clear_fake_data()
+ if clear_reals:
+ self.features_of_reals = []
+ self.num_real_feeded = 0
+
+ def load_to_gpu(self):
+ """Move models to GPU."""
+ if torch.cuda.is_available():
+ if self.vgg16 is not None:
+ self.vgg16 = self.vgg16.cuda()
+ elif self.inception_net is not None:
+ self.inception_net.cuda()
+ self.device = 'cuda'
+
+ def offload_to_cpu(self):
+ """Move models to CPU."""
+ if self.vgg16 is not None:
+ self.vgg16 = self.vgg16.cpu()
+ elif self.inception_net is not None:
+ self.inception_net.cpu()
+ self.device = 'cpu'
+
+
+@METRICS.register_module(force=True)
+class FID(_FID):
+
+ def __init__(self,
+ num_images=None,
+ image_shape=None,
+ inception_pkl=None,
+ bgr2rgb=True,
+ inception_args=dict(normalize_input=False)):
+ super().__init__(
+ num_images,
+ image_shape=image_shape,
+ inception_pkl=inception_pkl,
+ bgr2rgb=bgr2rgb,
+ inception_args=inception_args)
+
+ def prepare(self):
+ if self.inception_pkl is not None:
+ assert mmcv.is_filepath(self.inception_pkl)
+ if self.inception_pkl.startswith('huggingface://'):
+ self.inception_pkl = download_from_huggingface(self.inception_pkl)
+ elif self.inception_pkl.startswith(('http://', 'https://')):
+ self.inception_pkl = download_from_url(self.inception_pkl)
+ with open(self.inception_pkl, 'rb') as f:
+ reference = pickle.load(f)
+ self.real_mean = reference['mean']
+ self.real_cov = reference['cov']
+ mmcv.print_log(
+ f'Load reference inception pkl from {self.inception_pkl}',
+ 'mmgen')
+ self.num_real_feeded = self.num_images
+
+ @torch.no_grad()
+ def summary(self):
+ # calculate reference inception stat
+ if self.real_mean is None:
+ feats = torch.cat(self.real_feats, dim=0)
+ if self.num_images is not None:
+ assert feats.shape[0] >= self.num_images
+ feats = feats[:self.num_images]
+ feats_np = feats.numpy()
+ self.real_mean = np.mean(feats_np, 0)
+ self.real_cov = np.cov(feats_np, rowvar=False)
+
+ # calculate fake inception stat
+ fake_feats = torch.cat(self.fake_feats, dim=0)
+ if self.num_images is not None:
+ assert fake_feats.shape[0] >= self.num_images
+ fake_feats = fake_feats[:self.num_images]
+ fake_feats_np = fake_feats.numpy()
+ fake_mean = np.mean(fake_feats_np, 0)
+ fake_cov = np.cov(fake_feats_np, rowvar=False)
+
+ # calculate distance between real and fake statistics
+ fid, mean, cov = self._calc_fid(fake_mean, fake_cov, self.real_mean, self.real_cov)
+
+ # results for print/table
+ self._result_str = (f'{fid:.4f} ({mean:.5f}/{cov:.5f})')
+ # results for log_buffer
+ self._result_dict = dict(fid=fid, fid_mean=mean, fid_cov=cov)
+
+ return fid, mean, cov
+
+ def feed(self, batch, mode):
+ if self.num_images is not None:
+ return super().feed(batch, mode)
+ else:
+ self.feed_op(batch, mode)
+
+
+@METRICS.register_module()
+class FIDKID(FID):
+ name = 'FIDKID'
+
+ def __init__(self,
+ num_images=None,
+ num_subsets=100,
+ max_subset_size=1000,
+ **kwargs):
+ super().__init__(num_images=num_images, **kwargs)
+ self.num_subsets = num_subsets
+ self.max_subset_size = max_subset_size
+ self.real_feats_np = None
+
+ def prepare(self):
+ if self.inception_pkl is not None:
+ assert mmcv.is_filepath(self.inception_pkl)
+ with open(self.inception_pkl, 'rb') as f:
+ reference = pickle.load(f)
+ self.real_mean = reference['mean']
+ self.real_cov = reference['cov']
+ self.real_feats_np = reference['feats_np']
+ mmcv.print_log(
+ f'Load reference inception pkl from {self.inception_pkl}',
+ 'mmgen')
+ self.num_real_feeded = self.num_images
+
+ @staticmethod
+ def _calc_kid(real_feat, fake_feat, num_subsets, max_subset_size):
+ """Refer to the implementation from:
+ https://github.com/NVlabs/stylegan2-ada-pytorch/blob/main/metrics/kernel_inception_distance.py#L18 # noqa
+ Args:
+ real_feat (np.array): Features of the real samples.
+ fake_feat (np.array): Features of the fake samples.
+ num_subsets (int): Number of subsets to calculate KID.
+ max_subset_size (int): The max size of each subset.
+ Returns:
+ float: The calculated kid metric.
+ """
+ n = real_feat.shape[1]
+ m = min(min(real_feat.shape[0], fake_feat.shape[0]), max_subset_size)
+ t = 0
+ for _ in range(num_subsets):
+ x = fake_feat[np.random.choice(
+ fake_feat.shape[0], m, replace=False)]
+ y = real_feat[np.random.choice(
+ real_feat.shape[0], m, replace=False)]
+ a = (x @ x.T / n + 1)**3 + (y @ y.T / n + 1)**3
+ b = (x @ y.T / n + 1)**3
+ t += (a.sum() - np.diag(a).sum()) / (m - 1) - b.sum() * 2 / m
+
+ kid = t / num_subsets / m
+ return float(kid)
+
+ @torch.no_grad()
+ def summary(self):
+ if self.real_feats_np is None:
+ feats = torch.cat(self.real_feats, dim=0)
+ if self.num_images is not None:
+ assert feats.shape[0] >= self.num_images
+ feats = feats[:self.num_images]
+ feats_np = feats.numpy()
+ self.real_feats_np = feats_np
+ self.real_mean = np.mean(feats_np, 0)
+ self.real_cov = np.cov(feats_np, rowvar=False)
+
+ fake_feats = torch.cat(self.fake_feats, dim=0)
+ if self.num_images is not None:
+ assert fake_feats.shape[0] >= self.num_images
+ fake_feats = fake_feats[:self.num_images]
+ fake_feats_np = fake_feats.numpy()
+ fake_mean = np.mean(fake_feats_np, 0)
+ fake_cov = np.cov(fake_feats_np, rowvar=False)
+
+ fid, mean, cov = self._calc_fid(fake_mean, fake_cov, self.real_mean,
+ self.real_cov)
+ kid = self._calc_kid(self.real_feats_np, fake_feats_np, self.num_subsets,
+ self.max_subset_size) * 1000
+
+ self._result_str = f'{fid:.4f} ({mean:.5f}/{cov:.5f}), {kid:.4f}'
+ self._result_dict = dict(fid=fid, fid_mean=mean, fid_cov=cov, kid=kid)
+
+ return fid, mean, cov, kid
+
+
+@METRICS.register_module()
+class InceptionMetrics(Metric):
+ name = 'InceptionMetrics'
+
+ def __init__(self,
+ num_images=None,
+ reference_pkl=None,
+ bgr2rgb=False,
+ center_crop=False, # SDXL-Lightning patch FID
+ resize=True,
+ inception_args=dict(
+ type='StyleGAN',
+ inception_path=TERO_INCEPTION_URL),
+ use_kid=False,
+ use_pr=True,
+ use_is=True,
+ kid_num_subsets=100,
+ kid_max_subset_size=1000,
+ pr_k=3,
+ pr_row_batch_size=10000,
+ pr_col_batch_size=10000,
+ is_splits=10,
+ prefix=''):
+ super().__init__(num_images)
+ self.reference_pkl = reference_pkl
+ self.real_feats = []
+ self.fake_feats = []
+ self.preds = []
+ self.real_mean = None
+ self.real_cov = None
+ self.bgr2rgb = bgr2rgb
+ self.center_crop = center_crop
+ self.resize = resize
+ self.device = 'cpu'
+
+ if self.center_crop and self.resize:
+ warnings.warn('`center_crop` is set to True, `resize` will be ignored.')
+
+ logger = get_root_logger()
+ ori_level = logger.level
+ logger.setLevel('ERROR')
+ self.inception_net, self.inception_style = load_inception(
+ inception_args, 'FID', map_location=self.device)
+ logger.setLevel(ori_level)
+
+ self.inception_net.eval()
+
+ self.use_kid = use_kid
+ self.use_pr = use_pr
+ self.use_is = use_is
+ self.kid_num_subsets = kid_num_subsets
+ self.kid_max_subset_size = kid_max_subset_size
+ self.real_feats_np = None
+
+ self.pr_k = pr_k
+ self.pr_row_batch_size = pr_row_batch_size
+ self.pr_col_batch_size = pr_col_batch_size
+
+ self.is_splits = is_splits
+ self.prefix = prefix
+
+ def prepare(self):
+ self.real_feats = []
+ self.real_feats_np = None
+ self.fake_feats = []
+ self.preds = []
+ if self.reference_pkl is not None:
+ assert mmcv.is_filepath(self.reference_pkl)
+ if self.reference_pkl.startswith('huggingface://'):
+ self.reference_pkl = download_from_huggingface(self.reference_pkl)
+ elif self.reference_pkl.startswith(('http://', 'https://')):
+ self.reference_pkl = download_from_url(self.reference_pkl)
+ with open(self.reference_pkl, 'rb') as f:
+ reference = pickle.load(f)
+ self.real_mean = reference['mean']
+ self.real_cov = reference['cov']
+ self.real_feats_np = reference['real_feats_np']
+ self.real_feats = [torch.from_numpy(reference['real_feats_np'])]
+ self.num_real_feeded = reference['num_real_feeded']
+
+ @staticmethod
+ def _calc_fid(sample_mean, sample_cov, real_mean, real_cov, eps=1e-6):
+ """Refer to the implementation from:
+
+ https://github.com/rosinality/stylegan2-pytorch/blob/master/fid.py#L34
+ """
+ cov_sqrt, _ = linalg.sqrtm(sample_cov @ real_cov, disp=False)
+
+ if not np.isfinite(cov_sqrt).all():
+ print('product of cov matrices is singular')
+ offset = np.eye(sample_cov.shape[0]) * eps
+ cov_sqrt = linalg.sqrtm(
+ (sample_cov + offset) @ (real_cov + offset))
+
+ if np.iscomplexobj(cov_sqrt):
+ if not np.allclose(np.diagonal(cov_sqrt).imag, 0, atol=1e-3):
+ m = np.max(np.abs(cov_sqrt.imag))
+
+ raise ValueError(f'Imaginary component {m}')
+
+ cov_sqrt = cov_sqrt.real
+
+ mean_diff = sample_mean - real_mean
+ mean_norm = mean_diff @ mean_diff
+
+ trace = np.trace(sample_cov) + np.trace(
+ real_cov) - 2 * np.trace(cov_sqrt)
+
+ fid = mean_norm + trace
+
+ return fid, mean_norm, trace
+
+ @staticmethod
+ def _calc_kid(real_feat, fake_feat, num_subsets, max_subset_size):
+ """Refer to the implementation from:
+ https://github.com/NVlabs/stylegan2-ada-pytorch/blob/main/metrics/kernel_inception_distance.py#L18 # noqa
+ Args:
+ real_feat (np.array): Features of the real samples.
+ fake_feat (np.array): Features of the fake samples.
+ num_subsets (int): Number of subsets to calculate KID.
+ max_subset_size (int): The max size of each subset.
+ Returns:
+ float: The calculated kid metric.
+ """
+ n = real_feat.shape[1]
+ m = min(min(real_feat.shape[0], fake_feat.shape[0]), max_subset_size)
+ t = 0
+ for _ in range(num_subsets):
+ x = fake_feat[np.random.choice(
+ fake_feat.shape[0], m, replace=False)]
+ y = real_feat[np.random.choice(
+ real_feat.shape[0], m, replace=False)]
+ a = (x @ x.T / n + 1)**3 + (y @ y.T / n + 1)**3
+ b = (x @ y.T / n + 1)**3
+ t += (a.sum() - np.diag(a).sum()) / (m - 1) - b.sum() * 2 / m
+
+ kid = t / num_subsets / m
+ return float(kid)
+
+ def extract_features(self, batch):
+ if self.center_crop:
+ crop_size = 299
+ h, w = batch.shape[2], batch.shape[3]
+ assert h >= crop_size and w >= crop_size
+ h_offset = (h - crop_size) // 2
+ w_offset = (w - crop_size) // 2
+ batch = batch[:, :, h_offset:h_offset + crop_size, w_offset:w_offset + crop_size]
+ elif self.resize:
+ batch = F.interpolate(
+ batch, size=(299, 299), mode='bicubic', align_corners=False, antialias=True).clamp(min=-1, max=1)
+ assert self.inception_style == 'StyleGAN'
+ batch = (batch * 127.5 + 128).clamp(0, 255).to(torch.uint8)
+ feat = self.inception_net(batch, return_features=True)
+ pred = F.linear(feat, self.inception_net.output.weight).softmax(dim=1)
+ return feat, pred
+
+ @torch.no_grad()
+ def feed_op(self, batch, mode):
+ if self.bgr2rgb:
+ batch = batch[:, [2, 1, 0]]
+ batch = batch.to(self.device)
+
+ feat, pred = self.extract_features(batch)
+
+ if dist.is_initialized():
+ ws = dist.get_world_size()
+ placeholder = [torch.zeros_like(feat) for _ in range(ws)]
+ dist.all_gather(placeholder, feat)
+ feat = torch.stack(placeholder, dim=1).reshape(feat.size(0) * ws, *feat.shape[1:])
+ if mode == 'fakes':
+ placeholder = [torch.zeros_like(pred) for _ in range(ws)]
+ dist.all_gather(placeholder, pred)
+ pred = torch.stack(placeholder, dim=1).reshape(pred.size(0) * ws, *pred.shape[1:])
+
+ # in distributed training, we only collect features at rank-0.
+ if (dist.is_initialized() and dist.get_rank() == 0) or not dist.is_initialized():
+ if mode == 'reals':
+ self.real_feats.append(feat.cpu())
+ elif mode == 'fakes':
+ self.fake_feats.append(feat.cpu())
+ self.preds.append(pred.cpu().numpy())
+ else:
+ raise ValueError(
+ f"The expected mode should be set to 'reals' or 'fakes,\
+ but got '{mode}'")
+
+ def feed(self, batch, mode):
+ if self.num_images is None:
+ self.feed_op(batch, mode)
+
+ else:
+ _, ws = get_dist_info()
+ if mode == 'reals':
+ if self.num_real_feeded == self.num_real_need:
+ return 0
+
+ if isinstance(batch, dict):
+ batch_size = len(list(batch.values())[0])
+ end = min(batch_size, self.num_real_need - self.num_real_feeded)
+ batch_to_feed = {k: v[:end] for k, v in batch.items()}
+ else:
+ batch_size = batch.shape[0]
+ end = min(batch_size, self.num_real_need - self.num_real_feeded)
+ batch_to_feed = batch[:end]
+
+ global_end = min(batch_size * ws,
+ self.num_real_need - self.num_real_feeded)
+ self.feed_op(batch_to_feed, mode)
+ self.num_real_feeded += global_end
+ return end
+
+ elif mode == 'fakes':
+ if self.num_fake_feeded == self.num_fake_need:
+ return 0
+
+ if isinstance(batch, dict):
+ batch_size = len(list(batch.values())[0])
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = {k: v[:end] for k, v in batch.items()}
+ else:
+ batch_size = batch.shape[0]
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = batch[:end]
+
+ global_end = min(batch_size * ws,
+ self.num_fake_need - self.num_fake_feeded)
+ self.feed_op(batch_to_feed, mode)
+ self.num_fake_feeded += global_end
+ return end
+ else:
+ raise ValueError(
+ 'The expected mode should be set to \'reals\' or \'fakes\','
+ f'but got \'{mode}\'')
+
+ @torch.no_grad()
+ def summary(self):
+ real_feats = torch.cat(self.real_feats, dim=0)
+ fake_feats = torch.cat(self.fake_feats, dim=0)
+ if self.num_images is not None:
+ assert fake_feats.shape[0] >= self.num_images
+ fake_feats = fake_feats[:self.num_images]
+ if self.reference_pkl is None: # real feats not pre-calculated
+ assert real_feats.shape[0] >= self.num_images
+ real_feats = real_feats[:self.num_images]
+
+ if self.real_feats_np is None:
+ real_feats_np = real_feats.numpy()
+ self.real_feats_np = real_feats_np
+ self.real_mean = np.mean(real_feats_np, 0)
+ self.real_cov = np.cov(real_feats_np, rowvar=False)
+
+ self._result_dict = dict()
+
+ prefix = self.prefix + '_' if len(self.prefix) > 0 else ''
+
+ # FID
+ fake_feats_np = fake_feats.numpy()
+ fake_mean = np.mean(fake_feats_np, 0)
+ fake_cov = np.cov(fake_feats_np, rowvar=False)
+ fid, mean, cov = self._calc_fid(fake_mean, fake_cov, self.real_mean,
+ self.real_cov)
+ self._result_dict.update({f'{prefix}fid': fid})
+ _result_str = f'{prefix}FID: {fid:.4f} ({mean:.4f}/{cov:.4f})'
+
+ # KID
+ if self.use_kid:
+ kid = self._calc_kid(self.real_feats_np, fake_feats_np, self.kid_num_subsets,
+ self.kid_max_subset_size) * 1000
+ self._result_dict.update({f'{prefix}kid': kid})
+ _result_str += f', {prefix}KID: {kid:.4f}'
+ else:
+ kid = None
+
+ # PR
+ if self.use_pr:
+ for name, manifold, probes in [
+ (f'{prefix}precision', real_feats, fake_feats),
+ (f'{prefix}recall', fake_feats, real_feats)
+ ]:
+ kth = []
+ for manifold_batch in manifold.split(self.pr_row_batch_size):
+ distance = compute_pr_distances(
+ row_features=manifold_batch,
+ col_features=manifold,
+ col_batch_size=self.pr_col_batch_size)
+ kth.append(
+ distance.to(torch.float32).kthvalue(self.pr_k + 1).values.to(torch.float16))
+ kth = torch.cat(kth)
+ pred = []
+ for probes_batch in probes.split(self.pr_row_batch_size):
+ distance = compute_pr_distances(
+ row_features=probes_batch,
+ col_features=manifold,
+ col_batch_size=self.pr_col_batch_size)
+ pred.append((distance <= kth).any(dim=1))
+ self._result_dict[name] = float(torch.cat(pred).to(torch.float32).mean())
+ precision = self._result_dict[f'{prefix}precision']
+ recall = self._result_dict[f'{prefix}recall']
+ _result_str += f', {prefix}Precision: {precision:.5f}, {prefix}Recall:{recall:.5f}'
+ else:
+ precision = recall = None
+
+ # IS
+ if self.use_is:
+ split_scores = []
+ self.preds = np.concatenate(self.preds, axis=0)
+ if self.num_images is not None:
+ assert self.preds.shape[0] >= self.num_images
+ self.preds = self.preds[:self.num_images]
+ num_preds = self.preds.shape[0]
+ for k in range(self.is_splits):
+ part = self.preds[k * (num_preds // self.is_splits):(k + 1) * (num_preds // self.is_splits), :]
+ py = np.mean(part, axis=0)
+ scores = []
+ for i in range(part.shape[0]):
+ pyx = part[i, :]
+ scores.append(entropy(pyx, py))
+ split_scores.append(np.exp(np.mean(scores)))
+ is_mean = np.mean(split_scores)
+ self._result_dict.update({f'{prefix}is': is_mean})
+ _result_str += f', {prefix}IS: {is_mean:.2f}'
+ else:
+ is_mean = None
+
+ self._result_str = _result_str
+
+ return fid, kid, precision, recall, is_mean
+
+ def clear_fake_data(self):
+ self.fake_feats = []
+ self.preds = []
+ self.num_fake_feeded = 0
+
+ def clear(self, clear_reals=False):
+ self.clear_fake_data()
+ if clear_reals:
+ self.real_feats = []
+ self.real_feats_np = None
+ self.num_real_feeded = 0
+
+ def load_to_gpu(self):
+ """Move models to GPU."""
+ if torch.cuda.is_available():
+ self.inception_net.cuda()
+ self.device = 'cuda'
+
+ def offload_to_cpu(self):
+ """Move models to CPU."""
+ self.inception_net.cpu()
+ self.device = 'cpu'
+
+
+@METRICS.register_module()
+class ColorStats(Metric):
+ name = 'ColorStats'
+
+ def __init__(self,
+ num_images=None):
+ super().__init__(num_images)
+
+ def prepare(self):
+ self.stats = []
+
+ @staticmethod
+ def srgb_to_linear(c):
+ threshold = 0.04045
+ below = c <= threshold
+ out = torch.where(
+ below, c / 12.92, ((c + 0.055) / 1.055) ** 2.4)
+ return out
+
+ @staticmethod
+ def linear_to_srgb(c):
+ threshold = 0.0031308
+ below = c <= threshold
+ out = torch.where(
+ below, 12.92 * c, 1.055 * c ** (1.0 / 2.4) - 0.055)
+ return out
+
+ def rgb_to_grayscale_srgb(self, img_srgb):
+ img_lin = self.srgb_to_linear(img_srgb)
+ R_lin, G_lin, B_lin = img_lin.unbind(dim=1)
+ Y_lin = 0.2126 * R_lin + 0.7152 * G_lin + 0.0722 * B_lin
+ gray_srgb = self.linear_to_srgb(Y_lin)
+ return gray_srgb
+
+ @staticmethod
+ def srgb_to_hsv_saturation(img_srgb):
+ c_max = torch.amax(img_srgb, dim=1)
+ c_min = torch.amin(img_srgb, dim=1)
+ delta = c_max - c_min
+ sat = delta / c_max.clamp(min=1e-5)
+ return sat
+
+ def compute_stats(self, batch):
+ batch = (batch / 2 + 0.5).clamp(0, 1)
+ gray = self.rgb_to_grayscale_srgb(batch).flatten(1)
+ contrast, brightness = torch.std_mean(gray, dim=1)
+ saturation = self.srgb_to_hsv_saturation(batch).flatten(1).mean(dim=1)
+ return torch.stack([brightness, contrast, saturation], dim=-1)
+
+ @torch.no_grad()
+ def feed_op(self, batch, mode):
+ stats = self.compute_stats(batch)
+
+ if dist.is_initialized():
+ ws = dist.get_world_size()
+ placeholder = [torch.zeros_like(stats) for _ in range(ws)]
+ dist.all_gather(placeholder, stats)
+ stats = torch.stack(placeholder, dim=1).reshape(stats.size(0) * ws, *stats.shape[1:])
+
+ # in distributed training, we only collect features at rank-0.
+ if (dist.is_initialized() and dist.get_rank() == 0) or not dist.is_initialized():
+ self.stats.append(stats.cpu())
+
+ def feed(self, batch, mode):
+ if mode == 'reals':
+ return 0
+
+ if self.num_images is None:
+ self.feed_op(batch, mode)
+
+ else:
+ _, ws = get_dist_info()
+
+ if self.num_fake_feeded == self.num_fake_need:
+ return 0
+
+ if isinstance(batch, dict):
+ batch_size = len(list(batch.values())[0])
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = {k: v[:end] for k, v in batch.items()}
+ else:
+ batch_size = batch.shape[0]
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = batch[:end]
+
+ global_end = min(batch_size * ws,
+ self.num_fake_need - self.num_fake_feeded)
+ self.feed_op(batch_to_feed, mode)
+ self.num_fake_feeded += global_end
+ return end
+
+ @torch.no_grad()
+ def summary(self):
+ stats = torch.cat(self.stats, dim=0)
+ if self.num_images is not None:
+ assert stats.shape[0] >= self.num_images
+ stats = stats[:self.num_images]
+ stats = stats.mean(dim=0)
+ brightness, contrast, saturation = stats.tolist()
+ self._result_dict = dict(
+ brightness=brightness, contrast=contrast, saturation=saturation)
+ self._result_str = f'Brightness: {brightness:.4f}, Contrast: {contrast:.4f}, Saturation: {saturation:.4f}'
+ return brightness, contrast, saturation
+
+ def clear_fake_data(self):
+ self.stats = []
+ self.num_fake_feeded = 0
+
+ def clear(self, clear_reals=False):
+ self.clear_fake_data()
+
+
+@METRICS.register_module()
+class HPSv2(Metric):
+ name = 'HPSv2'
+ requires_prompt = True
+
+ def __init__(self,
+ num_images=None,
+ hps_version='v2.1'):
+ super().__init__(num_images)
+ self.hps_version = hps_version
+ self.device = 'cpu' # Initialize on CPU
+ self.dtype = torch.float16
+ self.model, self.tokenizer = load_hpsv2(hps_version, device=self.device, precision='fp16')
+ self.model.eval()
+ image_size = self.model.visual.image_size
+ if isinstance(image_size, tuple):
+ assert len(image_size) == 2 and image_size[0] == image_size[1]
+ image_size = image_size[0]
+ self.image_size = image_size
+ self.image_mean = torch.tensor(self.model.visual.image_mean, device=self.device).view(3, 1, 1)
+ self.image_std = torch.tensor(self.model.visual.image_std, device=self.device).view(3, 1, 1)
+
+ def prepare(self):
+ self.scores = []
+
+ def resize(self, imgs):
+ h, w = imgs.shape[2:]
+ scale = self.image_size / float(max(h, w))
+ if scale != 1.0:
+ h = int(round(h * scale))
+ w = int(round(w * scale))
+ imgs = F.interpolate(imgs, size=(h, w), mode='bicubic', align_corners=False, antialias=True).clamp(0, 1)
+ if h != w:
+ pad_h = self.image_size - h
+ pad_w = self.image_size - w
+ imgs = F.pad(
+ imgs, (pad_w // 2, pad_w - pad_w // 2, pad_h // 2, pad_h - pad_h // 2), mode='constant', value=0)
+ return imgs
+
+ @torch.no_grad()
+ def feed_op(self, batch, mode):
+ imgs = batch['imgs']
+ prompts = batch['prompts']
+
+ imgs = (imgs.to(device=self.device, dtype=torch.float32) / 2 + 0.5).clamp(0, 1)
+ imgs = ((self.resize(imgs) - self.image_mean) / self.image_std).to(dtype=self.dtype)
+ prompts = self.tokenizer(prompts).to(device=self.device)
+
+ outputs = self.model(imgs, prompts)
+ image_features, text_features = outputs['image_features'], outputs['text_features']
+ hps_scores = (image_features * text_features).sum(dim=-1) # (bs, )
+
+ if dist.is_initialized():
+ ws = dist.get_world_size()
+ placeholder = [torch.empty_like(hps_scores) for _ in range(ws)]
+ dist.all_gather(placeholder, hps_scores)
+ hps_scores = torch.stack(placeholder, dim=1).reshape(hps_scores.size(0) * ws)
+
+ if (dist.is_initialized() and dist.get_rank() == 0) or not dist.is_initialized():
+ self.scores.append(hps_scores.float().cpu())
+
+ def feed(self, batch, mode):
+ if mode == 'reals':
+ return 0
+
+ if self.num_images is None:
+ self.feed_op(batch, mode)
+
+ else:
+ _, ws = get_dist_info()
+
+ if self.num_fake_feeded == self.num_fake_need:
+ return 0
+
+ if isinstance(batch, dict):
+ batch_size = len(list(batch.values())[0])
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = {k: v[:end] for k, v in batch.items()}
+ else:
+ batch_size = batch.shape[0]
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = batch[:end]
+
+ global_end = min(batch_size * ws,
+ self.num_fake_need - self.num_fake_feeded)
+ self.feed_op(batch_to_feed, mode)
+ self.num_fake_feeded += global_end
+ return end
+
+ @torch.no_grad()
+ def summary(self):
+ scores = torch.cat(self.scores, dim=0)
+ if self.num_images is not None:
+ assert scores.shape[0] >= self.num_images
+ scores = scores[:self.num_images]
+ mean_score = scores.mean().item()
+ self._result_dict = dict(hpsv2=mean_score)
+ self._result_str = f'HPSv2: {mean_score:.4f}'
+ return mean_score
+
+ def clear_fake_data(self):
+ self.scores = []
+ self.num_fake_feeded = 0
+
+ def clear(self, clear_reals=False):
+ self.clear_fake_data()
+
+ def load_to_gpu(self):
+ if torch.cuda.is_available():
+ self.model.cuda()
+ self.image_mean = self.image_mean.cuda()
+ self.image_std = self.image_std.cuda()
+ self.device = 'cuda'
+
+ def offload_to_cpu(self):
+ self.model.cpu()
+ self.image_mean = self.image_mean.cpu()
+ self.image_std = self.image_std.cpu()
+ self.device = 'cpu'
+
+
+@METRICS.register_module()
+class CLIPSimilarity(Metric):
+ """
+ Average image–text CLIP cosine similarity (↑ better).
+ Preprocess emulates OpenAI CLIP for ViT-L/14@336:
+ - Resize so min(H, W) = 336 (bicubic, antialias), keep aspect ratio
+ - Center crop to 336x336
+ - Normalize with model.visual.image_mean/std
+ Expects batch = {'imgs': (B,3,H,W) in [-1,1], 'prompts': List[str]}
+ """
+ name = 'CLIPSimilarity'
+ requires_prompt = True
+
+ def __init__(
+ self,
+ num_images=None,
+ model_name='ViT-L-14-336-quickgelu',
+ pretrained='openai',
+ precision='fp16', # 'fp16' | 'fp32' | 'bf16'
+ ):
+ super().__init__(num_images)
+ self.model_name = model_name
+ self.pretrained = pretrained
+ self.precision = precision
+
+ self.device = 'cpu'
+ self.dtype = {
+ 'fp16': torch.float16,
+ 'bf16': torch.bfloat16,
+ 'fp32': torch.float32
+ }.get(precision, torch.float16)
+
+ self.model, self.tokenizer = load_openclip(
+ model_name=model_name,
+ pretrained=pretrained,
+ device=self.device,
+ precision=precision,
+ )
+ self.model.eval()
+
+ # OpenAI ViT-L/14@336 uses square 336 input
+ image_size = self.model.visual.image_size
+ if isinstance(image_size, tuple):
+ assert len(image_size) == 2 and image_size[0] == image_size[1]
+ image_size = image_size[0]
+ self.image_size = int(image_size) # 336
+
+ # Use the model's own stats for normalization
+ self.image_mean = torch.tensor(self.model.visual.image_mean, device=self.device).view(3, 1, 1)
+ self.image_std = torch.tensor(self.model.visual.image_std, device=self.device).view(3, 1, 1)
+
+ def prepare(self):
+ self.scores = []
+
+ def _resize_min_side_then_center_crop(self, imgs):
+ """
+ imgs: (B,3,H,W) in [0,1], float32, on self.device
+ 1) Resize so min(H,W) == self.image_size, preserve AR (bicubic, antialias)
+ 2) Center-crop to (self.image_size, self.image_size)
+ 3) Normalize with model mean/std
+ 4) Cast to self.dtype
+ """
+ _, _, H, W = imgs.shape
+ target = self.image_size
+
+ # Scale factor so that the shorter side becomes 'target'
+ short, long = (H, W) if H < W else (W, H)
+ if short == 0:
+ raise ValueError("Invalid image with zero dimension.")
+ scale = target / float(short)
+
+ new_h = max(1, int(round(H * scale)))
+ new_w = max(1, int(round(W * scale)))
+ if new_h != H or new_w != W:
+ imgs = F.interpolate(
+ imgs, size=(new_h, new_w),
+ mode='bicubic', align_corners=False, antialias=True
+ ).clamp(0, 1)
+
+ # Center crop to target x target
+ top = max(0, (new_h - target) // 2)
+ left = max(0, (new_w - target) // 2)
+ imgs = imgs[:, :, top:top + target, left:left + target]
+
+ imgs = (imgs - self.image_mean) / self.image_std
+ return imgs.to(dtype=self.dtype)
+
+ @torch.no_grad()
+ def feed_op(self, batch, mode):
+ if mode == 'reals':
+ return 0
+
+ imgs = batch['imgs']
+ prompts = batch['prompts']
+
+ # [-1,1] -> [0,1]
+ imgs = (imgs.to(device=self.device, dtype=torch.float32) / 2 + 0.5).clamp(0, 1)
+ imgs = self._resize_min_side_then_center_crop(imgs)
+
+ # Tokenize on device
+ text = self.tokenizer(prompts).to(device=self.device)
+
+ # Forward (create_model(..., output_dict=True)) => dict w/ features
+ out = self.model(imgs, text)
+ if isinstance(out, dict) and ('image_features' in out and 'text_features' in out):
+ img_feat = out['image_features']
+ txt_feat = out['text_features']
+ else:
+ img_feat = self.model.encode_image(imgs)
+ txt_feat = self.model.encode_text(text)
+
+ # Cosine similarity per pair
+ img_feat = F.normalize(img_feat, dim=-1)
+ txt_feat = F.normalize(txt_feat, dim=-1)
+ sim = (img_feat * txt_feat).sum(dim=-1).to(torch.float32) # (B,)
+
+ # DDP gather
+ if dist.is_initialized():
+ ws = dist.get_world_size()
+ bucket = [torch.empty_like(sim) for _ in range(ws)]
+ dist.all_gather(bucket, sim)
+ sim = torch.stack(bucket, dim=1).reshape(sim.size(0) * ws)
+
+ if (dist.is_initialized() and dist.get_rank() == 0) or not dist.is_initialized():
+ self.scores.append(sim.cpu())
+
+ def feed(self, batch, mode):
+ if mode == 'reals':
+ return 0
+
+ if self.num_images is None:
+ self.feed_op(batch, mode)
+
+ else:
+ _, ws = get_dist_info()
+
+ if self.num_fake_feeded == self.num_fake_need:
+ return 0
+
+ if isinstance(batch, dict):
+ batch_size = len(list(batch.values())[0])
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = {k: v[:end] for k, v in batch.items()}
+ else:
+ batch_size = batch.shape[0]
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = batch[:end]
+
+ global_end = min(batch_size * ws, self.num_fake_need - self.num_fake_feeded)
+ self.feed_op(batch_to_feed, mode)
+ self.num_fake_feeded += global_end
+ return end
+
+ @torch.no_grad()
+ def summary(self):
+ sims = torch.cat(self.scores, dim=0)
+ if self.num_images is not None:
+ assert sims.shape[0] >= self.num_images
+ sims = sims[:self.num_images]
+ mean_sim = sims.mean().item()
+
+ self._result_dict = dict(clipsim=mean_sim) # raw cosine in [-1,1]
+ self._result_str = f'CLIPSim: {mean_sim:.4f}'
+ return mean_sim
+
+ def clear_fake_data(self):
+ self.scores = []
+ self.num_fake_feeded = 0
+
+ def clear(self, clear_reals=False):
+ self.clear_fake_data()
+
+ def load_to_gpu(self):
+ if torch.cuda.is_available():
+ self.model.cuda()
+ self.image_mean = self.image_mean.cuda()
+ self.image_std = self.image_std.cuda()
+ self.device = 'cuda'
+
+ def offload_to_cpu(self):
+ self.model.cpu()
+ self.image_mean = self.image_mean.cpu()
+ self.image_std = self.image_std.cpu()
+ self.device = 'cpu'
diff --git a/lakonlab/evaluation/vqa_score.py b/lakonlab/evaluation/vqa_score.py
new file mode 100644
index 0000000000000000000000000000000000000000..02435ffbe72b570fb026b43ca4c7499fda81443c
--- /dev/null
+++ b/lakonlab/evaluation/vqa_score.py
@@ -0,0 +1,667 @@
+# Modified from https://github.com/linzhiqiu/t2v_metrics
+# Copyright 2023 Zhiqiu Lin
+
+import os
+import re
+import torch
+import torch.distributed as dist
+import torch.nn as nn
+import torch.nn.functional as F
+import mmcv
+
+from typing import List, Optional, Tuple, Union
+from dataclasses import dataclass, field
+from torch.distributed.fsdp import MixedPrecision, ShardingStrategy, FullyShardedDataParallel
+from torch.distributed.fsdp.wrap import ModuleWrapPolicy
+from transformers import (
+ AutoConfig, AutoTokenizer, AutoModelForSeq2SeqLM, T5Config, T5ForConditionalGeneration,
+ CLIPVisionModel, CLIPImageProcessor, CLIPVisionConfig)
+from transformers.models.t5.modeling_t5 import T5Block
+from transformers.modeling_outputs import Seq2SeqLMOutput
+from mmcv.runner import get_dist_info
+from mmgen.core.registry import METRICS
+from mmgen.core.evaluation.metrics import Metric
+
+IMAGE_TOKEN_INDEX = -200
+CONTEXT_LEN = 2048
+SYSTEM_MSG = "A chat between a curious user and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the user's questions."
+IGNORE_INDEX = -100
+DEFAULT_IMAGE_TOKEN = ""
+
+default_question_template = 'Does this figure show "{}"? Please answer yes or no.'
+default_answer_template = "Yes"
+
+
+def t5_tokenizer_image_token(prompt, tokenizer, image_token_index=IMAGE_TOKEN_INDEX, return_tensors=None):
+ prompt_chunks = [tokenizer(chunk).input_ids for chunk in prompt.split('')]
+
+ def insert_separator(X, sep):
+ return [ele for sublist in zip(X, [sep] * len(X)) for ele in sublist][:-1]
+
+ input_ids = []
+ # Since there's no bos_token_id, simply concatenate the tokenized prompt_chunks with the image_token_index
+ for x in insert_separator(prompt_chunks, [image_token_index]):
+ input_ids.extend(x)
+
+ if return_tensors is not None:
+ if return_tensors == 'pt':
+ return torch.tensor(input_ids, dtype=torch.long)
+ raise ValueError(f'Unsupported tensor type: {return_tensors}')
+ return input_ids
+
+
+def format_question(question, conversation_style='plain'):
+ if conversation_style == 't5_plain': # for 1st stage t5 model
+ question = DEFAULT_IMAGE_TOKEN + question
+ elif conversation_style == 't5_chat': # for 2nd stage t5 model
+ question = SYSTEM_MSG + " USER: " + DEFAULT_IMAGE_TOKEN + "\n" + question + " ASSISTANT: "
+ elif conversation_style == 't5_chat_no_system': # for 2nd stage t5 model
+ question = "USER: " + DEFAULT_IMAGE_TOKEN + "\n" + question + " ASSISTANT: "
+ elif conversation_style == 't5_chat_no_system_no_user': # for 2nd stage t5 model
+ question = "" + DEFAULT_IMAGE_TOKEN + "\n" + question + " : "
+ # elif conversation_style == 't5_chat_ood_system': # for 2nd stage t5 model
+ # question = SYSTEM_MSG + " HUMAN: " + DEFAULT_IMAGE_TOKEN + "\n" + question + " GPT: "
+ else:
+ raise NotImplementedError()
+ return question
+
+
+def format_answer(answer, conversation_style='plain'):
+ return answer
+
+
+class CLIPVisionTower(nn.Module):
+ def __init__(self, vision_tower, args, delay_load=False):
+ super().__init__()
+
+ self.is_loaded = False
+
+ self.vision_tower_name = vision_tower
+ self.select_layer = args.mm_vision_select_layer
+ self.select_feature = getattr(args, 'mm_vision_select_feature', 'patch')
+
+ if not delay_load:
+ self.load_model()
+ else:
+ self.cfg_only = CLIPVisionConfig.from_pretrained(self.vision_tower_name)
+
+ def load_model(self):
+ self.image_processor = CLIPImageProcessor.from_pretrained(self.vision_tower_name)
+ self.vision_tower = CLIPVisionModel.from_pretrained(self.vision_tower_name)
+ self.vision_tower.requires_grad_(False)
+
+ self.is_loaded = True
+
+ def feature_select(self, image_forward_outs):
+ image_features = image_forward_outs.hidden_states[self.select_layer]
+ if self.select_feature == 'patch':
+ image_features = image_features[:, 1:]
+ elif self.select_feature == 'cls_patch':
+ image_features = image_features
+ else:
+ raise ValueError(f'Unexpected select feature: {self.select_feature}')
+ return image_features
+
+ @torch.no_grad()
+ def forward(self, images):
+ if type(images) is list:
+ image_features = []
+ for image in images:
+ image_forward_out = self.vision_tower(image.to(device=self.device, dtype=self.dtype).unsqueeze(0),
+ output_hidden_states=True)
+ image_feature = self.feature_select(image_forward_out).to(image.dtype)
+ image_features.append(image_feature)
+ else:
+ image_forward_outs = self.vision_tower(images.to(device=self.device, dtype=self.dtype),
+ output_hidden_states=True)
+ image_features = self.feature_select(image_forward_outs).to(images.dtype)
+
+ return image_features
+
+ @property
+ def dummy_feature(self):
+ return torch.zeros(1, self.hidden_size, device=self.device, dtype=self.dtype)
+
+ @property
+ def dtype(self):
+ return self.vision_tower.dtype
+
+ @property
+ def device(self):
+ return self.vision_tower.device
+
+ @property
+ def config(self):
+ if self.is_loaded:
+ return self.vision_tower.config
+ else:
+ return self.cfg_only
+
+ @property
+ def hidden_size(self):
+ return self.config.hidden_size
+
+ @property
+ def num_patches(self):
+ return (self.config.image_size // self.config.patch_size) ** 2
+
+
+class IdentityMap(nn.Module):
+ def __init__(self):
+ super().__init__()
+
+ def forward(self, x, *args, **kwargs):
+ return x
+
+ @property
+ def config(self):
+ return {"mm_projector_type": 'identity'}
+
+
+def build_vision_tower(vision_tower_cfg, **kwargs):
+ vision_tower = getattr(vision_tower_cfg, 'mm_vision_tower', getattr(vision_tower_cfg, 'vision_tower', None))
+ is_absolute_path_exists = os.path.exists(vision_tower)
+ if is_absolute_path_exists or vision_tower.startswith("openai") or vision_tower.startswith("laion"):
+ return CLIPVisionTower(vision_tower, args=vision_tower_cfg, **kwargs)
+
+ raise ValueError(f'Unknown vision tower: {vision_tower}')
+
+
+def build_vision_projector(config, delay_load=False, **kwargs):
+ projector_type = getattr(config, 'mm_projector_type', 'linear')
+
+ if projector_type == 'linear':
+ return nn.Linear(config.mm_hidden_size, config.hidden_size)
+
+ mlp_gelu_match = re.match(r'^mlp(\d+)x_gelu$', projector_type)
+ if mlp_gelu_match:
+ mlp_depth = int(mlp_gelu_match.group(1))
+ modules = [nn.Linear(config.mm_hidden_size, config.hidden_size)]
+ for _ in range(1, mlp_depth):
+ modules.append(nn.GELU())
+ modules.append(nn.Linear(config.hidden_size, config.hidden_size))
+ return nn.Sequential(*modules)
+
+ if projector_type == 'identity':
+ return IdentityMap()
+
+ raise ValueError(f'Unknown projector type: {projector_type}')
+
+
+@dataclass
+class ModelArguments:
+ tune_mm_mlp_adapter: bool = field(default=False)
+ vision_tower: Optional[str] = field(default='openai/clip-vit-large-patch14-336')
+ mm_vision_select_layer: Optional[int] = field(default=-2) # default to the second last layer in llava1.5
+ pretrain_mm_mlp_adapter: Optional[str] = field(default=None)
+ mm_projector_type: Optional[str] = field(default='mlp2x_gelu')
+ mm_vision_select_feature: Optional[str] = field(default="patch")
+
+
+class CLIPT5Config(T5Config):
+ model_type = "clip_t5"
+
+
+class CLIPT5ForConditionalGeneration(T5ForConditionalGeneration):
+ # This class supports both T5 and FlanT5
+ config_class = CLIPT5Config
+
+ def __init__(self, config):
+ super(CLIPT5ForConditionalGeneration, self).__init__(config)
+ self.embed_tokens = self.encoder.embed_tokens
+ if hasattr(config, "mm_vision_tower"):
+ self.vision_tower = build_vision_tower(config, delay_load=False)
+ self.mm_projector = build_vision_projector(config)
+
+ def get_vision_tower(self):
+ vision_tower = getattr(self, 'vision_tower', None)
+ if type(vision_tower) is list:
+ vision_tower = vision_tower[0]
+ return vision_tower
+
+ def get_model(self):
+ return self # for compatibility with LlavaMetaForCausalLM
+
+ def prepare_inputs_labels_for_multimodal(
+ self, input_ids, attention_mask, decoder_attention_mask, past_key_values, labels, images
+ ):
+ # The labels are now separated from the input_ids.
+ vision_tower = self.get_vision_tower()
+ if vision_tower is None or images is None or input_ids.shape[1] == 1:
+ raise NotImplementedError()
+
+ if type(images) is list or images.ndim == 5:
+ concat_images = torch.cat([image for image in images], dim=0)
+ image_features = self.encode_images(concat_images)
+ split_sizes = [image.shape[0] for image in images]
+ image_features = torch.split(image_features, split_sizes, dim=0)
+ image_features = [x.flatten(0, 1) for x in image_features]
+ else:
+ image_features = self.encode_images(images)
+
+ new_input_embeds = []
+ cur_image_idx = 0
+ for _, cur_input_ids in enumerate(input_ids):
+ if (cur_input_ids == IMAGE_TOKEN_INDEX).sum() == 0:
+ # multimodal LLM, but the current sample is not multimodal
+ raise NotImplementedError()
+ image_token_indices = torch.where(cur_input_ids == IMAGE_TOKEN_INDEX)[0]
+ cur_new_input_embeds = []
+ while image_token_indices.numel() > 0:
+ cur_image_features = image_features[cur_image_idx]
+ image_token_start = image_token_indices[0]
+ cur_new_input_embeds.append(self.embed_tokens(cur_input_ids[:image_token_start]))
+ cur_new_input_embeds.append(cur_image_features)
+ cur_image_idx += 1
+ cur_input_ids = cur_input_ids[image_token_start + 1:]
+ image_token_indices = torch.where(cur_input_ids == IMAGE_TOKEN_INDEX)[0]
+ if cur_input_ids.numel() > 0:
+ cur_new_input_embeds.append(self.embed_tokens(cur_input_ids))
+ cur_new_input_embeds = [x.to(device=self.device) for x in cur_new_input_embeds]
+ cur_new_input_embeds = torch.cat(cur_new_input_embeds, dim=0)
+ new_input_embeds.append(cur_new_input_embeds)
+
+ if any(x.shape != new_input_embeds[0].shape for x in new_input_embeds):
+ max_len = max(x.shape[0] for x in new_input_embeds)
+
+ new_input_embeds_align = []
+ _input_embeds_lengths = []
+ for cur_new_embed in new_input_embeds:
+ _input_embeds_lengths.append(cur_new_embed.shape[0])
+ cur_new_embed = torch.cat((cur_new_embed,
+ torch.zeros((max_len - cur_new_embed.shape[0], cur_new_embed.shape[1]),
+ dtype=cur_new_embed.dtype, device=cur_new_embed.device)), dim=0)
+ new_input_embeds_align.append(cur_new_embed)
+ new_input_embeds = torch.stack(new_input_embeds_align, dim=0)
+
+ if attention_mask is not None:
+ new_attention_mask = []
+ for cur_attention_mask, _input_embeds_length in zip(attention_mask, _input_embeds_lengths):
+ new_attn_mask_pad_left = torch.full((_input_embeds_length - input_ids.shape[1],), True,
+ dtype=attention_mask.dtype, device=attention_mask.device)
+ new_attn_mask_pad_right = torch.full((new_input_embeds.shape[1] - _input_embeds_length,), False,
+ dtype=attention_mask.dtype, device=attention_mask.device)
+ cur_new_attention_mask = torch.cat(
+ (new_attn_mask_pad_left, cur_attention_mask, new_attn_mask_pad_right), dim=0)
+ new_attention_mask.append(cur_new_attention_mask)
+ attention_mask = torch.stack(new_attention_mask, dim=0)
+ assert attention_mask.shape == new_input_embeds.shape[:2]
+ else:
+ new_input_embeds = torch.stack(new_input_embeds, dim=0)
+
+ if attention_mask is not None:
+ new_attn_mask_pad_left = torch.full(
+ (attention_mask.shape[0], new_input_embeds.shape[1] - input_ids.shape[1]), True,
+ dtype=attention_mask.dtype, device=attention_mask.device)
+ attention_mask = torch.cat((new_attn_mask_pad_left, attention_mask), dim=1)
+ assert attention_mask.shape == new_input_embeds.shape[:2]
+
+ return None, attention_mask, decoder_attention_mask, past_key_values, new_input_embeds, labels
+
+ def encode_images(self, images):
+ image_features = self.get_vision_tower()(images)
+ image_features = self.mm_projector(image_features)
+ return image_features
+
+ def initialize_vision_modules(self, model_args, fsdp=None):
+ vision_tower = model_args.vision_tower
+ mm_vision_select_layer = model_args.mm_vision_select_layer
+ mm_vision_select_feature = model_args.mm_vision_select_feature
+ pretrain_mm_mlp_adapter = model_args.pretrain_mm_mlp_adapter
+
+ self.config.mm_vision_tower = vision_tower
+ self.config.pretrain_mm_mlp_adapter = pretrain_mm_mlp_adapter
+
+ if self.get_vision_tower() is None:
+ vision_tower = build_vision_tower(model_args)
+
+ if fsdp is not None and len(fsdp) > 0:
+ self.vision_tower = [vision_tower]
+ else:
+ self.vision_tower = vision_tower
+ else:
+ if fsdp is not None and len(fsdp) > 0:
+ vision_tower = self.vision_tower[0]
+ else:
+ vision_tower = self.vision_tower
+ if not vision_tower.is_loaded:
+ vision_tower.load_model()
+
+ self.config.use_mm_proj = True
+ self.config.mm_projector_type = getattr(model_args, 'mm_projector_type', 'mlp2x_gelu')
+ self.config.mm_hidden_size = vision_tower.hidden_size
+ self.config.mm_vision_select_layer = mm_vision_select_layer
+ self.config.mm_vision_select_feature = mm_vision_select_feature
+
+ if getattr(self, 'mm_projector', None) is None:
+ self.mm_projector = build_vision_projector(self.config)
+
+ if pretrain_mm_mlp_adapter is not None:
+ mm_projector_weights = torch.load(pretrain_mm_mlp_adapter, map_location='cpu')
+
+ def get_w(weights, keyword):
+ return {k.split(keyword + '.')[1]: v for k, v in weights.items() if keyword in k}
+
+ self.mm_projector.load_state_dict(get_w(mm_projector_weights, 'mm_projector'))
+
+ def forward(
+ self,
+ input_ids: torch.LongTensor = None,
+ attention_mask: Optional[torch.Tensor] = None,
+ decoder_attention_mask: Optional[torch.Tensor] = None,
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
+ inputs_embeds: Optional[torch.FloatTensor] = None,
+ labels: Optional[torch.LongTensor] = None,
+ use_cache: Optional[bool] = None,
+ output_attentions: Optional[bool] = None,
+ output_hidden_states: Optional[bool] = None,
+ images: Optional[torch.FloatTensor] = None,
+ return_dict: Optional[bool] = None,
+ **kwargs,
+ ) -> Union[Tuple[torch.FloatTensor], Seq2SeqLMOutput]:
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
+ output_hidden_states = (
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
+ )
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
+
+ if inputs_embeds is None:
+ _, attention_mask, decoder_attention_mask, past_key_values, inputs_embeds, labels = \
+ self.prepare_inputs_labels_for_multimodal(input_ids, attention_mask, decoder_attention_mask,
+ past_key_values, labels, images)
+
+ # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)
+ outputs = super(CLIPT5ForConditionalGeneration, self).forward(
+ input_ids=None, # will be None if inputs_embeds is not None
+ attention_mask=attention_mask,
+ decoder_attention_mask=decoder_attention_mask,
+ labels=labels,
+ past_key_values=past_key_values,
+ inputs_embeds=inputs_embeds,
+ use_cache=use_cache,
+ output_attentions=output_attentions,
+ output_hidden_states=output_hidden_states,
+ return_dict=return_dict,
+ **kwargs,
+ )
+
+ return outputs
+
+ @torch.no_grad()
+ def generate(
+ self,
+ inputs: Optional[torch.Tensor] = None,
+ attention_mask: Optional[torch.Tensor] = None,
+ images: Optional[torch.Tensor] = None,
+ **kwargs,
+ ):
+ assert images is not None, "images must be provided"
+ assert inputs is not None, "inputs must be provided"
+ assert attention_mask is not None, "attention_mask must be provided"
+ _, attention_mask, _, _, inputs_embeds, _ = \
+ self.prepare_inputs_labels_for_multimodal(inputs, attention_mask, None, None, None, images)
+ # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)
+ outputs = super(CLIPT5ForConditionalGeneration, self).generate(
+ input_ids=None, # will be None if inputs_embeds is not None
+ attention_mask=attention_mask,
+ inputs_embeds=inputs_embeds,
+ )
+ return outputs
+
+ def prepare_inputs_for_generation(
+ self,
+ input_ids,
+ past_key_values=None,
+ attention_mask=None,
+ head_mask=None,
+ decoder_head_mask=None,
+ decoder_attention_mask=None,
+ cross_attn_head_mask=None,
+ use_cache=None,
+ encoder_outputs=None,
+ inputs_embeds=None,
+ **kwargs,
+ ):
+ # cut decoder_input_ids if past_key_values is used
+ if past_key_values is not None:
+ past_length = past_key_values[0][0].shape[2]
+
+ # Some generation methods already pass only the last input ID
+ if input_ids.shape[1] > past_length:
+ remove_prefix_length = past_length
+ else:
+ # Default to old behavior: keep only final ID
+ remove_prefix_length = input_ids.shape[1] - 1
+
+ input_ids = input_ids[:, remove_prefix_length:]
+
+ # if `inputs_embeds` are passed, we only want to use them in the 1st generation step
+ if inputs_embeds is not None and past_key_values is None:
+ model_inputs = {"inputs_embeds": inputs_embeds}
+ else:
+ model_inputs = {"input_ids": input_ids}
+
+ model_inputs.update({
+ "decoder_input_ids": input_ids,
+ "past_key_values": past_key_values,
+ "encoder_outputs": encoder_outputs,
+ "attention_mask": attention_mask,
+ "head_mask": head_mask,
+ "decoder_head_mask": decoder_head_mask,
+ "decoder_attention_mask": decoder_attention_mask,
+ "cross_attn_head_mask": cross_attn_head_mask,
+ "use_cache": use_cache,
+ })
+ return model_inputs
+
+
+AutoConfig.register("clip_t5", CLIPT5Config)
+AutoModelForSeq2SeqLM.register(CLIPT5Config, CLIPT5ForConditionalGeneration)
+
+
+_vqascore_cache = {}
+
+
+def load_vqascore(device, dtype, use_fsdp=True):
+ # Create cache key from arguments
+ cache_key = f"{device}_{dtype}_{use_fsdp}"
+
+ # Check if model is already cached
+ if cache_key in _vqascore_cache:
+ return _vqascore_cache[cache_key]
+
+ tokenizer = AutoTokenizer.from_pretrained(
+ 'google/flan-t5-xxl', use_fast=False, model_max_length=2048)
+ model = CLIPT5ForConditionalGeneration.from_pretrained(
+ 'zhiqiulin/clip-flant5-xxl',
+ torch_dtype=dtype,
+ use_cache=False,
+ freeze_mm_mlp_adapter=True)
+ model.requires_grad_(False)
+ model.resize_token_embeddings(len(tokenizer))
+
+ if use_fsdp:
+ mmcv.print_log('Wrapping VQAScore model with FSDP.')
+ model = FullyShardedDataParallel(
+ model,
+ device_id=torch.cuda.current_device(),
+ use_orig_params=False,
+ mixed_precision=MixedPrecision(
+ param_dtype=dtype,
+ reduce_dtype=dtype,
+ buffer_dtype=dtype,
+ cast_root_forward_inputs=False),
+ sharding_strategy=ShardingStrategy.HYBRID_SHARD,
+ auto_wrap_policy=ModuleWrapPolicy([T5Block]))
+ else:
+ model.to(device)
+
+ result = model, tokenizer
+ _vqascore_cache[cache_key] = result
+ return result
+
+
+@METRICS.register_module()
+class VQAScore(Metric):
+ name = 'VQAScore'
+ requires_prompt = True
+
+ def __init__(self,
+ num_images=None,
+ use_fsdp=True):
+ super().__init__(num_images)
+ use_fsdp = use_fsdp and torch.cuda.is_available() and dist.is_initialized() and dist.get_world_size() > 0
+
+ self.use_fsdp = use_fsdp
+ self.dtype = torch.bfloat16
+ self.device = 'cuda' if use_fsdp else 'cpu'
+
+ self.model, self.tokenizer = load_vqascore(device=self.device, dtype=self.dtype, use_fsdp=use_fsdp)
+ self.model.eval()
+ image_processor = self.model.get_vision_tower().image_processor
+ image_size = tuple(image_processor.crop_size.values())
+ assert len(image_size) == 2 and image_size[0] == image_size[1]
+ self.image_size = image_size[0]
+ self.image_mean = torch.tensor(image_processor.image_mean, device=self.device).view(3, 1, 1)
+ self.image_std = torch.tensor(image_processor.image_std, device=self.device).view(3, 1, 1)
+ self.clamp_high = (1 - self.image_mean) / self.image_std
+ self.clamp_low = -self.image_mean / self.image_std
+
+ def prepare(self):
+ self.scores = []
+
+ def resize(self, imgs):
+ h, w = imgs.shape[2:]
+ if h != w:
+ pad_size = max(h, w)
+ pad_h = pad_size - h
+ pad_w = pad_size - w
+ imgs = F.pad(
+ imgs, (pad_w // 2, pad_w - pad_w // 2, pad_h // 2, pad_h - pad_h // 2), mode='constant', value=0)
+ h = w = pad_size
+ if h != self.image_size:
+ imgs = F.interpolate(imgs, size=self.image_size, mode='bicubic', align_corners=False, antialias=True)
+ imgs = torch.maximum(torch.minimum(imgs, self.clamp_high), self.clamp_low)
+ return imgs
+
+ @torch.no_grad()
+ def feed_op(self, batch, mode):
+ imgs = batch['imgs']
+ prompts = batch['prompts']
+
+ imgs = (imgs.to(device=self.device, dtype=torch.float32) / 2 + 0.5).clamp(0, 1)
+ imgs = self.resize((imgs - self.image_mean) / self.image_std).to(dtype=self.dtype)
+
+ # ========= preprocess prompts =========
+ questions = [default_question_template.format(prompt) for prompt in prompts]
+ answers = [default_answer_template.format(prompt) for prompt in prompts]
+
+ questions = [format_question(question, conversation_style='t5_chat') for question in questions]
+ answers = [format_answer(answer, conversation_style='t5_chat') for answer in answers]
+
+ input_ids = [t5_tokenizer_image_token(question, self.tokenizer, return_tensors='pt') for question in questions]
+ labels = [t5_tokenizer_image_token(answer, self.tokenizer, return_tensors='pt') for answer in answers]
+
+ input_ids = torch.nn.utils.rnn.pad_sequence(
+ input_ids, batch_first=True, padding_value=0)[:, :self.tokenizer.model_max_length]
+ labels = torch.nn.utils.rnn.pad_sequence(
+ labels, batch_first=True, padding_value=IGNORE_INDEX)[:, :self.tokenizer.model_max_length]
+
+ input_ids = input_ids.to(device=self.device)
+ labels = labels.to(device=self.device)
+
+ attention_mask = input_ids.ne(self.tokenizer.pad_token_id).to(device=self.device)
+ decoder_attention_mask = labels.ne(IGNORE_INDEX).to(device=self.device)
+
+ outputs = self.model(
+ input_ids=input_ids,
+ attention_mask=attention_mask,
+ decoder_attention_mask=decoder_attention_mask,
+ labels=labels,
+ images=imgs,
+ past_key_values=None,
+ inputs_embeds=None,
+ use_cache=None,
+ output_attentions=None,
+ output_hidden_states=None,
+ return_dict=True,
+ )
+
+ logits = outputs.logits
+ bs, seq_len, vocab_size = logits.size()
+ vqa_score = (-F.cross_entropy(
+ logits.reshape(bs * seq_len, vocab_size), labels.reshape(bs * seq_len), reduction='none'
+ ).reshape(bs, 2).mean(dim=1)).exp() # (bs, )
+
+ if dist.is_initialized():
+ ws = dist.get_world_size()
+ placeholder = [torch.empty_like(vqa_score) for _ in range(ws)]
+ dist.all_gather(placeholder, vqa_score)
+ vqa_score = torch.stack(placeholder, dim=1).reshape(vqa_score.size(0) * ws)
+
+ if (dist.is_initialized() and dist.get_rank() == 0) or not dist.is_initialized():
+ self.scores.append(vqa_score.float().cpu())
+
+ def feed(self, batch, mode):
+ if mode == 'reals':
+ return 0
+
+ if self.num_images is None:
+ self.feed_op(batch, mode)
+
+ else:
+ _, ws = get_dist_info()
+
+ if self.num_fake_feeded == self.num_fake_need:
+ return 0
+
+ if isinstance(batch, dict):
+ batch_size = len(list(batch.values())[0])
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = {k: v[:end] for k, v in batch.items()}
+ else:
+ batch_size = batch.shape[0]
+ end = min(batch_size, self.num_fake_need - self.num_fake_feeded)
+ batch_to_feed = batch[:end]
+
+ global_end = min(batch_size * ws,
+ self.num_fake_need - self.num_fake_feeded)
+ self.feed_op(batch_to_feed, mode)
+ self.num_fake_feeded += global_end
+ return end
+
+ @torch.no_grad()
+ def summary(self):
+ scores = torch.cat(self.scores, dim=0)
+ if self.num_images is not None:
+ assert scores.shape[0] >= self.num_images
+ scores = scores[:self.num_images]
+ mean_score = scores.mean().item()
+ self._result_dict = dict(vqascore=mean_score)
+ self._result_str = f'VQAScore: {mean_score:.4f}'
+ return mean_score
+
+ def clear_fake_data(self):
+ self.scores = []
+ self.num_fake_feeded = 0
+
+ def clear(self, clear_reals=False):
+ self.clear_fake_data()
+
+ def load_to_gpu(self):
+ if torch.cuda.is_available() and not isinstance(self.model, FullyShardedDataParallel):
+ self.model.cuda()
+ self.image_mean = self.image_mean.cuda()
+ self.image_std = self.image_std.cuda()
+ self.clamp_high = self.clamp_high.cuda()
+ self.clamp_low = self.clamp_low.cuda()
+ self.device = 'cuda'
+
+ def offload_to_cpu(self):
+ if not isinstance(self.model, FullyShardedDataParallel):
+ self.model.cpu()
+ self.image_mean = self.image_mean.cpu()
+ self.image_std = self.image_std.cpu()
+ self.clamp_high = self.clamp_high.cpu()
+ self.clamp_low = self.clamp_low.cpu()
+ self.device = 'cpu'
diff --git a/lakonlab/models/__init__.py b/lakonlab/models/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..0612681144fda37af0cd03dc096fbef6b4b6ebfb
--- /dev/null
+++ b/lakonlab/models/__init__.py
@@ -0,0 +1,6 @@
+from .losses import *
+from .architecture import *
+from .diffusions import *
+from .diffusion_2d import Diffusion2D
+from .latent_diffusion_class_image import LatentDiffusionClassImage
+from .latent_diffusion_text_image import LatentDiffusionTextImage
diff --git a/lakonlab/models/__pycache__/__init__.cpython-310.pyc b/lakonlab/models/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..2d140ef495903e73b290bcd6d1b9b8c413758640
Binary files /dev/null and b/lakonlab/models/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/models/__pycache__/base.cpython-310.pyc b/lakonlab/models/__pycache__/base.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..89f4c83e64cd30a56777201aebe183b034f3beb4
Binary files /dev/null and b/lakonlab/models/__pycache__/base.cpython-310.pyc differ
diff --git a/lakonlab/models/__pycache__/base_diffusion.cpython-310.pyc b/lakonlab/models/__pycache__/base_diffusion.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..857448316b8a066bceb9fa1835fa9ec9604c5ca4
Binary files /dev/null and b/lakonlab/models/__pycache__/base_diffusion.cpython-310.pyc differ
diff --git a/lakonlab/models/__pycache__/diffusion_2d.cpython-310.pyc b/lakonlab/models/__pycache__/diffusion_2d.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..175a4c1749ba702341c979f1c520419b3139ba57
Binary files /dev/null and b/lakonlab/models/__pycache__/diffusion_2d.cpython-310.pyc differ
diff --git a/lakonlab/models/__pycache__/latent_diffusion_class_image.cpython-310.pyc b/lakonlab/models/__pycache__/latent_diffusion_class_image.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..6c268caf4ec3aecfc47f48603659455a1a1b2f20
Binary files /dev/null and b/lakonlab/models/__pycache__/latent_diffusion_class_image.cpython-310.pyc differ
diff --git a/lakonlab/models/__pycache__/latent_diffusion_text_image.cpython-310.pyc b/lakonlab/models/__pycache__/latent_diffusion_text_image.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..c5c7ee4f538407f894a7a88e1babd4e31dcc6e99
Binary files /dev/null and b/lakonlab/models/__pycache__/latent_diffusion_text_image.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/__init__.py b/lakonlab/models/architecture/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..4cb42386a07c97340d148691f4c94b058f4b689b
--- /dev/null
+++ b/lakonlab/models/architecture/__init__.py
@@ -0,0 +1,4 @@
+from .ddpm import *
+from .diffusers import *
+from .gmflow import *
+from .dxflow import *
diff --git a/lakonlab/models/architecture/__pycache__/__init__.cpython-310.pyc b/lakonlab/models/architecture/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..85052ce7a1693c46388b9355d7b2f0fe5d899506
Binary files /dev/null and b/lakonlab/models/architecture/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/__pycache__/utils.cpython-310.pyc b/lakonlab/models/architecture/__pycache__/utils.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..db705c7001bc468c6a85a78deff2189414e83ef0
Binary files /dev/null and b/lakonlab/models/architecture/__pycache__/utils.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/ddpm/__init__.py b/lakonlab/models/architecture/ddpm/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..c0c1181d378514695c220f96b3992a9ac9f512e4
--- /dev/null
+++ b/lakonlab/models/architecture/ddpm/__init__.py
@@ -0,0 +1,6 @@
+from .denoising import DenoisingUnetMod
+from .modules import (
+ MultiHeadAttentionMod, DenoisingResBlockMod, DenoisingDownsampleMod, DenoisingUpsampleMod)
+
+__all__ = ['DenoisingUnetMod', 'MultiHeadAttentionMod', 'DenoisingResBlockMod',
+ 'DenoisingDownsampleMod', 'DenoisingUpsampleMod']
diff --git a/lakonlab/models/architecture/ddpm/__pycache__/__init__.cpython-310.pyc b/lakonlab/models/architecture/ddpm/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..8c6efb80cbc2df3405b273379c79cae989c24f82
Binary files /dev/null and b/lakonlab/models/architecture/ddpm/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/ddpm/__pycache__/denoising.cpython-310.pyc b/lakonlab/models/architecture/ddpm/__pycache__/denoising.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..8917f58341ba2b39af89bc2fd05d3e54d2863fd5
Binary files /dev/null and b/lakonlab/models/architecture/ddpm/__pycache__/denoising.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/ddpm/__pycache__/modules.cpython-310.pyc b/lakonlab/models/architecture/ddpm/__pycache__/modules.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..8269c6cfcf032c5324a253e279a63180b6513110
Binary files /dev/null and b/lakonlab/models/architecture/ddpm/__pycache__/modules.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/ddpm/denoising.py b/lakonlab/models/architecture/ddpm/denoising.py
new file mode 100644
index 0000000000000000000000000000000000000000..f206d15540a3c2ed94ad0ae8e85c0a4a355b0a97
--- /dev/null
+++ b/lakonlab/models/architecture/ddpm/denoising.py
@@ -0,0 +1,220 @@
+from copy import deepcopy
+
+import torch
+import torch.nn as nn
+from mmcv.cnn.bricks.conv_module import ConvModule
+
+from mmgen.models.architectures.ddpm.modules import TimeEmbedding, EmbedSequential
+from mmgen.models.architectures.ddpm.denoising import DenoisingUnet
+from mmgen.models.builder import MODULES, build_module
+
+
+@MODULES.register_module()
+class DenoisingUnetMod(DenoisingUnet):
+
+ def __init__(self,
+ image_size,
+ in_channels=3,
+ concat_cond_channels=0,
+ base_channels=128,
+ resblocks_per_downsample=3,
+ num_timesteps=1000,
+ use_rescale_timesteps=True,
+ dropout=0,
+ embedding_channels=-1,
+ num_classes=0,
+ channels_cfg=None,
+ groups=1,
+ norm_cfg=dict(type='GN', num_groups=32),
+ act_cfg=dict(type='SiLU', inplace=False),
+ shortcut_kernel_size=1,
+ use_scale_shift_norm=False,
+ num_heads=4,
+ time_embedding_mode='sin',
+ time_embedding_cfg=None,
+ resblock_cfg=dict(type='DenoisingResBlockMod'),
+ attention_cfg=dict(type='MultiHeadAttentionMod'),
+ downsample_conv=True,
+ upsample_conv=True,
+ downsample_cfg=dict(type='DenoisingDownsampleMod'),
+ upsample_cfg=dict(type='DenoisingUpsampleMod'),
+ attention_res=[16, 8],
+ pretrained=None):
+ super(DenoisingUnet, self).__init__()
+
+ self.num_classes = num_classes
+ self.num_timesteps = num_timesteps
+ self.use_rescale_timesteps = use_rescale_timesteps
+
+ out_channels = in_channels
+ self.out_channels = out_channels
+ self.concat_cond_channels = concat_cond_channels
+
+ # check type of image_size
+ if isinstance(image_size, list) or isinstance(image_size, tuple):
+ assert len(image_size) == 2, 'The length of `image_size` should be 2.'
+ elif isinstance(image_size, int):
+ image_size = [image_size, image_size]
+ else:
+ raise TypeError('Only support `int` and `list[int]` for `image_size`.')
+ self.image_size = image_size
+
+ if isinstance(channels_cfg, list):
+ self.channel_factor_list = channels_cfg
+ else:
+ raise ValueError('Only support list or dict for `channels_cfg`, '
+ f'receive {type(channels_cfg)}')
+
+ embedding_channels = base_channels * 4 \
+ if embedding_channels == -1 else embedding_channels
+ self.time_embedding = TimeEmbedding(
+ base_channels,
+ embedding_channels=embedding_channels,
+ embedding_mode=time_embedding_mode,
+ embedding_cfg=time_embedding_cfg,
+ act_cfg=act_cfg)
+
+ if self.num_classes != 0:
+ self.label_embedding = nn.Embedding(self.num_classes,
+ embedding_channels)
+
+ self.resblock_cfg = deepcopy(resblock_cfg)
+ self.resblock_cfg.setdefault('dropout', dropout)
+ self.resblock_cfg.setdefault('groups', groups)
+ self.resblock_cfg.setdefault('norm_cfg', norm_cfg)
+ self.resblock_cfg.setdefault('act_cfg', act_cfg)
+ self.resblock_cfg.setdefault('embedding_channels', embedding_channels)
+ self.resblock_cfg.setdefault('use_scale_shift_norm',
+ use_scale_shift_norm)
+ self.resblock_cfg.setdefault('shortcut_kernel_size',
+ shortcut_kernel_size)
+
+ # get scales of ResBlock to apply attention
+ attention_scale = [min(image_size) // int(res) for res in attention_res]
+ self.attention_cfg = deepcopy(attention_cfg)
+ self.attention_cfg.setdefault('num_heads', num_heads)
+ self.attention_cfg.setdefault('groups', groups)
+ self.attention_cfg.setdefault('norm_cfg', norm_cfg)
+
+ self.downsample_cfg = deepcopy(downsample_cfg)
+ self.downsample_cfg.setdefault('groups', groups)
+ self.downsample_cfg.setdefault('with_conv', downsample_conv)
+ self.upsample_cfg = deepcopy(upsample_cfg)
+ self.upsample_cfg.setdefault('groups', groups)
+ self.upsample_cfg.setdefault('with_conv', upsample_conv)
+
+ # init the channel scale factor
+ scale = 1
+ self.in_blocks = nn.ModuleList([
+ EmbedSequential(
+ nn.Conv2d(in_channels + concat_cond_channels, base_channels, 3, 1, padding=1, groups=groups))
+ ])
+ self.in_channels_list = [base_channels]
+
+ # construct the encoder part of Unet
+ for level, factor in enumerate(self.channel_factor_list):
+ in_channels_ = base_channels if level == 0 \
+ else base_channels * self.channel_factor_list[level - 1]
+ out_channels_ = base_channels * factor
+
+ for _ in range(resblocks_per_downsample):
+ layers = [
+ build_module(self.resblock_cfg, {
+ 'in_channels': in_channels_,
+ 'out_channels': out_channels_
+ })
+ ]
+ in_channels_ = out_channels_
+
+ if scale in attention_scale:
+ layers.append(
+ build_module(self.attention_cfg,
+ {'in_channels': in_channels_}))
+
+ self.in_channels_list.append(in_channels_)
+ self.in_blocks.append(EmbedSequential(*layers))
+
+ if level != len(self.channel_factor_list) - 1:
+ self.in_blocks.append(
+ EmbedSequential(
+ build_module(self.downsample_cfg,
+ {'in_channels': in_channels_})))
+ self.in_channels_list.append(in_channels_)
+ scale *= 2
+
+ # construct the bottom part of Unet
+ self.mid_blocks = EmbedSequential(
+ build_module(self.resblock_cfg, {'in_channels': in_channels_}),
+ build_module(self.attention_cfg, {'in_channels': in_channels_}),
+ build_module(self.resblock_cfg, {'in_channels': in_channels_}),
+ )
+
+ # construct the decoder part of Unet
+ in_channels_list = deepcopy(self.in_channels_list)
+ self.out_blocks = nn.ModuleList()
+ for level, factor in enumerate(self.channel_factor_list[::-1]):
+ for idx in range(resblocks_per_downsample + 1):
+ layers = [
+ build_module(
+ self.resblock_cfg, {
+ 'in_channels':
+ in_channels_ + in_channels_list.pop(),
+ 'out_channels': base_channels * factor
+ })
+ ]
+ in_channels_ = base_channels * factor
+ if scale in attention_scale:
+ layers.append(
+ build_module(self.attention_cfg,
+ {'in_channels': in_channels_}))
+ if (level != len(self.channel_factor_list) - 1
+ and idx == resblocks_per_downsample):
+ layers.append(
+ build_module(self.upsample_cfg,
+ {'in_channels': in_channels_}))
+ scale //= 2
+ self.out_blocks.append(EmbedSequential(*layers))
+
+ self.out = ConvModule(
+ in_channels=in_channels_,
+ out_channels=out_channels,
+ kernel_size=3,
+ padding=1,
+ groups=groups,
+ act_cfg=act_cfg,
+ norm_cfg=norm_cfg,
+ bias=True,
+ order=('norm', 'act', 'conv'))
+
+ self.init_weights(pretrained)
+
+ def forward(self, x_t, t, label=None, concat_cond=None, return_noise=False):
+ if self.use_rescale_timesteps:
+ t = t.float() * (1000.0 / self.num_timesteps)
+ with torch.autocast(
+ device_type='cuda',
+ enabled=True,
+ dtype=x_t.dtype):
+ embedding = self.time_embedding(t)
+
+ if label is not None:
+ assert hasattr(self, 'label_embedding')
+ embedding = self.label_embedding(label) + embedding
+
+ h, hs = x_t, []
+ if self.concat_cond_channels > 0:
+ h = torch.cat([h, concat_cond], dim=1)
+ # forward downsample blocks
+ for block in self.in_blocks:
+ h = block(h, embedding)
+ hs.append(h)
+
+ # forward middle blocks
+ h = self.mid_blocks(h, embedding)
+
+ # forward upsample blocks
+ for block in self.out_blocks:
+ h = block(torch.cat([h, hs.pop()], dim=1), embedding)
+ outputs = self.out(h)
+
+ return outputs
diff --git a/lakonlab/models/architecture/ddpm/modules.py b/lakonlab/models/architecture/ddpm/modules.py
new file mode 100644
index 0000000000000000000000000000000000000000..94f47ed074adf74e082a70754f978b384e62e8ca
--- /dev/null
+++ b/lakonlab/models/architecture/ddpm/modules.py
@@ -0,0 +1,141 @@
+from copy import deepcopy
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from mmcv.cnn.bricks import build_activation_layer, build_norm_layer
+
+from mmgen.models.builder import MODULES, build_module
+from mmgen.models.architectures.ddpm.modules import (
+ MultiHeadAttention, DenoisingResBlock, DenoisingDownsample, DenoisingUpsample)
+
+
+@MODULES.register_module()
+class MultiHeadAttentionMod(MultiHeadAttention):
+
+ def __init__(self,
+ in_channels,
+ num_heads=1,
+ groups=1,
+ norm_cfg=dict(type='GN', num_groups=32)):
+ super(MultiHeadAttention, self).__init__()
+ self.num_heads = num_heads
+ self.groups = groups
+ _, self.norm = build_norm_layer(norm_cfg, in_channels)
+ self.qkv = nn.Conv1d(in_channels, in_channels * 3, 1, groups=groups)
+ self.proj = nn.Conv1d(in_channels, in_channels, 1, groups=groups)
+
+ if hasattr(F, 'scaled_dot_product_attention'):
+ self.attn_proc = self.efficient_attn
+ else:
+ self.attn_proc = self.QKVAttention
+
+ self.init_weights()
+
+ @staticmethod
+ def efficient_attn(qkv):
+ q, k, v = torch.chunk(qkv, 3, dim=1)
+ return F.scaled_dot_product_attention(q, k, v)
+
+ def forward(self, x):
+ """Forward function for multi head attention.
+ Args:
+ x (torch.Tensor): Input feature map.
+
+ Returns:
+ torch.Tensor: Feature map after attention.
+ """
+ b, c, *spatial = x.shape
+ x = x.reshape(b, c, -1)
+ spatial_numel = x.size(-1)
+ qkv = self.qkv(self.norm(x))
+ qkv = qkv.reshape(
+ b, self.groups, -1, spatial_numel
+ ).transpose(1, 2).reshape(b * self.num_heads, -1, self.groups * spatial_numel)
+ h = self.attn_proc(qkv)
+ h = h.reshape(
+ b, -1, self.groups, spatial_numel
+ ).transpose(1, 2).reshape(b, -1, spatial_numel)
+ h = self.proj(h)
+ return (h + x).reshape(b, c, *spatial)
+
+
+@MODULES.register_module()
+class DenoisingResBlockMod(DenoisingResBlock):
+ def __init__(self,
+ in_channels,
+ embedding_channels,
+ use_scale_shift_norm,
+ dropout,
+ groups=1,
+ out_channels=None,
+ norm_cfg=dict(type='GN', num_groups=32),
+ act_cfg=dict(type='SiLU', inplace=False),
+ shortcut_kernel_size=1):
+ super(DenoisingResBlock, self).__init__()
+ out_channels = in_channels if out_channels is None else out_channels
+
+ _norm_cfg = deepcopy(norm_cfg)
+
+ _, norm_1 = build_norm_layer(_norm_cfg, in_channels)
+ conv_1 = [
+ norm_1,
+ build_activation_layer(act_cfg),
+ nn.Conv2d(in_channels, out_channels, 3, padding=1, groups=groups)
+ ]
+ self.conv_1 = nn.Sequential(*conv_1)
+
+ norm_with_embedding_cfg = dict(
+ in_channels=out_channels,
+ embedding_channels=embedding_channels,
+ use_scale_shift=use_scale_shift_norm,
+ norm_cfg=_norm_cfg)
+ self.norm_with_embedding = build_module(
+ dict(type='NormWithEmbedding'),
+ default_args=norm_with_embedding_cfg)
+
+ conv_2 = [
+ build_activation_layer(act_cfg),
+ nn.Dropout(dropout),
+ nn.Conv2d(out_channels, out_channels, 3, padding=1, groups=groups)
+ ] if dropout > 0 else [
+ build_activation_layer(act_cfg),
+ nn.Conv2d(out_channels, out_channels, 3, padding=1, groups=groups)
+ ]
+ self.conv_2 = nn.Sequential(*conv_2)
+
+ assert shortcut_kernel_size in [
+ 1, 3
+ ], ('Only support `1` and `3` for `shortcut_kernel_size`, but '
+ f'receive {shortcut_kernel_size}.')
+
+ self.learnable_shortcut = out_channels != in_channels
+
+ if self.learnable_shortcut:
+ shortcut_padding = 1 if shortcut_kernel_size == 3 else 0
+ self.shortcut = nn.Conv2d(
+ in_channels,
+ out_channels,
+ shortcut_kernel_size,
+ padding=shortcut_padding,
+ groups=groups)
+ self.init_weights()
+
+
+@MODULES.register_module()
+class DenoisingDownsampleMod(DenoisingDownsample):
+ def __init__(self, in_channels, groups=1, with_conv=True):
+ super(DenoisingDownsample, self).__init__()
+ if with_conv:
+ self.downsample = nn.Conv2d(in_channels, in_channels, 3, 2, 1, groups=groups)
+ else:
+ self.downsample = nn.AvgPool2d(stride=2)
+
+
+@MODULES.register_module()
+class DenoisingUpsampleMod(DenoisingUpsample):
+ def __init__(self, in_channels, groups=1, with_conv=True):
+ super(DenoisingUpsample, self).__init__()
+ if with_conv:
+ self.with_conv = True
+ self.conv = nn.Conv2d(in_channels, in_channels, 3, 1, 1, groups=groups)
diff --git a/lakonlab/models/architecture/dxflow/__init__.py b/lakonlab/models/architecture/dxflow/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..0dee64f106871d2a21c0cab1ffcff2c6dead92e6
--- /dev/null
+++ b/lakonlab/models/architecture/dxflow/__init__.py
@@ -0,0 +1,5 @@
+from .dxqwen import DXQwenImageTransformer2DModel
+from .dxdit import DXDiTTransformer2DModel
+from .dxflux import DXFluxTransformer2DModel
+
+__all__ = ['DXQwenImageTransformer2DModel', 'DXDiTTransformer2DModel', 'DXFluxTransformer2DModel']
diff --git a/lakonlab/models/architecture/dxflow/__pycache__/__init__.cpython-310.pyc b/lakonlab/models/architecture/dxflow/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..c7edc68fcc1a9c0ae2e93f4d08562d03613315ea
Binary files /dev/null and b/lakonlab/models/architecture/dxflow/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/dxflow/__pycache__/dxdit.cpython-310.pyc b/lakonlab/models/architecture/dxflow/__pycache__/dxdit.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..3296d31ddf4cf27e1144de1ab95411b09e2b161a
Binary files /dev/null and b/lakonlab/models/architecture/dxflow/__pycache__/dxdit.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/dxflow/__pycache__/dxflux.cpython-310.pyc b/lakonlab/models/architecture/dxflow/__pycache__/dxflux.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..9300bae1c2fac2d717626942f6be4720f1860e3d
Binary files /dev/null and b/lakonlab/models/architecture/dxflow/__pycache__/dxflux.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/dxflow/__pycache__/dxqwen.cpython-310.pyc b/lakonlab/models/architecture/dxflow/__pycache__/dxqwen.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..72c657b5931fbaa5e0dfdaa6b78b8de5c26d1518
Binary files /dev/null and b/lakonlab/models/architecture/dxflow/__pycache__/dxqwen.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/dxflow/dxdit.py b/lakonlab/models/architecture/dxflow/dxdit.py
new file mode 100644
index 0000000000000000000000000000000000000000..b398eabe922e81b0ed6b5126c5651427f2d9f767
--- /dev/null
+++ b/lakonlab/models/architecture/dxflow/dxdit.py
@@ -0,0 +1,114 @@
+import torch
+
+from typing import Optional
+from mmcv.runner import load_checkpoint, _load_checkpoint, load_state_dict
+from mmgen.models.builder import MODULES
+from mmgen.utils import get_root_logger
+from ..diffusers.dit import _DiTTransformer2DModelMod
+from ..utils import flex_freeze
+
+
+@MODULES.register_module()
+class DXDiTTransformer2DModel(_DiTTransformer2DModelMod):
+
+ def __init__(
+ self,
+ *args,
+ n_grid=16,
+ in_channels=4,
+ out_channels=None,
+ freeze=False,
+ freeze_exclude=[],
+ pretrained=None,
+ torch_dtype='float32',
+ autocast_dtype=None,
+ freeze_exclude_fp32=True,
+ freeze_exclude_autocast_dtype='float32',
+ checkpointing=True,
+ **kwargs):
+ out_channels = out_channels or in_channels
+ super().__init__(
+ *args, in_channels=in_channels, out_channels=n_grid * out_channels, **kwargs)
+ self.n_grid = n_grid
+
+ self.init_weights(pretrained)
+
+ if autocast_dtype is not None:
+ assert torch_dtype == 'float32'
+ self.autocast_dtype = autocast_dtype
+
+ if torch_dtype is not None:
+ self.to(getattr(torch, torch_dtype))
+
+ self.freeze = freeze
+ if self.freeze:
+ flex_freeze(
+ self,
+ exclude_keys=freeze_exclude,
+ exclude_fp32=freeze_exclude_fp32,
+ exclude_autocast_dtype=freeze_exclude_autocast_dtype)
+
+ if checkpointing:
+ self.enable_gradient_checkpointing()
+
+ def init_weights(self, pretrained=None):
+ super().init_weights()
+ if pretrained is not None:
+ logger = get_root_logger()
+ # load_checkpoint(self, pretrained, map_location='cpu', strict=False, logger=logger)
+ checkpoint = _load_checkpoint(pretrained, map_location='cpu', logger=logger)
+ if 'state_dict' in checkpoint:
+ state_dict = checkpoint['state_dict']
+ else:
+ state_dict = checkpoint
+ # expand the output channels
+ p2 = self.config.patch_size * self.config.patch_size
+ ori_out_channels = p2 * self.out_channels // self.n_grid
+ if 'proj_out_2.weight' in state_dict:
+ # if this is GMDiT V1 model with 1 Gaussian
+ if state_dict['proj_out_2.weight'].size(0) == p2 * (
+ self.out_channels // self.n_grid + 1):
+ state_dict['proj_out_2.weight'] = state_dict['proj_out_2.weight'].reshape(
+ p2, self.out_channels // self.n_grid + 1, -1
+ )[:, :-1].reshape(ori_out_channels, -1)
+ if state_dict['proj_out_2.weight'].size(0) == ori_out_channels:
+ state_dict['proj_out_2.weight'] = state_dict['proj_out_2.weight'].reshape(
+ p2, 1, self.out_channels // self.n_grid, -1
+ ).expand(-1, self.n_grid, -1, -1).reshape(
+ self.n_grid * ori_out_channels, -1)
+ if 'proj_out_2.bias' in state_dict:
+ # if this is GMDiT V1 model with 1 Gaussian
+ if state_dict['proj_out_2.bias'].size(0) == p2 * (
+ self.out_channels // self.n_grid + 1):
+ state_dict['proj_out_2.bias'] = state_dict['proj_out_2.bias'].reshape(
+ p2, self.out_channels // self.n_grid + 1
+ )[:, :-1].reshape(ori_out_channels)
+ if state_dict['proj_out_2.bias'].size(0) == ori_out_channels:
+ state_dict['proj_out_2.bias'] = state_dict['proj_out_2.bias'].reshape(
+ p2, 1, self.out_channels // self.n_grid
+ ).expand(-1, self.n_grid, -1).reshape(
+ self.n_grid * ori_out_channels)
+ load_state_dict(self, state_dict, logger=logger)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ timestep: Optional[torch.LongTensor] = None,
+ class_labels: Optional[torch.LongTensor] = None,
+ **kwargs):
+ if self.autocast_dtype is not None:
+ dtype = getattr(torch, self.autocast_dtype)
+ else:
+ dtype = hidden_states.dtype
+ with torch.autocast(
+ device_type='cuda',
+ enabled=self.autocast_dtype is not None,
+ dtype=dtype if self.autocast_dtype is not None else None):
+ output = super().forward(
+ hidden_states.to(dtype),
+ timestep=timestep,
+ class_labels=class_labels,
+ **kwargs)
+ bs, _, h, w = output.shape
+ output = output.reshape(bs, self.n_grid, self.out_channels // self.n_grid, h, w)
+ return output
diff --git a/lakonlab/models/architecture/dxflow/dxflux.py b/lakonlab/models/architecture/dxflow/dxflux.py
new file mode 100644
index 0000000000000000000000000000000000000000..346b8aa98342b2e1595a88c130d5b251acf438c0
--- /dev/null
+++ b/lakonlab/models/architecture/dxflow/dxflux.py
@@ -0,0 +1,177 @@
+import torch
+
+from typing import Optional
+from accelerate import init_empty_weights
+from diffusers.models import FluxTransformer2DModel
+from peft import LoraConfig
+from mmgen.models.builder import MODULES
+from mmgen.utils import get_root_logger
+from ..utils import flex_freeze
+from lakonlab.runner.checkpoint import _load_checkpoint, load_full_state_dict
+
+
+@MODULES.register_module()
+class DXFluxTransformer2DModel(FluxTransformer2DModel):
+
+ def __init__(
+ self,
+ n_grid=16,
+ patch_size=2,
+ in_channels=64,
+ out_channels=None,
+ freeze=False,
+ freeze_exclude=[],
+ pretrained=None,
+ pretrained_adapter=None,
+ torch_dtype='float32',
+ freeze_exclude_fp32=True,
+ freeze_exclude_autocast_dtype='float32',
+ checkpointing=True,
+ use_lora=False,
+ lora_target_modules=None,
+ lora_rank=16,
+ lora_dropout=0.0,
+ **kwargs):
+ out_channels = out_channels or in_channels
+ with init_empty_weights():
+ super().__init__(
+ patch_size=1, in_channels=in_channels, out_channels=n_grid * out_channels, **kwargs)
+ self.n_grid = n_grid
+ self.patch_size = patch_size
+
+ self.init_weights(pretrained, pretrained_adapter)
+
+ self.use_lora = use_lora
+ self.lora_target_modules = lora_target_modules
+ self.lora_rank = lora_rank
+ if self.use_lora:
+ transformer_lora_config = LoraConfig(
+ r=lora_rank,
+ lora_alpha=lora_rank,
+ init_lora_weights='gaussian',
+ target_modules=lora_target_modules,
+ lora_dropout=lora_dropout,
+ )
+ self.add_adapter(transformer_lora_config)
+
+ if torch_dtype is not None:
+ self.to(getattr(torch, torch_dtype))
+
+ self.freeze = freeze
+ if self.freeze:
+ flex_freeze(
+ self,
+ exclude_keys=freeze_exclude,
+ exclude_fp32=freeze_exclude_fp32,
+ exclude_autocast_dtype=freeze_exclude_autocast_dtype)
+
+ if checkpointing:
+ self.enable_gradient_checkpointing()
+
+ def init_weights(self, pretrained=None, pretrained_adapter=None):
+ if pretrained is not None:
+ logger = get_root_logger()
+ checkpoint = _load_checkpoint(pretrained, map_location='cpu', logger=logger)
+ if 'state_dict' in checkpoint:
+ state_dict = checkpoint['state_dict']
+ else:
+ state_dict = checkpoint
+ # expand the output channels
+ if 'proj_out.weight' in state_dict and \
+ state_dict['proj_out.weight'].size(0) == self.out_channels // self.n_grid:
+ state_dict['proj_out.weight'] = state_dict['proj_out.weight'][None].expand(
+ self.n_grid, -1, -1).reshape(self.out_channels, -1)
+ if 'proj_out.bias' in state_dict and \
+ state_dict['proj_out.bias'].size(0) == self.out_channels // self.n_grid:
+ state_dict['proj_out.bias'] = state_dict['proj_out.bias'][None].expand(
+ self.n_grid, -1).reshape(self.out_channels)
+ if pretrained_adapter is not None:
+ adapter_state_dict = _load_checkpoint(
+ pretrained_adapter, map_location='cpu', logger=logger)
+ lora_state_dict = dict()
+ for k, v in adapter_state_dict.items():
+ if 'lora' in k:
+ lora_state_dict[k] = v
+ else:
+ state_dict[k] = v
+ load_full_state_dict(self, state_dict, logger=logger, assign=True)
+ if len(lora_state_dict) > 0:
+ self.load_lora_adapter(lora_state_dict, prefix=None)
+ self.fuse_lora()
+ self.unload_lora()
+ else:
+ load_full_state_dict(self, state_dict, logger=logger, assign=True)
+
+ @staticmethod
+ def _prepare_latent_image_ids(height, width, device, dtype):
+ """
+ Copied from Diffusers
+ """
+ latent_image_ids = torch.zeros(height, width, 3)
+ latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height)[:, None]
+ latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width)[None, :]
+
+ latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape
+
+ latent_image_ids = latent_image_ids.reshape(
+ latent_image_id_height * latent_image_id_width, latent_image_id_channels)
+
+ return latent_image_ids.to(device=device, dtype=dtype)
+
+ def patchify(self, latents):
+ if self.patch_size > 1:
+ bs, c, h, w = latents.size()
+ latents = latents.reshape(
+ bs, c, h // self.patch_size, self.patch_size, w // self.patch_size, self.patch_size
+ ).permute(
+ 0, 1, 3, 5, 2, 4
+ ).reshape(
+ bs, c * self.patch_size * self.patch_size, h // self.patch_size, w // self.patch_size)
+ return latents
+
+ def unpatchify(self, latents):
+ if self.patch_size > 1:
+ bs, k, c, h, w = latents.size()
+ latents = latents.reshape(
+ bs, k, c // (self.patch_size * self.patch_size), self.patch_size, self.patch_size, h, w
+ ).permute(
+ 0, 1, 2, 5, 3, 6, 4
+ ).reshape(
+ bs, k, c // (self.patch_size * self.patch_size), h * self.patch_size, w * self.patch_size)
+ return latents
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ timestep: torch.Tensor,
+ encoder_hidden_states: torch.Tensor = None,
+ pooled_projections: torch.Tensor = None,
+ mask: Optional[torch.Tensor] = None,
+ masked_image_latents: Optional[torch.Tensor] = None,
+ **kwargs):
+ hidden_states = self.patchify(hidden_states)
+ bs, c, h, w = hidden_states.size()
+ dtype = hidden_states.dtype
+ device = hidden_states.device
+ hidden_states = hidden_states.reshape(bs, c, h * w).permute(0, 2, 1)
+ img_ids = self._prepare_latent_image_ids(
+ h, w, device, dtype)
+ txt_ids = img_ids.new_zeros((encoder_hidden_states.shape[-2], 3))
+
+ # Flux fill
+ if mask is not None and masked_image_latents is not None:
+ hidden_states = torch.cat(
+ (hidden_states, masked_image_latents.to(dtype=dtype), mask.to(dtype=dtype)), dim=-1)
+
+ output = super().forward(
+ hidden_states=hidden_states,
+ encoder_hidden_states=encoder_hidden_states.to(dtype),
+ pooled_projections=pooled_projections.to(dtype),
+ timestep=timestep,
+ img_ids=img_ids,
+ txt_ids=txt_ids,
+ return_dict=False,
+ **kwargs)[0]
+
+ output = output.permute(0, 2, 1).reshape(bs, self.n_grid, self.out_channels // self.n_grid, h, w)
+ return self.unpatchify(output)
diff --git a/lakonlab/models/architecture/dxflow/dxqwen.py b/lakonlab/models/architecture/dxflow/dxqwen.py
new file mode 100644
index 0000000000000000000000000000000000000000..e23eacda895c2fae70f5d77465181530a946bdca
--- /dev/null
+++ b/lakonlab/models/architecture/dxflow/dxqwen.py
@@ -0,0 +1,185 @@
+import torch
+
+from accelerate import init_empty_weights
+from diffusers.models import QwenImageTransformer2DModel
+from peft import LoraConfig
+from mmgen.models.builder import MODULES
+from mmgen.utils import get_root_logger
+from ..utils import flex_freeze
+from lakonlab.runner.checkpoint import _load_checkpoint, load_full_state_dict
+
+
+@MODULES.register_module()
+class DXQwenImageTransformer2DModel(QwenImageTransformer2DModel):
+
+ def __init__(
+ self,
+ n_grid=1,
+ p_order=1,
+ patch_size=2,
+ in_channels=64,
+ out_channels=None,
+ freeze=False,
+ freeze_exclude=[],
+ pretrained=None,
+ pretrained_adapter=None,
+ torch_dtype='float32',
+ freeze_exclude_fp32=True,
+ freeze_exclude_autocast_dtype='float32',
+ checkpointing=True,
+ use_lora=False,
+ lora_target_modules=None,
+ lora_rank=16,
+ lora_dropout=0.0,
+ **kwargs):
+ assert n_grid > 0 and p_order > 0
+ assert n_grid == 1 or p_order == 1, \
+ "Only one of n_grid and p_order can be greater than 1."
+
+ out_channels = out_channels or in_channels
+ with init_empty_weights():
+ super().__init__(
+ patch_size=1, in_channels=in_channels, out_channels=n_grid * p_order * out_channels, **kwargs)
+ self.n_grid = n_grid
+ self.p_order = p_order
+ self.patch_size = patch_size
+
+ self.init_weights(pretrained, pretrained_adapter, mode='grid' if n_grid > 1 else 'polynomial')
+
+ self.use_lora = use_lora
+ self.lora_target_modules = lora_target_modules
+ self.lora_rank = lora_rank
+ if self.use_lora:
+ transformer_lora_config = LoraConfig(
+ r=lora_rank,
+ lora_alpha=lora_rank,
+ init_lora_weights='gaussian',
+ target_modules=lora_target_modules,
+ lora_dropout=lora_dropout,
+ )
+ self.add_adapter(transformer_lora_config)
+
+ if torch_dtype is not None:
+ self.to(getattr(torch, torch_dtype))
+
+ self.freeze = freeze
+ if self.freeze:
+ flex_freeze(
+ self,
+ exclude_keys=freeze_exclude,
+ exclude_fp32=freeze_exclude_fp32,
+ exclude_autocast_dtype=freeze_exclude_autocast_dtype)
+
+ if checkpointing:
+ self.enable_gradient_checkpointing()
+
+ def init_weights(self, pretrained=None, pretrained_adapter=None, mode='grid'):
+ if pretrained is not None:
+ logger = get_root_logger()
+ checkpoint = _load_checkpoint(pretrained, map_location='cpu', logger=logger)
+ if 'state_dict' in checkpoint:
+ state_dict = checkpoint['state_dict']
+ else:
+ state_dict = checkpoint
+ # expand the output channels
+ if mode == 'grid':
+ if 'proj_out.weight' in state_dict and \
+ state_dict['proj_out.weight'].size(0) == self.out_channels // self.n_grid:
+ state_dict['proj_out.weight'] = state_dict['proj_out.weight'][None].expand(
+ self.n_grid, -1, -1).reshape(self.out_channels, -1)
+ if 'proj_out.bias' in state_dict and \
+ state_dict['proj_out.bias'].size(0) == self.out_channels // self.n_grid:
+ state_dict['proj_out.bias'] = state_dict['proj_out.bias'][None].expand(
+ self.n_grid, -1).reshape(self.out_channels)
+ elif mode == 'polynomial':
+ if 'proj_out.weight' in state_dict and \
+ state_dict['proj_out.weight'].size(0) == self.out_channels // self.p_order:
+ state_dict['proj_out.weight'] = torch.cat(
+ [state_dict['proj_out.weight'][None],
+ torch.zeros(
+ (self.p_order - 1, *state_dict['proj_out.weight'].size()),
+ device=state_dict['proj_out.weight'].device, dtype=state_dict['proj_out.weight'].dtype)],
+ dim=0).reshape(self.out_channels, -1)
+ if 'proj_out.bias' in state_dict and \
+ state_dict['proj_out.bias'].size(0) == self.out_channels // self.p_order:
+ state_dict['proj_out.bias'] = torch.cat(
+ [state_dict['proj_out.bias'][None],
+ torch.zeros(
+ (self.p_order - 1, *state_dict['proj_out.bias'].size()),
+ device=state_dict['proj_out.bias'].device, dtype=state_dict['proj_out.bias'].dtype)],
+ dim=0).reshape(self.out_channels)
+ else:
+ raise ValueError(f"Unknown mode: {mode}")
+ if pretrained_adapter is not None:
+ adapter_state_dict = _load_checkpoint(
+ pretrained_adapter, map_location='cpu', logger=logger)
+ lora_state_dict = dict()
+ for k, v in adapter_state_dict.items():
+ if 'lora' in k:
+ lora_state_dict[k] = v
+ else:
+ state_dict[k] = v
+ load_full_state_dict(self, state_dict, logger=logger, assign=True)
+ if len(lora_state_dict) > 0:
+ self.load_lora_adapter(lora_state_dict, prefix=None)
+ self.fuse_lora()
+ self.unload_lora()
+ else:
+ load_full_state_dict(self, state_dict, logger=logger, assign=True)
+
+ def patchify(self, latents):
+ if self.patch_size > 1:
+ bs, c, h, w = latents.size()
+ latents = latents.reshape(
+ bs, c, h // self.patch_size, self.patch_size, w // self.patch_size, self.patch_size
+ ).permute(
+ 0, 1, 3, 5, 2, 4
+ ).reshape(
+ bs, c * self.patch_size * self.patch_size, h // self.patch_size, w // self.patch_size)
+ return latents
+
+ def unpatchify(self, latents):
+ if self.patch_size > 1:
+ bs, k, c, h, w = latents.size()
+ latents = latents.reshape(
+ bs, k, c // (self.patch_size * self.patch_size), self.patch_size, self.patch_size, h, w
+ ).permute(
+ 0, 1, 2, 5, 3, 6, 4
+ ).reshape(
+ bs, k, c // (self.patch_size * self.patch_size), h * self.patch_size, w * self.patch_size)
+ return latents
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ timestep: torch.Tensor,
+ encoder_hidden_states: torch.Tensor = None,
+ encoder_hidden_states_mask: torch.Tensor = None,
+ **kwargs):
+ hidden_states = self.patchify(hidden_states)
+ bs, c, h, w = hidden_states.size()
+ dtype = hidden_states.dtype
+ hidden_states = hidden_states.reshape(bs, c, h * w).permute(0, 2, 1)
+ img_shapes = [[(1, h, w)]]
+ if encoder_hidden_states_mask is not None:
+ txt_seq_lens = encoder_hidden_states_mask.sum(dim=1)
+ max_txt_seq_len = txt_seq_lens.max()
+ encoder_hidden_states = encoder_hidden_states[:, :max_txt_seq_len]
+ encoder_hidden_states_mask = encoder_hidden_states_mask[:, :max_txt_seq_len]
+ txt_seq_lens = txt_seq_lens.tolist()
+ else:
+ txt_seq_lens = None
+
+ output = super().forward(
+ hidden_states=hidden_states,
+ encoder_hidden_states=encoder_hidden_states.to(dtype),
+ encoder_hidden_states_mask=encoder_hidden_states_mask,
+ timestep=timestep,
+ img_shapes=img_shapes,
+ txt_seq_lens=txt_seq_lens,
+ return_dict=False,
+ **kwargs)[0]
+
+ extra_dim = max(self.n_grid, self.p_order)
+ output = output.permute(0, 2, 1).reshape(bs, extra_dim, self.out_channels // extra_dim, h, w)
+ return self.unpatchify(output)
diff --git a/lakonlab/models/architecture/gmflow/__init__.py b/lakonlab/models/architecture/gmflow/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..4a90a63b57f7900c86329e4cab68e85f08d7830b
--- /dev/null
+++ b/lakonlab/models/architecture/gmflow/__init__.py
@@ -0,0 +1,13 @@
+from .gmflux import GMFluxTransformer2DModel
+from .gmdit import GMDiTTransformer2DModel
+from .gmdit_v2 import GMDiTTransformer2DModelV2
+from .toymodels import GMFlowMLP2DDenoiser
+from .spectrum_mlp import SpectrumMLP
+from .gmunet_ddpm import GMUnet
+from .gmsd3 import GMSD3Transformer2DModel
+from .gmqwen import GMQwenImageTransformer2DModel
+
+__all__ = [
+ 'GMDiTTransformer2DModel', 'GMDiTTransformer2DModelV2',
+ 'GMFluxTransformer2DModel', 'GMFlowMLP2DDenoiser', 'SpectrumMLP', 'GMUnet',
+ 'GMSD3Transformer2DModel', 'GMQwenImageTransformer2DModel']
diff --git a/lakonlab/models/architecture/gmflow/__pycache__/__init__.cpython-310.pyc b/lakonlab/models/architecture/gmflow/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..bf7a48f46167eb6e6b604d825bd3fa1dc06557b9
Binary files /dev/null and b/lakonlab/models/architecture/gmflow/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/gmflow/__pycache__/gm_output.cpython-310.pyc b/lakonlab/models/architecture/gmflow/__pycache__/gm_output.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..532d7687db8f327d15194a0d6be5cf0df3c99a3b
Binary files /dev/null and b/lakonlab/models/architecture/gmflow/__pycache__/gm_output.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/gmflow/__pycache__/gmdit.cpython-310.pyc b/lakonlab/models/architecture/gmflow/__pycache__/gmdit.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..684e0e56493c3afb18b9cd6b62a7e7afc6216ba4
Binary files /dev/null and b/lakonlab/models/architecture/gmflow/__pycache__/gmdit.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/gmflow/__pycache__/gmdit_v2.cpython-310.pyc b/lakonlab/models/architecture/gmflow/__pycache__/gmdit_v2.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..93b775350f0e5a06c4d89df0daa2ff48617261fb
Binary files /dev/null and b/lakonlab/models/architecture/gmflow/__pycache__/gmdit_v2.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/gmflow/__pycache__/gmflux.cpython-310.pyc b/lakonlab/models/architecture/gmflow/__pycache__/gmflux.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..e89571b198dac655a946b4777f5f3c236e2daf05
Binary files /dev/null and b/lakonlab/models/architecture/gmflow/__pycache__/gmflux.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/gmflow/__pycache__/gmqwen.cpython-310.pyc b/lakonlab/models/architecture/gmflow/__pycache__/gmqwen.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..8fcd4de06a64527af23eb0276cc7cd16ef2d6871
Binary files /dev/null and b/lakonlab/models/architecture/gmflow/__pycache__/gmqwen.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/gmflow/__pycache__/gmsd3.cpython-310.pyc b/lakonlab/models/architecture/gmflow/__pycache__/gmsd3.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..7b811f0e80c62c80d0c69ee92706dc693ea5f56b
Binary files /dev/null and b/lakonlab/models/architecture/gmflow/__pycache__/gmsd3.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/gmflow/__pycache__/gmunet_ddpm.cpython-310.pyc b/lakonlab/models/architecture/gmflow/__pycache__/gmunet_ddpm.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..a4ae32f45064d6a843adf643ef6a70345ad6a142
Binary files /dev/null and b/lakonlab/models/architecture/gmflow/__pycache__/gmunet_ddpm.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/gmflow/__pycache__/spectrum_mlp.cpython-310.pyc b/lakonlab/models/architecture/gmflow/__pycache__/spectrum_mlp.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..b4b4fdebeff2387d6ff82633fac79c2692ea955f
Binary files /dev/null and b/lakonlab/models/architecture/gmflow/__pycache__/spectrum_mlp.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/gmflow/__pycache__/toymodels.cpython-310.pyc b/lakonlab/models/architecture/gmflow/__pycache__/toymodels.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..2764e464ccf50f3d81df897721f67123f0aaf0ac
Binary files /dev/null and b/lakonlab/models/architecture/gmflow/__pycache__/toymodels.cpython-310.pyc differ
diff --git a/lakonlab/models/architecture/gmflow/gm_output.py b/lakonlab/models/architecture/gmflow/gm_output.py
new file mode 100644
index 0000000000000000000000000000000000000000..fe45c63321d5d52d6195d213fe81385b140b6d59
--- /dev/null
+++ b/lakonlab/models/architecture/gmflow/gm_output.py
@@ -0,0 +1,90 @@
+import torch
+import torch.nn as nn
+from functools import partial
+from dataclasses import dataclass
+from mmcv.cnn import constant_init, xavier_init
+from diffusers.utils import BaseOutput
+
+
+@dataclass
+class GMFlowModelOutput(BaseOutput):
+ """
+ The output of GMFlow models.
+
+ Args:
+ means (`torch.Tensor` of shape `(batch_size, num_gaussians, num_channels, height, width)` or
+ `(batch_size, num_gaussians, num_channels, frame, height, width)`):
+ Gaussian mixture means.
+ logweights (`torch.Tensor` of shape `(batch_size, num_gaussians, 1, height, width)` or
+ `(batch_size, num_gaussians, 1, frame, height, width)`):
+ Gaussian mixture log-weights (logits).
+ logstds (`torch.Tensor` of shape `(batch_size, 1, 1, 1, 1)` or `(batch_size, 1, 1, 1, 1, 1)`):
+ Gaussian mixture log-standard-deviations (logstds are shared across all Gaussians and channels).
+ """
+
+ means: torch.Tensor
+ logweights: torch.Tensor
+ logstds: torch.Tensor
+
+
+class GMOutput2D(nn.Module):
+
+ def __init__(self,
+ num_gaussians,
+ out_channels,
+ embed_dim,
+ constant_logstd=None,
+ logstd_inner_dim=1024,
+ num_logstd_layers=2,
+ activation_fn='silu'):
+ super(GMOutput2D, self).__init__()
+ self.num_gaussians = num_gaussians
+ self.out_channels = out_channels
+ self.embed_dim = embed_dim
+ self.constant_logstd = constant_logstd
+
+ if constant_logstd is None:
+ if activation_fn == 'gelu-approximate':
+ act = partial(nn.GELU, approximate='tanh')
+ elif activation_fn == 'silu':
+ act = nn.SiLU
+ else:
+ raise ValueError(f'Unsupported activation function: {activation_fn}')
+
+ assert num_logstd_layers >= 1
+ in_dim = self.embed_dim
+ logstd_layers = []
+ for _ in range(num_logstd_layers - 1):
+ logstd_layers.extend([
+ act(),
+ nn.Linear(in_dim, logstd_inner_dim)])
+ in_dim = logstd_inner_dim
+ self.logstd_layers = nn.Sequential(
+ *logstd_layers,
+ act(),
+ nn.Linear(in_dim, 1))
+
+ self.init_weights()
+
+ def init_weights(self):
+ if self.constant_logstd is None:
+ for m in self.modules():
+ if isinstance(m, nn.Linear):
+ xavier_init(m, distribution='uniform')
+ constant_init(self.logstd_layers[-1], val=0)
+
+ def forward(self, x, emb):
+ bs, c, h, w = x.size()
+ means, logweights = x.split([self.num_gaussians * self.out_channels, self.num_gaussians], dim=1)
+ means = means.view(bs, self.num_gaussians, self.out_channels, h, w)
+ logweights = logweights.view(bs, self.num_gaussians, 1, h, w).log_softmax(dim=1)
+ if self.constant_logstd is None:
+ logstds = self.logstd_layers(emb).view(bs, 1, 1, 1, 1)
+ else:
+ logstds = torch.full(
+ (bs, 1, 1, 1, 1), float(self.constant_logstd),
+ dtype=x.dtype, device=x.device)
+ return GMFlowModelOutput(
+ means=means,
+ logweights=logweights,
+ logstds=logstds)
diff --git a/lakonlab/models/architecture/gmflow/gmdit.py b/lakonlab/models/architecture/gmflow/gmdit.py
new file mode 100644
index 0000000000000000000000000000000000000000..9da392ef2d0cac6f1803f0cf44f36634cf8bcdd2
--- /dev/null
+++ b/lakonlab/models/architecture/gmflow/gmdit.py
@@ -0,0 +1,268 @@
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from typing import Any, Dict, Optional
+from diffusers.models import DiTTransformer2DModel, ModelMixin
+from diffusers.models.embeddings import PatchEmbed
+from diffusers.models.normalization import AdaLayerNormZero
+from diffusers.configuration_utils import register_to_config
+from mmcv.runner import load_checkpoint
+from mmcv.cnn import constant_init, xavier_init
+from mmgen.models.builder import MODULES
+from mmgen.utils import get_root_logger
+from ..diffusers.dit import CombinedTimestepLabelEmbeddingsMod, BasicTransformerBlockMod
+from ..utils import flex_freeze
+from .gm_output import GMOutput2D
+
+
+class _GMDiTTransformer2DModel(DiTTransformer2DModel):
+
+ @register_to_config
+ def __init__(
+ self,
+ num_gaussians=16,
+ constant_logstd=None,
+ logstd_inner_dim=1024,
+ gm_num_logstd_layers=2,
+ class_dropout_prob=0.0,
+ num_attention_heads: int = 16,
+ attention_head_dim: int = 72,
+ in_channels: int = 4,
+ out_channels: Optional[int] = None,
+ num_layers: int = 28,
+ dropout: float = 0.0,
+ norm_num_groups: int = 32,
+ attention_bias: bool = True,
+ sample_size: int = 32,
+ patch_size: int = 2,
+ activation_fn: str = 'gelu-approximate',
+ num_embeds_ada_norm: Optional[int] = 1000,
+ upcast_attention: bool = False,
+ norm_type: str = 'ada_norm_zero',
+ norm_elementwise_affine: bool = False,
+ norm_eps: float = 1e-5):
+
+ super(DiTTransformer2DModel, self).__init__()
+
+ # Validate inputs.
+ if norm_type != "ada_norm_zero":
+ raise NotImplementedError(
+ f"Forward pass is not implemented when `patch_size` is not None and `norm_type` is '{norm_type}'."
+ )
+ elif norm_type == "ada_norm_zero" and num_embeds_ada_norm is None:
+ raise ValueError(
+ f"When using a `patch_size` and this `norm_type` ({norm_type}), `num_embeds_ada_norm` cannot be None."
+ )
+
+ # Set some common variables used across the board.
+ self.attention_head_dim = attention_head_dim
+ self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim
+ self.out_channels = in_channels if out_channels is None else out_channels
+ self.gm_channels = num_gaussians * (self.out_channels + 1)
+ self.gradient_checkpointing = False
+
+ # 2. Initialize the position embedding and transformer blocks.
+ self.height = self.config.sample_size
+ self.width = self.config.sample_size
+
+ self.patch_size = self.config.patch_size
+ self.pos_embed = PatchEmbed(
+ height=self.config.sample_size,
+ width=self.config.sample_size,
+ patch_size=self.config.patch_size,
+ in_channels=self.config.in_channels,
+ embed_dim=self.inner_dim)
+ self.emb = CombinedTimestepLabelEmbeddingsMod(
+ num_embeds_ada_norm, self.inner_dim, class_dropout_prob=0.0)
+
+ self.transformer_blocks = nn.ModuleList([
+ BasicTransformerBlockMod(
+ self.inner_dim,
+ self.config.num_attention_heads,
+ self.config.attention_head_dim,
+ dropout=self.config.dropout,
+ activation_fn=self.config.activation_fn,
+ num_embeds_ada_norm=None,
+ attention_bias=self.config.attention_bias,
+ upcast_attention=self.config.upcast_attention,
+ norm_type=norm_type,
+ norm_elementwise_affine=self.config.norm_elementwise_affine,
+ norm_eps=self.config.norm_eps)
+ for _ in range(self.config.num_layers)])
+
+ # 3. Output blocks.
+ self.norm_out = nn.LayerNorm(self.inner_dim, elementwise_affine=False, eps=1e-6)
+ self.proj_out_1 = nn.Linear(self.inner_dim, 2 * self.inner_dim)
+ self.proj_out_2 = nn.Linear(
+ self.inner_dim, self.config.patch_size * self.config.patch_size * self.gm_channels)
+
+ self.gm_out = GMOutput2D(
+ num_gaussians,
+ self.out_channels,
+ self.inner_dim,
+ constant_logstd=constant_logstd,
+ logstd_inner_dim=logstd_inner_dim,
+ num_logstd_layers=gm_num_logstd_layers)
+
+ # https://github.com/facebookresearch/DiT/blob/main/models.py
+ def init_weights(self):
+ for m in self.modules():
+ if isinstance(m, nn.Linear):
+ xavier_init(m, distribution='uniform')
+ elif isinstance(m, nn.Embedding):
+ torch.nn.init.normal_(m.weight, mean=0.0, std=0.02)
+
+ # Initialize patch_embed like nn.Linear (instead of nn.Conv2d)
+ w = self.pos_embed.proj.weight.data
+ nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
+ nn.init.constant_(self.pos_embed.proj.bias, 0)
+
+ # Zero-out adaLN modulation layers in DiT blocks
+ for m in self.modules():
+ if isinstance(m, AdaLayerNormZero):
+ constant_init(m.linear, val=0)
+
+ # Zero-out output layers
+ constant_init(self.proj_out_1, val=0)
+
+ self.gm_out.init_weights()
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ timestep: Optional[torch.LongTensor] = None,
+ class_labels: Optional[torch.LongTensor] = None,
+ cross_attention_kwargs: Dict[str, Any] = None):
+ # 1. Input
+ bs, _, h, w = hidden_states.size()
+ height, width = h // self.patch_size, w // self.patch_size
+ hidden_states = self.pos_embed(hidden_states)
+
+ cond_emb = self.emb(
+ timestep, class_labels, hidden_dtype=hidden_states.dtype)
+ dropout_enabled = self.config.class_dropout_prob > 0 and self.training
+ if dropout_enabled:
+ uncond_emb = self.emb(timestep, torch.full_like(
+ class_labels, self.config.num_embeds_ada_norm), hidden_dtype=hidden_states.dtype)
+
+ # 2. Blocks
+ for block in self.transformer_blocks:
+ if dropout_enabled:
+ dropout_mask = torch.rand((bs, 1), device=hidden_states.device) < self.config.class_dropout_prob
+ emb = torch.where(dropout_mask, uncond_emb, cond_emb)
+ else:
+ emb = cond_emb
+
+ if torch.is_grad_enabled() and self.gradient_checkpointing:
+
+ def create_custom_forward(module, return_dict=None):
+ def custom_forward(*inputs):
+ if return_dict is not None:
+ return module(*inputs, return_dict=return_dict)
+ else:
+ return module(*inputs)
+
+ return custom_forward
+
+ hidden_states = torch.utils.checkpoint.checkpoint(
+ create_custom_forward(block),
+ hidden_states,
+ None,
+ None,
+ None,
+ timestep,
+ cross_attention_kwargs,
+ class_labels,
+ emb,
+ use_reentrant=False)
+
+ else:
+ hidden_states = block(
+ hidden_states,
+ attention_mask=None,
+ encoder_hidden_states=None,
+ encoder_attention_mask=None,
+ timestep=timestep,
+ cross_attention_kwargs=cross_attention_kwargs,
+ class_labels=class_labels,
+ emb=emb)
+
+ # 3. Output
+ if dropout_enabled:
+ dropout_mask = torch.rand((bs, 1), device=hidden_states.device) < self.config.class_dropout_prob
+ emb = torch.where(dropout_mask, uncond_emb, cond_emb)
+ else:
+ emb = cond_emb
+ shift, scale = self.proj_out_1(F.silu(emb)).chunk(2, dim=1)
+ hidden_states = self.norm_out(hidden_states) * (1 + scale[:, None]) + shift[:, None]
+ hidden_states = self.proj_out_2(hidden_states).reshape(
+ bs, height, width, self.patch_size, self.patch_size, self.gm_channels
+ ).permute(0, 5, 1, 3, 2, 4).reshape(
+ bs, self.gm_channels, height * self.patch_size, width * self.patch_size)
+
+ return self.gm_out(hidden_states, cond_emb.detach())
+
+
+@MODULES.register_module()
+class GMDiTTransformer2DModel(_GMDiTTransformer2DModel):
+
+ def __init__(
+ self,
+ *args,
+ freeze=False,
+ freeze_exclude=[],
+ pretrained=None,
+ torch_dtype='float32',
+ autocast_dtype=None,
+ freeze_exclude_fp32=True,
+ freeze_exclude_autocast_dtype='float32',
+ checkpointing=True,
+ **kwargs):
+ super().__init__(*args, **kwargs)
+
+ self.init_weights(pretrained)
+
+ if autocast_dtype is not None:
+ assert torch_dtype == 'float32'
+ self.autocast_dtype = autocast_dtype
+
+ if torch_dtype is not None:
+ self.to(getattr(torch, torch_dtype))
+
+ self.freeze = freeze
+ if self.freeze:
+ flex_freeze(
+ self,
+ exclude_keys=freeze_exclude,
+ exclude_fp32=freeze_exclude_fp32,
+ exclude_autocast_dtype=freeze_exclude_autocast_dtype)
+
+ if checkpointing:
+ self.enable_gradient_checkpointing()
+
+ def init_weights(self, pretrained=None):
+ super().init_weights()
+ if pretrained is not None:
+ logger = get_root_logger()
+ load_checkpoint(self, pretrained, map_location='cpu', strict=False, logger=logger)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ timestep: Optional[torch.LongTensor] = None,
+ class_labels: Optional[torch.LongTensor] = None,
+ **kwargs):
+ if self.autocast_dtype is not None:
+ dtype = getattr(torch, self.autocast_dtype)
+ else:
+ dtype = hidden_states.dtype
+ with torch.autocast(
+ device_type='cuda',
+ enabled=self.autocast_dtype is not None,
+ dtype=dtype if self.autocast_dtype is not None else None):
+ return super().forward(
+ hidden_states.to(dtype),
+ timestep=timestep,
+ class_labels=class_labels,
+ **kwargs)
diff --git a/lakonlab/models/architecture/gmflow/gmdit_v2.py b/lakonlab/models/architecture/gmflow/gmdit_v2.py
new file mode 100644
index 0000000000000000000000000000000000000000..0f87c1b3e4ffce56ff8df29fb0abe682bc1a075b
--- /dev/null
+++ b/lakonlab/models/architecture/gmflow/gmdit_v2.py
@@ -0,0 +1,340 @@
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from typing import Any, Dict, Optional
+from diffusers.models import DiTTransformer2DModel, ModelMixin
+from diffusers.models.embeddings import PatchEmbed
+from diffusers.models.normalization import AdaLayerNormZero
+from diffusers.configuration_utils import register_to_config
+from mmcv.runner import load_checkpoint, _load_checkpoint, load_state_dict
+from mmcv.cnn import constant_init, xavier_init
+from mmgen.models.builder import MODULES
+from mmgen.utils import get_root_logger
+from ..diffusers.dit import CombinedTimestepLabelEmbeddingsMod, BasicTransformerBlockMod
+from ..utils import flex_freeze
+from .gm_output import GMFlowModelOutput
+
+
+class _GMDiTTransformer2DModelV2(DiTTransformer2DModel):
+
+ @register_to_config
+ def __init__(
+ self,
+ num_gaussians=16,
+ constant_logstd=None,
+ logstd_inner_dim=1024,
+ gm_num_logstd_layers=2,
+ class_dropout_prob=0.0,
+ num_attention_heads: int = 16,
+ attention_head_dim: int = 72,
+ in_channels: int = 4,
+ out_channels: Optional[int] = None,
+ num_layers: int = 28,
+ dropout: float = 0.0,
+ norm_num_groups: int = 32,
+ attention_bias: bool = True,
+ sample_size: int = 32,
+ patch_size: int = 2,
+ activation_fn: str = 'gelu-approximate',
+ num_embeds_ada_norm: Optional[int] = 1000,
+ upcast_attention: bool = False,
+ norm_type: str = 'ada_norm_zero',
+ norm_elementwise_affine: bool = False,
+ norm_eps: float = 1e-5):
+
+ super(DiTTransformer2DModel, self).__init__()
+
+ # Validate inputs.
+ if norm_type != "ada_norm_zero":
+ raise NotImplementedError(
+ f"Forward pass is not implemented when `patch_size` is not None and `norm_type` is '{norm_type}'."
+ )
+ elif norm_type == "ada_norm_zero" and num_embeds_ada_norm is None:
+ raise ValueError(
+ f"When using a `patch_size` and this `norm_type` ({norm_type}), `num_embeds_ada_norm` cannot be None."
+ )
+
+ # Set some common variables used across the board.
+ self.attention_head_dim = attention_head_dim
+ self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim
+ self.out_channels = in_channels if out_channels is None else out_channels
+ self.gradient_checkpointing = False
+
+ # 2. Initialize the position embedding and transformer blocks.
+ self.height = self.config.sample_size
+ self.width = self.config.sample_size
+
+ self.patch_size = self.config.patch_size
+ self.pos_embed = PatchEmbed(
+ height=self.config.sample_size,
+ width=self.config.sample_size,
+ patch_size=self.config.patch_size,
+ in_channels=self.config.in_channels,
+ embed_dim=self.inner_dim)
+ self.emb = CombinedTimestepLabelEmbeddingsMod(
+ num_embeds_ada_norm, self.inner_dim, class_dropout_prob=0.0)
+
+ self.transformer_blocks = nn.ModuleList([
+ BasicTransformerBlockMod(
+ self.inner_dim,
+ self.config.num_attention_heads,
+ self.config.attention_head_dim,
+ dropout=self.config.dropout,
+ activation_fn=self.config.activation_fn,
+ num_embeds_ada_norm=None,
+ attention_bias=self.config.attention_bias,
+ upcast_attention=self.config.upcast_attention,
+ norm_type=norm_type,
+ norm_elementwise_affine=self.config.norm_elementwise_affine,
+ norm_eps=self.config.norm_eps)
+ for _ in range(self.config.num_layers)])
+
+ # 3. Output blocks.
+ self.norm_out = nn.LayerNorm(self.inner_dim, elementwise_affine=False, eps=1e-6)
+ self.proj_out_1 = nn.Linear(self.inner_dim, 2 * self.inner_dim)
+ self.proj_out_means = nn.Linear(
+ self.inner_dim,
+ self.config.patch_size * self.config.patch_size * self.config.num_gaussians * self.out_channels)
+ self.proj_out_logweights = nn.Linear(
+ self.inner_dim,
+ self.config.patch_size * self.config.patch_size * self.config.num_gaussians)
+ self.constant_logstd = constant_logstd
+
+ if self.constant_logstd is None:
+ assert gm_num_logstd_layers >= 1
+ in_dim = self.inner_dim
+ logstd_layers = []
+ for _ in range(gm_num_logstd_layers - 1):
+ logstd_layers.extend([
+ nn.SiLU(),
+ nn.Linear(in_dim, logstd_inner_dim)])
+ in_dim = logstd_inner_dim
+ self.proj_out_logstds = nn.Sequential(
+ *logstd_layers,
+ nn.SiLU(),
+ nn.Linear(in_dim, 1))
+
+ # https://github.com/facebookresearch/DiT/blob/main/models.py
+ def init_weights(self):
+ for m in self.modules():
+ if isinstance(m, nn.Linear):
+ xavier_init(m, distribution='uniform')
+ elif isinstance(m, nn.Embedding):
+ torch.nn.init.normal_(m.weight, mean=0.0, std=0.02)
+
+ # Initialize patch_embed like nn.Linear (instead of nn.Conv2d)
+ w = self.pos_embed.proj.weight.data
+ nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
+ nn.init.constant_(self.pos_embed.proj.bias, 0)
+
+ # Zero-out adaLN modulation layers in DiT blocks
+ for m in self.modules():
+ if isinstance(m, AdaLayerNormZero):
+ constant_init(m.linear, val=0)
+
+ # Output layers
+ constant_init(self.proj_out_1, val=0)
+ constant_init(self.proj_out_means, val=0)
+ rand_noise = torch.randn((self.config.num_gaussians * self.out_channels)) * 0.1
+ self.proj_out_means.bias.data.copy_(rand_noise[None, :].expand(
+ self.config.patch_size * self.config.patch_size, -1).flatten())
+ constant_init(self.proj_out_logweights, val=0)
+ if self.constant_logstd is None:
+ constant_init(self.proj_out_logstds[-1], val=0)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ timestep: Optional[torch.LongTensor] = None,
+ class_labels: Optional[torch.LongTensor] = None,
+ cross_attention_kwargs: Dict[str, Any] = None):
+ # 1. Input
+ bs, _, h, w = hidden_states.size()
+ height, width = h // self.patch_size, w // self.patch_size
+ hidden_states = self.pos_embed(hidden_states)
+
+ cond_emb = self.emb(
+ timestep, class_labels, hidden_dtype=hidden_states.dtype)
+ dropout_enabled = self.config.class_dropout_prob > 0 and self.training
+ if dropout_enabled:
+ uncond_emb = self.emb(timestep, torch.full_like(
+ class_labels, self.config.num_embeds_ada_norm), hidden_dtype=hidden_states.dtype)
+
+ # 2. Blocks
+ for block in self.transformer_blocks:
+ if dropout_enabled:
+ dropout_mask = torch.rand((bs, 1), device=hidden_states.device) < self.config.class_dropout_prob
+ emb = torch.where(dropout_mask, uncond_emb, cond_emb)
+ else:
+ emb = cond_emb
+
+ if torch.is_grad_enabled() and self.gradient_checkpointing:
+
+ def create_custom_forward(module, return_dict=None):
+ def custom_forward(*inputs):
+ if return_dict is not None:
+ return module(*inputs, return_dict=return_dict)
+ else:
+ return module(*inputs)
+
+ return custom_forward
+
+ hidden_states = torch.utils.checkpoint.checkpoint(
+ create_custom_forward(block),
+ hidden_states,
+ None,
+ None,
+ None,
+ timestep,
+ cross_attention_kwargs,
+ class_labels,
+ emb,
+ use_reentrant=False)
+
+ else:
+ hidden_states = block(
+ hidden_states,
+ attention_mask=None,
+ encoder_hidden_states=None,
+ encoder_attention_mask=None,
+ timestep=timestep,
+ cross_attention_kwargs=cross_attention_kwargs,
+ class_labels=class_labels,
+ emb=emb)
+
+ # 3. Output
+ if dropout_enabled:
+ dropout_mask = torch.rand((bs, 1), device=hidden_states.device) < self.config.class_dropout_prob
+ emb = torch.where(dropout_mask, uncond_emb, cond_emb)
+ else:
+ emb = cond_emb
+ shift, scale = self.proj_out_1(F.silu(emb)).chunk(2, dim=1)
+ hidden_states = self.norm_out(hidden_states) * (1 + scale[:, None]) + shift[:, None]
+
+ out_means = self.proj_out_means(hidden_states).reshape(
+ bs, height, width, self.patch_size, self.patch_size, self.config.num_gaussians * self.out_channels
+ ).permute(0, 5, 1, 3, 2, 4).reshape(
+ bs, self.config.num_gaussians, self.out_channels, height * self.patch_size, width * self.patch_size)
+ out_logweights = self.proj_out_logweights(hidden_states).reshape(
+ bs, height, width, self.patch_size, self.patch_size, self.config.num_gaussians
+ ).permute(0, 5, 1, 3, 2, 4).reshape(
+ bs, self.config.num_gaussians, 1, height * self.patch_size, width * self.patch_size
+ ).log_softmax(dim=1)
+ if self.constant_logstd is None:
+ out_logstds = self.proj_out_logstds(cond_emb.detach()).reshape(bs, 1, 1, 1, 1)
+ else:
+ out_logstds = hidden_states.new_full((bs, 1, 1, 1, 1), float(self.constant_logstd))
+
+ return GMFlowModelOutput(
+ means=out_means,
+ logweights=out_logweights,
+ logstds=out_logstds)
+
+
+@MODULES.register_module()
+class GMDiTTransformer2DModelV2(_GMDiTTransformer2DModelV2):
+
+ def __init__(
+ self,
+ *args,
+ freeze=False,
+ freeze_exclude=[],
+ pretrained=None,
+ torch_dtype='float32',
+ autocast_dtype=None,
+ freeze_exclude_fp32=True,
+ freeze_exclude_autocast_dtype='float32',
+ checkpointing=True,
+ **kwargs):
+ super().__init__(*args, **kwargs)
+
+ self.init_weights(pretrained)
+
+ if autocast_dtype is not None:
+ assert torch_dtype == 'float32'
+ self.autocast_dtype = autocast_dtype
+
+ if torch_dtype is not None:
+ self.to(getattr(torch, torch_dtype))
+
+ self.freeze = freeze
+ if self.freeze:
+ flex_freeze(
+ self,
+ exclude_keys=freeze_exclude,
+ exclude_fp32=freeze_exclude_fp32,
+ exclude_autocast_dtype=freeze_exclude_autocast_dtype)
+
+ if checkpointing:
+ self.enable_gradient_checkpointing()
+
+ def init_weights(self, pretrained=None):
+ super().init_weights()
+ if pretrained is not None:
+ logger = get_root_logger()
+ # load_checkpoint(self, pretrained, map_location='cpu', strict=False, logger=logger)
+ checkpoint = _load_checkpoint(pretrained, map_location='cpu', logger=logger)
+ if 'state_dict' in checkpoint:
+ state_dict = checkpoint['state_dict']
+ else:
+ state_dict = checkpoint
+ # expand the output channels
+ p2 = self.config.patch_size * self.config.patch_size
+ ori_out_channels = p2 * self.out_channels
+ if 'proj_out_2.weight' in state_dict:
+ # if this is GMDiT V1 model with 1 Gaussian
+ if state_dict['proj_out_2.weight'].size(0) == p2 * (self.out_channels + 1):
+ state_dict['proj_out_2.weight'] = state_dict['proj_out_2.weight'].reshape(
+ p2, self.out_channels + 1, -1
+ )[:, :-1].reshape(ori_out_channels, -1)
+ if state_dict['proj_out_2.weight'].size(0) == ori_out_channels:
+ state_dict['proj_out_means.weight'] = state_dict['proj_out_2.weight'].reshape(
+ p2, 1, self.out_channels, -1
+ ).expand(-1, self.config.num_gaussians, -1, -1).reshape(
+ self.config.num_gaussians * ori_out_channels, -1)
+ del state_dict['proj_out_2.weight']
+ if 'proj_out_2.bias' in state_dict:
+ # if this is GMDiT V1 model with 1 Gaussian
+ if state_dict['proj_out_2.bias'].size(0) == p2 * (self.out_channels + 1):
+ state_dict['proj_out_2.bias'] = state_dict['proj_out_2.bias'].reshape(
+ p2, self.out_channels + 1
+ )[:, :-1].reshape(ori_out_channels)
+ if state_dict['proj_out_2.bias'].size(0) == ori_out_channels:
+ state_dict['proj_out_means.bias'] = state_dict['proj_out_2.bias'].reshape(
+ p2, 1, self.out_channels
+ ).expand(-1, self.config.num_gaussians, -1).reshape(
+ self.config.num_gaussians * ori_out_channels)
+ rand_noise = torch.randn(
+ (self.config.num_gaussians * self.out_channels),
+ dtype=state_dict['proj_out_means.bias'].dtype,
+ device=state_dict['proj_out_means.bias'].device) * 0.05
+ state_dict['proj_out_means.bias'] += rand_noise[None, :].expand(p2, -1).flatten()
+ del state_dict['proj_out_2.bias']
+ if (self.constant_logstd is None
+ and 'proj_out_means.weight' in state_dict
+ and 'proj_out_means.bias' in state_dict):
+ self.proj_out_logstds[-1].bias.data = torch.full_like(
+ self.proj_out_logstds[-1].bias.data, np.log(0.05)) # reduce the initial logstd
+ load_state_dict(self, state_dict, logger=logger)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ timestep: Optional[torch.LongTensor] = None,
+ class_labels: Optional[torch.LongTensor] = None,
+ **kwargs):
+ if self.autocast_dtype is not None:
+ dtype = getattr(torch, self.autocast_dtype)
+ else:
+ dtype = hidden_states.dtype
+ with torch.autocast(
+ device_type='cuda',
+ enabled=self.autocast_dtype is not None,
+ dtype=dtype if self.autocast_dtype is not None else None):
+ return super().forward(
+ hidden_states.to(dtype),
+ timestep=timestep,
+ class_labels=class_labels,
+ **kwargs)
diff --git a/lakonlab/models/architecture/gmflow/gmflux.py b/lakonlab/models/architecture/gmflow/gmflux.py
new file mode 100644
index 0000000000000000000000000000000000000000..a4d1eadfa3d318b892a515ceea205997c41b447b
--- /dev/null
+++ b/lakonlab/models/architecture/gmflow/gmflux.py
@@ -0,0 +1,452 @@
+import numpy as np
+import torch
+import torch.nn as nn
+
+from typing import Any, Dict, Optional, Tuple
+from accelerate import init_empty_weights
+from diffusers.models import ModelMixin
+from diffusers.models.transformers.transformer_flux import (
+ FluxTransformer2DModel, FluxPosEmbed, FluxTransformerBlock, FluxSingleTransformerBlock)
+from diffusers.models.embeddings import (
+ CombinedTimestepGuidanceTextProjEmbeddings, CombinedTimestepTextProjEmbeddings)
+from diffusers.models.normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle
+from diffusers.configuration_utils import register_to_config
+from diffusers.utils import USE_PEFT_BACKEND, scale_lora_layers, unscale_lora_layers
+from peft import LoraConfig
+from mmcv.cnn import constant_init, xavier_init
+from mmgen.models.builder import MODULES
+from mmgen.utils import get_root_logger
+from lakonlab.runner.checkpoint import _load_checkpoint, load_full_state_dict
+from ..utils import flex_freeze
+from .gm_output import GMFlowModelOutput
+
+
+class _GMFluxTransformer2DModel(FluxTransformer2DModel):
+
+ @register_to_config
+ def __init__(
+ self,
+ num_gaussians=16,
+ constant_logstd=None,
+ logstd_inner_dim=1024,
+ gm_num_logstd_layers=2,
+ logweights_channels=1,
+ in_channels: int = 64,
+ out_channels: Optional[int] = None,
+ num_layers: int = 19,
+ num_single_layers: int = 38,
+ attention_head_dim: int = 128,
+ num_attention_heads: int = 24,
+ joint_attention_dim: int = 4096,
+ pooled_projection_dim: int = 768,
+ guidance_embeds: bool = False,
+ axes_dims_rope: Tuple[int, int, int] = (16, 56, 56)):
+ super(FluxTransformer2DModel, self).__init__()
+
+ self.num_gaussians = num_gaussians
+ self.logweights_channels = logweights_channels
+
+ self.out_channels = out_channels or in_channels
+ self.inner_dim = num_attention_heads * attention_head_dim
+
+ self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_dims_rope)
+
+ text_time_guidance_cls = (
+ CombinedTimestepGuidanceTextProjEmbeddings if guidance_embeds else CombinedTimestepTextProjEmbeddings
+ )
+ self.time_text_embed = text_time_guidance_cls(
+ embedding_dim=self.inner_dim, pooled_projection_dim=pooled_projection_dim
+ )
+
+ self.context_embedder = nn.Linear(joint_attention_dim, self.inner_dim)
+ self.x_embedder = nn.Linear(in_channels, self.inner_dim)
+
+ self.transformer_blocks = nn.ModuleList(
+ [
+ FluxTransformerBlock(
+ dim=self.inner_dim,
+ num_attention_heads=num_attention_heads,
+ attention_head_dim=attention_head_dim,
+ )
+ for _ in range(num_layers)
+ ]
+ )
+
+ self.single_transformer_blocks = nn.ModuleList(
+ [
+ FluxSingleTransformerBlock(
+ dim=self.inner_dim,
+ num_attention_heads=num_attention_heads,
+ attention_head_dim=attention_head_dim,
+ )
+ for _ in range(num_single_layers)
+ ]
+ )
+
+ self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6)
+ self.proj_out_means = nn.Linear(self.inner_dim, self.num_gaussians * self.out_channels)
+ self.proj_out_logweights = nn.Linear(self.inner_dim, self.num_gaussians * self.logweights_channels)
+ self.constant_logstd = constant_logstd
+
+ if self.constant_logstd is None:
+ assert gm_num_logstd_layers >= 1
+ in_dim = self.inner_dim
+ logstd_layers = []
+ for _ in range(gm_num_logstd_layers - 1):
+ logstd_layers.extend([
+ nn.SiLU(),
+ nn.Linear(in_dim, logstd_inner_dim)])
+ in_dim = logstd_inner_dim
+ self.proj_out_logstds = nn.Sequential(
+ *logstd_layers,
+ nn.SiLU(),
+ nn.Linear(in_dim, 1))
+
+ self.gradient_checkpointing = False
+
+ def init_weights(self):
+ # for m in self.modules():
+ # if isinstance(m, nn.Linear):
+ # xavier_init(m.to_empty(device='cpu'), distribution='uniform')
+ #
+ # # Zero-out adaLN modulation layers in DiT blocks
+ # for m in self.modules():
+ # if isinstance(m, (AdaLayerNormZero, AdaLayerNormZeroSingle, AdaLayerNormContinuous)):
+ # constant_init(m.linear, val=0)
+
+ # Output layers
+ constant_init(self.proj_out_means.to_empty(device='cpu'), val=0)
+ rand_noise = torch.randn((self.num_gaussians * self.out_channels // self.logweights_channels)) * 0.1
+ self.proj_out_means.bias.data.copy_(rand_noise[:, None].expand(-1, self.logweights_channels).flatten())
+ constant_init(self.proj_out_logweights.to_empty(device='cpu'), val=0)
+ if self.constant_logstd is None:
+ # logstd layers
+ for m in self.proj_out_logstds:
+ if isinstance(m, nn.Linear):
+ xavier_init(m.to_empty(device='cpu'), distribution='uniform')
+ constant_init(self.proj_out_logstds[-1], val=0)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor = None,
+ pooled_projections: torch.Tensor = None,
+ timestep: torch.Tensor = None,
+ img_ids: torch.Tensor = None,
+ txt_ids: torch.Tensor = None,
+ guidance: torch.Tensor = None,
+ joint_attention_kwargs: Optional[Dict[str, Any]] = None,
+ controlnet_block_samples=None,
+ controlnet_single_block_samples=None,
+ controlnet_blocks_repeat: bool = False):
+ if joint_attention_kwargs is not None:
+ joint_attention_kwargs = joint_attention_kwargs.copy()
+ lora_scale = joint_attention_kwargs.pop("scale", 1.0)
+ else:
+ lora_scale = 1.0
+
+ if USE_PEFT_BACKEND:
+ scale_lora_layers(self, lora_scale)
+ else:
+ assert joint_attention_kwargs is None or joint_attention_kwargs.get('scale', None) is None
+
+ hidden_states = self.x_embedder(hidden_states)
+
+ timestep = timestep.to(hidden_states.dtype) * 1000
+ if guidance is not None:
+ guidance = guidance.to(hidden_states.dtype) * 1000
+
+ temb = (
+ self.time_text_embed(timestep, pooled_projections)
+ if guidance is None
+ else self.time_text_embed(timestep, guidance, pooled_projections)
+ )
+ encoder_hidden_states = self.context_embedder(encoder_hidden_states)
+
+ ids = torch.cat((txt_ids, img_ids), dim=0)
+ image_rotary_emb = self.pos_embed(ids)
+ image_rotary_emb = tuple([x.to(hidden_states.dtype) for x in image_rotary_emb])
+
+ if joint_attention_kwargs is not None and "ip_adapter_image_embeds" in joint_attention_kwargs:
+ ip_adapter_image_embeds = joint_attention_kwargs.pop("ip_adapter_image_embeds")
+ ip_hidden_states = self.encoder_hid_proj(ip_adapter_image_embeds)
+ joint_attention_kwargs.update({"ip_hidden_states": ip_hidden_states})
+
+ for index_block, block in enumerate(self.transformer_blocks):
+ if torch.is_grad_enabled() and self.gradient_checkpointing:
+ encoder_hidden_states, hidden_states = self._gradient_checkpointing_func(
+ block,
+ hidden_states,
+ encoder_hidden_states,
+ temb,
+ image_rotary_emb,
+ joint_attention_kwargs,
+ )
+
+ else:
+ encoder_hidden_states, hidden_states = block(
+ hidden_states=hidden_states,
+ encoder_hidden_states=encoder_hidden_states,
+ temb=temb,
+ image_rotary_emb=image_rotary_emb,
+ joint_attention_kwargs=joint_attention_kwargs,
+ )
+
+ # controlnet residual
+ if controlnet_block_samples is not None:
+ interval_control = len(self.transformer_blocks) / len(controlnet_block_samples)
+ interval_control = int(np.ceil(interval_control))
+ # For Xlabs ControlNet.
+ if controlnet_blocks_repeat:
+ hidden_states = (
+ hidden_states + controlnet_block_samples[index_block % len(controlnet_block_samples)]
+ )
+ else:
+ hidden_states = hidden_states + controlnet_block_samples[index_block // interval_control]
+
+ for index_block, block in enumerate(self.single_transformer_blocks):
+ if torch.is_grad_enabled() and self.gradient_checkpointing:
+ encoder_hidden_states, hidden_states = self._gradient_checkpointing_func(
+ block,
+ hidden_states,
+ encoder_hidden_states,
+ temb,
+ image_rotary_emb,
+ joint_attention_kwargs,
+ )
+
+ else:
+ encoder_hidden_states, hidden_states = block(
+ hidden_states=hidden_states,
+ encoder_hidden_states=encoder_hidden_states,
+ temb=temb,
+ image_rotary_emb=image_rotary_emb,
+ joint_attention_kwargs=joint_attention_kwargs,
+ )
+
+ # controlnet residual
+ if controlnet_single_block_samples is not None:
+ interval_control = len(self.single_transformer_blocks) / len(controlnet_single_block_samples)
+ interval_control = int(np.ceil(interval_control))
+ hidden_states[:, encoder_hidden_states.shape[1] :, ...] = (
+ hidden_states[:, encoder_hidden_states.shape[1] :, ...]
+ + controlnet_single_block_samples[index_block // interval_control]
+ )
+
+ hidden_states = self.norm_out(hidden_states, temb)
+
+ bs, seq_len, _ = hidden_states.size()
+ out_means = self.proj_out_means(hidden_states).reshape(
+ bs, seq_len, self.num_gaussians, self.out_channels)
+ out_logweights = self.proj_out_logweights(hidden_states).reshape(
+ bs, seq_len, self.num_gaussians, self.logweights_channels).log_softmax(dim=-2)
+ if self.constant_logstd is None:
+ out_logstds = self.proj_out_logstds(temb.detach()).reshape(bs, 1, 1, 1)
+ else:
+ out_logstds = hidden_states.new_full((bs, 1, 1, 1), float(self.constant_logstd))
+
+ if USE_PEFT_BACKEND:
+ unscale_lora_layers(self, lora_scale)
+
+ return GMFlowModelOutput(
+ means=out_means,
+ logweights=out_logweights,
+ logstds=out_logstds)
+
+
+@MODULES.register_module()
+class GMFluxTransformer2DModel(_GMFluxTransformer2DModel):
+
+ def __init__(
+ self,
+ *args,
+ patch_size=2,
+ freeze=False,
+ freeze_exclude=[],
+ pretrained=None,
+ pretrained_adapter=None,
+ torch_dtype='float32',
+ autocast_dtype=None,
+ freeze_exclude_fp32=True,
+ freeze_exclude_autocast_dtype='float32',
+ checkpointing=True,
+ use_lora=False,
+ lora_target_modules=None,
+ lora_rank=16,
+ lora_dropout=0.0,
+ **kwargs):
+ with init_empty_weights():
+ super().__init__(*args, **kwargs)
+ self.patch_size = patch_size
+ assert self.patch_size * self.patch_size == self.logweights_channels
+
+ self.init_weights(pretrained, pretrained_adapter)
+
+ if autocast_dtype is not None:
+ assert torch_dtype == 'float32'
+ self.autocast_dtype = autocast_dtype
+
+ self.use_lora = use_lora
+ self.lora_target_modules = lora_target_modules
+ self.lora_rank = lora_rank
+ if self.use_lora:
+ transformer_lora_config = LoraConfig(
+ r=lora_rank,
+ lora_alpha=lora_rank,
+ init_lora_weights='gaussian',
+ target_modules=lora_target_modules,
+ lora_dropout=lora_dropout,
+ )
+ self.add_adapter(transformer_lora_config)
+
+ if torch_dtype is not None:
+ self.to(getattr(torch, torch_dtype))
+
+ self.freeze = freeze
+ if self.freeze:
+ flex_freeze(
+ self,
+ exclude_keys=freeze_exclude,
+ exclude_fp32=freeze_exclude_fp32,
+ exclude_autocast_dtype=freeze_exclude_autocast_dtype)
+
+ if checkpointing:
+ self.enable_gradient_checkpointing()
+
+ def init_weights(self, pretrained=None, pretrained_adapter=None):
+ super().init_weights()
+ if pretrained is not None:
+ logger = get_root_logger()
+ checkpoint = _load_checkpoint(pretrained, map_location='cpu', logger=logger)
+ if 'state_dict' in checkpoint:
+ state_dict = checkpoint['state_dict']
+ else:
+ state_dict = checkpoint
+ # expand the output channels
+ if 'proj_out.weight' in state_dict and state_dict['proj_out.weight'].size(0) == self.out_channels:
+ state_dict['proj_out_means.weight'] = state_dict['proj_out.weight'][None].expand(
+ self.num_gaussians, -1, -1).reshape(self.num_gaussians * self.out_channels, -1)
+ del state_dict['proj_out.weight']
+ if 'proj_out.bias' in state_dict and state_dict['proj_out.bias'].size(0) == self.out_channels:
+ state_dict['proj_out_means.bias'] = state_dict['proj_out.bias'][None].expand(
+ self.num_gaussians, -1).reshape(self.num_gaussians * self.out_channels)
+ p2 = self.patch_size * self.patch_size
+ rand_noise = torch.randn(
+ (self.num_gaussians * self.out_channels // p2),
+ dtype=state_dict['proj_out_means.bias'].dtype,
+ device=state_dict['proj_out_means.bias'].device) * 0.05
+ state_dict['proj_out_means.bias'] += rand_noise[:, None].expand(-1, p2).flatten()
+ del state_dict['proj_out.bias']
+ if (self.constant_logstd is None
+ and 'proj_out_means.weight' in state_dict
+ and 'proj_out_means.bias' in state_dict):
+ self.proj_out_logstds[-1].bias.data = torch.full_like(
+ self.proj_out_logstds[-1].bias.data, np.log(0.05)) # reduce the initial logstd
+ if pretrained_adapter is not None:
+ adapter_state_dict = _load_checkpoint(
+ pretrained_adapter, map_location='cpu', logger=logger)
+ lora_state_dict = dict()
+ for k, v in adapter_state_dict.items():
+ if 'lora' in k:
+ lora_state_dict[k] = v
+ else:
+ state_dict[k] = v
+ load_full_state_dict(self, state_dict, logger=logger, assign=True)
+ if len(lora_state_dict) > 0:
+ self.load_lora_adapter(lora_state_dict, prefix=None)
+ self.fuse_lora()
+ self.unload_lora()
+ else:
+ load_full_state_dict(self, state_dict, logger=logger, assign=True)
+
+ @staticmethod
+ def _prepare_latent_image_ids(height, width, device, dtype):
+ """
+ Copied from Diffusers
+ """
+ latent_image_ids = torch.zeros(height, width, 3)
+ latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height)[:, None]
+ latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width)[None, :]
+
+ latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape
+
+ latent_image_ids = latent_image_ids.reshape(
+ latent_image_id_height * latent_image_id_width, latent_image_id_channels)
+
+ return latent_image_ids.to(device=device, dtype=dtype)
+
+ def patchify(self, latents):
+ if self.patch_size > 1:
+ bs, c, h, w = latents.size()
+ latents = latents.reshape(
+ bs, c, h // self.patch_size, self.patch_size, w // self.patch_size, self.patch_size
+ ).permute(
+ 0, 1, 3, 5, 2, 4
+ ).reshape(
+ bs, c * self.patch_size * self.patch_size, h // self.patch_size, w // self.patch_size)
+ return latents
+
+ def unpatchify(self, gm):
+ if self.patch_size > 1:
+ bs, k, c, h, w = gm['means'].size()
+ gm['means'] = gm['means'].reshape(
+ bs, k, c // (self.patch_size * self.patch_size), self.patch_size, self.patch_size, h, w
+ ).permute(
+ 0, 1, 2, 5, 3, 6, 4
+ ).reshape(
+ bs, k, c // (self.patch_size * self.patch_size), h * self.patch_size, w * self.patch_size)
+ gm['logweights'] = gm['logweights'].reshape(
+ bs, k, 1, self.patch_size, self.patch_size, h, w
+ ).permute(
+ 0, 1, 2, 5, 3, 6, 4
+ ).reshape(
+ bs, k, 1, h * self.patch_size, w * self.patch_size)
+ return gm
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ timestep: torch.Tensor,
+ encoder_hidden_states: torch.Tensor = None,
+ pooled_projections: torch.Tensor = None,
+ mask: Optional[torch.Tensor] = None,
+ masked_image_latents: Optional[torch.Tensor] = None,
+ **kwargs):
+ hidden_states = self.patchify(hidden_states)
+ bs, c, h, w = hidden_states.size()
+ if self.autocast_dtype is not None:
+ dtype = getattr(torch, self.autocast_dtype)
+ else:
+ dtype = hidden_states.dtype
+ device = hidden_states.device
+ hidden_states = hidden_states.reshape(bs, c, h * w).permute(0, 2, 1)
+ img_ids = self._prepare_latent_image_ids(
+ h, w, device, dtype)
+ txt_ids = img_ids.new_zeros((encoder_hidden_states.shape[-2], 3))
+
+ # Flux fill
+ if mask is not None and masked_image_latents is not None:
+ hidden_states = torch.cat(
+ (hidden_states.to(dtype=dtype),
+ masked_image_latents.to(dtype=dtype),
+ mask.to(dtype=dtype)), dim=-1)
+
+ with torch.autocast(
+ device_type='cuda',
+ enabled=self.autocast_dtype is not None,
+ dtype=dtype if self.autocast_dtype is not None else None):
+ output = super().forward(
+ hidden_states=hidden_states.to(dtype),
+ encoder_hidden_states=encoder_hidden_states.to(dtype),
+ pooled_projections=pooled_projections.to(dtype),
+ timestep=timestep,
+ img_ids=img_ids,
+ txt_ids=txt_ids,
+ **kwargs)
+
+ output['means'] = output['means'].permute(0, 2, 3, 1).reshape(
+ bs, self.num_gaussians, self.out_channels, h, w)
+ output['logweights'] = output['logweights'].permute(0, 2, 3, 1).reshape(
+ bs, self.num_gaussians, self.logweights_channels, h, w)
+ output['logstds'] = output['logstds'].unsqueeze(-1) # (bs, 1, 1, 1, 1)
+ return self.unpatchify(output)
diff --git a/lakonlab/models/architecture/gmflow/gmqwen.py b/lakonlab/models/architecture/gmflow/gmqwen.py
new file mode 100644
index 0000000000000000000000000000000000000000..c54cd9f27c61468c9ce88b48e8792e0851a35121
--- /dev/null
+++ b/lakonlab/models/architecture/gmflow/gmqwen.py
@@ -0,0 +1,348 @@
+import numpy as np
+import torch
+import torch.nn as nn
+
+from typing import Any, Dict, Optional, Tuple, List
+from accelerate import init_empty_weights
+from diffusers.models import ModelMixin
+from diffusers.models.transformers.transformer_qwenimage import (
+ QwenImageTransformer2DModel, QwenEmbedRope, QwenImageTransformerBlock, QwenTimestepProjEmbeddings)
+from diffusers.models.normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle, RMSNorm
+from diffusers.configuration_utils import register_to_config
+from diffusers.utils import USE_PEFT_BACKEND, scale_lora_layers, unscale_lora_layers
+from peft import LoraConfig
+from mmcv.cnn import constant_init, xavier_init
+from mmgen.models.builder import MODULES
+from mmgen.utils import get_root_logger
+from lakonlab.runner.checkpoint import _load_checkpoint, load_full_state_dict
+from ..utils import flex_freeze
+from .gm_output import GMFlowModelOutput
+
+
+class _GMQwenImageTransformer2DModel(QwenImageTransformer2DModel):
+
+ @register_to_config
+ def __init__(
+ self,
+ num_gaussians=16,
+ constant_logstd=None,
+ logstd_inner_dim=1024,
+ gm_num_logstd_layers=2,
+ logweights_channels=1,
+ in_channels: int = 64,
+ out_channels: Optional[int] = None,
+ num_layers: int = 60,
+ attention_head_dim: int = 128,
+ num_attention_heads: int = 24,
+ joint_attention_dim: int = 3584,
+ axes_dims_rope: Tuple[int, int, int] = (16, 56, 56)):
+ super(QwenImageTransformer2DModel, self).__init__()
+
+ self.num_gaussians = num_gaussians
+ self.logweights_channels = logweights_channels
+
+ self.out_channels = out_channels or in_channels
+ self.inner_dim = num_attention_heads * attention_head_dim
+
+ self.pos_embed = QwenEmbedRope(theta=10000, axes_dim=list(axes_dims_rope), scale_rope=True)
+
+ self.time_text_embed = QwenTimestepProjEmbeddings(embedding_dim=self.inner_dim)
+
+ self.txt_norm = RMSNorm(joint_attention_dim, eps=1e-6)
+
+ self.img_in = nn.Linear(in_channels, self.inner_dim)
+ self.txt_in = nn.Linear(joint_attention_dim, self.inner_dim)
+
+ self.transformer_blocks = nn.ModuleList(
+ [
+ QwenImageTransformerBlock(
+ dim=self.inner_dim,
+ num_attention_heads=num_attention_heads,
+ attention_head_dim=attention_head_dim,
+ )
+ for _ in range(num_layers)
+ ]
+ )
+
+ self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6)
+ self.proj_out_means = nn.Linear(self.inner_dim, self.num_gaussians * self.out_channels)
+ self.proj_out_logweights = nn.Linear(self.inner_dim, self.num_gaussians * self.logweights_channels)
+ self.constant_logstd = constant_logstd
+
+ if self.constant_logstd is None:
+ assert gm_num_logstd_layers >= 1
+ in_dim = self.inner_dim
+ logstd_layers = []
+ for _ in range(gm_num_logstd_layers - 1):
+ logstd_layers.extend([
+ nn.SiLU(),
+ nn.Linear(in_dim, logstd_inner_dim)])
+ in_dim = logstd_inner_dim
+ self.proj_out_logstds = nn.Sequential(
+ *logstd_layers,
+ nn.SiLU(),
+ nn.Linear(in_dim, 1))
+
+ self.gradient_checkpointing = False
+
+ def init_weights(self):
+ # Output layers
+ constant_init(self.proj_out_means.to_empty(device='cpu'), val=0)
+ rand_noise = torch.randn((self.num_gaussians * self.out_channels // self.logweights_channels)) * 0.1
+ self.proj_out_means.bias.data.copy_(rand_noise[:, None].expand(-1, self.logweights_channels).flatten())
+ constant_init(self.proj_out_logweights.to_empty(device='cpu'), val=0)
+ if self.constant_logstd is None:
+ # logstd layers
+ for m in self.proj_out_logstds:
+ if isinstance(m, nn.Linear):
+ xavier_init(m.to_empty(device='cpu'), distribution='uniform')
+ constant_init(self.proj_out_logstds[-1], val=0)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor = None,
+ encoder_hidden_states_mask: torch.Tensor = None,
+ timestep: torch.LongTensor = None,
+ img_shapes: Optional[List[Tuple[int, int, int]]] = None,
+ txt_seq_lens: Optional[List[int]] = None,
+ attention_kwargs: Optional[Dict[str, Any]] = None):
+ if attention_kwargs is not None:
+ attention_kwargs = attention_kwargs.copy()
+ lora_scale = attention_kwargs.pop("scale", 1.0)
+ else:
+ lora_scale = 1.0
+
+ if USE_PEFT_BACKEND:
+ scale_lora_layers(self, lora_scale)
+ else:
+ assert attention_kwargs is None or attention_kwargs.get('scale', None) is None
+
+ hidden_states = self.img_in(hidden_states)
+
+ timestep = timestep.to(hidden_states.dtype)
+ encoder_hidden_states = self.txt_norm(encoder_hidden_states)
+ encoder_hidden_states = self.txt_in(encoder_hidden_states)
+
+ temb = self.time_text_embed(timestep, hidden_states)
+
+ image_rotary_emb = self.pos_embed(img_shapes, txt_seq_lens, device=hidden_states.device)
+
+ for index_block, block in enumerate(self.transformer_blocks):
+ if torch.is_grad_enabled() and self.gradient_checkpointing:
+ encoder_hidden_states, hidden_states = self._gradient_checkpointing_func(
+ block,
+ hidden_states,
+ encoder_hidden_states,
+ encoder_hidden_states_mask,
+ temb,
+ image_rotary_emb,
+ )
+
+ else:
+ encoder_hidden_states, hidden_states = block(
+ hidden_states=hidden_states,
+ encoder_hidden_states=encoder_hidden_states,
+ encoder_hidden_states_mask=encoder_hidden_states_mask,
+ temb=temb,
+ image_rotary_emb=image_rotary_emb,
+ joint_attention_kwargs=attention_kwargs,
+ )
+
+ hidden_states = self.norm_out(hidden_states, temb)
+
+ bs, seq_len, _ = hidden_states.size()
+ out_means = self.proj_out_means(hidden_states).reshape(
+ bs, seq_len, self.num_gaussians, self.out_channels)
+ out_logweights = self.proj_out_logweights(hidden_states).reshape(
+ bs, seq_len, self.num_gaussians, self.logweights_channels).log_softmax(dim=-2)
+ if self.constant_logstd is None:
+ out_logstds = self.proj_out_logstds(temb.detach()).reshape(bs, 1, 1, 1)
+ else:
+ out_logstds = hidden_states.new_full((bs, 1, 1, 1), float(self.constant_logstd))
+
+ if USE_PEFT_BACKEND:
+ unscale_lora_layers(self, lora_scale)
+
+ return GMFlowModelOutput(
+ means=out_means,
+ logweights=out_logweights,
+ logstds=out_logstds)
+
+
+@MODULES.register_module()
+class GMQwenImageTransformer2DModel(_GMQwenImageTransformer2DModel):
+
+ def __init__(
+ self,
+ *args,
+ patch_size=2,
+ freeze=False,
+ freeze_exclude=[],
+ pretrained=None,
+ pretrained_adapter=None,
+ torch_dtype='float32',
+ autocast_dtype=None,
+ freeze_exclude_fp32=True,
+ freeze_exclude_autocast_dtype='float32',
+ checkpointing=True,
+ use_lora=False,
+ lora_target_modules=None,
+ lora_rank=16,
+ lora_dropout=0.0,
+ **kwargs):
+ with init_empty_weights():
+ super().__init__(*args, **kwargs)
+ self.patch_size = patch_size
+ assert self.patch_size * self.patch_size == self.logweights_channels
+
+ self.init_weights(pretrained, pretrained_adapter)
+
+ if autocast_dtype is not None:
+ assert torch_dtype == 'float32'
+ self.autocast_dtype = autocast_dtype
+
+ self.use_lora = use_lora
+ self.lora_target_modules = lora_target_modules
+ self.lora_rank = lora_rank
+ if self.use_lora:
+ transformer_lora_config = LoraConfig(
+ r=lora_rank,
+ lora_alpha=lora_rank,
+ init_lora_weights='gaussian',
+ target_modules=lora_target_modules,
+ lora_dropout=lora_dropout,
+ )
+ self.add_adapter(transformer_lora_config)
+
+ if torch_dtype is not None:
+ self.to(getattr(torch, torch_dtype))
+
+ self.freeze = freeze
+ if self.freeze:
+ flex_freeze(
+ self,
+ exclude_keys=freeze_exclude,
+ exclude_fp32=freeze_exclude_fp32,
+ exclude_autocast_dtype=freeze_exclude_autocast_dtype)
+
+ if checkpointing:
+ self.enable_gradient_checkpointing()
+
+ def init_weights(self, pretrained=None, pretrained_adapter=None):
+ super().init_weights()
+ if pretrained is not None:
+ logger = get_root_logger()
+ checkpoint = _load_checkpoint(pretrained, map_location='cpu', logger=logger)
+ if 'state_dict' in checkpoint:
+ state_dict = checkpoint['state_dict']
+ else:
+ state_dict = checkpoint
+ # expand the output channels
+ if 'proj_out.weight' in state_dict and state_dict['proj_out.weight'].size(0) == self.out_channels:
+ state_dict['proj_out_means.weight'] = state_dict['proj_out.weight'][None].expand(
+ self.num_gaussians, -1, -1).reshape(self.num_gaussians * self.out_channels, -1)
+ del state_dict['proj_out.weight']
+ if 'proj_out.bias' in state_dict and state_dict['proj_out.bias'].size(0) == self.out_channels:
+ state_dict['proj_out_means.bias'] = state_dict['proj_out.bias'][None].expand(
+ self.num_gaussians, -1).reshape(self.num_gaussians * self.out_channels)
+ p2 = self.patch_size * self.patch_size
+ rand_noise = torch.randn(
+ (self.num_gaussians * self.out_channels // p2),
+ dtype=state_dict['proj_out_means.bias'].dtype,
+ device=state_dict['proj_out_means.bias'].device) * 0.05
+ state_dict['proj_out_means.bias'] += rand_noise[:, None].expand(-1, p2).flatten()
+ del state_dict['proj_out.bias']
+ if (self.constant_logstd is None
+ and 'proj_out_means.weight' in state_dict
+ and 'proj_out_means.bias' in state_dict):
+ self.proj_out_logstds[-1].bias.data = torch.full_like(
+ self.proj_out_logstds[-1].bias.data, np.log(0.05)) # reduce the initial logstd
+ if pretrained_adapter is not None:
+ adapter_state_dict = _load_checkpoint(
+ pretrained_adapter, map_location='cpu', logger=logger)
+ lora_state_dict = dict()
+ for k, v in adapter_state_dict.items():
+ if 'lora' in k:
+ lora_state_dict[k] = v
+ else:
+ state_dict[k] = v
+ load_full_state_dict(self, state_dict, logger=logger, assign=True)
+ if len(lora_state_dict) > 0:
+ self.load_lora_adapter(lora_state_dict, prefix=None)
+ self.fuse_lora()
+ self.unload_lora()
+ else:
+ load_full_state_dict(self, state_dict, logger=logger, assign=True)
+
+ def patchify(self, latents):
+ if self.patch_size > 1:
+ bs, c, h, w = latents.size()
+ latents = latents.reshape(
+ bs, c, h // self.patch_size, self.patch_size, w // self.patch_size, self.patch_size
+ ).permute(
+ 0, 1, 3, 5, 2, 4
+ ).reshape(
+ bs, c * self.patch_size * self.patch_size, h // self.patch_size, w // self.patch_size)
+ return latents
+
+ def unpatchify(self, gm):
+ if self.patch_size > 1:
+ bs, k, c, h, w = gm['means'].size()
+ gm['means'] = gm['means'].reshape(
+ bs, k, c // (self.patch_size * self.patch_size), self.patch_size, self.patch_size, h, w
+ ).permute(
+ 0, 1, 2, 5, 3, 6, 4
+ ).reshape(
+ bs, k, c // (self.patch_size * self.patch_size), h * self.patch_size, w * self.patch_size)
+ gm['logweights'] = gm['logweights'].reshape(
+ bs, k, 1, self.patch_size, self.patch_size, h, w
+ ).permute(
+ 0, 1, 2, 5, 3, 6, 4
+ ).reshape(
+ bs, k, 1, h * self.patch_size, w * self.patch_size)
+ return gm
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ timestep: torch.Tensor,
+ encoder_hidden_states: torch.Tensor = None,
+ encoder_hidden_states_mask: torch.Tensor = None,
+ **kwargs):
+ hidden_states = self.patchify(hidden_states)
+ bs, c, h, w = hidden_states.size()
+ if self.autocast_dtype is not None:
+ dtype = getattr(torch, self.autocast_dtype)
+ else:
+ dtype = hidden_states.dtype
+ hidden_states = hidden_states.reshape(bs, c, h * w).permute(0, 2, 1)
+ img_shapes = [[(1, h, w)]]
+ if encoder_hidden_states_mask is not None:
+ txt_seq_lens = encoder_hidden_states_mask.sum(dim=1)
+ max_txt_seq_len = txt_seq_lens.max()
+ encoder_hidden_states = encoder_hidden_states[:, :max_txt_seq_len]
+ encoder_hidden_states_mask = encoder_hidden_states_mask[:, :max_txt_seq_len]
+ txt_seq_lens = txt_seq_lens.tolist()
+ else:
+ txt_seq_lens = None
+
+ with torch.autocast(
+ device_type='cuda',
+ enabled=self.autocast_dtype is not None,
+ dtype=dtype if self.autocast_dtype is not None else None):
+ output = super().forward(
+ hidden_states=hidden_states,
+ encoder_hidden_states=encoder_hidden_states.to(dtype),
+ encoder_hidden_states_mask=encoder_hidden_states_mask,
+ timestep=timestep,
+ img_shapes=img_shapes,
+ txt_seq_lens=txt_seq_lens,
+ **kwargs)
+
+ output['means'] = output['means'].permute(0, 2, 3, 1).reshape(
+ bs, self.num_gaussians, self.out_channels, h, w)
+ output['logweights'] = output['logweights'].permute(0, 2, 3, 1).reshape(
+ bs, self.num_gaussians, self.logweights_channels, h, w)
+ output['logstds'] = output['logstds'].unsqueeze(-1) # (bs, 1, 1, 1, 1)
+ return self.unpatchify(output)
diff --git a/lakonlab/models/architecture/gmflow/gmsd3.py b/lakonlab/models/architecture/gmflow/gmsd3.py
new file mode 100644
index 0000000000000000000000000000000000000000..bd276f958d10d1d1e683d4c3327923823d739bdb
--- /dev/null
+++ b/lakonlab/models/architecture/gmflow/gmsd3.py
@@ -0,0 +1,330 @@
+import numpy as np
+import torch
+import torch.nn as nn
+
+from typing import Any, Dict, Optional, Tuple, List
+from accelerate import init_empty_weights
+from diffusers.models import ModelMixin
+from diffusers.models.transformers.transformer_sd3 import (
+ SD3Transformer2DModel, JointTransformerBlock)
+from diffusers.models.embeddings import PatchEmbed, CombinedTimestepTextProjEmbeddings
+from diffusers.models.normalization import AdaLayerNormContinuous, AdaLayerNormZero, SD35AdaLayerNormZeroX
+from diffusers.configuration_utils import register_to_config
+from diffusers.utils import USE_PEFT_BACKEND, scale_lora_layers, unscale_lora_layers
+from peft import LoraConfig
+from mmcv.cnn import constant_init, xavier_init
+from mmgen.models.builder import MODULES
+from mmgen.utils import get_root_logger
+from lakonlab.runner.checkpoint import _load_checkpoint, load_full_state_dict
+from ..utils import flex_freeze
+from .gm_output import GMFlowModelOutput
+
+
+class _GMSD3Transformer2DModel(SD3Transformer2DModel):
+
+ @register_to_config
+ def __init__(
+ self,
+ num_gaussians=16,
+ constant_logstd=None,
+ logstd_inner_dim=1024,
+ gm_num_logstd_layers=2,
+ sample_size: int = 128,
+ patch_size: int = 2,
+ in_channels: int = 16,
+ num_layers: int = 18,
+ attention_head_dim: int = 64,
+ num_attention_heads: int = 18,
+ joint_attention_dim: int = 4096,
+ caption_projection_dim: int = 1152,
+ pooled_projection_dim: int = 2048,
+ out_channels: int = 16,
+ pos_embed_max_size: int = 96,
+ dual_attention_layers: Tuple[
+ int, ...
+ ] = (), # () for sd3.0; (0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12) for sd3.5
+ qk_norm: Optional[str] = None):
+ super(SD3Transformer2DModel, self).__init__()
+
+ self.num_gaussians = num_gaussians
+
+ self.out_channels = out_channels if out_channels is not None else in_channels
+ self.inner_dim = num_attention_heads * attention_head_dim
+
+ self.pos_embed = PatchEmbed(
+ height=sample_size,
+ width=sample_size,
+ patch_size=patch_size,
+ in_channels=in_channels,
+ embed_dim=self.inner_dim,
+ pos_embed_max_size=pos_embed_max_size, # hard-code for now.
+ )
+ self.time_text_embed = CombinedTimestepTextProjEmbeddings(
+ embedding_dim=self.inner_dim, pooled_projection_dim=pooled_projection_dim
+ )
+ self.context_embedder = nn.Linear(joint_attention_dim, caption_projection_dim)
+
+ self.transformer_blocks = nn.ModuleList(
+ [
+ JointTransformerBlock(
+ dim=self.inner_dim,
+ num_attention_heads=num_attention_heads,
+ attention_head_dim=attention_head_dim,
+ context_pre_only=i == num_layers - 1,
+ qk_norm=qk_norm,
+ use_dual_attention=True if i in dual_attention_layers else False,
+ )
+ for i in range(num_layers)
+ ]
+ )
+
+ self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6)
+ self.proj_out_means = nn.Linear(
+ self.inner_dim,
+ self.config.patch_size * self.config.patch_size * self.num_gaussians * self.out_channels)
+ self.proj_out_logweights = nn.Linear(
+ self.inner_dim,
+ self.config.patch_size * self.config.patch_size * self.num_gaussians)
+ self.constant_logstd = constant_logstd
+
+ if self.constant_logstd is None:
+ assert gm_num_logstd_layers >= 1
+ in_dim = self.inner_dim
+ logstd_layers = []
+ for _ in range(gm_num_logstd_layers - 1):
+ logstd_layers.extend([
+ nn.SiLU(),
+ nn.Linear(in_dim, logstd_inner_dim)])
+ in_dim = logstd_inner_dim
+ self.proj_out_logstds = nn.Sequential(
+ *logstd_layers,
+ nn.SiLU(),
+ nn.Linear(in_dim, 1))
+
+ self.gradient_checkpointing = False
+
+ def init_weights(self):
+ # for m in self.modules():
+ # if isinstance(m, nn.Linear):
+ # xavier_init(m, distribution='uniform')
+
+ # # Initialize patch_embed like nn.Linear (instead of nn.Conv2d)
+ # w = self.pos_embed.proj.weight.data
+ # nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
+ # nn.init.constant_(self.pos_embed.proj.bias, 0)
+
+ # # Zero-out adaLN modulation layers in DiT blocks
+ # for m in self.modules():
+ # if isinstance(m, (AdaLayerNormZero, SD35AdaLayerNormZeroX, AdaLayerNormContinuous)):
+ # constant_init(m.linear, val=0)
+
+ # Output layers
+ constant_init(self.proj_out_means.to_empty(device='cpu'), val=0)
+ rand_noise = torch.randn((self.config.num_gaussians * self.out_channels)) * 0.1
+ self.proj_out_means.bias.data.copy_(rand_noise[None, :].expand(
+ self.config.patch_size * self.config.patch_size, -1).flatten())
+ constant_init(self.proj_out_logweights.to_empty(device='cpu'), val=0)
+ if self.constant_logstd is None:
+ # logstd layers
+ for m in self.proj_out_logstds:
+ if isinstance(m, nn.Linear):
+ xavier_init(m.to_empty(device='cpu'), distribution='uniform')
+ constant_init(self.proj_out_logstds[-1], val=0)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor = None,
+ pooled_projections: torch.Tensor = None,
+ timestep: torch.LongTensor = None,
+ block_controlnet_hidden_states: List = None,
+ joint_attention_kwargs: Optional[Dict[str, Any]] = None,
+ skip_layers: Optional[List[int]] = None):
+ if joint_attention_kwargs is not None:
+ joint_attention_kwargs = joint_attention_kwargs.copy()
+ lora_scale = joint_attention_kwargs.pop("scale", 1.0)
+ else:
+ lora_scale = 1.0
+
+ if USE_PEFT_BACKEND:
+ scale_lora_layers(self, lora_scale)
+ else:
+ assert joint_attention_kwargs is None or joint_attention_kwargs.get('scale', None) is None
+
+ bs, _, h, w = hidden_states.size()
+ height, width = h // self.patch_size, w // self.patch_size
+
+ hidden_states = self.pos_embed(hidden_states) # takes care of adding positional embeddings too.
+ temb = self.time_text_embed(timestep, pooled_projections)
+ encoder_hidden_states = self.context_embedder(encoder_hidden_states)
+
+ if joint_attention_kwargs is not None and "ip_adapter_image_embeds" in joint_attention_kwargs:
+ ip_adapter_image_embeds = joint_attention_kwargs.pop("ip_adapter_image_embeds")
+ ip_hidden_states, ip_temb = self.image_proj(ip_adapter_image_embeds, timestep)
+
+ joint_attention_kwargs.update(ip_hidden_states=ip_hidden_states, temb=ip_temb)
+
+ for index_block, block in enumerate(self.transformer_blocks):
+ # Skip specified layers
+ is_skip = True if skip_layers is not None and index_block in skip_layers else False
+
+ if torch.is_grad_enabled() and self.gradient_checkpointing and not is_skip:
+ encoder_hidden_states, hidden_states = self._gradient_checkpointing_func(
+ block,
+ hidden_states,
+ encoder_hidden_states,
+ temb,
+ joint_attention_kwargs,
+ )
+ elif not is_skip:
+ encoder_hidden_states, hidden_states = block(
+ hidden_states=hidden_states,
+ encoder_hidden_states=encoder_hidden_states,
+ temb=temb,
+ joint_attention_kwargs=joint_attention_kwargs,
+ )
+
+ # controlnet residual
+ if block_controlnet_hidden_states is not None and block.context_pre_only is False:
+ interval_control = len(self.transformer_blocks) / len(block_controlnet_hidden_states)
+ hidden_states = hidden_states + block_controlnet_hidden_states[int(index_block / interval_control)]
+
+ hidden_states = self.norm_out(hidden_states, temb)
+
+ num_gaussians = self.config.num_gaussians
+ patch_size = self.config.patch_size
+ out_means = self.proj_out_means(hidden_states).reshape(
+ bs, height, width, patch_size, patch_size, num_gaussians * self.out_channels
+ ).permute(0, 5, 1, 3, 2, 4).reshape(
+ bs, num_gaussians, self.out_channels, height * patch_size, width * patch_size)
+ out_logweights = self.proj_out_logweights(hidden_states).reshape(
+ bs, height, width, patch_size, patch_size, num_gaussians
+ ).permute(0, 5, 1, 3, 2, 4).reshape(
+ bs, num_gaussians, 1, height * patch_size, width * patch_size
+ ).log_softmax(dim=1)
+ if self.constant_logstd is None:
+ out_logstds = self.proj_out_logstds(temb.detach()).reshape(bs, 1, 1, 1, 1)
+ else:
+ out_logstds = hidden_states.new_full((bs, 1, 1, 1, 1), float(self.constant_logstd))
+
+ if USE_PEFT_BACKEND:
+ unscale_lora_layers(self, lora_scale)
+
+ return GMFlowModelOutput(
+ means=out_means,
+ logweights=out_logweights,
+ logstds=out_logstds)
+
+
+@MODULES.register_module()
+class GMSD3Transformer2DModel(_GMSD3Transformer2DModel):
+
+ def __init__(
+ self,
+ *args,
+ freeze=False,
+ freeze_exclude=[],
+ pretrained=None,
+ torch_dtype='float32',
+ autocast_dtype=None,
+ freeze_exclude_fp32=True,
+ freeze_exclude_autocast_dtype='float32',
+ checkpointing=True,
+ use_lora=False,
+ lora_target_modules=None,
+ lora_rank=16,
+ **kwargs):
+ with init_empty_weights():
+ super().__init__(*args, **kwargs)
+ self.init_weights(pretrained)
+
+ if autocast_dtype is not None:
+ assert torch_dtype == 'float32'
+ self.autocast_dtype = autocast_dtype
+
+ self.use_lora = use_lora
+ self.lora_target_modules = lora_target_modules
+ self.lora_rank = lora_rank
+ if self.use_lora:
+ transformer_lora_config = LoraConfig(
+ r=lora_rank,
+ lora_alpha=lora_rank,
+ init_lora_weights='gaussian',
+ target_modules=lora_target_modules,
+ )
+ self.add_adapter(transformer_lora_config)
+
+ if torch_dtype is not None:
+ self.to(getattr(torch, torch_dtype))
+
+ self.freeze = freeze
+ if self.freeze:
+ flex_freeze(
+ self,
+ exclude_keys=freeze_exclude,
+ exclude_fp32=freeze_exclude_fp32,
+ exclude_autocast_dtype=freeze_exclude_autocast_dtype)
+
+ if checkpointing:
+ self.enable_gradient_checkpointing()
+
+ def init_weights(self, pretrained=None):
+ super().init_weights()
+ if pretrained is not None:
+ logger = get_root_logger()
+ # load_checkpoint(self, pretrained, map_location='cpu', strict=False, logger=logger)
+ checkpoint = _load_checkpoint(pretrained, map_location='cpu', logger=logger)
+ if 'state_dict' in checkpoint:
+ state_dict = checkpoint['state_dict']
+ else:
+ state_dict = checkpoint
+ # expand the output channels
+ p2 = self.config.patch_size * self.config.patch_size
+ ori_out_channels = p2 * self.out_channels
+ if 'proj_out.weight' in state_dict:
+ if state_dict['proj_out.weight'].size(0) == ori_out_channels:
+ state_dict['proj_out_means.weight'] = state_dict['proj_out.weight'].reshape(
+ p2, 1, self.out_channels, -1
+ ).expand(-1, self.config.num_gaussians, -1, -1).reshape(
+ self.config.num_gaussians * ori_out_channels, -1)
+ del state_dict['proj_out.weight']
+ if 'proj_out.bias' in state_dict:
+ if state_dict['proj_out.bias'].size(0) == ori_out_channels:
+ state_dict['proj_out_means.bias'] = state_dict['proj_out.bias'].reshape(
+ p2, 1, self.out_channels
+ ).expand(-1, self.config.num_gaussians, -1).reshape(
+ self.config.num_gaussians * ori_out_channels)
+ rand_noise = torch.randn(
+ (self.config.num_gaussians * self.out_channels),
+ dtype=state_dict['proj_out_means.bias'].dtype,
+ device=state_dict['proj_out_means.bias'].device) * 0.05
+ state_dict['proj_out_means.bias'] += rand_noise[None, :].expand(p2, -1).flatten()
+ del state_dict['proj_out.bias']
+ if (self.constant_logstd is None
+ and 'proj_out_means.weight' in state_dict
+ and 'proj_out_means.bias' in state_dict):
+ self.proj_out_logstds[-1].bias.data = torch.full_like(
+ self.proj_out_logstds[-1].bias.data, np.log(0.05)) # reduce the initial logstd
+ load_full_state_dict(self, state_dict, logger=logger, assign=True)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ timestep: torch.Tensor,
+ encoder_hidden_states: torch.Tensor = None,
+ pooled_projections: torch.Tensor = None,
+ **kwargs):
+ if self.autocast_dtype is not None:
+ dtype = getattr(torch, self.autocast_dtype)
+ else:
+ dtype = hidden_states.dtype
+ with torch.autocast(
+ device_type='cuda',
+ enabled=self.autocast_dtype is not None,
+ dtype=dtype if self.autocast_dtype is not None else None):
+ return super().forward(
+ hidden_states=hidden_states.to(dtype),
+ encoder_hidden_states=encoder_hidden_states.to(dtype),
+ pooled_projections=pooled_projections.to(dtype),
+ timestep=timestep,
+ **kwargs)
diff --git a/lakonlab/models/architecture/gmflow/gmunet_ddpm.py b/lakonlab/models/architecture/gmflow/gmunet_ddpm.py
new file mode 100644
index 0000000000000000000000000000000000000000..11bfdcf9d049190d556d5a19b8e23bf87125cac4
--- /dev/null
+++ b/lakonlab/models/architecture/gmflow/gmunet_ddpm.py
@@ -0,0 +1,254 @@
+from copy import deepcopy
+
+import torch
+import torch.nn as nn
+
+from mmcv.cnn.bricks.conv_module import ConvModule
+from mmcv.cnn import constant_init, kaiming_init
+from mmcv.runner import load_checkpoint
+from mmgen.models.architectures.ddpm.modules import TimeEmbedding, EmbedSequential
+from mmgen.models.architectures.ddpm.denoising import DenoisingUnet
+from mmgen.models.builder import MODULES, build_module
+from mmgen.utils import get_root_logger
+from .gm_output import GMOutput2D
+
+
+@MODULES.register_module()
+class GMUnet(DenoisingUnet):
+
+ def __init__(self,
+ image_size,
+ num_gaussians=16,
+ constant_logstd=None,
+ logstd_inner_dim=1024,
+ gm_num_logstd_layers=2,
+ in_channels=3,
+ base_channels=128,
+ resblocks_per_downsample=3,
+ num_timesteps=1000,
+ use_rescale_timesteps=True,
+ dropout=0,
+ embedding_channels=-1,
+ num_classes=0,
+ channels_cfg=None,
+ groups=1,
+ norm_cfg=dict(type='GN', num_groups=32),
+ act_cfg=dict(type='SiLU', inplace=False),
+ shortcut_kernel_size=1,
+ use_scale_shift_norm=False,
+ num_heads=4,
+ time_embedding_mode='sin',
+ time_embedding_cfg=None,
+ resblock_cfg=dict(type='DenoisingResBlockMod'),
+ attention_cfg=dict(type='MultiHeadAttentionMod'),
+ downsample_conv=True,
+ upsample_conv=True,
+ downsample_cfg=dict(type='DenoisingDownsampleMod'),
+ upsample_cfg=dict(type='DenoisingUpsampleMod'),
+ attention_res=[16, 8],
+ pretrained=None):
+ super(DenoisingUnet, self).__init__()
+
+ self.num_gaussians = num_gaussians
+ self.constant_logstd = constant_logstd
+ self.logstd_inner_dim = logstd_inner_dim
+ self.gm_num_logstd_layers = gm_num_logstd_layers
+
+ self.num_classes = num_classes
+ self.num_timesteps = num_timesteps
+ self.use_rescale_timesteps = use_rescale_timesteps
+
+ out_channels = num_gaussians * (in_channels + 1)
+ self.out_channels = out_channels
+
+ # check type of image_size
+ if isinstance(image_size, list) or isinstance(image_size, tuple):
+ assert len(image_size) == 2, 'The length of `image_size` should be 2.'
+ elif isinstance(image_size, int):
+ image_size = [image_size, image_size]
+ else:
+ raise TypeError('Only support `int` and `list[int]` for `image_size`.')
+ self.image_size = image_size
+
+ if isinstance(channels_cfg, list):
+ self.channel_factor_list = channels_cfg
+ else:
+ raise ValueError('Only support list or dict for `channels_cfg`, '
+ f'receive {type(channels_cfg)}')
+
+ embedding_channels = base_channels * 4 \
+ if embedding_channels == -1 else embedding_channels
+ self.time_embedding = TimeEmbedding(
+ base_channels,
+ embedding_channels=embedding_channels,
+ embedding_mode=time_embedding_mode,
+ embedding_cfg=time_embedding_cfg,
+ act_cfg=act_cfg)
+
+ if self.num_classes != 0:
+ self.label_embedding = nn.Embedding(self.num_classes,
+ embedding_channels)
+
+ self.resblock_cfg = deepcopy(resblock_cfg)
+ self.resblock_cfg.setdefault('dropout', dropout)
+ self.resblock_cfg.setdefault('groups', groups)
+ self.resblock_cfg.setdefault('norm_cfg', norm_cfg)
+ self.resblock_cfg.setdefault('act_cfg', act_cfg)
+ self.resblock_cfg.setdefault('embedding_channels', embedding_channels)
+ self.resblock_cfg.setdefault('use_scale_shift_norm',
+ use_scale_shift_norm)
+ self.resblock_cfg.setdefault('shortcut_kernel_size',
+ shortcut_kernel_size)
+
+ # get scales of ResBlock to apply attention
+ attention_scale = [min(image_size) // int(res) for res in attention_res]
+ self.attention_cfg = deepcopy(attention_cfg)
+ self.attention_cfg.setdefault('num_heads', num_heads)
+ self.attention_cfg.setdefault('groups', groups)
+ self.attention_cfg.setdefault('norm_cfg', norm_cfg)
+
+ self.downsample_cfg = deepcopy(downsample_cfg)
+ self.downsample_cfg.setdefault('groups', groups)
+ self.downsample_cfg.setdefault('with_conv', downsample_conv)
+ self.upsample_cfg = deepcopy(upsample_cfg)
+ self.upsample_cfg.setdefault('groups', groups)
+ self.upsample_cfg.setdefault('with_conv', upsample_conv)
+
+ # init the channel scale factor
+ scale = 1
+ self.in_blocks = nn.ModuleList([
+ EmbedSequential(
+ nn.Conv2d(in_channels, base_channels, 3, 1, padding=1, groups=groups))
+ ])
+ self.in_channels_list = [base_channels]
+
+ # construct the encoder part of Unet
+ for level, factor in enumerate(self.channel_factor_list):
+ in_channels_ = base_channels if level == 0 \
+ else base_channels * self.channel_factor_list[level - 1]
+ out_channels_ = base_channels * factor
+
+ for _ in range(resblocks_per_downsample):
+ layers = [
+ build_module(self.resblock_cfg, {
+ 'in_channels': in_channels_,
+ 'out_channels': out_channels_
+ })
+ ]
+ in_channels_ = out_channels_
+
+ if scale in attention_scale:
+ layers.append(
+ build_module(self.attention_cfg,
+ {'in_channels': in_channels_}))
+
+ self.in_channels_list.append(in_channels_)
+ self.in_blocks.append(EmbedSequential(*layers))
+
+ if level != len(self.channel_factor_list) - 1:
+ self.in_blocks.append(
+ EmbedSequential(
+ build_module(self.downsample_cfg,
+ {'in_channels': in_channels_})))
+ self.in_channels_list.append(in_channels_)
+ scale *= 2
+
+ # construct the bottom part of Unet
+ self.mid_blocks = EmbedSequential(
+ build_module(self.resblock_cfg, {'in_channels': in_channels_}),
+ build_module(self.attention_cfg, {'in_channels': in_channels_}),
+ build_module(self.resblock_cfg, {'in_channels': in_channels_}),
+ )
+
+ # construct the decoder part of Unet
+ in_channels_list = deepcopy(self.in_channels_list)
+ self.out_blocks = nn.ModuleList()
+ for level, factor in enumerate(self.channel_factor_list[::-1]):
+ for idx in range(resblocks_per_downsample + 1):
+ layers = [
+ build_module(
+ self.resblock_cfg, {
+ 'in_channels':
+ in_channels_ + in_channels_list.pop(),
+ 'out_channels': base_channels * factor
+ })
+ ]
+ in_channels_ = base_channels * factor
+ if scale in attention_scale:
+ layers.append(
+ build_module(self.attention_cfg,
+ {'in_channels': in_channels_}))
+ if (level != len(self.channel_factor_list) - 1
+ and idx == resblocks_per_downsample):
+ layers.append(
+ build_module(self.upsample_cfg,
+ {'in_channels': in_channels_}))
+ scale //= 2
+ self.out_blocks.append(EmbedSequential(*layers))
+
+ self.out = ConvModule(
+ in_channels=in_channels_,
+ out_channels=out_channels,
+ kernel_size=3,
+ padding=1,
+ groups=groups,
+ act_cfg=act_cfg,
+ norm_cfg=norm_cfg,
+ bias=True,
+ order=('norm', 'act', 'conv'))
+
+ self.gm_out = GMOutput2D(
+ num_gaussians,
+ in_channels,
+ embedding_channels,
+ constant_logstd=constant_logstd,
+ logstd_inner_dim=logstd_inner_dim,
+ num_logstd_layers=gm_num_logstd_layers)
+
+ self.init_weights(pretrained)
+
+ def init_weights(self, pretrained=None):
+ if isinstance(pretrained, str):
+ logger = get_root_logger()
+ load_checkpoint(self, pretrained, strict=False, logger=logger)
+ elif pretrained is None:
+ for n, m in self.named_modules():
+ if isinstance(m, nn.Conv2d) and 'conv_2' in n:
+ constant_init(m, 0)
+ if isinstance(m, nn.Conv1d) and 'proj' in n:
+ constant_init(m, 0)
+ # Note: the output layer of the Unet cannot be zero-initialized.
+ kaiming_init(self.out.conv, nonlinearity='linear')
+ self.out.conv.weight.data *= 0.1
+ else:
+ raise TypeError('pretrained must be a str or None but'
+ f' got {type(pretrained)} instead.')
+
+ def forward(self, x_t, t, label=None):
+ if self.use_rescale_timesteps:
+ t = t.float() * (1000.0 / self.num_timesteps)
+ with torch.autocast(
+ device_type='cuda',
+ enabled=True,
+ dtype=x_t.dtype):
+ embedding = self.time_embedding(t)
+
+ if label is not None:
+ assert hasattr(self, 'label_embedding')
+ embedding = self.label_embedding(label) + embedding
+
+ h, hs = x_t, []
+ # forward downsample blocks
+ for block in self.in_blocks:
+ h = block(h, embedding)
+ hs.append(h)
+
+ # forward middle blocks
+ h = self.mid_blocks(h, embedding)
+
+ # forward upsample blocks
+ for block in self.out_blocks:
+ h = block(torch.cat([h, hs.pop()], dim=1), embedding)
+ outputs = self.out(h)
+
+ return self.gm_out(outputs, embedding.detach())
diff --git a/lakonlab/models/architecture/gmflow/spectrum_mlp.py b/lakonlab/models/architecture/gmflow/spectrum_mlp.py
new file mode 100644
index 0000000000000000000000000000000000000000..92716b8eac70d73ec7b5fa585ec396d3bbd4a078
--- /dev/null
+++ b/lakonlab/models/architecture/gmflow/spectrum_mlp.py
@@ -0,0 +1,91 @@
+import math
+import torch
+import torch.nn as nn
+from diffusers.models.modeling_utils import ModelMixin
+from diffusers.configuration_utils import register_to_config, ConfigMixin
+
+from mmcv.runner import load_checkpoint
+from mmcv.cnn import constant_init, xavier_init
+from mmgen.models.builder import MODULES
+from mmgen.utils import get_root_logger
+
+
+class _SpectrumMLP(ModelMixin, ConfigMixin):
+
+ @register_to_config
+ def __init__(
+ self,
+ base_size=(4, 32, 32), # (c, h, w)
+ layers=[64, 8]):
+ super().__init__()
+ assert len(base_size) == 3
+ mlp = []
+ in_chn = 2
+ for i, out_chn in enumerate(layers):
+ mlp.append(nn.Linear(in_chn, out_chn))
+ mlp.append(nn.SiLU())
+ in_chn = out_chn
+ mlp.append(nn.Linear(in_chn, base_size[0] * base_size[1] * base_size[2]))
+ self.mlp = nn.Sequential(*mlp)
+
+ def init_weights(self):
+ for m in self.modules():
+ if isinstance(m, nn.Linear):
+ xavier_init(m, distribution='uniform')
+ constant_init(self.mlp[-1], val=0)
+
+ def forward(self, gaussian_output):
+
+ ori_dtype = gaussian_output['mean'].dtype
+ spectral_mlp_dtype = self.dtype
+
+ output_stats = torch.stack(
+ [gaussian_output['var'].mean(dim=(-3, -2, -1)),
+ gaussian_output['mean'].var(dim=(-2, -1)).mean(dim=-1)],
+ dim=-1).to(spectral_mlp_dtype)
+
+ c, base_h, base_w = self.config.base_size
+ h, w = gaussian_output['mean'].shape[-2:]
+ batch_shape = output_stats.shape[:-1]
+ bs = batch_shape.numel()
+
+ output_stats = output_stats.reshape(bs, 2)
+ spectrum = self.mlp(output_stats).reshape(bs, c, base_h, base_w)
+
+ if h != base_h or w != base_w:
+ assert h <= base_h and w <= base_w
+ h1 = (h + 1) // 2
+ h2 = h - h1
+ w1 = (w + 1) // 2
+ w2 = w - w1
+ spectrum = torch.cat(
+ [torch.cat([spectrum[..., :h1, :w1], spectrum[..., :h1, -w2:]], dim=-1),
+ torch.cat([spectrum[..., -h2:, :w1], spectrum[..., -h2:, -w2:]], dim=-1)], dim=-2)
+
+ log_var = spectrum.flatten(-2).log_softmax(dim=-1) + math.log(h * w)
+ log_var = log_var.reshape(*batch_shape, c, h, w)
+
+ return log_var.to(ori_dtype)
+
+
+@MODULES.register_module()
+class SpectrumMLP(_SpectrumMLP):
+
+ def __init__(
+ self,
+ *args,
+ pretrained=None,
+ torch_dtype='float32',
+ **kwargs):
+
+ super().__init__(*args, **kwargs)
+
+ self.init_weights(pretrained)
+ if torch_dtype is not None:
+ self.to(getattr(torch, torch_dtype))
+
+ def init_weights(self, pretrained=None):
+ super().init_weights()
+ if pretrained is not None:
+ logger = get_root_logger()
+ load_checkpoint(self, pretrained, map_location='cpu', strict=False, logger=logger)
diff --git a/lakonlab/models/architecture/gmflow/toymodels.py b/lakonlab/models/architecture/gmflow/toymodels.py
new file mode 100644
index 0000000000000000000000000000000000000000..6bdbd03d406a023d42caa78a3d421297246a3eb9
--- /dev/null
+++ b/lakonlab/models/architecture/gmflow/toymodels.py
@@ -0,0 +1,122 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import math
+import torch
+import torch.nn as nn
+
+from mmgen.models.builder import MODULES
+from diffusers.models.embeddings import Timesteps
+
+from .gm_output import GMFlowModelOutput
+
+
+def get_1d_sincos_pos_embed(
+ embed_dim, pos, min_period=1e-3, max_period=10):
+ if embed_dim % 2 != 0:
+ raise ValueError('embed_dim must be divisible by 2')
+ half_dim = embed_dim // 2
+ period = torch.logspace(
+ math.log(min_period), math.log(max_period), half_dim, base=math.e,
+ dtype=pos.dtype, device=pos.device)
+ out = pos.unsqueeze(-1) * (2 * math.pi / period)
+ emb = torch.cat(
+ [torch.sin(out), torch.cos(out)],
+ dim=-1)
+ return emb
+
+
+def get_2d_sincos_pos_embed(
+ embed_dim, crd, min_period=1e-3, max_period=10):
+ """
+ Args:
+ embed_dim (int)
+ crd (torch.Tensor): Shape (bs, 2)
+ """
+ if embed_dim % 2 != 0:
+ raise ValueError('embed_dim must be divisible by 2')
+ bs = crd.size(0)
+ emb = get_1d_sincos_pos_embed(
+ embed_dim // 2, crd.flatten(), min_period, max_period) # (bs * 2, embed_dim // 2)
+ return emb.reshape(bs, embed_dim)
+
+
+class SinCos2DPosEmbed(nn.Module):
+ def __init__(self, num_channels=256, min_period=1e-3, max_period=10):
+ super().__init__()
+ self.num_channels = num_channels
+ self.min_period = min_period
+ self.max_period = max_period
+
+ def forward(self, hidden_states):
+ """
+ Args:
+ hidden_states (torch.Tensor): Shape (B, 2)
+
+ Returns:
+ torch.Tensor: Shape (B, num_channels)
+ """
+ return get_2d_sincos_pos_embed(self.num_channels, hidden_states, self.min_period, self.max_period)
+
+
+@MODULES.register_module()
+class GMFlowMLP2DDenoiser(nn.Module):
+ def __init__(
+ self,
+ num_gaussians=32,
+ pos_min_period=5e-3,
+ pos_max_period=50,
+ embed_dim=256,
+ hidden_dim=512,
+ constant_logstd=None,
+ num_layers=5):
+ super().__init__()
+ self.num_gaussians = num_gaussians
+ self.constant_logstd = constant_logstd
+ self.time_proj = Timesteps(num_channels=embed_dim, flip_sin_to_cos=True, downscale_freq_shift=1)
+ self.pos_emb = SinCos2DPosEmbed(num_channels=embed_dim, min_period=pos_min_period, max_period=pos_max_period)
+
+ in_dim = embed_dim * 2
+ mlp = []
+ for _ in range(num_layers):
+ mlp.append(nn.Linear(in_dim, hidden_dim))
+ mlp.append(nn.SiLU())
+ in_dim = hidden_dim
+ self.net = nn.Sequential(*mlp)
+ self.out_means = nn.Linear(hidden_dim, num_gaussians * 2)
+ self.out_logweights = nn.Linear(hidden_dim, num_gaussians)
+ if constant_logstd is None:
+ self.out_logstds = nn.Linear(hidden_dim, 1)
+
+ self.init_weights()
+
+ def init_weights(self):
+ for m in self.modules():
+ if isinstance(m, nn.Linear):
+ nn.init.kaiming_normal_(m.weight, mode='fan_in', nonlinearity='relu')
+ nn.init.zeros_(m.bias)
+ nn.init.zeros_(self.out_logweights.weight)
+ if self.constant_logstd is None:
+ nn.init.zeros_(self.out_logstds.weight)
+
+ def forward(self, hidden_states, timestep):
+ shape = hidden_states.shape
+ bs = shape[0]
+ extra_dims = shape[2:]
+
+ t_emb = self.time_proj(timestep).to(hidden_states)
+ pos_emb = self.pos_emb(hidden_states.reshape(bs, 2)).to(hidden_states)
+ embeddings = torch.cat([t_emb, pos_emb], dim=-1)
+
+ feat = self.net(embeddings)
+ means = self.out_means(feat).reshape(bs, self.num_gaussians, 2, *extra_dims)
+ logweights = self.out_logweights(feat).log_softmax(dim=-1).reshape(bs, self.num_gaussians, 1, *extra_dims)
+ if self.constant_logstd is None:
+ logstds = self.out_logstds(feat).reshape(bs, 1, 1, *extra_dims)
+ else:
+ logstds = torch.full(
+ (bs, 1, 1, *extra_dims), float(self.constant_logstd),
+ dtype=hidden_states.dtype, device=hidden_states.device)
+ return GMFlowModelOutput(
+ means=means,
+ logweights=logweights,
+ logstds=logstds)
diff --git a/lakonlab/models/architecture/utils.py b/lakonlab/models/architecture/utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..8d57299b0d018ef3845d50c28803fc35e67b9651
--- /dev/null
+++ b/lakonlab/models/architecture/utils.py
@@ -0,0 +1,81 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import torch
+import torch.nn as nn
+from mmgen.utils import get_root_logger
+from lakonlab.utils import rgetattr
+
+
+def autocast_patch(module, dtype=None, enabled=True):
+
+ def make_new_forward(old_forward, dtype, enabled):
+ def new_forward(*args, **kwargs):
+ with torch.autocast(device_type='cuda', dtype=dtype, enabled=enabled):
+ result = old_forward(*args, **kwargs)
+ return result
+
+ return new_forward
+
+ module.forward = make_new_forward(module.forward, dtype, enabled)
+
+
+def flex_freeze(module, exclude_keys=None, exclude_fp32=True, exclude_autocast_dtype='float32'):
+ module.requires_grad_(False)
+
+ if exclude_keys is not None and len(exclude_keys) > 0:
+
+ logger = get_root_logger()
+
+ # find modules
+ excluded_module_keys = set()
+ exclude_modules_names = []
+ for name, _ in module.named_modules():
+ for exclude_key in exclude_keys:
+ if exclude_key.startswith('self.'): # use full name matching
+ if exclude_key[5:] == name:
+ exclude_modules_names.append(name)
+ excluded_module_keys.add(exclude_key)
+ break
+ elif exclude_key in name: # use partial name matching
+ exclude_modules_names.append(name)
+ excluded_module_keys.add(exclude_key)
+ break
+
+ for name in exclude_modules_names:
+ m = rgetattr(module, name)
+ if exclude_fp32:
+ m.to(torch.float32)
+ autocast_patch(m, dtype=getattr(torch, exclude_autocast_dtype))
+ m.requires_grad_(True)
+
+ exclude_keys = set(exclude_keys) - excluded_module_keys
+
+ if len(exclude_keys) > 0:
+ # find parameters
+ excluded_parameter_keys = set()
+ exclude_parameters_names = []
+ for name, _ in module.named_parameters():
+ for exclude_key in exclude_keys:
+ if exclude_key.startswith('self.'): # use full name matching
+ if exclude_key[5:] == name:
+ exclude_parameters_names.append(name)
+ excluded_parameter_keys.add(exclude_key)
+ break
+ elif exclude_key in name: # use partial name matching
+ exclude_parameters_names.append(name)
+ excluded_parameter_keys.add(exclude_key)
+ break
+
+ for name in exclude_parameters_names:
+ p = rgetattr(module, name)
+ if exclude_fp32:
+ logger.warning(
+ f'Parameter autocast patching is not supported yet. '
+ f'Please ensure that parameter {name} is used in fp32 context.')
+ p.data = p.data.to(torch.float32)
+ p.requires_grad_(True)
+
+ exclude_keys = exclude_keys - excluded_parameter_keys
+
+ if len(exclude_keys) > 0:
+ logger.warning(f'Exclusion keys not found: {exclude_keys}')
diff --git a/lakonlab/models/base.py b/lakonlab/models/base.py
new file mode 100644
index 0000000000000000000000000000000000000000..37952cabacd2ee49328c85389e9fe2336a21e192
--- /dev/null
+++ b/lakonlab/models/base.py
@@ -0,0 +1,170 @@
+# Copyright (c) 2025 Hansheng Chen
+
+from abc import ABCMeta, abstractmethod
+import torch
+import torch.nn as nn
+from torch.distributed.fsdp import FullyShardedDataParallel
+try:
+ from torch.distributed.fsdp import FSDPModule
+except:
+ FSDPModule = None
+
+from lakonlab.utils import kai_zhang_clip_grad
+
+
+def chunk_list(input_list, chunks):
+ """
+ Splits a list into a specified number of chunks, similar to torch.chunk.
+
+ Args:
+ input_list (list): The list to be chunked.
+ chunks (int): The desired number of chunks.
+
+ Returns:
+ list: A list of sub-lists (chunks).
+ """
+ list_len = len(input_list)
+ assert list_len % chunks == 0
+ chunk_size = list_len // chunks
+
+ result_chunks = []
+
+ for i in range(chunks):
+ result_chunks.append(input_list[i * chunk_size:(i + 1) * chunk_size])
+
+ return result_chunks
+
+
+def chunk_data_dict(data, chunks):
+ data_splits = [dict() for _ in range(chunks)]
+ for k, v in data.items():
+ if isinstance(v, torch.Tensor):
+ assert v.size(0) % chunks == 0
+ v_splits = torch.chunk(v, chunks, dim=0)
+ elif isinstance(v, list):
+ v_splits = chunk_list(v, chunks)
+ elif isinstance(v, dict):
+ v_splits = chunk_data_dict(v, chunks)
+ else:
+ raise TypeError(
+ f'Unsupported data type {type(v)} for gradient accumulation. '
+ 'Only torch.Tensor and list are supported.')
+ for grad_step_id in range(chunks):
+ data_splits[grad_step_id][k] = v_splits[grad_step_id]
+ return data_splits
+
+
+def guess_bs(data):
+ for v in data.values():
+ if isinstance(v, torch.Tensor):
+ return v.size(0)
+ elif isinstance(v, list):
+ return len(v)
+ elif isinstance(v, dict):
+ bs = guess_bs(v)
+ if bs is not None:
+ return bs
+ return None
+
+
+class BaseModel(nn.Module, metaclass=ABCMeta):
+ """Base class for all models in the training framework. Optionally supports:
+ - Gradient accumulation
+ - Gradient clipping
+ """
+
+ def step_optimizer(self, optimizer, loss_scaler, running_status):
+ log_vars = dict()
+ for k, v in optimizer.items():
+ grad_clip = self.train_cfg.get(k + '_grad_clip', 0.0)
+ grad_clip_begin_iter = self.train_cfg.get(k + '_grad_clip_begin_iter', 0)
+ grad_clip_skip_ratio = self.train_cfg.get(k + '_grad_clip_skip_ratio', 0.0)
+ skip_step = False
+ if grad_clip > 0.0 and running_status['iteration'] >= grad_clip_begin_iter:
+ m = getattr(self, k)
+ if isinstance(m, FullyShardedDataParallel): # FSDP1
+ grad_norm = m.clip_grad_norm_(grad_clip)
+ elif FSDPModule is not None and isinstance(m, FSDPModule): # FSDP2
+ grad_norm = kai_zhang_clip_grad(m, grad_clip)
+ else:
+ grad_norm = torch.nn.utils.clip_grad_norm_(m.parameters(), grad_clip)
+ if torch.logical_or(grad_norm.isnan(), grad_norm.isinf()).item() or (
+ grad_clip_skip_ratio > 0 and grad_norm > grad_clip * grad_clip_skip_ratio):
+ grad_norm = float('nan')
+ v.zero_grad()
+ skip_step = True
+ log_vars.update({k + '_grad_norm': float(grad_norm)})
+ if not skip_step:
+ if loss_scaler is None:
+ v.step()
+ else:
+ loss_scaler.unscale_(v)
+ loss_scaler.step(v)
+ return log_vars
+
+ @abstractmethod
+ def train_minibatch(self, data, loss_scaler=None, running_status=None):
+ """Training forward/backward inside a single gradient accumulation minibatch.
+ """
+
+ def train_grad_accum(
+ self, train_minibatch_func, data_splits, optimizer, grad_accum_steps, loss_scaler=None, running_status=None):
+ log_vars = dict()
+
+ bs = 0
+ for grad_step_id in range(grad_accum_steps):
+ log_vars_single, bs_single = train_minibatch_func(
+ data_splits[grad_step_id], loss_scaler, running_status)
+ for k, v in log_vars_single.items():
+ if k in log_vars:
+ log_vars[k] += float(v)
+ else:
+ log_vars[k] = float(v)
+ bs += bs_single
+
+ if grad_accum_steps > 1:
+ norm_factor = 1 / grad_accum_steps
+ log_vars = {k: v * norm_factor for k, v in log_vars.items()}
+ for v in optimizer.values():
+ for group in v.param_groups:
+ for p in group['params']:
+ if p.grad is None:
+ continue
+ if p.grad.is_sparse:
+ p.grad._values().mul_(norm_factor)
+ else:
+ p.grad.mul_(norm_factor)
+
+ log_vars_optim = self.step_optimizer(optimizer, loss_scaler, running_status)
+ log_vars.update(log_vars_optim)
+
+ return log_vars, bs
+
+ def train_step(self, data, optimizer, loss_scaler=None, running_status=None):
+ for v in optimizer.values():
+ v.zero_grad()
+
+ _bs = guess_bs(data)
+ grad_accum_batch_size = self.train_cfg.get('grad_accum_batch_size', None)
+ if grad_accum_batch_size is not None and _bs is not None:
+ grad_accum_batch_size = max(min(grad_accum_batch_size, _bs), 1)
+ assert _bs % grad_accum_batch_size == 0, \
+ f'Data batch size {_bs} is not divisible by `grad_accum_batch_size` {grad_accum_batch_size}.'
+ grad_accum_steps = _bs // grad_accum_batch_size
+ else:
+ grad_accum_steps = 1
+ data_splits = chunk_data_dict(data, grad_accum_steps)
+
+ log_vars, bs = self.train_grad_accum(
+ self.train_minibatch,
+ data_splits,
+ optimizer,
+ grad_accum_steps,
+ loss_scaler=loss_scaler,
+ running_status=running_status)
+ if _bs is not None:
+ assert bs == _bs
+
+ outputs_dict = dict(log_vars=log_vars, num_samples=bs)
+
+ return outputs_dict
diff --git a/lakonlab/models/base_diffusion.py b/lakonlab/models/base_diffusion.py
new file mode 100644
index 0000000000000000000000000000000000000000..3f4f0a8ac2c0dae0511c686a8fca8d8480c4fa4c
--- /dev/null
+++ b/lakonlab/models/base_diffusion.py
@@ -0,0 +1,188 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import torch
+
+from abc import abstractmethod
+from copy import deepcopy
+from accelerate import init_empty_weights
+from mmgen.models.builder import build_module
+
+from .base import BaseModel
+from lakonlab.utils import clone_params, rgetattr, tie_untrained_submodules
+
+
+def train_fwd_bwd(model, args, kwargs, loss_scaler=None):
+ is_multistep = rgetattr(model, 'is_multistep', False)
+
+ if is_multistep:
+ step_states, log_vars = model(*args, return_step_states=True, **kwargs)
+ loss = 0
+ step_id = 0
+ while not step_states['terminate']:
+ step_loss, step_log_vars, step_states = model(
+ *args, return_loss=True, step_states=step_states, **kwargs)
+ if step_states['detachable']:
+ step_loss.backward() if loss_scaler is None else loss_scaler.scale(step_loss).backward()
+ step_loss.detach_()
+ loss = loss + step_loss
+ for k, v in step_log_vars.items():
+ if k in log_vars:
+ log_vars[k] += v
+ else:
+ log_vars[k] = v
+ step_id += 1
+
+ else:
+ loss, log_vars = model(*args, return_loss=True, **kwargs)
+
+ if isinstance(loss, torch.Tensor) and loss.requires_grad:
+ loss.backward() if loss_scaler is None else loss_scaler.scale(loss).backward()
+
+ return log_vars
+
+
+class BaseDiffusion(BaseModel):
+ """Base class providing the common training interface for diffusion models. Optionally supports:
+ - Teacher model for distillation training
+ - EMA version of the diffusion model
+ - Multi-step diffusion training
+ - Image/video patching for patch-wise GMFlow
+ """
+
+ def __init__(self,
+ diffusion=dict(type='GaussianFlow'),
+ diffusion_use_ema=False,
+ tie_ema=True,
+ teacher=None,
+ tie_teacher=False,
+ patch_size=1,
+ inference_only=False,
+ train_cfg=None,
+ test_cfg=None):
+ super().__init__()
+ # order matters: teacher must be built before diffusion for FSDP tying
+ if teacher is not None and not inference_only:
+ teacher.update(train_cfg=train_cfg, test_cfg=test_cfg)
+ self.teacher = build_module(teacher)
+ else:
+ self.teacher = None
+
+ diffusion.update(train_cfg=train_cfg, test_cfg=test_cfg)
+ self.diffusion = build_module(diffusion)
+ if self.teacher is not None and tie_teacher:
+ tie_untrained_submodules(self.diffusion, self.teacher, tie_tgt_lora_base_layer=True)
+
+ self.patch_size = patch_size
+
+ self.diffusion_use_ema = diffusion_use_ema
+ if self.diffusion_use_ema:
+ if inference_only:
+ self.diffusion_ema = self.diffusion
+ else:
+ diffusion_ema = deepcopy(diffusion)
+ if isinstance(diffusion_ema.get('denoising', None), dict):
+ diffusion_ema['denoising'].pop('pretrained', None)
+ with init_empty_weights():
+ self.diffusion_ema = build_module(diffusion_ema)
+ if tie_ema:
+ tie_untrained_submodules(self.diffusion_ema, self.diffusion)
+ clone_params(self.diffusion_ema, self.diffusion)
+
+ self.train_cfg = dict() if train_cfg is None else deepcopy(train_cfg)
+ self.test_cfg = dict() if test_cfg is None else deepcopy(test_cfg)
+
+ def patchify(self, x):
+ if isinstance(self.patch_size, int) and self.patch_size == 1:
+ return x
+ if x.dim() == 4:
+ if isinstance(self.patch_size, int):
+ ph = pw = self.patch_size
+ else:
+ assert len(self.patch_size) == 2
+ ph, pw = self.patch_size
+ bs, c, h, w = x.size()
+ x = x.reshape(
+ bs, c, h // ph, ph, w // pw, pw
+ ).permute(
+ 0, 1, 3, 5, 2, 4
+ ).reshape(
+ bs, c * ph * pw, h // ph, w // pw)
+ elif x.dim() == 5:
+ if isinstance(self.patch_size, int):
+ pt = ph = pw = self.patch_size
+ else:
+ assert len(self.patch_size) == 3
+ pt, ph, pw = self.patch_size
+ bs, c, t, h, w = x.size()
+ x = x.reshape(
+ bs, c, t // pt, pt, h // ph, ph, w // pw, pw
+ ).permute(
+ 0, 1, 3, 5, 7, 2, 4, 6
+ ).reshape(
+ bs, c * pt * ph * pw, t // pt, h // ph, w // pw)
+ else:
+ raise ValueError(f'Unsupported input dimension {x.dim()}. Expected 4 or 5 dimensions.')
+ return x
+
+ def unpatchify(self, x):
+ if isinstance(self.patch_size, int) and self.patch_size == 1:
+ return x
+ if x.dim() == 4:
+ if isinstance(self.patch_size, int):
+ ph = pw = self.patch_size
+ else:
+ assert len(self.patch_size) == 2
+ ph, pw = self.patch_size
+ bs, c, h, w = x.size()
+ x = x.reshape(
+ bs, c // (ph * pw), ph, pw, h, w
+ ).permute(
+ 0, 1, 4, 2, 5, 3
+ ).reshape(
+ bs, c // (ph * pw), h * ph, w * pw)
+ elif x.dim() == 5:
+ if isinstance(self.patch_size, int):
+ pt = ph = pw = self.patch_size
+ else:
+ assert len(self.patch_size) == 3
+ pt, ph, pw = self.patch_size
+ bs, c, t, h, w = x.size()
+ x = x.reshape(
+ bs, c // (pt * ph * pw), pt, ph, pw, t, h, w
+ ).permute(
+ 0, 1, 5, 2, 6, 3, 7, 4
+ ).reshape(
+ bs, c // (pt * ph * pw), t * pt, h * ph, w * pw)
+ else:
+ raise ValueError(f'Unsupported input dimension {x.dim()}. Expected 4 or 5 dimensions.')
+ return x
+
+ @abstractmethod
+ def _prepare_train_minibatch_args(self, data, running_status=None):
+ """
+ Prepare the arguments for the training minibatch.
+
+ Args:
+ data (dict): The input data for the training step.
+ running_status (dict): The running status for the training step.
+
+ Returns:
+ tuple: A tuple containing the batch size, diffusion arguments, and diffusion keyword arguments.
+ """
+
+ def train_minibatch(self, data, loss_scaler=None, running_status=None):
+ bs, diffusion_args, diffusion_kwargs = self._prepare_train_minibatch_args(data, running_status)
+ log_vars = train_fwd_bwd(self.diffusion, diffusion_args, diffusion_kwargs, loss_scaler)
+ return log_vars, bs
+
+ @abstractmethod
+ def val_step(self, data, test_cfg_override=dict(), **kwargs):
+ """Perform a validation step.
+
+ Args:
+ data (dict): The input data for the validation step.
+ test_cfg_override (dict): Override configuration for the test.
+
+ Returns:
+ dict: A dictionary containing the number of samples and predicted outputs.
+ """
diff --git a/lakonlab/models/diffusion_2d.py b/lakonlab/models/diffusion_2d.py
new file mode 100644
index 0000000000000000000000000000000000000000..5af669776a0aeab586615b5deeb139da70bf6545
--- /dev/null
+++ b/lakonlab/models/diffusion_2d.py
@@ -0,0 +1,69 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import torch
+
+from copy import deepcopy
+from accelerate import init_empty_weights
+from mmgen.models.builder import MODELS, build_module
+
+from .base import BaseModel
+from lakonlab.utils import clone_params, tie_untrained_submodules
+
+
+@MODELS.register_module()
+class Diffusion2D(BaseModel):
+
+ def __init__(self,
+ diffusion=dict(type='GMFlow'),
+ diffusion_use_ema=False,
+ tie_ema=True,
+ inference_only=False,
+ train_cfg=None,
+ test_cfg=None):
+ super().__init__()
+ diffusion.update(train_cfg=train_cfg, test_cfg=test_cfg)
+ self.diffusion = build_module(diffusion)
+ self.diffusion_use_ema = diffusion_use_ema
+ if self.diffusion_use_ema:
+ if inference_only:
+ self.diffusion_ema = self.diffusion
+ else:
+ diffusion_ema = deepcopy(diffusion)
+ if isinstance(diffusion_ema.get('denoising', None), dict):
+ diffusion_ema['denoising'].pop('pretrained', None)
+ with init_empty_weights():
+ self.diffusion_ema = build_module(diffusion_ema)
+ if tie_ema:
+ tie_untrained_submodules(self.diffusion_ema, self.diffusion)
+ clone_params(self.diffusion_ema, self.diffusion)
+
+ self.train_cfg = dict() if train_cfg is None else deepcopy(train_cfg)
+ self.test_cfg = dict() if test_cfg is None else deepcopy(test_cfg)
+
+ def train_minibatch(self, data, loss_scaler=None, running_status=None):
+ bs = data['x'].size(0)
+
+ loss, log_vars = self.diffusion(
+ data['x'].reshape(bs, 2, 1, 1),
+ return_loss=True)
+
+ loss.backward() if loss_scaler is None else loss_scaler.scale(loss).backward()
+
+ return log_vars, bs
+
+ def val_step(self, data, test_cfg_override=dict(), **kwargs):
+ bs = data['x'].size(0)
+ cfg = deepcopy(self.test_cfg)
+ cfg.update(test_cfg_override)
+ diffusion = self.diffusion_ema if self.diffusion_use_ema else self.diffusion
+
+ with torch.no_grad():
+ if 'noise' in data:
+ noise = data['noise'].reshape(bs, 2, 1, 1)
+ else:
+ noise = torch.randn((bs, 2, 1, 1), device=data['x'].device)
+ x_out = diffusion(
+ noise=noise,
+ test_cfg_override=test_cfg_override)
+
+ return dict(num_samples=bs, pred_x=x_out)
diff --git a/lakonlab/models/diffusions/__init__.py b/lakonlab/models/diffusions/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..5083ecd21d3ff286855bd316daec4aca8b9c7d94
--- /dev/null
+++ b/lakonlab/models/diffusions/__init__.py
@@ -0,0 +1,6 @@
+from .sampler import ContinuousTimeStepSampler
+from .gaussian_flow import GaussianFlow
+from .gmflow import GMFlow
+from .piflow import PiFlowImitation, PiFlowImitationDataFree
+
+__all__ = ['ContinuousTimeStepSampler', 'GaussianFlow', 'GMFlow', 'PiFlowImitation', 'PiFlowImitationDataFree']
diff --git a/lakonlab/models/diffusions/__pycache__/__init__.cpython-310.pyc b/lakonlab/models/diffusions/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..31361bad6b0778fb117dd55bbf4c99e1885c0531
Binary files /dev/null and b/lakonlab/models/diffusions/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/__pycache__/gaussian_flow.cpython-310.pyc b/lakonlab/models/diffusions/__pycache__/gaussian_flow.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..414f6150709ad28ea2421f094571f76ee147af13
Binary files /dev/null and b/lakonlab/models/diffusions/__pycache__/gaussian_flow.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/__pycache__/gmflow.cpython-310.pyc b/lakonlab/models/diffusions/__pycache__/gmflow.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..1855aa25c3753f8d0a8fccc2d896fea690e8cbe6
Binary files /dev/null and b/lakonlab/models/diffusions/__pycache__/gmflow.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/__pycache__/piflow.cpython-310.pyc b/lakonlab/models/diffusions/__pycache__/piflow.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..37db196c2474f7fc39c30940b460845f4b620fa7
Binary files /dev/null and b/lakonlab/models/diffusions/__pycache__/piflow.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/__pycache__/sampler.cpython-310.pyc b/lakonlab/models/diffusions/__pycache__/sampler.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..c6516f1e18cd95102d04f2a76b23129ca10a7fe2
Binary files /dev/null and b/lakonlab/models/diffusions/__pycache__/sampler.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/gaussian_flow.py b/lakonlab/models/diffusions/gaussian_flow.py
new file mode 100644
index 0000000000000000000000000000000000000000..a24d0a120a1a1891cbe8081f8b0f89d1384fb11f
--- /dev/null
+++ b/lakonlab/models/diffusions/gaussian_flow.py
@@ -0,0 +1,270 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import sys
+import inspect
+import torch
+import torch.nn as nn
+import mmcv
+import diffusers
+
+from copy import deepcopy
+from mmcv.runner.fp16_utils import force_fp32
+from mmgen.models.architectures.common import get_module_device
+from mmgen.models.builder import MODULES, build_module
+
+from . import schedulers
+
+
+@torch.jit.script
+def guidance_jit(pos_mean, neg_mean, guidance_scale: float, orthogonal: bool = False):
+ bias = (pos_mean - neg_mean) * (guidance_scale - 1)
+ if orthogonal:
+ dim = list(range(1, pos_mean.dim()))
+ bias = bias - (bias * pos_mean).mean(
+ dim=dim, keepdim=True
+ ) / (pos_mean * pos_mean).mean(dim=dim, keepdim=True).clamp(min=1e-6) * pos_mean
+ return bias
+
+
+@MODULES.register_module()
+class GaussianFlow(nn.Module):
+
+ def __init__(self,
+ denoising=None,
+ flow_loss=None,
+ num_timesteps=1000,
+ timestep_sampler=dict(type='ContinuousTimeStepSampler', shift=1.0),
+ flip_model_timesteps=False,
+ denoising_mean_mode='U',
+ train_cfg=None,
+ test_cfg=None):
+ super().__init__()
+ # build denoising module in this function
+ self.num_timesteps = num_timesteps
+ self.denoising = build_module(denoising) if isinstance(denoising, dict) else denoising
+ self.denoising_mean_mode = denoising_mean_mode
+
+ self.flip_model_timesteps = flip_model_timesteps
+ self.train_cfg = deepcopy(train_cfg) if train_cfg is not None else dict()
+ self.test_cfg = deepcopy(test_cfg) if test_cfg is not None else dict()
+
+ # build sampler
+ self.timestep_sampler = build_module(
+ timestep_sampler,
+ default_args=dict(num_timesteps=num_timesteps))
+ self.flow_loss = build_module(flow_loss) if flow_loss is not None else None
+
+ def forward_transition(
+ self, x_t_src, t_src=None, t_tgt=None, sigma_src=None, sigma_tgt=None, eps=1e-6):
+ if sigma_src is None:
+ if not isinstance(t_src, torch.Tensor):
+ t_src = torch.tensor(t_src, device=x_t_src.device)
+ t_src = t_src.reshape(*t_src.size(), *((x_t_src.dim() - t_src.dim()) * [1]))
+ sigma_src = t_src / self.num_timesteps
+
+ if sigma_tgt is None:
+ if not isinstance(t_tgt, torch.Tensor):
+ t_tgt = torch.tensor(t_tgt, device=x_t_src.device)
+ t_tgt = t_tgt.reshape(*t_tgt.size(), *((x_t_src.dim() - t_tgt.dim()) * [1]))
+ sigma_tgt = t_tgt / self.num_timesteps
+
+ alpha_src = 1 - sigma_src
+ alpha_tgt = 1 - sigma_tgt
+
+ scale_trans = alpha_tgt / alpha_src.clamp(min=eps)
+ var_trans = sigma_tgt ** 2 - (scale_trans * sigma_src) ** 2
+ return dict(mean=x_t_src * scale_trans, var=var_trans), scale_trans
+
+ def sample_forward_transition(self, x_t_src, noise, t_src=None, t_tgt=None, sigma_src=None, sigma_tgt=None):
+ trans_g = self.forward_transition(
+ x_t_src, t_src=t_src, t_tgt=t_tgt, sigma_src=sigma_src, sigma_tgt=sigma_tgt)[0]
+ return trans_g['mean'] + noise * trans_g['var'].sqrt()
+
+ def sample_forward_diffusion(self, x_0, t, noise):
+ if t.dim() == 0:
+ t = t.expand(x_0.size(0))
+ std = t.reshape(*t.size(), *((x_0.dim() - t.dim()) * [1])) / self.num_timesteps
+ mean = 1 - std
+ return x_0 * mean + noise * std, mean, std
+
+ def pred(self, x_t=None, t=None, **kwargs):
+ ori_dtype = x_t.dtype
+ if hasattr(self.denoising, 'dtype'):
+ denoising_dtype = self.denoising.dtype
+ else:
+ denoising_dtype = next(self.denoising.parameters()).dtype
+ x_t = x_t.to(denoising_dtype)
+ num_batches = x_t.size(0)
+ if t.dim() == 0 or len(t) != num_batches:
+ t = t.expand(num_batches)
+ if self.flip_model_timesteps:
+ t = self.num_timesteps - t
+ output = self.denoising(x_t, t, **kwargs)
+ if isinstance(output, dict):
+ output = {k: v.to(ori_dtype) for k, v in output.items()}
+ else:
+ output = output.to(ori_dtype)
+ return output
+
+ @force_fp32()
+ def loss(self, denoising_output, x_0, noise, t, pred_mask=None):
+ if self.denoising_mean_mode.upper() == 'U':
+ if isinstance(denoising_output, dict):
+ loss_kwargs = denoising_output
+ elif isinstance(denoising_output, torch.Tensor):
+ loss_kwargs = dict(u_t_pred=denoising_output)
+ else:
+ raise AttributeError('Unknown denoising output type '
+ f'[{type(denoising_output)}].')
+ loss_kwargs.update(u_t=noise - x_0)
+ else:
+ raise AttributeError('Unknown denoising mean output type '
+ f'[{self.denoising_mean_mode}].')
+ loss_kwargs.update(
+ x_0=x_0,
+ noise=noise,
+ timesteps=t,
+ weight=pred_mask.float() if pred_mask is not None else None)
+
+ return self.flow_loss(loss_kwargs)
+
+ def forward_train(self, x_0, **kwargs):
+ device = get_module_device(self)
+
+ num_batches = x_0.size(0)
+ seq_len = x_0.shape[2:].numel() # h * w or t * h * w
+
+ t = self.timestep_sampler(num_batches, seq_len=seq_len, device=device)
+
+ noise = torch.randn_like(x_0)
+ x_t, _, _ = self.sample_forward_diffusion(x_0, t, noise)
+
+ denoising_output = self.pred(x_t, t, **kwargs)
+ loss = self.loss(denoising_output, x_0, noise, t)
+ log_vars = self.flow_loss.log_vars
+ log_vars.update(loss_diffusion=float(loss))
+
+ return loss, log_vars
+
+ def forward_test(
+ self, x_0=None, noise=None, guidance_scale=1.0,
+ test_cfg_override=dict(), show_pbar=False, **kwargs):
+ x_t = torch.randn_like(x_0) if noise is None else noise
+ num_batches = x_t.size(0)
+ ori_dtype = x_t.dtype
+ x_t = x_t.float()
+
+ cfg = deepcopy(self.test_cfg)
+ cfg.update(test_cfg_override)
+
+ sampler = cfg.get('sampler', 'FlowEulerODE')
+ sampler_class = getattr(diffusers.schedulers, sampler + 'Scheduler', None)
+ if sampler_class is None:
+ sampler_class = getattr(schedulers, sampler + 'Scheduler', None)
+ if sampler_class is None:
+ raise AttributeError(f'Cannot find sampler [{sampler}].')
+
+ sampler_kwargs = cfg.get('sampler_kwargs', {})
+ signatures = inspect.signature(sampler_class).parameters.keys()
+ for key in ['shift', 'use_dynamic_shifting', 'base_seq_len', 'max_seq_len', 'base_logshift', 'max_logshift']:
+ if key in signatures and key not in sampler_kwargs:
+ sampler_kwargs[key] = cfg.get(key, getattr(self.timestep_sampler, key))
+ if 'flow_shift' in signatures and 'use_flow_sigmas' in signatures:
+ sampler_kwargs['prediction_type'] = 'flow_prediction'
+ sampler_kwargs['use_flow_sigmas'] = True
+ if 'flow_shift' not in sampler_kwargs:
+ sampler_kwargs['flow_shift'] = cfg.get('shift', self.timestep_sampler.shift)
+ sampler = sampler_class(self.num_timesteps, **sampler_kwargs)
+
+ num_timesteps = cfg.get('num_timesteps', self.num_timesteps)
+ guidance_interval = cfg.get('guidance_interval', [0, self.num_timesteps])
+ orthogonal_guidance = cfg.get('orthogonal_guidance', False)
+ use_guidance = guidance_scale > 1.0
+
+ set_timesteps_signatures = inspect.signature(sampler.set_timesteps).parameters.keys()
+ if 'seq_len' in set_timesteps_signatures:
+ seq_len = x_t.shape[2:].numel() # h * w or t * h * w
+ sampler.set_timesteps(num_timesteps, seq_len=seq_len, device=x_t.device)
+ else:
+ sampler.set_timesteps(num_timesteps, device=x_t.device)
+
+ timesteps = sampler.timesteps
+
+ if show_pbar:
+ pbar = mmcv.ProgressBar(len(timesteps))
+
+ for t in timesteps:
+ x_t_input = x_t
+ _kwargs = kwargs
+ if use_guidance:
+ guidance_active = guidance_interval[0] <= t <= guidance_interval[1]
+ if guidance_active:
+ x_t_input = torch.cat([x_t_input, x_t_input], dim=0)
+ else:
+ _kwargs = {
+ k: v[num_batches:] if isinstance(v, torch.Tensor) and v.size(0) == 2 * num_batches else v
+ for k, v in kwargs.items()}
+
+ denoising_output = self.pred(x_t_input, t, **_kwargs)
+
+ if use_guidance and guidance_active:
+ mean_neg, mean_pos = denoising_output.chunk(2, dim=0)
+ bias = guidance_jit(mean_pos, mean_neg, guidance_scale, orthogonal_guidance)
+ denoising_output = mean_pos + bias
+
+ x_t = sampler.step(denoising_output, t, x_t, return_dict=False)[0]
+ if show_pbar:
+ pbar.update()
+
+ if show_pbar:
+ sys.stdout.write('\n')
+
+ return x_t.to(ori_dtype)
+
+ def forward_u(self, x_t=None, t=None, guidance_scale=1.0, test_cfg_override=dict(), **kwargs):
+ ori_dtype = x_t.dtype
+ x_t = x_t.float()
+ num_batches = x_t.size(0)
+
+ cfg = deepcopy(self.test_cfg)
+ cfg.update(test_cfg_override)
+
+ orthogonal_guidance = cfg.get('orthogonal_guidance', False)
+ guidance_interval = cfg.get('guidance_interval', [0, self.num_timesteps])
+
+ use_guidance = guidance_scale > 1.0
+
+ x_t_input = x_t
+ t_input = t
+ if use_guidance:
+ x_t_input = torch.cat([x_t_input, x_t_input], dim=0)
+ t_input = torch.cat([t_input, t_input], dim=0)
+
+ denoising_output = self.pred(x_t_input, t_input, **kwargs)
+
+ if use_guidance:
+ mean_neg, mean_pos = denoising_output.chunk(2, dim=0)
+ bias = guidance_jit(mean_pos, mean_neg, guidance_scale, orthogonal_guidance)
+ if guidance_interval[0] > 0 or guidance_interval[1] < self.num_timesteps:
+ guidance_active = ((t >= guidance_interval[0]) & (t <= guidance_interval[1])).reshape(
+ [num_batches] + [1] * (bias.dim() - 1))
+ bias = bias.masked_fill(~guidance_active, 0.0)
+ denoising_output = mean_pos + bias
+
+ return denoising_output.to(ori_dtype)
+
+ def forward(
+ self,
+ x_0=None,
+ return_loss=False,
+ return_u=False,
+ return_denoising_output=False,
+ **kwargs):
+ if return_loss:
+ return self.forward_train(x_0, **kwargs)
+ elif return_u:
+ return self.forward_u(**kwargs)
+ elif return_denoising_output:
+ return self.pred(**kwargs)
+ else:
+ return self.forward_test(x_0, **kwargs)
diff --git a/lakonlab/models/diffusions/gmflow.py b/lakonlab/models/diffusions/gmflow.py
new file mode 100644
index 0000000000000000000000000000000000000000..7f1c2d89e00b867e400fb364c357bc991c9a2cf3
--- /dev/null
+++ b/lakonlab/models/diffusions/gmflow.py
@@ -0,0 +1,677 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import sys
+import inspect
+import torch
+import diffusers
+import mmcv
+
+from typing import Optional
+from copy import deepcopy
+from mmgen.models.architectures.common import get_module_device
+from mmgen.models.builder import MODULES, build_module
+
+from . import GaussianFlow, schedulers
+from lakonlab.ops.gmflow_ops.gmflow_ops import (
+ gm_to_sample, gm_to_mean, gaussian_samples_to_gm_samples, gm_samples_to_gaussian_samples,
+ iso_gaussian_mul_iso_gaussian, gm_mul_iso_gaussian, gm_to_iso_gaussian, gm_transpose_t_first)
+
+
+@torch.jit.script
+def probabilistic_guidance_jit(
+ cond_mean, total_var, uncond_mean, guidance_scale: float,
+ orthogonal: float = 1.0, orthogonal_axis: Optional[torch.Tensor] = None):
+ dim = list(range(1, cond_mean.dim()))
+ bias = cond_mean - uncond_mean
+ if orthogonal > 0.0:
+ if orthogonal_axis is None:
+ orthogonal_axis = cond_mean
+ bias = bias - ((bias * orthogonal_axis).mean(
+ dim=dim, keepdim=True
+ ) / (orthogonal_axis * orthogonal_axis).mean(
+ dim=dim, keepdim=True
+ ).clamp(min=1e-6) * orthogonal_axis).mul(orthogonal)
+ bias_power = (bias * bias).mean(dim=dim, keepdim=True)
+ avg_var = total_var.mean(dim=dim, keepdim=True)
+ bias = bias * ((avg_var / bias_power.clamp(min=1e-6)).sqrt() * guidance_scale)
+ gaussian_output = dict(
+ mean=cond_mean + bias,
+ var=total_var * (1 - (guidance_scale * guidance_scale)))
+ return gaussian_output, bias, avg_var
+
+
+@torch.jit.script
+def gmflow_posterior_jit(
+ sigma_t_src, sigma_t, x_t_src, x_t,
+ gm_means, gm_vars, gm_logstds, gm_logweights,
+ eps: float, gm_dim: int = -4, channel_dim: int = -3):
+ alpha_t_src = 1 - sigma_t_src
+ alpha_t = 1 - sigma_t
+
+ sigma_t_src_sq = sigma_t_src.square()
+ sigma_t_sq = sigma_t.square()
+
+ # compute gaussian params
+ denom = (alpha_t.square() * sigma_t_src_sq - alpha_t_src.square() * sigma_t_sq).clamp(min=eps) # ζ
+ g_mean = (alpha_t * sigma_t_src_sq * x_t - alpha_t_src * sigma_t_sq * x_t_src) / denom # ν / ζ
+ g_var = sigma_t_sq * sigma_t_src_sq / denom
+
+ # gm_mul_iso_gaussian
+ g_mean = g_mean.unsqueeze(gm_dim) # (bs, *, 1, out_channels, h, w)
+ g_var = g_var.unsqueeze(gm_dim) # (bs, *, 1, 1, 1, 1)
+ g_logstd = g_var.clamp(min=eps).log() / 2
+
+ gm_diffs = gm_means - g_mean # (bs, *, num_gaussians, out_channels, h, w)
+ norm_factor = (g_var + gm_vars).clamp(min=eps)
+
+ out_means = (g_var * gm_means + gm_vars * g_mean) / norm_factor
+ # (bs, *, num_gaussians, 1, h, w)
+ logweights_delta = gm_diffs.square().sum(dim=channel_dim, keepdim=True) * (-0.5 / norm_factor)
+
+ out_logweights = (gm_logweights + logweights_delta).log_softmax(dim=gm_dim)
+ out_logstds = gm_logstds + g_logstd - 0.5 * torch.log(norm_factor)
+
+ return out_means, out_logstds, out_logweights
+
+
+@torch.jit.script
+def gmflow_posterior_mean_jit(
+ sigma_t_src, sigma_t, x_t_src, x_t,
+ gm_means, gm_vars, gm_logweights,
+ eps: float, gm_dim: int = -4, channel_dim: int = -3):
+ alpha_t_src = 1 - sigma_t_src
+ alpha_t = 1 - sigma_t
+
+ sigma_t_src_sq = sigma_t_src.square()
+ sigma_t_sq = sigma_t.square()
+
+ # compute gaussian params
+ denom = (alpha_t.square() * sigma_t_src_sq - alpha_t_src.square() * sigma_t_sq).clamp(min=eps) # ζ
+ g_mean = (alpha_t * sigma_t_src_sq * x_t - alpha_t_src * sigma_t_sq * x_t_src) / denom # ν / ζ
+ g_var = sigma_t_sq * sigma_t_src_sq / denom
+
+ # gm_mul_iso_gaussian
+ g_mean = g_mean.unsqueeze(gm_dim) # (bs, *, 1, out_channels, h, w)
+ g_var = g_var.unsqueeze(gm_dim) # (bs, *, 1, 1, 1, 1)
+
+ gm_diffs = gm_means - g_mean # (bs, *, num_gaussians, out_channels, h, w)
+ norm_factor = (g_var + gm_vars).clamp(min=eps)
+
+ out_means = (g_var * gm_means + gm_vars * g_mean) / norm_factor
+ # (bs, *, num_gaussians, 1, h, w)
+ logweights_delta = gm_diffs.square().sum(dim=channel_dim, keepdim=True) * (-0.5 / norm_factor)
+ out_weights = (gm_logweights + logweights_delta).softmax(dim=gm_dim)
+
+ out_mean = (out_means * out_weights).sum(dim=gm_dim)
+
+ return out_mean
+
+
+class GMFlowMixin:
+
+ @property
+ def time_scaling(self):
+ if hasattr(self, 'scheduler'): # for diffusers pipelines
+ return self.scheduler.config.num_train_timesteps
+ elif hasattr(self, 'num_timesteps'):
+ return self.num_timesteps
+ else:
+ raise ValueError('num_timesteps or scheduler is not defined.')
+
+ def u_to_x_0(self, denoising_output, x_t, t=None, sigma=None, eps=1e-6):
+ if isinstance(denoising_output, dict) and 'logweights' in denoising_output:
+ x_t = x_t.unsqueeze(-4)
+
+ if sigma is None:
+ if not isinstance(t, torch.Tensor):
+ t = torch.tensor(t, device=x_t.device)
+ t = t.reshape(*t.size(), *((x_t.dim() - t.dim()) * [1]))
+ sigma = t / self.time_scaling
+ else:
+ assert sigma.dim() == x_t.dim() - 1
+ sigma = sigma.unsqueeze(-4)
+
+ if isinstance(denoising_output, dict):
+ if 'logweights' in denoising_output:
+ means_x_0 = x_t - sigma * denoising_output['means']
+ logstds_x_0 = denoising_output['logstds'] + torch.log(sigma.clamp(min=eps))
+ return dict(
+ means=means_x_0,
+ logstds=logstds_x_0,
+ logweights=denoising_output['logweights'])
+ elif 'var' in denoising_output:
+ mean = x_t - sigma * denoising_output['mean']
+ var = denoising_output['var'] * sigma.square()
+ return dict(mean=mean, var=var)
+ else:
+ raise ValueError('Invalid denoising_output.')
+
+ else: # sample mode
+ x_0 = x_t - sigma * denoising_output
+ return x_0
+
+ def gmflow_posterior_mean(
+ self, gm, x_t, x_t_src, t=None, t_src=None,
+ sigma_t_src=None, sigma_t=None, eps=1e-6, prediction_type='x0',
+ checkpointing=False):
+ """
+ Fuse gmflow_posterior and gm_to_mean to avoid redundant computation.
+ """
+ assert isinstance(gm, dict)
+
+ if sigma_t_src is None:
+ if not isinstance(t_src, torch.Tensor):
+ t_src = torch.tensor(t_src, device=x_t_src.device)
+ t_src = t_src.reshape(*t_src.size(), *((x_t_src.dim() - t_src.dim()) * [1]))
+ sigma_t_src = t_src / self.time_scaling
+
+ if sigma_t is None:
+ if not isinstance(t, torch.Tensor):
+ t = torch.tensor(t, device=x_t_src.device)
+ t = t.reshape(*t.size(), *((x_t_src.dim() - t.dim()) * [1]))
+ sigma_t = t / self.time_scaling
+
+ if prediction_type == 'u':
+ gm = self.u_to_x_0(gm, x_t_src, sigma=sigma_t_src)
+ else:
+ assert prediction_type == 'x0'
+
+ gm_means = gm['means'] # (bs, *, num_gaussians, out_channels, h, w)
+ gm_logweights = gm['logweights'] # (bs, *, num_gaussians, 1, h, w)
+ if 'gm_vars' in gm:
+ gm_vars = gm['gm_vars']
+ else:
+ gm_vars = (gm['logstds'] * 2).exp() # (bs, *, 1, 1, 1, 1)
+ gm['gm_vars'] = gm_vars
+
+ if checkpointing and torch.is_grad_enabled():
+ return torch.utils.checkpoint.checkpoint(
+ gmflow_posterior_mean_jit,
+ sigma_t_src, sigma_t, x_t_src, x_t,
+ gm_means, gm_vars, gm_logweights, eps,
+ use_reentrant=True) # use_reentrant=False does not work with jit
+ else:
+ return gmflow_posterior_mean_jit(
+ sigma_t_src, sigma_t, x_t_src, x_t,
+ gm_means, gm_vars, gm_logweights, eps)
+
+ def reverse_transition(self, denoising_output, x_t_high, t_low, t_high, eps=1e-6, prediction_type='u'):
+ if isinstance(denoising_output, dict):
+ x_t_high = x_t_high.unsqueeze(-4)
+
+ bs = x_t_high.size(0)
+ if not isinstance(t_low, torch.Tensor):
+ t_low = torch.tensor(t_low, device=x_t_high.device)
+ if not isinstance(t_high, torch.Tensor):
+ t_high = torch.tensor(t_high, device=x_t_high.device)
+ if t_low.dim() == 0:
+ t_low = t_low.expand(bs)
+ if t_high.dim() == 0:
+ t_high = t_high.expand(bs)
+ t_low = t_low.reshape(*t_low.size(), *((x_t_high.dim() - t_low.dim()) * [1]))
+ t_high = t_high.reshape(*t_high.size(), *((x_t_high.dim() - t_high.dim()) * [1]))
+
+ sigma = t_high / self.time_scaling
+ sigma_to = t_low / self.time_scaling
+ alpha = 1 - sigma
+ alpha_to = 1 - sigma_to
+
+ sigma_to_over_sigma = sigma_to / sigma.clamp(min=eps)
+ alpha_over_alpha_to = alpha / alpha_to.clamp(min=eps)
+ beta_over_sigma_sq = 1 - (sigma_to_over_sigma * alpha_over_alpha_to) ** 2
+
+ c1 = sigma_to_over_sigma ** 2 * alpha_over_alpha_to
+ c2 = beta_over_sigma_sq * alpha_to
+
+ if isinstance(denoising_output, dict):
+ c3 = beta_over_sigma_sq * sigma_to ** 2
+ if prediction_type == 'u':
+ means_x_0 = x_t_high - sigma * denoising_output['means']
+ logstds_x_t_low = torch.logaddexp(
+ (denoising_output['logstds'] + torch.log((sigma * c2).clamp(min=eps))) * 2,
+ torch.log(c3.clamp(min=eps))
+ ) / 2
+ elif prediction_type == 'x0':
+ means_x_0 = denoising_output['means']
+ logstds_x_t_low = torch.logaddexp(
+ (denoising_output['logstds'] + torch.log(c2.clamp(min=eps))) * 2,
+ torch.log(c3.clamp(min=eps))
+ ) / 2
+ else:
+ raise ValueError('Invalid prediction_type.')
+ means_x_t_low = c1 * x_t_high + c2 * means_x_0
+ return dict(
+ means=means_x_t_low,
+ logstds=logstds_x_t_low,
+ logweights=denoising_output['logweights'])
+
+ else: # sample mode
+ c3_sqrt = beta_over_sigma_sq ** 0.5 * sigma_to
+ noise = torch.randn_like(denoising_output)
+ if prediction_type == 'u':
+ x_0 = x_t_high - sigma * denoising_output
+ elif prediction_type == 'x0':
+ x_0 = denoising_output
+ else:
+ raise ValueError('Invalid prediction_type.')
+ x_t_low = c1 * x_t_high + c2 * x_0 + c3_sqrt * noise
+ return x_t_low
+
+ @staticmethod
+ def gm_sample(gm, power_spectrum=None, n_samples=1, generator=None):
+ device = gm['means'].device
+ if power_spectrum is not None:
+ power_spectrum = power_spectrum.to(dtype=torch.float32).unsqueeze(-4)
+ shape = list(gm['means'].size())
+ shape[-4] = n_samples
+ half_size = shape[-1] // 2 + 1
+ spectral_samples = torch.randn(
+ shape, dtype=torch.float32, device=device, generator=generator) * (power_spectrum / 2).exp()
+ z_1 = spectral_samples.roll((-1, -1), dims=(-2, -1)).flip((-2, -1))[..., :half_size]
+ z_0 = spectral_samples[..., :half_size]
+ z_kr = torch.complex(z_0 + z_1, z_0 - z_1) / 2
+ gaussian_samples = torch.fft.irfft2(z_kr, norm='ortho')
+ samples = gaussian_samples_to_gm_samples(gm, gaussian_samples, axis_aligned=True)
+ else:
+ samples = gm_to_sample(gm, n_samples=n_samples)
+ spectral_samples = None
+ return samples, spectral_samples
+
+ def gm_to_model_output(self, gm, output_mode, power_spectrum=None, generator=None):
+ assert output_mode in ['mean', 'sample']
+ if output_mode == 'mean':
+ output = gm_to_mean(gm)
+ else: # sample
+ output = self.gm_sample(gm, power_spectrum, generator=generator)[0].squeeze(-4)
+ return output
+
+ def gm_2nd_order(
+ self, gm_output, gaussian_output, x_t, t, h,
+ guidance_scale=0.0, gm_cond=None, gaussian_cond=None, avg_var=None, cfg_bias=None, ca=0.005, cb=1.0,
+ gm2_correction_steps=0):
+ if self.prev_gm is not None:
+ dim = list(range(1, x_t.dim()))
+
+ if cfg_bias is not None:
+ gm_mean = gm_to_mean(gm_output)
+ base_gaussian = gaussian_cond
+ base_gm = gm_cond
+ else:
+ gm_mean = gaussian_output['mean']
+ base_gaussian = gaussian_output
+ base_gaussian['var'] = base_gaussian['var'].mean(dim=dim[:-3] + dim[-2:], keepdim=True) # exclude channel dim
+ avg_var = base_gaussian['var'].mean(dim=dim, keepdim=True)
+ base_gm = gm_output
+
+ mean_from_prev = self.gmflow_posterior_mean(
+ self.prev_gm, x_t, self.prev_x_t, t, self.prev_t, prediction_type='x0')
+ self.prev_gm = gm_output
+
+ # Compute rescaled 2nd-order mean difference
+ k = 0.5 * h / self.prev_h
+ prev_h_norm = self.prev_h / self.time_scaling
+ _guidance_scale = guidance_scale * cb
+ err_power = avg_var * (_guidance_scale * _guidance_scale + ca)
+ mean_diff = (gm_mean - mean_from_prev) * (
+ (1 - err_power / (prev_h_norm * prev_h_norm)).clamp(min=0).sqrt() * k)
+
+ bias = mean_diff
+ # Here we fuse probabilistic guidance bias and 2nd-order mean difference to perform one single
+ # update to the base GM, which avoids cumulative errors.
+ if cfg_bias is not None:
+ bias = mean_diff + cfg_bias
+ bias_power = bias.square().mean(dim=dim, keepdim=True)
+ bias = bias * (avg_var / bias_power.clamp(min=1e-6)).clamp(max=1).sqrt()
+
+ gaussian_output = dict(
+ mean=base_gaussian['mean'] + bias,
+ var=base_gaussian['var'] * (1 - bias_power / avg_var.clamp(min=1e-6)).clamp(min=1e-6))
+ gm_output = gm_mul_iso_gaussian(
+ base_gm, iso_gaussian_mul_iso_gaussian(gaussian_output, base_gaussian, 1, -1),
+ 1, 1)[0]
+
+ # Additional correction steps for strictly matching the 2nd order mean difference
+ if gm2_correction_steps > 0:
+ adjusted_bias = bias
+ tgt_bias = mean_diff + gm_mean - base_gaussian['mean']
+ for _ in range(gm2_correction_steps):
+ out_bias = gm_to_mean(gm_output) - base_gaussian['mean']
+ err = out_bias - tgt_bias
+ adjusted_bias = adjusted_bias - err * (
+ adjusted_bias.norm(dim=-3, keepdim=True) / out_bias.norm(dim=-3, keepdim=True).clamp(min=1e-6)
+ ).clamp(max=1)
+ adjusted_bias_power = adjusted_bias.square().mean(dim=dim, keepdim=True)
+ adjusted_bias = adjusted_bias * (avg_var / adjusted_bias_power.clamp(min=1e-6)).clamp(max=1).sqrt()
+ adjusted_gaussian_output = dict(
+ mean=base_gaussian['mean'] + adjusted_bias,
+ var=base_gaussian['var'] * (1 - adjusted_bias_power / avg_var.clamp(min=1e-6)).clamp(min=1e-6))
+ gm_output = gm_mul_iso_gaussian(
+ base_gm, iso_gaussian_mul_iso_gaussian(adjusted_gaussian_output, base_gaussian, 1, -1),
+ 1, 1)[0]
+
+ else:
+ self.prev_gm = gm_output
+
+ self.prev_x_t = x_t
+ self.prev_t = t
+ self.prev_h = h
+
+ return gm_output, gaussian_output
+
+ def init_gm_cache(self):
+ self.prev_gm = None
+ self.prev_x_t = None
+ self.prev_t = None
+ self.prev_h = None
+
+
+@MODULES.register_module()
+class GMFlow(GaussianFlow, GMFlowMixin):
+
+ def __init__(
+ self,
+ *args,
+ spectrum_net=None,
+ spectral_loss_weight=1.0,
+ **kwargs):
+ super().__init__(*args, **kwargs)
+ self.spectrum_net = build_module(spectrum_net) if spectrum_net is not None else None
+ self.spectral_loss_weight = spectral_loss_weight
+ self.intermediate_x_t = []
+ self.intermediate_x_0 = []
+
+ def loss(self, denoising_output, x_t_low, x_t_high, t_low, t_high):
+ """
+ GMFlow transition loss.
+ """
+ x_t_low = x_t_low.float()
+ x_t_high = x_t_high.float()
+ t_low = t_low.float()
+ t_high = t_high.float()
+
+ x_t_low_gm = self.reverse_transition(denoising_output, x_t_high, t_low, t_high)
+ loss_kwargs = {k: v for k, v in x_t_low_gm.items()}
+
+ loss_kwargs.update(x_t_low=x_t_low, timesteps=t_high)
+ return self.flow_loss(loss_kwargs)
+
+ def spectral_loss(self, denoising_output, x_0, x_t, t, eps=1e-6):
+ x_0 = x_0.float()
+ x_t = x_t.float()
+ t = t.float()
+
+ t = t.reshape(*t.size(), *((x_t.dim() - t.dim()) * [1]))
+ inv_sigma = self.num_timesteps / t.clamp(min=eps)
+
+ with torch.no_grad():
+ output_g = self.u_to_x_0(gm_to_iso_gaussian(denoising_output)[0], x_t, t)
+ u = (x_t - x_0) * inv_sigma
+ z_kr = gm_samples_to_gaussian_samples(
+ denoising_output, u.unsqueeze(-4), axis_aligned=True).squeeze(-4)
+ z_kr_fft = torch.fft.fft2(z_kr, norm='ortho')
+ z_kr_fft = z_kr_fft.real + z_kr_fft.imag
+
+ log_var = self.spectrum_net(output_g)
+
+ loss = z_kr_fft.square() * (torch.exp(-log_var) - 1) + log_var
+ loss = loss.mean() * (0.5 * self.spectral_loss_weight)
+ return loss
+
+ def pred(self, x_t=None, t=None, **kwargs):
+ ndim = x_t.dim()
+ assert ndim in [4, 5], f'Invalid x_t shape: {x_t.shape}. Expected 4D or 5D tensor.'
+ if ndim == 5: # (bs, t, c, h, w)
+ x_t = x_t.permute(0, 2, 1, 3, 4) # (bs, c, t, h, w)
+ output = super().pred(x_t=x_t, t=t, **kwargs)
+ if ndim == 5:
+ output = gm_transpose_t_first(output) # (bs, t, c, h, w)
+ return output
+
+ def forward_train(self, x_0, **kwargs):
+ device = get_module_device(self)
+
+ num_batches = x_0.size(0)
+ seq_len = x_0.shape[2:].numel() # h * w or t * h * w
+ ndim = x_0.dim()
+ assert ndim in [4, 5], f'Invalid x_0 shape: {x_0.shape}. Expected 4D or 5D tensor.'
+ if ndim == 5: # (bs, c, t, h, w)
+ x_0 = x_0.permute(0, 2, 1, 3, 4) # (bs, t, c, h, w)
+
+ trans_ratio = self.train_cfg.get('trans_ratio', 1.0)
+ eps = self.train_cfg.get('eps', 1e-4)
+
+ t_high = self.timestep_sampler(
+ num_batches, seq_len=seq_len).to(device).clamp(min=eps, max=self.num_timesteps)
+ t_low = t_high * (1 - trans_ratio)
+ t_low = torch.minimum(t_low, t_high - eps).clamp(min=0)
+
+ noise = torch.randn((num_batches * 2, *x_0.shape[1:]), device=device, dtype=x_0.dtype)
+ noise_0, noise_1 = torch.chunk(noise, 2, dim=0)
+
+ x_t_low, _, _ = self.sample_forward_diffusion(x_0, t_low, noise_0)
+ x_t_high = self.sample_forward_transition(x_t_low, noise_1, t_src=t_low, t_tgt=t_high)
+
+ denoising_output = self.pred(x_t_high, t_high, **kwargs)
+ loss = self.loss(denoising_output, x_t_low, x_t_high, t_low, t_high)
+ log_vars = self.flow_loss.log_vars
+ log_vars.update(loss_transition=float(loss))
+
+ if self.spectrum_net is not None:
+ # Note: only support 2D power spectrum for now.
+ loss_spectral = self.spectral_loss(denoising_output, x_0, x_t_high, t_high)
+ log_vars.update(loss_spectral=float(loss_spectral))
+ loss = loss + loss_spectral
+
+ return loss, log_vars
+
+ def forward_test(
+ self, x_0=None, noise=None, guidance_scale=0.0,
+ test_cfg_override=dict(), show_pbar=False, **kwargs):
+ x_t = torch.randn_like(x_0) if noise is None else noise
+ num_batches = x_t.size(0)
+ seq_len = x_t.shape[2:].numel() # h * w or t * h * w
+ ori_dtype = x_t.dtype
+ x_t = x_t.float()
+ ndim = x_t.dim()
+ assert ndim in [4, 5], f'Invalid x_t shape: {x_t.shape}. Expected 4D or 5D tensor.'
+ if ndim == 5: # (bs, c, t, h, w)
+ x_t = x_t.permute(0, 2, 1, 3, 4) # (bs, t, c, h, w)
+
+ cfg = deepcopy(self.test_cfg)
+ cfg.update(test_cfg_override)
+
+ output_mode = cfg.get('output_mode', 'mean')
+ assert output_mode in ['mean', 'sample']
+
+ sampler = cfg['sampler']
+ sampler_class = getattr(diffusers.schedulers, sampler + 'Scheduler', None)
+ if sampler_class is None:
+ sampler_class = getattr(schedulers, sampler + 'Scheduler', None)
+ if sampler_class is None:
+ raise AttributeError(f'Cannot find sampler [{sampler}].')
+
+ sampler_kwargs = cfg.get('sampler_kwargs', {})
+ signatures = inspect.signature(sampler_class).parameters.keys()
+ for key in ['shift', 'use_dynamic_shifting', 'base_seq_len', 'max_seq_len', 'base_logshift', 'max_logshift']:
+ if key in signatures and key not in sampler_kwargs:
+ sampler_kwargs[key] = cfg.get(key, getattr(self.timestep_sampler, key))
+ sampler = sampler_class(self.num_timesteps, **sampler_kwargs)
+
+ num_timesteps = cfg.get('num_timesteps', self.num_timesteps)
+ num_substeps = cfg.get('num_substeps', 1)
+ guidance_interval = cfg.get('guidance_interval', [0, self.num_timesteps])
+ orthogonal_guidance = cfg.get('orthogonal_guidance', 1.0)
+ save_intermediate = cfg.get('save_intermediate', False)
+ order = cfg.get('order', 1)
+ gm2_coefs = cfg.get('gm2_coefs', [0.005, 1.0])
+ gm2_correction_steps = cfg.get('gm2_correction_steps', 0)
+
+ set_timesteps_signatures = inspect.signature(sampler.set_timesteps).parameters.keys()
+ if 'seq_len' in set_timesteps_signatures:
+ sampler.set_timesteps(num_timesteps * num_substeps, seq_len=seq_len, device=x_t.device)
+ else:
+ sampler.set_timesteps(num_timesteps * num_substeps, device=x_t.device)
+
+ timesteps = sampler.timesteps
+ self.intermediate_x_t = []
+ self.intermediate_x_0 = []
+ self.intermediate_gm_x_0 = []
+ self.intermediate_gm_trans = []
+ self.intermediate_t = []
+ use_guidance = 0.0 < guidance_scale < 1.0
+ assert order in [1, 2]
+
+ if show_pbar:
+ pbar = mmcv.ProgressBar(num_timesteps)
+
+ # ========== Main sampling loop ==========
+ self.init_gm_cache()
+
+ for timestep_id in range(num_timesteps):
+ t = timesteps[timestep_id * num_substeps]
+
+ if save_intermediate:
+ self.intermediate_x_t.append(x_t)
+ self.intermediate_t.append(t)
+
+ x_t_input = x_t
+ _kwargs = kwargs
+ if use_guidance:
+ guidance_active = guidance_interval[0] <= t <= guidance_interval[1]
+ if guidance_active:
+ x_t_input = torch.cat([x_t_input, x_t_input], dim=0)
+ else:
+ _kwargs = {
+ k: v[num_batches:] if isinstance(v, torch.Tensor) and v.size(0) == 2 * num_batches else v
+ for k, v in kwargs.items()}
+
+ gm_output = self.pred(x_t_input, t, **_kwargs)
+ assert isinstance(gm_output, dict)
+ gm_output = self.u_to_x_0(gm_output, x_t_input, t)
+
+ # ========== Probabilistic CFG ==========
+ if use_guidance and guidance_active:
+ gm_cond = {k: v[num_batches:] for k, v in gm_output.items()}
+ gm_uncond = {k: v[:num_batches] for k, v in gm_output.items()}
+ uncond_mean = gm_to_mean(gm_uncond)
+ gaussian_cond = gm_to_iso_gaussian(gm_cond)[0]
+ if ndim == 5:
+ gaussian_cond['var'] = gaussian_cond['var'].mean(dim=(-4, -2, -1), keepdim=True) # exclude channel dim
+ else:
+ gaussian_cond['var'] = gaussian_cond['var'].mean(dim=(-2, -1), keepdim=True) # exclude channel dim
+ gaussian_output, cfg_bias, avg_var = probabilistic_guidance_jit(
+ gaussian_cond['mean'], gaussian_cond['var'], uncond_mean, guidance_scale,
+ orthogonal=orthogonal_guidance)
+ gm_output = gm_mul_iso_gaussian(
+ gm_cond, iso_gaussian_mul_iso_gaussian(gaussian_output, gaussian_cond, 1, -1),
+ 1, 1)[0]
+ else:
+ gaussian_output = gm_to_iso_gaussian(gm_output)[0]
+ gm_cond = gaussian_cond = avg_var = cfg_bias = None
+
+ # ========== 2nd order GM ==========
+ if order == 2:
+ if timestep_id < num_timesteps - 1:
+ h = t - timesteps[(timestep_id + 1) * num_substeps]
+ else:
+ h = t
+ gm_output, gaussian_output = self.gm_2nd_order(
+ gm_output, gaussian_output, x_t, t, h,
+ guidance_scale if guidance_active else 0.0, gm_cond, gaussian_cond, avg_var, cfg_bias,
+ ca=gm2_coefs[0], cb=gm2_coefs[1], gm2_correction_steps=gm2_correction_steps)
+
+ if save_intermediate:
+ self.intermediate_gm_x_0.append(gm_output)
+ if timestep_id < num_timesteps - 1:
+ t_next = timesteps[(timestep_id + 1) * num_substeps]
+ else:
+ t_next = 0
+ gm_trans = self.reverse_transition(gm_output, x_t, t_next, t, prediction_type='x0')
+ self.intermediate_gm_trans.append(gm_trans)
+
+ # ========== GM SDE step or GM ODE substeps ==========
+ x_t_base = x_t
+ t_base = t
+ for substep_id in range(num_substeps):
+ if substep_id == 0:
+ if self.spectrum_net is not None and output_mode == 'sample':
+ # Note: only support 2D power spectrum for now.
+ power_spectrum = self.spectrum_net(gaussian_output)
+ else:
+ power_spectrum = None
+ model_output = self.gm_to_model_output(gm_output, output_mode, power_spectrum=power_spectrum)
+ else:
+ assert output_mode == 'mean'
+ t = timesteps[timestep_id * num_substeps + substep_id]
+ model_output = self.gmflow_posterior_mean(
+ gm_output, x_t, x_t_base, t, t_base, prediction_type='x0')
+ x_t = sampler.step(model_output, t, x_t, return_dict=False, prediction_type='x0')[0]
+
+ if save_intermediate:
+ self.intermediate_x_0.append(model_output)
+
+ if show_pbar:
+ pbar.update()
+
+ if show_pbar:
+ sys.stdout.write('\n')
+
+ if ndim == 5: # (bs, t, c, h, w)
+ x_t = x_t.permute(0, 2, 1, 3, 4) # (bs, c, t, h, w)
+
+ return x_t.to(ori_dtype)
+
+ def forward_u(self, x_t, t, guidance_scale=0.0, test_cfg_override=dict(), **kwargs):
+ ori_dtype = x_t.dtype
+ x_t = x_t.float()
+ ndim = x_t.dim()
+ assert ndim in [4, 5], f'Invalid x_t shape: {x_t.shape}. Expected 4D or 5D tensor.'
+ if ndim == 5: # (bs, c, t, h, w)
+ x_t = x_t.permute(0, 2, 1, 3, 4) # (bs, t, c, h, w)
+
+ cfg = deepcopy(self.test_cfg)
+ cfg.update(test_cfg_override)
+
+ orthogonal_guidance = cfg.get('orthogonal_guidance', 1.0)
+ guidance_interval = cfg.get('guidance_interval', [0, self.num_timesteps])
+
+ use_guidance = 0.0 < guidance_scale < 1.0
+
+ x_t_input = x_t
+ t_input = t
+ if use_guidance:
+ x_t_input = torch.cat([x_t_input, x_t_input], dim=0)
+ t_input = torch.cat([t_input, t_input], dim=0)
+
+ gm_output = self.pred(x_t_input, t_input, **kwargs)
+ assert isinstance(gm_output, dict)
+
+ # ========== Probabilistic CFG ==========
+ if use_guidance:
+ num_batches = x_t.size(0)
+ gm_cond = {k: v[num_batches:] for k, v in gm_output.items()}
+ gm_uncond = {k: v[:num_batches] for k, v in gm_output.items()}
+ uncond_mean = gm_to_mean(gm_uncond)
+ gaussian_cond = gm_to_iso_gaussian(gm_cond)[0]
+ if ndim == 5:
+ gaussian_cond['var'] = gaussian_cond['var'].mean(dim=(-4, -2, -1), keepdim=True) # exclude channel dim
+ else:
+ gaussian_cond['var'] = gaussian_cond['var'].mean(dim=(-2, -1), keepdim=True)
+ gaussian_output = probabilistic_guidance_jit(
+ gaussian_cond['mean'], gaussian_cond['var'], uncond_mean, guidance_scale,
+ orthogonal=orthogonal_guidance,
+ orthogonal_axis=self.u_to_x_0(gaussian_cond['mean'], x_t, t))[0]
+ gm_output = gm_mul_iso_gaussian(
+ gm_cond, iso_gaussian_mul_iso_gaussian(gaussian_output, gaussian_cond, 1, -1),
+ 1, 1)[0]
+ if guidance_interval[0] > 0 or guidance_interval[1] < self.num_timesteps:
+ guidance_active = ((t >= guidance_interval[0]) & (t <= guidance_interval[1])).reshape(
+ [num_batches] + [1] * ndim)
+ gm_output = {k: torch.where(guidance_active, v, gm_cond[k]) for k, v in gm_output.items()}
+
+ u = gm_to_mean(gm_output)
+
+ if ndim == 5: # (bs, t, c, h, w)
+ u = u.permute(0, 2, 1, 3, 4) # (bs, c, t, h, w)
+
+ return u.to(ori_dtype)
diff --git a/lakonlab/models/diffusions/piflow.py b/lakonlab/models/diffusions/piflow.py
new file mode 100644
index 0000000000000000000000000000000000000000..0319c01a2a7b79aac415a63d63feb0f7894cac10
--- /dev/null
+++ b/lakonlab/models/diffusions/piflow.py
@@ -0,0 +1,410 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import sys
+import torch
+import mmcv
+
+from copy import deepcopy
+from functools import partial
+from mmgen.models.architectures.common import get_module_device
+from mmgen.models.builder import MODULES
+
+from . import GaussianFlow
+from .piflow_policies import POLICY_CLASSES, GMFlowPolicy
+from lakonlab.utils import module_eval
+
+
+class PiFlowImitationBase(GaussianFlow):
+
+ def __init__(self, *args, policy_type='GMFlow', policy_kwargs=None, **kwargs):
+ super().__init__(*args, **kwargs)
+ assert policy_type in POLICY_CLASSES, \
+ f'Invalid policy: {policy_type}. Supported policies are {list(POLICY_CLASSES.keys())}.'
+ self.policy_type = policy_type
+ self.policy_class = partial(
+ POLICY_CLASSES[policy_type], **policy_kwargs
+ ) if policy_kwargs else POLICY_CLASSES[policy_type]
+
+ def policy_rollout(
+ self,
+ x_t_start: torch.Tensor, # (B, C, *, H, W)
+ sigma_t_start: torch.Tensor, # (B, 1, *, 1, 1)
+ raw_t_start: torch.Tensor, # (B, )
+ raw_t_end: torch.Tensor, # (B, )
+ total_substeps: int,
+ policy,
+ seq_len=None):
+ num_batches = x_t_start.size(0)
+ ndim = x_t_start.dim()
+ raw_t_start = raw_t_start.reshape(num_batches, *((ndim - 1) * [1]))
+ raw_t_end = raw_t_end.reshape(num_batches, *((ndim - 1) * [1]))
+
+ delta_raw_t = raw_t_start - raw_t_end
+ num_substeps = (delta_raw_t * total_substeps).round().to(torch.long).clamp(min=1)
+ substep_size = delta_raw_t / num_substeps
+ max_num_substeps = num_substeps.max()
+
+ raw_t = raw_t_start
+ sigma_t = sigma_t_start
+ x_t = x_t_start
+
+ for substep_id in range(max_num_substeps.item()):
+ u = policy.pi(x_t, sigma_t)
+
+ raw_t_minus = (raw_t - substep_size).clamp(min=0)
+ sigma_t_minus = self.timestep_sampler.warp_t(raw_t_minus, seq_len=seq_len)
+ x_t_minus = x_t + u * (sigma_t_minus - sigma_t)
+
+ active_mask = num_substeps > substep_id
+ x_t = torch.where(active_mask, x_t_minus, x_t)
+ sigma_t = torch.where(active_mask, sigma_t_minus, sigma_t)
+ raw_t = torch.where(active_mask, raw_t_minus, raw_t)
+
+ x_t_end = x_t
+ sigma_t_end = sigma_t
+ t_end = sigma_t_end.flatten() * self.num_timesteps
+ return x_t_end, sigma_t_end, t_end
+
+ def policy_average_u(
+ self,
+ x_t_start: torch.Tensor, # (B, C, *, H, W)
+ sigma_t_start: torch.Tensor, # (B, 1, *, 1, 1)
+ raw_t_start: torch.Tensor, # (B, )
+ raw_t_end: torch.Tensor, # (B, )
+ total_substeps: int,
+ policy,
+ seq_len=None,
+ eps=1e-4):
+ num_batches = x_t_start.size(0)
+ ndim = x_t_start.dim()
+ is_small_length = torch.round((raw_t_start - raw_t_end) * total_substeps) < 2
+ pred_mean_u = pred_local_u = None
+ if not is_small_length.all(): # mean velocity over the rollout length
+ x_t_end, sigma_t_end, _ = self.policy_rollout(
+ x_t_start, sigma_t_start, raw_t_start, raw_t_end, total_substeps,
+ policy, seq_len=seq_len)
+ pred_mean_u = (x_t_start - x_t_end) / (sigma_t_start - sigma_t_end).clamp(min=eps)
+ if is_small_length.any(): # numerically stable local velocity
+ pred_local_u = policy.pi(x_t_start, sigma_t_start)
+ if pred_mean_u is None:
+ pred_u = pred_local_u
+ elif pred_local_u is None:
+ pred_u = pred_mean_u
+ else:
+ pred_u = torch.where(
+ is_small_length.reshape(num_batches, *((ndim - 1) * [1])), pred_local_u, pred_mean_u)
+ return pred_u
+
+ @staticmethod
+ def get_shape_info(x):
+ x_t_dst_shape = x.size()
+ bs = x_t_dst_shape[0]
+ ndim = len(x_t_dst_shape)
+ seq_len = x.shape[2:].numel()
+ return ndim, bs, seq_len
+
+ def piid_segment(
+ self, teacher, policy, x_t_src, raw_t_src, sigma_t_src, teacher_ratio, segment_size,
+ teacher_kwargs, get_x_t_dst=False):
+ eps = self.train_cfg.get('eps', 1e-4)
+ total_substeps = self.train_cfg.get('total_substeps', 128)
+ num_intermediate_states = self.train_cfg.get('num_intermediate_states', 2)
+ window_substeps = self.train_cfg.get('window_substeps', 0)
+
+ device = x_t_src.device
+ ndim, bs, seq_len = self.get_shape_info(x_t_src)
+ if not isinstance(segment_size, torch.Tensor):
+ segment_size = torch.tensor(
+ [segment_size], dtype=torch.float32, device=device)
+
+ # window size ∆τ ≈ window_substeps / total_substeps
+ num_substeps = (segment_size * total_substeps).round().to(torch.long).clamp(min=1)
+ substep_size = segment_size / num_substeps
+ window_size = torch.minimum(window_substeps * substep_size, segment_size)
+
+ raw_t_dst = raw_t_src - segment_size
+
+ policy_detached = policy.detach()
+ if isinstance(policy_detached, GMFlowPolicy):
+ gm_dropout = self.train_cfg.get('gm_dropout', 0.0)
+ policy_detached.dropout_(gm_dropout)
+
+ # time sampling for scheduled trajectory mixing
+ assert not self.timestep_sampler.logit_normal_enable
+ student_intervals = torch.rand(
+ (bs, num_intermediate_states), device=device
+ ) * ((1 - teacher_ratio) * (segment_size - window_size).unsqueeze(-1))
+ student_intervals = torch.sort(student_intervals, dim=-1)[0]
+ student_intervals = torch.diff(student_intervals, dim=-1, prepend=torch.zeros((bs, 1), device=device))
+
+ teacher_intervals = torch.rand((bs, num_intermediate_states - 1), device=device)
+ teacher_intervals = torch.sort(teacher_intervals, dim=-1)[0]
+ teacher_intervals = torch.diff(
+ teacher_intervals, dim=-1,
+ prepend=torch.zeros((bs, 1), device=device),
+ append=torch.ones(
+ (bs, 1), device=device)
+ ) * (teacher_ratio * (segment_size - window_size).unsqueeze(-1))
+
+ x_t = x_t_src
+ raw_t = raw_t_src
+ sigma_t = sigma_t_src
+
+ all_pred_u = []
+ all_tgt_u = []
+ all_timesteps = []
+
+ for teacher_step_id in range(num_intermediate_states):
+ raw_t_a = (raw_t - student_intervals[:, teacher_step_id]).clamp(min=0)
+ raw_t_b = (raw_t_a - teacher_intervals[:, teacher_step_id]).clamp(min=0)
+
+ with torch.no_grad(), module_eval(teacher):
+ x_t_a, sigma_t_a, t_a = self.policy_rollout(
+ x_t, sigma_t, raw_t, raw_t_a, total_substeps,
+ policy_detached, seq_len=seq_len)
+ tgt_u = teacher(return_u=True, x_t=x_t_a, t=t_a, **teacher_kwargs)
+ all_tgt_u.append(tgt_u)
+ all_timesteps.append(t_a)
+
+ pred_u = self.policy_average_u(
+ x_t_a, sigma_t_a, raw_t_a, raw_t_b - window_size, total_substeps,
+ policy, seq_len=seq_len, eps=eps)
+ all_pred_u.append(pred_u)
+
+ sigma_t_b = self.timestep_sampler.warp_t(raw_t_b, seq_len=seq_len).reshape(bs, *((ndim - 1) * [1]))
+ x_t = x_t_a + tgt_u * (sigma_t_b - sigma_t_a)
+ raw_t = raw_t_b
+ sigma_t = sigma_t_b
+
+ loss_kwargs = dict(
+ u_t_pred=torch.cat(all_pred_u, dim=0),
+ u_t=torch.cat(all_tgt_u, dim=0),
+ timesteps=torch.cat(all_timesteps, dim=0)
+ )
+ loss = self.flow_loss(loss_kwargs)
+
+ if get_x_t_dst:
+ with torch.no_grad():
+ x_t_dst, _, _ = self.policy_rollout(
+ x_t, sigma_t, raw_t, raw_t_dst, total_substeps,
+ policy_detached, seq_len=seq_len)
+ else:
+ x_t_dst = None
+
+ return loss, x_t_dst, raw_t_dst
+
+ def forward_test(
+ self, x_0=None, noise=None, guidance_scale=None,
+ test_cfg_override=dict(), show_pbar=False, **kwargs):
+ x_t_src = torch.randn_like(x_0) if noise is None else noise
+ num_batches = x_t_src.size(0)
+ seq_len = x_t_src.shape[2:].numel() # h * w or t * h * w
+ ori_dtype = x_t_src.dtype
+ device = x_t_src.device
+ x_t_src = x_t_src.float()
+ ndim = x_t_src.dim()
+ assert ndim in [4, 5], f'Invalid x_t_src shape: {x_t_src.shape}. Expected 4D or 5D tensor.'
+
+ cfg = deepcopy(self.test_cfg)
+ cfg.update(test_cfg_override)
+
+ total_substeps = cfg.get('total_substeps', self.num_timesteps)
+ eps = cfg.get('eps', 1e-4)
+ nfe = cfg['nfe']
+ final_step_size_scale = max(cfg.get('final_step_size_scale', 1.0), eps)
+ base_segment_size = 1 / (nfe - 1 + final_step_size_scale)
+
+ raw_t_src = torch.ones((num_batches,), dtype=torch.float32, device=device)
+ sigma_t_src = self.timestep_sampler.warp_t(raw_t_src, seq_len=seq_len).reshape(
+ num_batches, *((ndim - 1) * [1]))
+ t_src = sigma_t_src.flatten() * self.num_timesteps
+
+ if show_pbar:
+ pbar = mmcv.ProgressBar(self.distill_steps)
+
+ # ========== Main sampling loop ==========
+ for step_id in range(nfe):
+ is_final_step = step_id == nfe - 1
+ if is_final_step:
+ segment_size = base_segment_size * final_step_size_scale
+ else:
+ segment_size = base_segment_size
+
+ raw_t_dst = raw_t_src - segment_size
+
+ denoising_output = self.pred(x_t_src, t_src, **kwargs)
+ policy = self.policy_class(
+ denoising_output, x_t_src, sigma_t_src, eps=eps)
+ if isinstance(policy, GMFlowPolicy) and not is_final_step:
+ temperature = cfg.get('temperature', 1.0)
+ policy.temperature_(temperature)
+
+ x_t_dst, sigma_t_dst, t_dst = self.policy_rollout(
+ x_t_src, sigma_t_src, raw_t_src, raw_t_dst, total_substeps,
+ policy, seq_len=seq_len)
+
+ x_t_src = x_t_dst
+ raw_t_src = raw_t_dst
+ sigma_t_src = sigma_t_dst
+ t_src = t_dst
+
+ if show_pbar:
+ pbar.update()
+
+ if show_pbar:
+ sys.stdout.write('\n')
+
+ return x_t_src.to(ori_dtype)
+
+
+@MODULES.register_module()
+class PiFlowImitation(PiFlowImitationBase):
+
+ def sample_t(self, num_batches, ndim, seq_len=None, device=None):
+ eps = self.train_cfg.get('eps', 1e-4)
+ nfe = self.train_cfg['nfe']
+
+ final_step_size_scale = max(self.train_cfg.get('final_step_size_scale', 1.0), eps)
+ one_minus_final_scale = 1 - final_step_size_scale
+ base_segment_size = 1 / (nfe - one_minus_final_scale)
+ final_step_size = final_step_size_scale * base_segment_size
+
+ raw_t = self.timestep_sampler(
+ num_batches, warp_t=False, scale_t=False, device=device).clamp(min=eps)
+ raw_t_src_idx = torch.ceil(
+ raw_t / base_segment_size + one_minus_final_scale).clamp(min=1)
+ if isinstance(nfe, torch.Tensor):
+ raw_t_src_idx = torch.minimum(raw_t_src_idx, nfe)
+ else:
+ raw_t_src_idx = raw_t_src_idx.clamp(max=nfe)
+ raw_t_src = ((raw_t_src_idx - one_minus_final_scale) * base_segment_size).clamp(min=eps, max=1)
+ is_final_step = raw_t_src_idx == 1
+ segment_size = torch.where(
+ is_final_step, final_step_size, base_segment_size)
+
+ sigma_t_src = self.timestep_sampler.warp_t(raw_t_src, seq_len=seq_len).reshape(
+ num_batches, *((ndim - 1) * [1]))
+ t_src = sigma_t_src.flatten() * self.num_timesteps
+ return raw_t_src, sigma_t_src, t_src, segment_size
+
+ def forward_train(self, x_0, teacher=None, teacher_kwargs=dict(), running_status=None, **kwargs):
+ device = get_module_device(self)
+ num_batches = x_0.size(0)
+ seq_len = x_0.shape[2:].numel() # h * w or t * h * w
+ ndim = x_0.dim()
+ assert ndim in [4, 5], f'Invalid x_0 shape: {x_0.shape}. Expected 4D or 5D tensor.'
+
+ num_decay_iters = self.train_cfg.get('num_decay_iters', 0)
+ if num_decay_iters > 0:
+ teacher_ratio = 1 - min(running_status['iteration'], num_decay_iters) / num_decay_iters
+ log_vars = dict(teacher_ratio=teacher_ratio)
+ else:
+ teacher_ratio = 0.0
+ log_vars = dict()
+
+ raw_t_src, sigma_t_src, t_src, segment_size = self.sample_t(
+ num_batches, ndim, seq_len=seq_len, device=device)
+ noise = torch.randn_like(x_0)
+ x_t_src, _, _ = self.sample_forward_diffusion(x_0, t_src, noise)
+
+ denoising_output = self.pred(x_t_src, t_src, **kwargs)
+ policy = self.policy_class(denoising_output, x_t_src, sigma_t_src)
+
+ loss_diffusion, _, _ = self.piid_segment(
+ teacher, policy, x_t_src, raw_t_src, sigma_t_src, teacher_ratio, segment_size,
+ teacher_kwargs)
+
+ loss = loss_diffusion
+ log_vars.update(self.flow_loss.log_vars)
+ log_vars.update(loss_diffusion=float(loss_diffusion))
+
+ return loss, log_vars
+
+
+@MODULES.register_module()
+class PiFlowImitationDataFree(PiFlowImitationBase):
+
+ is_multistep = True
+
+ def forward_initialize(
+ self, x_0, teacher=None, teacher_kwargs=dict(), running_status=None, **kwargs):
+ device = get_module_device(self)
+ num_batches = x_0.size(0) # x_0 is a dummy input
+
+ num_decay_iters = self.train_cfg.get('num_decay_iters', 0)
+ if num_decay_iters > 0:
+ teacher_ratio = 1 - min(running_status['iteration'], num_decay_iters) / num_decay_iters
+ log_vars = dict(teacher_ratio=teacher_ratio)
+ else:
+ teacher_ratio = 0.0
+ log_vars = dict()
+
+ x_t_src = torch.randn_like(x_0)
+ raw_t_src = torch.ones((num_batches,), dtype=torch.float32, device=device)
+ step_states = dict(
+ step_id=0,
+ terminate=False,
+ detachable=True,
+ teacher_ratio=teacher_ratio,
+ x_t_src=x_t_src,
+ raw_t_src=raw_t_src,
+ )
+
+ return step_states, log_vars
+
+ def forward_train(
+ self, x_0, step_states=None, teacher=None, teacher_kwargs=dict(), running_status=None, **kwargs):
+ step_id = step_states['step_id']
+ teacher_ratio = step_states['teacher_ratio']
+ x_t_src = step_states['x_t_src']
+ raw_t_src = step_states['raw_t_src']
+
+ num_batches = x_t_src.size(0)
+ seq_len = x_t_src.shape[2:].numel()
+ ndim = x_t_src.dim()
+ assert ndim in [4, 5], f'Invalid x_t_src shape: {x_t_src.shape}. Expected 4D or 5D tensor.'
+
+ eps = self.train_cfg.get('eps', 1e-4)
+ nfe = self.train_cfg['nfe']
+ final_step_size_scale = max(self.train_cfg.get('final_step_size_scale', 1.0), eps)
+ base_segment_size = 1 / (nfe - 1 + final_step_size_scale)
+ is_final_step = step_id == nfe - 1
+ if is_final_step:
+ segment_size = base_segment_size * final_step_size_scale
+ else:
+ segment_size = base_segment_size
+
+ sigma_t_src = self.timestep_sampler.warp_t(raw_t_src, seq_len=seq_len).reshape(
+ num_batches, *((ndim - 1) * [1]))
+ t_src = sigma_t_src.flatten() * self.num_timesteps
+
+ denoising_output = self.pred(x_t_src, t_src, **kwargs)
+ policy = self.policy_class(denoising_output, x_t_src, sigma_t_src)
+
+ step_loss_diffusion, x_t_dst, raw_t_dst = self.piid_segment(
+ teacher, policy, x_t_src, raw_t_src, sigma_t_src, teacher_ratio, segment_size,
+ teacher_kwargs, get_x_t_dst=True)
+
+ loss_diffusion = step_loss_diffusion * segment_size # Weighing by segment size
+ loss = loss_diffusion
+ log_vars = {k: v * segment_size for k, v in self.flow_loss.log_vars.items()}
+ log_vars.update({
+ 'loss_diffusion': float(loss_diffusion),
+ f'loss_diffusion_step{step_id}': float(step_loss_diffusion)
+ })
+
+ if step_id < nfe - 1:
+ step_states.update(
+ step_id=step_id + 1,
+ x_t_src=x_t_dst,
+ raw_t_src=raw_t_dst)
+ else:
+ step_states.update(terminate=True)
+
+ return loss, log_vars, step_states
+
+ def forward(self, x_0=None, return_step_states=False, **kwargs):
+ if return_step_states:
+ return self.forward_initialize(x_0=x_0, **kwargs)
+ else:
+ return super().forward(x_0=x_0, **kwargs)
diff --git a/lakonlab/models/diffusions/piflow_policies/__init__.py b/lakonlab/models/diffusions/piflow_policies/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..df0fa4f1d2aff8fc26eccd4012396b06dd68ff57
--- /dev/null
+++ b/lakonlab/models/diffusions/piflow_policies/__init__.py
@@ -0,0 +1,8 @@
+from .dx import DXPolicy
+from .gmflow import GMFlowPolicy
+
+
+POLICY_CLASSES = dict(
+ DX=DXPolicy,
+ GMFlow=GMFlowPolicy
+)
diff --git a/lakonlab/models/diffusions/piflow_policies/__pycache__/__init__.cpython-310.pyc b/lakonlab/models/diffusions/piflow_policies/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..71158b2cd39c8f91d712e6f5835a9f8a4d45e133
Binary files /dev/null and b/lakonlab/models/diffusions/piflow_policies/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/piflow_policies/__pycache__/base.cpython-310.pyc b/lakonlab/models/diffusions/piflow_policies/__pycache__/base.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..b02fb3431c0584cca4555773f09210ffe23f0e26
Binary files /dev/null and b/lakonlab/models/diffusions/piflow_policies/__pycache__/base.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/piflow_policies/__pycache__/dx.cpython-310.pyc b/lakonlab/models/diffusions/piflow_policies/__pycache__/dx.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..0434e2ce69fffe6013a858f68753c6ab20a94159
Binary files /dev/null and b/lakonlab/models/diffusions/piflow_policies/__pycache__/dx.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/piflow_policies/__pycache__/gmflow.cpython-310.pyc b/lakonlab/models/diffusions/piflow_policies/__pycache__/gmflow.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..da40acdd0962c39bd8ebe3845daf9514e49ce7d8
Binary files /dev/null and b/lakonlab/models/diffusions/piflow_policies/__pycache__/gmflow.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/piflow_policies/base.py b/lakonlab/models/diffusions/piflow_policies/base.py
new file mode 100644
index 0000000000000000000000000000000000000000..1aee3e7eb6404ac93c054c974acda6e93ccf257e
--- /dev/null
+++ b/lakonlab/models/diffusions/piflow_policies/base.py
@@ -0,0 +1,21 @@
+from abc import ABCMeta, abstractmethod
+
+
+class BasePolicy(metaclass=ABCMeta):
+
+ @abstractmethod
+ def pi(self, x_t, sigma_t):
+ """Compute the flow velocity at (x_t, t).
+
+ Args:
+ x_t (torch.Tensor): Noisy input at time t.
+ sigma_t (torch.Tensor): Noise level at time t.
+
+ Returns:
+ torch.Tensor: The computed flow velocity u_t.
+ """
+ pass
+
+ @abstractmethod
+ def detach(self):
+ pass
diff --git a/lakonlab/models/diffusions/piflow_policies/dx.py b/lakonlab/models/diffusions/piflow_policies/dx.py
new file mode 100644
index 0000000000000000000000000000000000000000..8d1e3c48a354cdbceab74df7a6ac6ca21f88c2ea
--- /dev/null
+++ b/lakonlab/models/diffusions/piflow_policies/dx.py
@@ -0,0 +1,123 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import torch
+from .base import BasePolicy
+
+
+class DXPolicy(BasePolicy):
+ """DX policy. The number of grid points N is inferred from the denoising output.
+
+ Note: segment_size and shift are intrinsic parameters of the DX policy. For elastic inference (i.e., changing
+ the number of function evaluations or noise schedule at test time), these parameters should be kept unchanged.
+
+ Args:
+ denoising_output (torch.Tensor): The output of the denoising model. Shape (B, N, C, H, W) or (B, N, C, T, H, W).
+ x_t_src (torch.Tensor): The initial noisy sample. Shape (B, C, H, W) or (B, C, T, H, W).
+ sigma_t_src (torch.Tensor): The initial noise level. Shape (B,).
+ segment_size (float): The size of each DX policy time segment. Defaults to 1.0.
+ shift (float): The shift parameter for the DX policy noise schedule. Defaults to 1.0.
+ mode (str): Either 'grid' or 'polynomial' mode for calculating x_0. Defaults to 'grid'.
+ eps (float): A small value to avoid numerical issues. Defaults to 1e-4.
+ """
+
+ def __init__(
+ self,
+ denoising_output: torch.Tensor,
+ x_t_src: torch.Tensor,
+ sigma_t_src: torch.Tensor,
+ segment_size: float = 1.0,
+ shift: float = 1.0,
+ mode: str = 'grid',
+ eps: float = 1e-4):
+ self.x_t_src = x_t_src
+ self.ndim = x_t_src.dim()
+ self.shift = shift
+ self.eps = eps
+
+ assert mode in ['grid', 'polynomial']
+ self.mode = mode
+
+ self.sigma_t_src = sigma_t_src.reshape(*sigma_t_src.size(), *((self.ndim - sigma_t_src.dim()) * [1]))
+ self.raw_t_src = self._unwarp_t(self.sigma_t_src)
+ self.raw_t_dst = (self.raw_t_src - segment_size).clamp(min=0)
+ self.segment_size = (self.raw_t_src - self.raw_t_dst).clamp(min=eps)
+
+ self.denoising_output_x_0 = self._u_to_x_0(
+ denoising_output, self.x_t_src, self.sigma_t_src)
+
+ def _unwarp_t(self, sigma_t):
+ return sigma_t / (self.shift + (1 - self.shift) * sigma_t)
+
+ @staticmethod
+ def _u_to_x_0(denoising_output, x_t, sigma_t):
+ x_0 = x_t.unsqueeze(1) - sigma_t.unsqueeze(1) * denoising_output
+ return x_0
+
+ @staticmethod
+ def _interpolate(x, t):
+ """
+ Args:
+ x (torch.Tensor): (B, N, *)
+ t (torch.Tensor): (B, *) in [0, 1]
+
+ Returns:
+ torch.Tensor: (B, *)
+ """
+ n = x.size(1)
+ if n < 2:
+ return x.squeeze(1)
+ t = t.clamp(min=0, max=1) * (n - 1)
+ t0 = t.floor().to(torch.long).clamp(min=0, max=n - 2)
+ t1 = t0 + 1
+ t0t1 = torch.stack([t0, t1], dim=1) # (B, 2, *)
+ x0x1 = torch.gather(x, dim=1, index=t0t1.expand(-1, -1, *x.shape[2:]))
+ x_interp = (t1 - t) * x0x1[:, 0] + (t - t0) * x0x1[:, 1]
+ return x_interp
+
+ def pi(self, x_t, sigma_t):
+ """Compute the flow velocity at (x_t, t).
+
+ Args:
+ x_t (torch.Tensor): Noisy input at time t.
+ sigma_t (torch.Tensor): Noise level at time t.
+
+ Returns:
+ torch.Tensor: The computed flow velocity u_t.
+ """
+ sigma_t = sigma_t.reshape(*sigma_t.size(), *((self.ndim - sigma_t.dim()) * [1]))
+ raw_t = self._unwarp_t(sigma_t)
+ if self.mode == 'grid':
+ x_0 = self._interpolate(
+ self.denoising_output_x_0, (raw_t - self.raw_t_dst) / self.segment_size)
+ elif self.mode == 'polynomial':
+ p_order = self.denoising_output_x_0.size(1)
+ diff_t = self.raw_t_src - raw_t # (B, 1, 1, 1)
+ basis = torch.stack(
+ [diff_t ** i for i in range(p_order)], dim=1) # (B, N, 1, 1, 1)
+ x_0 = torch.sum(basis * self.denoising_output_x_0, dim=1)
+ else:
+ raise ValueError(f"Unknown mode: {self.mode}")
+ u = (x_t - x_0) / sigma_t.clamp(min=self.eps)
+ return u
+
+ def copy(self):
+ new_policy = DXPolicy.__new__(DXPolicy)
+ new_policy.x_t_src = self.x_t_src
+ new_policy.ndim = self.ndim
+ new_policy.shift = self.shift
+ new_policy.eps = self.eps
+ new_policy.mode = self.mode
+ new_policy.sigma_t_src = self.sigma_t_src
+ new_policy.raw_t_src = self.raw_t_src
+ new_policy.raw_t_dst = self.raw_t_dst
+ new_policy.segment_size = self.segment_size
+ new_policy.denoising_output_x_0 = self.denoising_output_x_0
+ return new_policy
+
+ def detach_(self):
+ self.denoising_output_x_0 = self.denoising_output_x_0.detach()
+ return self
+
+ def detach(self):
+ new_policy = self.copy()
+ return new_policy.detach_()
diff --git a/lakonlab/models/diffusions/piflow_policies/gmflow.py b/lakonlab/models/diffusions/piflow_policies/gmflow.py
new file mode 100644
index 0000000000000000000000000000000000000000..53d249b82803fd7e93e12614f52f538b2854395d
--- /dev/null
+++ b/lakonlab/models/diffusions/piflow_policies/gmflow.py
@@ -0,0 +1,133 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import torch
+
+from typing import Dict
+from .base import BasePolicy
+from ..gmflow import gmflow_posterior_mean_jit
+from lakonlab.ops.gmflow_ops.gmflow_ops import gm_temperature
+
+
+class GMFlowPolicy(BasePolicy):
+ """GMFlow policy. The number of components K is inferred from the denoising output.
+
+ Args:
+ denoising_output (dict): The output of the denoising model, containing:
+ means (torch.Tensor): The means of the Gaussian components. Shape (B, K, C, H, W) or (B, K, C, T, H, W).
+ logstds (torch.Tensor): The log standard deviations of the Gaussian components. Shape (B, K, 1, 1, 1)
+ or (B, K, 1, 1, 1, 1).
+ logweights (torch.Tensor): The log weights of the Gaussian components. Shape (B, K, 1, H, W) or
+ (B, K, 1, T, H, W).
+ x_t_src (torch.Tensor): The initial noisy sample. Shape (B, C, H, W) or (B, C, T, H, W).
+ sigma_t_src (torch.Tensor): The initial noise level. Shape (B,).
+ checkpointing (bool): Whether to use gradient checkpointing to save memory. Defaults to True.
+ eps (float): A small value to avoid numerical issues. Defaults to 1e-4.
+ """
+
+ def __init__(
+ self,
+ denoising_output: Dict[str, torch.Tensor],
+ x_t_src: torch.Tensor,
+ sigma_t_src: torch.Tensor,
+ checkpointing: bool = True,
+ eps: float = 1e-4):
+ self.x_t_src = x_t_src
+ self.ndim = x_t_src.dim()
+ self.checkpointing = checkpointing
+ self.eps = eps
+
+ self.sigma_t_src = sigma_t_src.reshape(*sigma_t_src.size(), *((self.ndim - sigma_t_src.dim()) * [1]))
+ self.denoising_output_x_0 = self._u_to_x_0(
+ denoising_output, self.x_t_src, self.sigma_t_src)
+
+ @staticmethod
+ def _u_to_x_0(denoising_output, x_t, sigma_t):
+ x_t = x_t.unsqueeze(1)
+ sigma_t = sigma_t.unsqueeze(1)
+ means_x_0 = x_t - sigma_t * denoising_output['means']
+ gm_vars = (denoising_output['logstds'] * 2).exp() * sigma_t.square()
+ return dict(
+ means=means_x_0,
+ gm_vars=gm_vars,
+ logweights=denoising_output['logweights'])
+
+ def pi(self, x_t, sigma_t):
+ """Compute the flow velocity at (x_t, t).
+
+ Args:
+ x_t (torch.Tensor): Noisy input at time t.
+ sigma_t (torch.Tensor): Noise level at time t.
+
+ Returns:
+ torch.Tensor: The computed flow velocity u_t.
+ """
+ sigma_t = sigma_t.reshape(*sigma_t.size(), *((self.ndim - sigma_t.dim()) * [1]))
+ means = self.denoising_output_x_0['means']
+ gm_vars = self.denoising_output_x_0['gm_vars']
+ logweights = self.denoising_output_x_0['logweights']
+ if (sigma_t == self.sigma_t_src).all() and (x_t == self.x_t_src).all():
+ x_0 = (logweights.softmax(dim=1) * means).sum(dim=1)
+ else:
+ if self.checkpointing and torch.is_grad_enabled():
+ x_0 = torch.utils.checkpoint.checkpoint(
+ gmflow_posterior_mean_jit,
+ self.sigma_t_src, sigma_t, self.x_t_src, x_t,
+ means,
+ gm_vars,
+ logweights,
+ self.eps, 1, 2,
+ use_reentrant=True) # use_reentrant=False does not work with jit
+ else:
+ x_0 = gmflow_posterior_mean_jit(
+ self.sigma_t_src, sigma_t, self.x_t_src, x_t,
+ means,
+ gm_vars,
+ logweights,
+ self.eps, 1, 2)
+ u = (x_t - x_0) / sigma_t.clamp(min=self.eps)
+ return u
+
+ def copy(self):
+ new_policy = GMFlowPolicy.__new__(GMFlowPolicy)
+ new_policy.x_t_src = self.x_t_src
+ new_policy.ndim = self.ndim
+ new_policy.checkpointing = self.checkpointing
+ new_policy.eps = self.eps
+ new_policy.sigma_t_src = self.sigma_t_src
+ new_policy.denoising_output_x_0 = self.denoising_output_x_0.copy()
+ return new_policy
+
+ def detach_(self):
+ self.denoising_output_x_0 = {k: v.detach() for k, v in self.denoising_output_x_0.items()}
+ return self
+
+ def detach(self):
+ new_policy = self.copy()
+ return new_policy.detach_()
+
+ def dropout_(self, p):
+ if p <= 0 or p >= 1:
+ return self
+ logweights = self.denoising_output_x_0['logweights']
+ dropout_mask = torch.rand(
+ (*logweights.shape[:2], *((self.ndim - 1) * [1])), device=logweights.device) < p
+ is_all_dropout = dropout_mask.all(dim=1, keepdim=True)
+ dropout_mask &= ~is_all_dropout
+ self.denoising_output_x_0['logweights'] = logweights.masked_fill(
+ dropout_mask, float('-inf'))
+ return self
+
+ def dropout(self, p):
+ new_policy = self.copy()
+ return new_policy.dropout_(p)
+
+ def temperature_(self, temp):
+ if temp >= 1.0:
+ return self
+ self.denoising_output_x_0 = gm_temperature(
+ self.denoising_output_x_0, temp, gm_dim=1, eps=self.eps)
+ return self
+
+ def temperature(self, temp):
+ new_policy = self.copy()
+ return new_policy.temperature_(temp)
diff --git a/lakonlab/models/diffusions/sampler.py b/lakonlab/models/diffusions/sampler.py
new file mode 100644
index 0000000000000000000000000000000000000000..59193d2d0652ab2c9346375b9c058ce78792a1ba
--- /dev/null
+++ b/lakonlab/models/diffusions/sampler.py
@@ -0,0 +1,76 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import numpy as np
+import torch
+
+from mmgen.models.builder import MODULES
+
+
+@MODULES.register_module()
+class ContinuousTimeStepSampler:
+ def __init__(
+ self,
+ num_timesteps,
+ shift=1.0,
+ logit_normal_enable=False,
+ logit_normal_mean=0.0,
+ logit_normal_std=1.0,
+ use_dynamic_shifting=False,
+ base_seq_len=256,
+ max_seq_len=4096,
+ base_logshift=0.5,
+ max_logshift=1.15):
+ self.num_timesteps = num_timesteps
+ self.shift = shift
+ self.logit_normal_enable = logit_normal_enable
+ self.logit_normal_mean = logit_normal_mean
+ self.logit_normal_std = logit_normal_std
+ self.use_dynamic_shifting = use_dynamic_shifting
+ self.base_seq_len = base_seq_len
+ self.max_seq_len = max_seq_len
+ self.base_logshift = base_logshift
+ self.max_logshift = max_logshift
+
+ def get_shift(self, seq_len=None):
+ if self.use_dynamic_shifting and seq_len is not None:
+ m = (self.max_logshift - self.base_logshift) / (self.max_seq_len - self.base_seq_len)
+ logshift = (seq_len - self.base_seq_len) * m + self.base_logshift
+ if isinstance(logshift, torch.Tensor):
+ shift = torch.exp(logshift)
+ else:
+ shift = np.exp(logshift)
+ else:
+ shift = self.shift
+ return shift
+
+ def warp_t(self, t, seq_len=None):
+ shift = self.get_shift(seq_len=seq_len)
+ return shift * t / (1 + (shift - 1) * t)
+
+ def unwarp_t(self, t, seq_len=None):
+ shift = self.get_shift(seq_len=seq_len)
+ return t / (shift + (1 - shift) * t)
+
+ def sample(self, batch_size, warp_t=True, scale_t=True, seq_len=None,
+ raw_t_range=None, device=None):
+ if self.logit_normal_enable:
+ assert raw_t_range is None
+ t = torch.sigmoid(
+ self.logit_normal_mean + self.logit_normal_std * torch.randn(
+ (batch_size, ), dtype=torch.float, device=device))
+ else:
+ if raw_t_range is not None:
+ assert isinstance(raw_t_range, (tuple, list)) and len(raw_t_range) == 2
+ t = torch.rand(
+ (batch_size, ), dtype=torch.float, device=device
+ ) * (raw_t_range[0] - raw_t_range[1]) + raw_t_range[1]
+ else:
+ t = 1 - torch.rand((batch_size, ), dtype=torch.float, device=device)
+ if warp_t:
+ t = self.warp_t(t, seq_len=seq_len)
+ if scale_t:
+ t = t * self.num_timesteps
+ return t
+
+ def __call__(self, batch_size, **kwargs):
+ return self.sample(batch_size, **kwargs)
diff --git a/lakonlab/models/diffusions/schedulers/__init__.py b/lakonlab/models/diffusions/schedulers/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..a2380ce847713495355a4b1bc810a97e42a6d2cf
--- /dev/null
+++ b/lakonlab/models/diffusions/schedulers/__init__.py
@@ -0,0 +1,6 @@
+from .flow_euler_ode import FlowEulerODEScheduler
+from .flow_sde import FlowSDEScheduler
+from .flow_adapter import FlowAdapterScheduler
+
+
+__all__ = ['FlowEulerODEScheduler', 'FlowSDEScheduler', 'FlowAdapterScheduler']
diff --git a/lakonlab/models/diffusions/schedulers/__pycache__/__init__.cpython-310.pyc b/lakonlab/models/diffusions/schedulers/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..13b26977ae89ad84d085741b6a64d4d6efd32bfd
Binary files /dev/null and b/lakonlab/models/diffusions/schedulers/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/schedulers/__pycache__/flow_adapter.cpython-310.pyc b/lakonlab/models/diffusions/schedulers/__pycache__/flow_adapter.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..5ff3c06178d9c39e6ace26cef55a1f5a366a7136
Binary files /dev/null and b/lakonlab/models/diffusions/schedulers/__pycache__/flow_adapter.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/schedulers/__pycache__/flow_euler_ode.cpython-310.pyc b/lakonlab/models/diffusions/schedulers/__pycache__/flow_euler_ode.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..89298572f078e0a15c357c9832d0c839f6f1d2de
Binary files /dev/null and b/lakonlab/models/diffusions/schedulers/__pycache__/flow_euler_ode.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/schedulers/__pycache__/flow_sde.cpython-310.pyc b/lakonlab/models/diffusions/schedulers/__pycache__/flow_sde.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..a79fb826396849bdc509bf7651d8dec68d4ac920
Binary files /dev/null and b/lakonlab/models/diffusions/schedulers/__pycache__/flow_sde.cpython-310.pyc differ
diff --git a/lakonlab/models/diffusions/schedulers/flow_adapter.py b/lakonlab/models/diffusions/schedulers/flow_adapter.py
new file mode 100644
index 0000000000000000000000000000000000000000..f4bac2f5d9de8b7e00bfb324f14321f913e009af
--- /dev/null
+++ b/lakonlab/models/diffusions/schedulers/flow_adapter.py
@@ -0,0 +1,233 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import inspect
+import numpy as np
+import torch
+import diffusers
+
+from dataclasses import dataclass
+from typing import Optional, Tuple, Union
+from diffusers.configuration_utils import register_to_config
+from diffusers.utils import BaseOutput
+from diffusers.schedulers import SchedulerMixin
+from diffusers.configuration_utils import ConfigMixin
+
+
+@dataclass
+class FlowWrapperSchedulerOutput(BaseOutput):
+ prev_sample: torch.FloatTensor
+
+
+class FlowAdapterScheduler(SchedulerMixin, ConfigMixin):
+
+ order = 1
+
+ @register_to_config
+ def __init__(
+ self,
+ num_train_timesteps: int = 1000,
+ shift: float = 1.0,
+ use_dynamic_shifting=False,
+ base_seq_len=256,
+ max_seq_len=4096,
+ base_logshift=0.5,
+ max_logshift=1.15,
+ terminal_sigma=None,
+ base_scheduler='UniPCMultistep',
+ eps=1e-4,
+ **kwargs):
+
+ sigmas = torch.from_numpy(1 - np.linspace(
+ 0, 1, num_train_timesteps, dtype=np.float32, endpoint=False))
+ self.sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
+ self.timesteps = self.sigmas * num_train_timesteps
+ alphas = 1 - self.sigmas
+
+ base_scheduler_class = getattr(diffusers.schedulers, base_scheduler + 'Scheduler', None)
+
+ if base_scheduler_class is None:
+ raise AttributeError(f'Cannot find base_scheduler [{base_scheduler}].')
+ if base_scheduler in ['EulerDiscrete', 'EulerAncestralDiscrete']:
+ assert kwargs.get('prediction_type', 'epsilon') == 'epsilon'
+ kwargs['prediction_type'] = 'epsilon'
+ self.scales = ((alphas ** 2 + self.sigmas ** 2) / (
+ 1 + (self.sigmas / alphas.clamp(min=self.config.eps)) ** 2)).sqrt()
+ elif base_scheduler in [
+ 'DPMSolverSinglestep', 'DPMSolverMultistep', 'DEISMultistep', 'SASolver']:
+ assert kwargs.get('prediction_type', 'epsilon') == 'epsilon'
+ kwargs['prediction_type'] = 'epsilon'
+ self.scales = (alphas ** 2 + self.sigmas ** 2).sqrt()
+ elif base_scheduler in ['UniPCMultistep']:
+ self.scales = torch.ones_like(alphas)
+ assert kwargs.get('prediction_type', 'flow_prediction') == 'flow_prediction'
+ kwargs['prediction_type'] = 'flow_prediction'
+ kwargs['use_flow_sigmas'] = True
+ else:
+ raise AttributeError(f'Unsupported base_scheduler [{base_scheduler}].')
+
+ signatures = inspect.signature(base_scheduler_class).parameters.keys()
+ if 'final_sigmas_type' in signatures:
+ kwargs['final_sigmas_type'] = 'zero'
+ if 'lower_order_final' in signatures:
+ kwargs['lower_order_final'] = True
+
+ self.base_scheduler = base_scheduler_class(
+ num_train_timesteps=num_train_timesteps,
+ **kwargs)
+ self.base_scheduler.timesteps = self.timesteps
+ if self.config.base_scheduler in ['EulerDiscrete', 'EulerAncestralDiscrete']:
+ self.base_scheduler.sigmas = self.sigmas / alphas.clamp(min=self.config.eps)
+ elif self.config.base_scheduler in [
+ 'DPMSolverSinglestep', 'DPMSolverMultistep', 'DEISMultistep', 'SASolver']:
+ self.base_scheduler.sigmas = self.sigmas / alphas.clamp(min=self.config.eps)
+ elif self.config.base_scheduler in ['UniPCMultistep']:
+ self.base_scheduler.sigmas = self.sigmas
+ else:
+ raise AttributeError(f'Unsupported base_scheduler [{self.config.base_scheduler}].')
+
+ self._step_index = None
+ self._begin_index = None
+
+ @property
+ def step_index(self):
+ return self._step_index
+
+ @property
+ def begin_index(self):
+ return self._begin_index
+
+ def get_shift(self, seq_len=None):
+ if self.config.use_dynamic_shifting and seq_len is not None:
+ m = (self.config.max_logshift - self.config.base_logshift
+ ) / (self.config.max_seq_len - self.config.base_seq_len)
+ logshift = (seq_len - self.config.base_seq_len) * m + self.config.base_logshift
+ if isinstance(logshift, torch.Tensor):
+ shift = torch.exp(logshift)
+ else:
+ shift = np.exp(logshift)
+ else:
+ shift = self.config.shift
+ return shift
+
+ def stretch_to_terminal(self, sigma):
+ one_minus_sigma = 1 - sigma
+ stretched_sigma = 1 - (one_minus_sigma * (1 - self.config.terminal_sigma) / one_minus_sigma[-1])
+ return stretched_sigma
+
+ def set_timesteps(self, num_inference_steps: int, seq_len=None, device=None):
+ self.num_inference_steps = num_inference_steps
+
+ sigmas = torch.from_numpy(np.linspace(
+ 1, 0, num_inference_steps, dtype=np.float32, endpoint=False))
+ shift = self.get_shift(seq_len=seq_len)
+ sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
+
+ if self.config.terminal_sigma is not None:
+ sigmas = self.stretch_to_terminal(sigmas)
+
+ self.timesteps = (sigmas * self.config.num_train_timesteps).to(device)
+ if self.config.base_scheduler in ['DEISMultistep', 'SASolver']:
+ self.sigmas = torch.cat(
+ [sigmas, torch.tensor([self.config.eps], dtype=torch.float32, device=sigmas.device)])
+ else:
+ self.sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
+ alphas = 1 - self.sigmas
+
+ self.base_scheduler.set_timesteps(num_inference_steps, device=device)
+
+ self.base_scheduler.timesteps = self.timesteps
+ if self.config.base_scheduler in ['EulerDiscrete', 'EulerAncestralDiscrete']:
+ self.base_scheduler.sigmas = self.sigmas / alphas.clamp(min=self.config.eps)
+ self.scales = ((alphas ** 2 + self.sigmas ** 2) / (
+ 1 + (self.sigmas / alphas.clamp(min=self.config.eps)) ** 2)).sqrt()
+ elif self.config.base_scheduler in [
+ 'DPMSolverSinglestep', 'DPMSolverMultistep', 'DEISMultistep', 'SASolver']:
+ self.base_scheduler.sigmas = self.sigmas / alphas.clamp(min=self.config.eps)
+ self.scales = (alphas**2 + self.sigmas**2).sqrt()
+ elif self.config.base_scheduler in ['UniPCMultistep']:
+ self.base_scheduler.sigmas = self.sigmas.clamp(max=1 - self.config.eps)
+ self.scales = torch.ones_like(alphas)
+ else:
+ raise AttributeError(f'Unsupported base_scheduler [{self.config.base_scheduler}].')
+
+ self._step_index = None
+ self._begin_index = None
+
+ def index_for_timestep(self, timestep, schedule_timesteps=None):
+ if schedule_timesteps is None:
+ schedule_timesteps = self.timesteps
+
+ indices = (schedule_timesteps == timestep).nonzero()
+
+ # The sigma index that is taken for the **very** first `step`
+ # is always the second index (or the last index if there is only 1)
+ # This way we can ensure we don't accidentally skip a sigma in
+ # case we start in the middle of the denoising schedule (e.g. for image-to-image)
+ pos = 1 if len(indices) > 1 else 0
+
+ return indices[pos].item()
+
+ def _init_step_index(self, timestep):
+ if self.begin_index is None:
+ if isinstance(timestep, torch.Tensor):
+ timestep = timestep.to(self.timesteps.device)
+ self._step_index = self.index_for_timestep(timestep)
+ else:
+ self._step_index = self._begin_index
+
+ def step(
+ self,
+ model_output: torch.FloatTensor,
+ timestep: Union[float, torch.FloatTensor],
+ sample: torch.FloatTensor,
+ generator: Optional[torch.Generator] = None,
+ return_dict: bool = True,
+ prediction_type='u',
+ eps=1e-6) -> Union[FlowWrapperSchedulerOutput, Tuple]:
+ assert prediction_type in ['u', 'x0']
+
+ if self.step_index is None:
+ self._init_step_index(timestep)
+
+ # Upcast to avoid precision issues when computing prev_sample
+ ori_dtype = model_output.dtype
+ sample = sample.to(torch.float32)
+ model_output = model_output.to(torch.float32)
+
+ sigma = self.sigmas[self.step_index]
+ alpha = 1 - sigma
+ scale = self.scales[self.step_index]
+ next_scale = self.scales[self.step_index + 1]
+
+ if hasattr(self.base_scheduler, 'is_scale_input_called'):
+ self.base_scheduler.is_scale_input_called = True
+ kwargs = dict(return_dict=False)
+ if generator is not None:
+ kwargs.update(generator=generator)
+
+ if self.config.base_scheduler in ['UniPCMultistep']: # to u
+ if prediction_type == 'u':
+ model_output = model_output
+ else:
+ model_output = (sample - model_output) / sigma.clamp(min=eps)
+ else: # to epsilon
+ if prediction_type == 'u':
+ model_output = sample + alpha * model_output
+ else:
+ model_output = (sample - alpha * model_output) / sigma.clamp(min=eps)
+ prev_sample = self.base_scheduler.step(
+ model_output,
+ timestep,
+ sample / scale,
+ **kwargs
+ )[0] * next_scale
+
+ prev_sample = prev_sample.to(ori_dtype)
+
+ # upon completion increase step index by one
+ self._step_index += 1
+
+ if not return_dict:
+ return (prev_sample,)
+
+ return FlowWrapperSchedulerOutput(prev_sample=prev_sample)
diff --git a/lakonlab/models/diffusions/schedulers/flow_euler_ode.py b/lakonlab/models/diffusions/schedulers/flow_euler_ode.py
new file mode 100644
index 0000000000000000000000000000000000000000..35e83a67c85888381c4450e085d2ee72f2a936b2
--- /dev/null
+++ b/lakonlab/models/diffusions/schedulers/flow_euler_ode.py
@@ -0,0 +1,164 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import numpy as np
+import torch
+
+from dataclasses import dataclass
+from typing import Optional, Tuple, Union
+from diffusers.configuration_utils import ConfigMixin, register_to_config
+from diffusers.utils import BaseOutput, logging
+from diffusers.schedulers.scheduling_utils import SchedulerMixin
+
+logger = logging.get_logger(__name__) # pylint: disable=invalid-name
+
+
+@dataclass
+class FlowEulerODESchedulerOutput(BaseOutput):
+ prev_sample: torch.FloatTensor
+
+
+class FlowEulerODEScheduler(SchedulerMixin, ConfigMixin):
+
+ _compatibles = []
+ order = 1
+
+ @register_to_config
+ def __init__(
+ self,
+ num_train_timesteps: int = 1000,
+ shift: float = 1.0,
+ use_dynamic_shifting=False,
+ base_seq_len=256,
+ max_seq_len=4096,
+ base_logshift=0.5,
+ max_logshift=1.15,
+ terminal_sigma=None):
+ sigmas = torch.from_numpy(1 - np.linspace(
+ 0, 1, num_train_timesteps, dtype=np.float32, endpoint=False))
+ self.sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
+ self.timesteps = self.sigmas * num_train_timesteps
+
+ self._step_index = None
+ self._begin_index = None
+
+ self.sigma_min = self.sigmas[-1].item()
+ self.sigma_max = self.sigmas[0].item()
+
+ @property
+ def step_index(self):
+ return self._step_index
+
+ @property
+ def begin_index(self):
+ return self._begin_index
+
+ def set_begin_index(self, begin_index: int = 0):
+ self._begin_index = begin_index
+
+ def get_shift(self, seq_len=None):
+ if self.config.use_dynamic_shifting and seq_len is not None:
+ m = (self.config.max_logshift - self.config.base_logshift
+ ) / (self.config.max_seq_len - self.config.base_seq_len)
+ logshift = (seq_len - self.config.base_seq_len) * m + self.config.base_logshift
+ if isinstance(logshift, torch.Tensor):
+ shift = torch.exp(logshift)
+ else:
+ shift = np.exp(logshift)
+ else:
+ shift = self.config.shift
+ return shift
+
+ def stretch_to_terminal(self, sigma):
+ one_minus_sigma = 1 - sigma
+ stretched_sigma = 1 - (one_minus_sigma * (1 - self.config.terminal_sigma) / one_minus_sigma[-1])
+ return stretched_sigma
+
+ def set_timesteps(self, num_inference_steps: int, seq_len=None, device=None):
+ self.num_inference_steps = num_inference_steps
+
+ sigmas = torch.from_numpy(np.linspace(
+ 1, 0, num_inference_steps, dtype=np.float32, endpoint=False))
+ shift = self.get_shift(seq_len=seq_len)
+ sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
+
+ if self.config.terminal_sigma is not None:
+ sigmas = self.stretch_to_terminal(sigmas)
+
+ self.timesteps = (sigmas * self.config.num_train_timesteps).to(device)
+ self.sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
+
+ self._step_index = None
+ self._begin_index = None
+
+ def index_for_timestep(self, timestep, schedule_timesteps=None):
+ if schedule_timesteps is None:
+ schedule_timesteps = self.timesteps
+
+ indices = (schedule_timesteps == timestep).nonzero()
+
+ pos = 1 if len(indices) > 1 else 0
+
+ return indices[pos].item()
+
+ def _init_step_index(self, timestep):
+ if self.begin_index is None:
+ if isinstance(timestep, torch.Tensor):
+ timestep = timestep.to(self.timesteps.device)
+ self._step_index = self.index_for_timestep(timestep)
+ else:
+ self._step_index = self._begin_index
+
+ def step(
+ self,
+ model_output: torch.FloatTensor,
+ timestep: Union[float, torch.FloatTensor],
+ sample: torch.FloatTensor,
+ generator: Optional[torch.Generator] = None,
+ return_dict: bool = True,
+ prediction_type='u',
+ eps=1e-6) -> Union[FlowEulerODESchedulerOutput, Tuple]:
+ assert prediction_type in ['u', 'x0']
+
+ if isinstance(timestep, int) \
+ or isinstance(timestep, torch.IntTensor) \
+ or isinstance(timestep, torch.LongTensor):
+ raise ValueError(
+ (
+ 'Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to'
+ ' `EulerDiscreteScheduler.step()` is not supported. Make sure to pass'
+ ' one of the `scheduler.timesteps` as a timestep.'
+ ),
+ )
+
+ if self.step_index is None:
+ self._init_step_index(timestep)
+
+ # Upcast to avoid precision issues when computing prev_sample
+ ori_dtype = model_output.dtype
+ sample = sample.to(torch.float32)
+ model_output = model_output.to(torch.float32)
+
+ sigma = self.sigmas[self.step_index]
+ sigma_to = self.sigmas[self.step_index + 1]
+
+ if prediction_type == 'u':
+ derivative = model_output
+ else:
+ derivative = (sample - model_output) / sigma
+
+ dt = sigma_to - sigma
+ prev_sample = sample + derivative * dt
+
+ # Cast sample back to model compatible dtype
+ prev_sample = prev_sample.to(ori_dtype)
+
+ # upon completion increase step index by one
+ self._step_index += 1
+
+ if not return_dict:
+ return (prev_sample,)
+
+ return FlowEulerODESchedulerOutput(prev_sample=prev_sample)
+
+ def __len__(self):
+ return self.config.num_train_timesteps
diff --git a/lakonlab/models/diffusions/schedulers/flow_sde.py b/lakonlab/models/diffusions/schedulers/flow_sde.py
new file mode 100644
index 0000000000000000000000000000000000000000..af2ae88cdad8d3ca23a56f52d574677e8997e6f5
--- /dev/null
+++ b/lakonlab/models/diffusions/schedulers/flow_sde.py
@@ -0,0 +1,180 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import numpy as np
+import torch
+
+from dataclasses import dataclass
+from typing import Optional, Tuple, Union
+from diffusers.configuration_utils import ConfigMixin, register_to_config
+from diffusers.utils import BaseOutput, logging
+from diffusers.utils.torch_utils import randn_tensor
+from diffusers.schedulers.scheduling_utils import SchedulerMixin
+
+logger = logging.get_logger(__name__) # pylint: disable=invalid-name
+
+
+@dataclass
+class FlowSDESchedulerOutput(BaseOutput):
+ prev_sample: torch.FloatTensor
+
+
+class FlowSDEScheduler(SchedulerMixin, ConfigMixin):
+
+ _compatibles = []
+ order = 1
+
+ @register_to_config
+ def __init__(
+ self,
+ num_train_timesteps: int = 1000,
+ h: Union[float, str] = 1.0,
+ shift: float = 1.0,
+ use_dynamic_shifting=False,
+ base_seq_len=256,
+ max_seq_len=4096,
+ base_logshift=0.5,
+ max_logshift=1.15,
+ terminal_sigma=None):
+ sigmas = torch.from_numpy(1 - np.linspace(
+ 0, 1, num_train_timesteps, dtype=np.float32, endpoint=False))
+ self.sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
+ self.timesteps = self.sigmas * num_train_timesteps
+
+ self._step_index = None
+ self._begin_index = None
+
+ self.sigma_min = self.sigmas[-1].item()
+ self.sigma_max = self.sigmas[0].item()
+
+ @property
+ def step_index(self):
+ return self._step_index
+
+ @property
+ def begin_index(self):
+ return self._begin_index
+
+ def set_begin_index(self, begin_index: int = 0):
+ self._begin_index = begin_index
+
+ def get_shift(self, seq_len=None):
+ if self.config.use_dynamic_shifting and seq_len is not None:
+ m = (self.config.max_logshift - self.config.base_logshift
+ ) / (self.config.max_seq_len - self.config.base_seq_len)
+ logshift = (seq_len - self.config.base_seq_len) * m + self.config.base_logshift
+ if isinstance(logshift, torch.Tensor):
+ shift = torch.exp(logshift)
+ else:
+ shift = np.exp(logshift)
+ else:
+ shift = self.config.shift
+ return shift
+
+ def stretch_to_terminal(self, sigma):
+ one_minus_sigma = 1 - sigma
+ stretched_sigma = 1 - (one_minus_sigma * (1 - self.config.terminal_sigma) / one_minus_sigma[-1])
+ return stretched_sigma
+
+ def set_timesteps(self, num_inference_steps: int, seq_len=None, device=None):
+ self.num_inference_steps = num_inference_steps
+
+ sigmas = torch.from_numpy(np.linspace(
+ 1, 0, num_inference_steps, dtype=np.float32, endpoint=False))
+ shift = self.get_shift(seq_len=seq_len)
+ sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
+
+ if self.config.terminal_sigma is not None:
+ sigmas = self.stretch_to_terminal(sigmas)
+
+ self.timesteps = (sigmas * self.config.num_train_timesteps).to(device)
+ self.sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
+
+ self._step_index = None
+ self._begin_index = None
+
+ def index_for_timestep(self, timestep, schedule_timesteps=None):
+ if schedule_timesteps is None:
+ schedule_timesteps = self.timesteps
+
+ indices = (schedule_timesteps == timestep).nonzero()
+
+ pos = 1 if len(indices) > 1 else 0
+
+ return indices[pos].item()
+
+ def _init_step_index(self, timestep):
+ if self.begin_index is None:
+ if isinstance(timestep, torch.Tensor):
+ timestep = timestep.to(self.timesteps.device)
+ self._step_index = self.index_for_timestep(timestep)
+ else:
+ self._step_index = self._begin_index
+
+ def step(
+ self,
+ model_output: torch.FloatTensor,
+ timestep: Union[float, torch.FloatTensor],
+ sample: torch.FloatTensor,
+ generator: Optional[torch.Generator] = None,
+ return_dict: bool = True,
+ prediction_type='u',
+ eps=1e-6) -> Union[FlowSDESchedulerOutput, Tuple]:
+ assert prediction_type in ['u', 'x0']
+
+ if isinstance(timestep, int) \
+ or isinstance(timestep, torch.IntTensor) \
+ or isinstance(timestep, torch.LongTensor):
+ raise ValueError(
+ (
+ 'Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to'
+ ' `EulerDiscreteScheduler.step()` is not supported. Make sure to pass'
+ ' one of the `scheduler.timesteps` as a timestep.'
+ ),
+ )
+
+ if self.step_index is None:
+ self._init_step_index(timestep)
+
+ # Upcast to avoid precision issues when computing prev_sample
+ ori_dtype = model_output.dtype
+ sample = sample.to(torch.float32)
+ model_output = model_output.to(torch.float32)
+
+ sigma = self.sigmas[self.step_index]
+ sigma_to = self.sigmas[self.step_index + 1]
+ alpha = 1 - sigma
+ alpha_to = 1 - sigma_to
+
+ if prediction_type == 'u':
+ x0 = sample - sigma * model_output
+ epsilon = sample + alpha * model_output
+ else:
+ x0 = model_output
+ epsilon = (sample - alpha * x0) / sigma.clamp(min=eps)
+ noise = randn_tensor(
+ model_output.shape, dtype=torch.float32, device=model_output.device, generator=generator)
+
+ if self.config.h == 'inf':
+ m = torch.zeros_like(sigma)
+ elif self.config.h == 0.0:
+ m = torch.ones_like(sigma)
+ else:
+ assert self.config.h > 0.0
+ h2 = self.config.h * self.config.h
+ m = (sigma_to * alpha / (sigma * alpha_to).clamp(min=eps)) ** h2
+
+ prev_sample = alpha_to * x0 + sigma_to * (m * epsilon + (1 - m.square()).clamp(min=0).sqrt() * noise)
+
+ # Cast sample back to model compatible dtype
+ prev_sample = prev_sample.to(ori_dtype)
+
+ # upon completion increase step index by one
+ self._step_index += 1
+
+ if not return_dict:
+ return (prev_sample,)
+
+ return FlowSDESchedulerOutput(prev_sample=prev_sample)
+
+ def __len__(self):
+ return self.config.num_train_timesteps
diff --git a/lakonlab/models/latent_diffusion_class_image.py b/lakonlab/models/latent_diffusion_class_image.py
new file mode 100644
index 0000000000000000000000000000000000000000..31e1518194a6ec067aaadba7fc57c67936560c71
--- /dev/null
+++ b/lakonlab/models/latent_diffusion_class_image.py
@@ -0,0 +1,105 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import inspect
+import torch
+
+from copy import deepcopy
+from mmgen.models.builder import MODELS, build_module
+
+from .base_diffusion import BaseDiffusion
+from lakonlab.utils import rgetattr
+
+
+@MODELS.register_module()
+class LatentDiffusionClassImage(BaseDiffusion):
+
+ def __init__(self,
+ *args,
+ vae=dict(type='PretrainedVAE'),
+ **kwargs):
+ super().__init__(*args, **kwargs)
+ self.vae = build_module(vae) if vae is not None else None
+
+ def _prepare_train_minibatch_args(self, data, running_status=None):
+ if 'latents' in data:
+ latents = data['latents']
+ elif 'images' in data:
+ assert self.vae is not None, 'VAE must be provided for encoding images to latents.'
+ with torch.no_grad():
+ if hasattr(self.vae, 'dtype'):
+ vae_dtype = self.vae.dtype
+ else:
+ vae_dtype = next(self.vae.parameters()).dtype
+ latents = self.vae.encode((data['images'] * 2 - 1).to(vae_dtype)).float()
+ else:
+ raise ValueError('Either `latents` or `images` should be provided in the input data.')
+
+ labels = data['labels']
+ bs = labels.size(0)
+
+ diffusion_args = (self.patchify(latents), )
+
+ prob_class = self.train_cfg.get('prob_class', 1.0)
+ if prob_class < 1.0:
+ labels = torch.where(
+ torch.rand_like(labels, dtype=torch.float32) < prob_class,
+ labels, data['negative_labels'])
+
+ diffusion_kwargs = dict(class_labels=labels)
+
+ parameters = inspect.signature(rgetattr(self.diffusion, 'forward_train')).parameters
+
+ if 'running_status' in parameters:
+ diffusion_kwargs['running_status'] = running_status
+
+ if 'teacher' in parameters and 'teacher_kwargs' in parameters and self.teacher is not None:
+ teacher_guidance_scale = self.train_cfg.get('teacher_guidance_scale', None)
+ teacher_use_guidance = (teacher_guidance_scale is not None
+ and teacher_guidance_scale != 0.0 and teacher_guidance_scale != 1.0)
+ if teacher_use_guidance:
+ teacher_kwargs = dict(class_labels=torch.cat([data['negative_labels'], labels], dim=0))
+ teacher_kwargs.update(guidance_scale=teacher_guidance_scale)
+ else:
+ teacher_kwargs = dict(class_labels=labels)
+
+ diffusion_kwargs.update(
+ teacher=self.teacher,
+ teacher_kwargs=teacher_kwargs)
+
+ return bs, diffusion_args, diffusion_kwargs
+
+ def val_step(self, data, test_cfg_override=dict(), **kwargs):
+ bs = len(data['labels'])
+ cfg = deepcopy(self.test_cfg)
+ cfg.update(test_cfg_override)
+ guidance_scale = cfg.get('guidance_scale', 1.0)
+ diffusion = self.diffusion_ema if self.diffusion_use_ema else self.diffusion
+
+ with torch.no_grad():
+ class_labels = data['labels']
+
+ if guidance_scale != 0.0 and guidance_scale != 1.0:
+ class_labels = torch.cat([data['negative_labels'], class_labels], dim=0)
+
+ if 'noise' in data:
+ noise = data['noise']
+ else:
+ latent_size = cfg['latent_size']
+ noise = torch.randn((bs, *latent_size), device=data['labels'].device)
+ noise = self.patchify(noise)
+ latents_out = diffusion(
+ noise=noise,
+ class_labels=class_labels,
+ guidance_scale=guidance_scale,
+ test_cfg_override=test_cfg_override)
+ latents_out = self.unpatchify(latents_out)
+
+ if hasattr(self.vae, 'dtype'):
+ vae_dtype = self.vae.dtype
+ else:
+ vae_dtype = next(self.vae.parameters()).dtype
+ latents_out = latents_out.to(vae_dtype)
+
+ out_images = (self.vae.decode(latents_out).float() / 2 + 0.5).clamp(min=0, max=1)
+
+ return dict(num_samples=bs, pred_imgs=out_images)
diff --git a/lakonlab/models/latent_diffusion_text_image.py b/lakonlab/models/latent_diffusion_text_image.py
new file mode 100644
index 0000000000000000000000000000000000000000..0c83beb61572d062eb3759217e6426706e9fc9b0
--- /dev/null
+++ b/lakonlab/models/latent_diffusion_text_image.py
@@ -0,0 +1,170 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import torch
+import inspect
+
+from copy import deepcopy
+from mmgen.models.builder import MODELS, build_module
+
+from .base_diffusion import BaseDiffusion
+from lakonlab.utils import rgetattr
+
+
+@MODELS.register_module()
+class LatentDiffusionTextImage(BaseDiffusion):
+
+ def __init__(self,
+ *args,
+ vae=None,
+ text_encoder=None,
+ **kwargs):
+ super().__init__(*args, **kwargs)
+ self.vae = build_module(vae) if vae is not None else None
+ self.text_encoder = build_module(text_encoder) if text_encoder is not None else None
+
+ def _prepare_train_minibatch_diffusion_args(self, data):
+ if 'prompt_embed_kwargs' in data:
+ prompt_embed_kwargs = data['prompt_embed_kwargs']
+ elif 'prompt_kwargs' in data:
+ assert self.text_encoder is not None, 'Text encoder must be provided for encoding text to embeddings.'
+ prompt_embed_kwargs = self.text_encoder(**data['prompt_kwargs'])
+ else:
+ raise ValueError('Either `prompt_embed_kwargs` or `prompt_kwargs` should be provided in the input data.')
+
+ if 'latents' in data:
+ latents = data['latents']
+ elif 'images' in data:
+ assert self.vae is not None, 'VAE must be provided for encoding images to latents.'
+ with torch.no_grad():
+ if hasattr(self.vae, 'dtype'):
+ vae_dtype = self.vae.dtype
+ else:
+ vae_dtype = next(self.vae.parameters()).dtype
+ latents = self.vae.encode((data['images'] * 2 - 1).to(vae_dtype)).float()
+ else:
+ raise ValueError('Either `latents` or `images` should be provided in the input data.')
+
+ v = next(iter(prompt_embed_kwargs.values()))
+ bs = v.size(0)
+ device = v.device
+
+ diffusion_args = (self.patchify(latents), )
+ diffusion_kwargs = prompt_embed_kwargs.copy()
+
+ distilled_guidance_scale = self.train_cfg.get('distilled_guidance_scale', None)
+ if distilled_guidance_scale is not None:
+ distilled_guidance_scale = torch.full(
+ (bs,), distilled_guidance_scale, dtype=torch.float32, device=device)
+ diffusion_kwargs.update(guidance=distilled_guidance_scale)
+
+ return diffusion_args, diffusion_kwargs, prompt_embed_kwargs, bs, device
+
+ def _prepare_train_minibatch_teacher_args(self, data, prompt_embed_kwargs, bs, device):
+ teacher_guidance_scale = self.train_cfg.get('teacher_guidance_scale', None)
+ teacher_use_guidance = (teacher_guidance_scale is not None
+ and teacher_guidance_scale != 0.0 and teacher_guidance_scale != 1.0)
+ if teacher_use_guidance:
+ if 'negative_prompt_embed_kwargs' in data:
+ negative_prompt_embed_kwargs = data['negative_prompt_embed_kwargs']
+ elif 'negative_prompt_kwargs' in data:
+ negative_prompt_embed_kwargs = self.text_encoder(**data['negative_prompt_kwargs'])
+ else:
+ raise ValueError(
+ 'Either `negative_prompt_embed_kwargs` or `negative_prompt_kwargs` should be provided in the '
+ 'input data for classifier-free guidance.')
+ teacher_kwargs = {
+ k: torch.cat([negative_prompt_embed_kwargs[k], v], dim=0)
+ for k, v in prompt_embed_kwargs.items()}
+ teacher_kwargs.update(guidance_scale=teacher_guidance_scale)
+ else:
+ teacher_kwargs = prompt_embed_kwargs.copy()
+
+ teacher_distilled_guidance_scale = self.train_cfg.get('teacher_distilled_guidance_scale', None)
+ if teacher_distilled_guidance_scale is not None:
+ teacher_distilled_guidance_scale = torch.full(
+ (bs * 2,) if teacher_use_guidance else (bs,),
+ teacher_distilled_guidance_scale, dtype=torch.float32, device=device)
+ teacher_kwargs.update(guidance=teacher_distilled_guidance_scale)
+
+ return teacher_kwargs
+
+ def _prepare_train_minibatch_args(self, data, running_status=None):
+ diffusion_args, diffusion_kwargs, prompt_embed_kwargs, bs, device = \
+ self._prepare_train_minibatch_diffusion_args(data)
+ parameters = inspect.signature(rgetattr(self.diffusion, 'forward_train')).parameters
+ if 'running_status' in parameters:
+ diffusion_kwargs['running_status'] = running_status
+
+ if 'teacher' in parameters and 'teacher_kwargs' in parameters and self.teacher is not None:
+ teacher_kwargs = self._prepare_train_minibatch_teacher_args(
+ data, prompt_embed_kwargs, bs, device)
+
+ diffusion_kwargs.update(
+ teacher=self.teacher,
+ teacher_kwargs=teacher_kwargs)
+
+ return bs, diffusion_args, diffusion_kwargs
+
+ def val_step(self, data, test_cfg_override=dict(), **kwargs):
+ if 'prompt_embed_kwargs' in data:
+ prompt_embed_kwargs = data['prompt_embed_kwargs']
+ elif 'prompt_kwargs' in data:
+ assert self.text_encoder is not None, 'Text encoder must be provided for encoding text to embeddings.'
+ prompt_embed_kwargs = self.text_encoder(**data['prompt_kwargs'])
+ else:
+ raise ValueError('Either `prompt_embed_kwargs` or `prompt_kwargs` should be provided in the input data.')
+
+ v = next(iter(prompt_embed_kwargs.values()))
+ bs = v.size(0)
+ device = v.device
+
+ cfg = deepcopy(self.test_cfg)
+ cfg.update(test_cfg_override)
+ guidance_scale = cfg.get('guidance_scale', 1.0)
+ diffusion = self.diffusion_ema if self.diffusion_use_ema else self.diffusion
+
+ with torch.no_grad():
+ use_guidance = guidance_scale != 0.0 and guidance_scale != 1.0
+ if use_guidance:
+ if 'negative_prompt_embed_kwargs' in data:
+ negative_prompt_embed_kwargs = data['negative_prompt_embed_kwargs']
+ elif 'negative_prompt_kwargs' in data:
+ negative_prompt_embed_kwargs = self.text_encoder(**data['negative_prompt_kwargs'])
+ else:
+ raise ValueError(
+ 'Either `negative_prompt_embed_kwargs` or `negative_prompt_kwargs` should be provided in the '
+ 'input data for classifier-free guidance.')
+ kwargs = {
+ k: torch.cat([negative_prompt_embed_kwargs[k], v], dim=0)
+ for k, v in prompt_embed_kwargs.items()}
+ else:
+ kwargs = prompt_embed_kwargs.copy()
+ distilled_guidance_scale = cfg.get('distilled_guidance_scale', None)
+ if distilled_guidance_scale is not None:
+ distilled_guidance_scale = torch.full(
+ (bs * 2,) if use_guidance else (bs,),
+ distilled_guidance_scale, dtype=torch.float32, device=device)
+ kwargs.update(guidance=distilled_guidance_scale)
+
+ if 'noise' in data:
+ noise = data['noise']
+ else:
+ latent_size = cfg['latent_size']
+ noise = torch.randn((bs, *latent_size), device=device)
+ noise = self.patchify(noise)
+ latents_out = diffusion(
+ noise=noise,
+ guidance_scale=guidance_scale,
+ test_cfg_override=test_cfg_override,
+ **kwargs)
+ latents_out = self.unpatchify(latents_out)
+
+ if hasattr(self.vae, 'dtype'):
+ vae_dtype = self.vae.dtype
+ else:
+ vae_dtype = next(self.vae.parameters()).dtype
+ latents_out = latents_out.to(vae_dtype)
+
+ out_images = (self.vae.decode(latents_out).float() / 2 + 0.5).clamp(min=0, max=1)
+
+ return dict(num_samples=bs, pred_imgs=out_images)
diff --git a/lakonlab/models/losses/__init__.py b/lakonlab/models/losses/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..15439383ab36da7f2476fdf923729e7bb07d9778
--- /dev/null
+++ b/lakonlab/models/losses/__init__.py
@@ -0,0 +1,3 @@
+from .diffusion_loss import DiffusionMSELoss, DiffusionNLLLoss, GMFlowNLLLoss
+
+__all__ = ['DiffusionMSELoss', 'DiffusionNLLLoss', 'GMFlowNLLLoss']
diff --git a/lakonlab/models/losses/__pycache__/__init__.cpython-310.pyc b/lakonlab/models/losses/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..69175768ccfc7f2274ac3c233d0c1e4638f7bcaf
Binary files /dev/null and b/lakonlab/models/losses/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/models/losses/__pycache__/diffusion_loss.cpython-310.pyc b/lakonlab/models/losses/__pycache__/diffusion_loss.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..9f06f49fdb812f0a63521918151165b5e97d36b1
Binary files /dev/null and b/lakonlab/models/losses/__pycache__/diffusion_loss.cpython-310.pyc differ
diff --git a/lakonlab/models/losses/diffusion_loss.py b/lakonlab/models/losses/diffusion_loss.py
new file mode 100644
index 0000000000000000000000000000000000000000..b9095d7592eedc305f9473c4973c7e993b937fd4
--- /dev/null
+++ b/lakonlab/models/losses/diffusion_loss.py
@@ -0,0 +1,291 @@
+import math
+import torch
+import torch.distributed as dist
+
+from functools import partial
+from mmgen.models import MODULES
+from mmgen.models.losses.ddpm_loss import DDPMLoss, mse_loss, reduce_loss
+from mmgen.models.losses.utils import weighted_loss
+
+from lakonlab.ops.gmflow_ops.gmflow_ops import gm_logprob
+
+
+@weighted_loss
+def gaussian_nll_loss(pred, target, logstd, eps=1e-4):
+ inverse_std = torch.exp(-logstd).clamp(max=1 / eps)
+ diff_weighted = (pred - target) * inverse_std
+ loss = 0.5 * (diff_weighted.square() + math.log(2 * math.pi)) + logstd
+ return loss
+
+
+@weighted_loss
+def gaussian_mixture_nll_loss(
+ pred_means, target, pred_logstds, pred_logweights):
+ """
+ Args:
+ pred_means (torch.Tensor): Shape (bs, *, num_gaussians, c, h, w)
+ target (torch.Tensor): Shape (bs, *, c, h, w)
+ pred_logstds (torch.Tensor): Shape (bs, *, 1 or num_gaussians, 1 or c, 1 or h, 1 or w)
+ pred_logweights (torch.Tensor): Shape (bs, *, num_gaussians, 1, h, w)
+
+ Returns:
+ torch.Tensor: Shape (bs, *, h, w)
+ """
+ num_channels = pred_means.size(-3)
+ loss = -gm_logprob(
+ dict(
+ means=pred_means,
+ logstds=pred_logstds,
+ logweights=pred_logweights),
+ target.unsqueeze(-4))[0]
+ return loss.squeeze(-3) / num_channels
+
+
+@MODULES.register_module()
+class DiffusionMSELoss(DDPMLoss):
+ _default_data_info = dict(pred='eps_t_pred', target='noise')
+
+ def __init__(self,
+ rescale_mode='constant',
+ rescale_cfg=dict(scale=1.0),
+ sampler=None,
+ weight=None,
+ log_cfgs=None,
+ reduction='mean',
+ data_info=None,
+ loss_name='loss_mse'):
+ super().__init__(rescale_mode=rescale_mode,
+ rescale_cfg=rescale_cfg,
+ log_cfgs=log_cfgs,
+ weight=weight,
+ sampler=sampler,
+ reduction=reduction,
+ loss_name=loss_name)
+
+ self.data_info = self._default_data_info \
+ if data_info is None else data_info
+
+ self.loss_fn = partial(mse_loss, reduction='flatmean')
+
+ def _forward_loss(self, outputs_dict):
+ """Forward function for loss calculation.
+ Args:
+ outputs_dict (dict): Outputs of the model used to calculate losses.
+
+ Returns:
+ torch.Tensor: Calculated loss.
+ """
+ loss_input_dict = {
+ k: outputs_dict[v]
+ for k, v in self.data_info.items()
+ }
+ loss = self.loss_fn(**loss_input_dict) * 0.5
+ return loss
+
+
+@MODULES.register_module()
+class DiffusionNLLLoss(DDPMLoss):
+ _default_data_info = dict(pred='u_t_pred', target='u_t', logstd='logstd')
+
+ def __init__(self,
+ rescale_mode='constant',
+ rescale_cfg=dict(scale=1.0),
+ log_cfgs=None,
+ data_info=None,
+ reduction='mean',
+ loss_name='loss_nll'):
+ super().__init__(
+ rescale_mode=rescale_mode,
+ rescale_cfg=rescale_cfg,
+ log_cfgs=log_cfgs,
+ reduction=reduction,
+ loss_name=loss_name)
+ self.data_info = self._default_data_info \
+ if data_info is None else data_info
+ self.loss_fn = partial(gaussian_nll_loss, reduction='flatmean')
+ if log_cfgs is not None and log_cfgs.get('type', None) == 'quartile':
+ for i in range(4):
+ self.register_buffer(f'loss_quartile_{i}', torch.zeros((1, ), dtype=torch.float))
+ self.register_buffer(f'var_quartile_{i}', torch.ones((1, ), dtype=torch.float))
+ self.register_buffer(f'count_quartile_{i}', torch.zeros((1, ), dtype=torch.long))
+
+ @torch.no_grad()
+ def collect_log(self, loss, var, timesteps):
+ if not self.log_fn_list:
+ return
+
+ if dist.is_initialized():
+ ws = dist.get_world_size()
+ placeholder_l = [torch.zeros_like(loss) for _ in range(ws)]
+ placeholder_v = [torch.zeros_like(var) for _ in range(ws)]
+ placeholder_t = [torch.zeros_like(timesteps) for _ in range(ws)]
+ dist.all_gather(placeholder_l, loss)
+ dist.all_gather(placeholder_v, var)
+ dist.all_gather(placeholder_t, timesteps)
+ loss = torch.cat(placeholder_l, dim=0)
+ var = torch.cat(placeholder_v, dim=0)
+ timesteps = torch.cat(placeholder_t, dim=0)
+ log_vars = dict()
+
+ if (dist.is_initialized()
+ and dist.get_rank() == 0) or not dist.is_initialized():
+ for log_fn in self.log_fn_list:
+ log_vars.update(log_fn(loss, var, timesteps))
+ self.log_vars = log_vars
+
+ @torch.no_grad()
+ def quartile_log_collect(self,
+ loss,
+ var,
+ timesteps,
+ total_timesteps,
+ prefix_name,
+ reduction='mean',
+ momentum=0.1):
+ quartile = (timesteps / total_timesteps * 4)
+ quartile = quartile.to(torch.long).clamp(min=0, max=3)
+
+ log_vars = dict()
+
+ for idx in range(4):
+ quartile_mask = quartile == idx
+ quartile_count = torch.count_nonzero(quartile_mask).reshape(1)
+ if quartile_count > 0:
+ loss_quartile = reduce_loss(loss[quartile_mask], reduction).reshape(1)
+ var_quartile = reduce_loss(var[quartile_mask], reduction).reshape(1)
+
+ cur_weight = 1 - torch.exp(-momentum * quartile_count)
+ getattr(self, f'count_quartile_{idx}').add_(quartile_count)
+ total_weight = 1 - torch.exp(-momentum * getattr(self, f'count_quartile_{idx}'))
+ cur_weight /= total_weight.clamp(min=1e-4)
+ getattr(self, f'loss_quartile_{idx}').mul_(1 - cur_weight).add_(loss_quartile * cur_weight)
+ getattr(self, f'var_quartile_{idx}').mul_(1 - cur_weight).add_(var_quartile * cur_weight)
+
+ log_vars[f'{prefix_name}_quartile_{idx}'] = getattr(self, f'loss_quartile_{idx}').item()
+ log_vars[f'{prefix_name}_var_quartile_{idx}'] = getattr(self, f'var_quartile_{idx}').item()
+
+ return log_vars
+
+ def _forward_loss(self, outputs_dict):
+ loss_input_dict = {
+ k: outputs_dict[v]
+ for k, v in self.data_info.items()
+ }
+ loss = self.loss_fn(**loss_input_dict)
+ return loss
+
+ def forward(self, *args, **kwargs):
+ if len(args) == 1:
+ assert isinstance(args[0], dict), (
+ 'You should offer a dictionary containing network outputs '
+ 'for building up computational graph of this loss module.')
+ output_dict = args[0]
+ elif 'output_dict' in kwargs:
+ assert len(args) == 0, (
+ 'If the outputs dict is given in keyworded arguments, no'
+ ' further non-keyworded arguments should be offered.')
+ output_dict = kwargs.pop('outputs_dict')
+ else:
+ raise NotImplementedError(
+ 'Cannot parsing your arguments passed to this loss module.'
+ ' Please check the usage of this module')
+
+ # check keys in output_dict
+ assert 'timesteps' in output_dict, (
+ '\'timesteps\' is must for DDPM-based losses, but found'
+ f'{output_dict.keys()} in \'output_dict\'')
+
+ timesteps = output_dict['timesteps']
+ loss = self._forward_loss(output_dict)
+
+ with torch.no_grad():
+ var = torch.exp(output_dict['logstd'] * 2) # (bs, *)
+ if 'weight' in self.data_info:
+ weight = output_dict[self.data_info['weight']] # (bs, *)
+ weight_norm_factor = weight.flatten(1).mean(dim=1).clamp(min=1e-6)
+ _var = (var * weight).flatten(1).mean(dim=1) / weight_norm_factor
+ _loss = loss / weight_norm_factor
+ else:
+ _var = var.flatten(1).mean(dim=1)
+ _loss = loss
+
+ # update log_vars of this class
+ self.collect_log(_loss, _var, timesteps=timesteps) # Mod: log after rescaling
+
+ loss_rescaled = self.rescale_fn(loss, timesteps)
+ return reduce_loss(loss_rescaled, self.reduction)
+
+
+@MODULES.register_module()
+class GMFlowNLLLoss(DiffusionNLLLoss):
+ _default_data_info = dict(
+ pred_means='means',
+ target='u_t',
+ pred_logstds='logstds',
+ pred_logweights='logweights')
+
+ def __init__(self,
+ rescale_mode='constant',
+ rescale_cfg=dict(scale=1.0),
+ log_cfgs=None,
+ data_info=None,
+ reduction='mean',
+ loss_name='loss_nll'):
+ super().__init__(
+ rescale_mode=rescale_mode,
+ rescale_cfg=rescale_cfg,
+ log_cfgs=log_cfgs,
+ reduction=reduction,
+ loss_name=loss_name)
+ self.data_info = self._default_data_info \
+ if data_info is None else data_info
+ self.loss_fn = partial(gaussian_mixture_nll_loss, reduction='flatmean')
+ if log_cfgs is not None and log_cfgs.get('type', None) == 'quartile':
+ for i in range(4):
+ self.register_buffer(f'loss_quartile_{i}', torch.zeros((1,), dtype=torch.float))
+ self.register_buffer(f'var_quartile_{i}', torch.ones((1,), dtype=torch.float))
+ self.register_buffer(f'count_quartile_{i}', torch.zeros((1,), dtype=torch.long))
+
+ def forward(self, *args, **kwargs):
+ if len(args) == 1:
+ assert isinstance(args[0], dict), (
+ 'You should offer a dictionary containing network outputs '
+ 'for building up computational graph of this loss module.')
+ output_dict = args[0]
+ elif 'output_dict' in kwargs:
+ assert len(args) == 0, (
+ 'If the outputs dict is given in keyworded arguments, no'
+ ' further non-keyworded arguments should be offered.')
+ output_dict = kwargs.pop('outputs_dict')
+ else:
+ raise NotImplementedError(
+ 'Cannot parsing your arguments passed to this loss module.'
+ ' Please check the usage of this module')
+
+ # check keys in output_dict
+ assert 'timesteps' in output_dict, (
+ '\'timesteps\' is must for DDPM-based losses, but found'
+ f'{output_dict.keys()} in \'output_dict\'')
+
+ timesteps = output_dict['timesteps']
+ loss = self._forward_loss(output_dict)
+
+ with torch.no_grad():
+ weights = output_dict['logweights'].exp()
+ mean = (weights * output_dict['means']).sum(-4, keepdim=True) # (bs, *, 1, c, h, w)
+ var = (weights * ((output_dict['means'] - mean).square()
+ + (output_dict['logstds'] * 2).exp())).sum(-4) # (bs, *, c, h, w)
+ if 'weight' in self.data_info:
+ weight = output_dict[self.data_info['weight']].unsqueeze(-3) # (bs, *, 1, h, w)
+ weight_norm_factor = weight.flatten(1).mean(dim=1).clamp(min=1e-6)
+ _var = (var * weight).flatten(1).mean(dim=1) / weight_norm_factor
+ _loss = loss / weight_norm_factor
+ else:
+ _var = var.flatten(1).mean(dim=1)
+ _loss = loss
+
+ # update log_vars of this class
+ self.collect_log(_loss, _var, timesteps=timesteps) # Mod: log after rescaling
+
+ loss_rescaled = self.rescale_fn(loss, timesteps)
+ return reduce_loss(loss_rescaled, self.reduction)
diff --git a/lakonlab/ops/__init__.py b/lakonlab/ops/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/lakonlab/ops/__pycache__/__init__.cpython-310.pyc b/lakonlab/ops/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..f68b83f8efb5092e00068a3c9e47054359cdfd0f
Binary files /dev/null and b/lakonlab/ops/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/ops/gmflow_ops/__init__.py b/lakonlab/ops/gmflow_ops/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/lakonlab/ops/gmflow_ops/__pycache__/__init__.cpython-310.pyc b/lakonlab/ops/gmflow_ops/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..3e2db438859e6f5a8eb5826aea84b51829e4f0f9
Binary files /dev/null and b/lakonlab/ops/gmflow_ops/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/ops/gmflow_ops/__pycache__/gmflow_ops.cpython-310.pyc b/lakonlab/ops/gmflow_ops/__pycache__/gmflow_ops.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..c62efe6c3727f79635ebfec42572ee13bbeac766
Binary files /dev/null and b/lakonlab/ops/gmflow_ops/__pycache__/gmflow_ops.cpython-310.pyc differ
diff --git a/lakonlab/ops/gmflow_ops/backend.py b/lakonlab/ops/gmflow_ops/backend.py
new file mode 100644
index 0000000000000000000000000000000000000000..1e61381c8e7ab73c8e8efa0b4ae33380a1764021
--- /dev/null
+++ b/lakonlab/ops/gmflow_ops/backend.py
@@ -0,0 +1,41 @@
+import os
+from torch.utils.cpp_extension import load
+
+_src_path = os.path.dirname(os.path.abspath(__file__))
+
+nvcc_flags = [
+ '-O3', '-std=c++17',
+ '-U__CUDA_NO_HALF_OPERATORS__', '-U__CUDA_NO_HALF_CONVERSIONS__', '-U__CUDA_NO_HALF2_OPERATORS__',
+]
+
+if os.name == "posix":
+ c_flags = ['-O3', '-std=c++17']
+elif os.name == "nt":
+ c_flags = ['/O2', '/std:c++17']
+
+ # find cl.exe
+ def find_cl_path():
+ import glob
+ for program_files in [r"C:\\Program Files (x86)", r"C:\\Program Files"]:
+ for edition in ["Enterprise", "Professional", "BuildTools", "Community"]:
+ paths = sorted(glob.glob(r"%s\\Microsoft Visual Studio\\*\\%s\\VC\\Tools\\MSVC\\*\\bin\\Hostx64\\x64" % (program_files, edition)), reverse=True)
+ if paths:
+ return paths[0]
+
+ # If cl.exe is not on path, try to find it.
+ if os.system("where cl.exe >nul 2>nul") != 0:
+ cl_path = find_cl_path()
+ if cl_path is None:
+ raise RuntimeError("Could not locate a supported Microsoft Visual C++ installation")
+ os.environ["PATH"] += ";" + cl_path
+
+_backend = load(name='_gmflow_ops',
+ extra_cflags=c_flags,
+ extra_cuda_cflags=nvcc_flags,
+ sources=[os.path.join(_src_path, 'src', f) for f in [
+ 'gmflow_ops.cu',
+ 'bindings.cpp',
+ ]],
+ )
+
+__all__ = ['_backend']
diff --git a/lakonlab/ops/gmflow_ops/gmflow_ops.py b/lakonlab/ops/gmflow_ops/gmflow_ops.py
new file mode 100644
index 0000000000000000000000000000000000000000..a045327d5b3693368260856694f099df22845845
--- /dev/null
+++ b/lakonlab/ops/gmflow_ops/gmflow_ops.py
@@ -0,0 +1,1144 @@
+# Copyright (c) 2025 Hansheng Chen
+
+# A complete library for Gaussian mixture operations in PyTorch.
+# Some of the functions are not used in GMFlow but are kept for future reference.
+
+import math
+import torch
+
+from torch.autograd import Function
+from torch.amp import custom_fwd
+
+BACKEND = None
+
+
+def get_backend():
+ global BACKEND
+
+ if BACKEND is None:
+ try:
+ import _gmflow_ops as _backend
+ except ImportError:
+ from .backend import _backend
+
+ BACKEND = _backend
+
+ return BACKEND
+
+
+class _gm1d_inverse_cdf(Function):
+ @staticmethod
+ @custom_fwd(device_type='cuda', cast_inputs=torch.float32)
+ def forward(ctx, gm1d, scaled_cdfs, n_steps, eps, max_step_size, init_samples):
+ means = gm1d['means']
+ logstds = gm1d['logstds']
+ logweights = gm1d['logweights']
+ gm_weights = gm1d['gm_weights']
+
+ batch_shapes = max(means.shape[:-3], scaled_cdfs.shape[:-3])
+ num_gaussians, h, w = means.shape[-3:]
+ batch_numel = batch_shapes.numel()
+ n_samples = scaled_cdfs.size(-3)
+ assert n_steps >= 0
+
+ samples = init_samples.expand(*batch_shapes, n_samples, h, w).reshape(
+ batch_numel, n_samples, h, w).contiguous()
+ get_backend().gm1d_inverse_cdf(
+ means.expand(*batch_shapes, num_gaussians, h, w).reshape(
+ batch_numel, num_gaussians, h, w).contiguous(),
+ logstds.expand(*batch_shapes, 1, 1, 1).reshape(
+ batch_numel, 1, 1, 1).contiguous(),
+ logweights.expand(*batch_shapes, num_gaussians, h, w).reshape(
+ batch_numel, num_gaussians, h, w).contiguous(),
+ gm_weights.expand(*batch_shapes, num_gaussians, h, w).reshape(
+ batch_numel, num_gaussians, h, w).contiguous(),
+ scaled_cdfs.expand(*batch_shapes, n_samples, h, w).reshape(
+ batch_numel, n_samples, h, w).contiguous(),
+ samples,
+ n_steps, eps, max_step_size)
+
+ return samples.reshape(*batch_shapes, n_samples, h, w)
+
+
+def psd_inverse(x):
+ return torch.cholesky_inverse(torch.linalg.cholesky(x))
+
+
+def gm1d_pdf_cdf(gm1d, samples):
+ """
+ Args:
+ gm1d (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, h, w)
+ (Optional) gm_vars (torch.Tensor): (bs, *, 1, 1, 1)
+ (Optional) gm_weights (torch.Tensor): (bs, *, num_gaussians, h, w)
+ samples (torch.Tensor): (bs, *, n_samples, h, w)
+
+ Returns:
+ tuple[torch.Tensor, torch.Tensor]: (pdf, cdf) in shape (bs, *, n_samples, h, w)
+ Note that the CDF is scaled to [-1, 1].
+ """
+ gm_logstds = gm1d['logstds'].unsqueeze(-4) # (bs, *, 1, 1, 1, 1)
+ gm_stds = gm_logstds.exp() # (bs, *, 1, 1, 1, 1)
+ gm_logweights = gm1d['logweights'].unsqueeze(-4) # (bs, *, 1, num_gaussians, h, w)
+ if 'gm_weights' in gm1d:
+ gm_weights = gm1d['gm_weights'].unsqueeze(-4) # (bs, *, 1, num_gaussians, h, w)
+ else:
+ gm_weights = gm_logweights.exp() # (bs, *, 1, num_gaussians, h, w)
+ # (bs, *, n_samples, num_gaussians, h, w)
+ gm1d_norm_diffs = (samples.unsqueeze(-3) - gm1d['means'].unsqueeze(-4)) / gm_stds
+ pdf = (-0.5 * gm1d_norm_diffs.square() - gm_logstds + gm_logweights).exp().sum(dim=-3) / math.sqrt(2 * math.pi)
+ cdf = (gm_weights * torch.erf(gm1d_norm_diffs / math.sqrt(2))).sum(dim=-3)
+ return pdf, cdf
+
+
+@torch.jit.script
+def gm1d_inverse_cdf_init_jit(gaussian_samples, var, mean):
+ init_samples = gaussian_samples * var.sqrt() + mean
+ return init_samples
+
+
+def gm1d_inverse_cdf(
+ gm1d, scaled_cdfs, n_steps=8, eps=1e-6, max_step_size=1.5, gaussian_samples=None, backward_steps=2):
+ """
+ Inverse CDF for 1D Gaussian mixture using Newton-Raphson method.
+
+ Args:
+ gm1d (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, h, w)
+ (Optional) gm_vars (torch.Tensor): (bs, *, 1, 1, 1)
+ (Optional) gm_weights (torch.Tensor): (bs, *, num_gaussians, h, w)
+ scaled_cdfs (torch.Tensor): (bs, *, n_samples, h, w) scaled to range [-1, 1]
+ gaussian_samples (torch.Tensor | None): (bs, *, n_samples, h, w) standard Gaussian samples with scaled_cdfs
+ n_steps (int)
+ eps (float)
+ max_step_size (float)
+
+ Returns:
+ torch.Tensor: (bs, *, n_samples, h, w)
+ """
+ _gm1d = {k: v.unsqueeze(-3) for k, v in gm1d.items()}
+ gaussian_proxy = gm_to_iso_gaussian(_gm1d)[0]
+ gm1d['gm_weights'] = _gm1d['gm_weights'].squeeze(-3)
+ gm1d['gm_vars'] = _gm1d['gm_vars'].squeeze(-3)
+
+ with torch.no_grad():
+ if gaussian_samples is None:
+ gaussian_samples = (torch.erfinv(
+ scaled_cdfs.float().clamp(min=-1 + eps, max=1 - eps)
+ ) * math.sqrt(2)).to(dtype=scaled_cdfs.dtype) # (bs, *, n_samples, h, w)
+ # (bs, *, n_samples, h, w)
+ init_samples = gm1d_inverse_cdf_init_jit(gaussian_samples, gaussian_proxy['var'], gaussian_proxy['mean'])
+ samples = _gm1d_inverse_cdf.apply(
+ gm1d, scaled_cdfs, n_steps - backward_steps, eps, max_step_size, init_samples)
+
+ if backward_steps > 0: # fallback to pytorch implementation, which supports autograd but is slower
+ clamp_range = max_step_size * gm1d['logstds'].exp()
+ for i in range(n_steps):
+ cur_pdfs, cur_cdfs = gm1d_pdf_cdf(gm1d, samples) # (bs, *, n_samples, h, w)
+ delta = 0.5 * (cur_cdfs - scaled_cdfs) / cur_pdfs.clamp(min=eps) # (bs, *, n_samples, h, w)
+ delta = torch.maximum(torch.minimum(delta, clamp_range), -clamp_range)
+ samples = samples - delta # (bs, *, n_samples, h, w)
+ return samples
+
+
+@torch.jit.script
+def gm_to_iso_gaussian_jit(gm_weights, gm_means, gm_vars):
+ g_mean = (gm_weights * gm_means).sum(-4, keepdim=True) # (bs, *, 1, out_channels, h, w)
+ # (bs, *, num_gaussians, out_channels, h, w)
+ gm_diffs = gm_means - g_mean
+ g_var = (gm_weights * (gm_diffs * gm_diffs)).sum(-4, keepdim=True).mean(-3, keepdim=True) + gm_vars # (bs, *, 1, 1, h, w)
+ return g_mean, g_var, gm_diffs
+
+
+def gm_to_iso_gaussian(gm):
+ """
+ Args:
+ gm (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ (Optional) gm_vars (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ (Optional) gm_weights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ or
+ means (torch.Tensor): (bs, *, num_gaussians, h, w, out_channels)
+ covs (torch.Tensor): (bs, *, 1 or num_gaussians, h, w, out_channels, out_channels)
+ logweights (torch.Tensor): (bs, *, num_gaussians, h, w)
+
+ Returns:
+ tuple[dict, torch.Tensor, torch.Tensor]:
+ dict: Output gaussian
+ mean (torch.Tensor): (bs, *, out_channels, h, w)
+ var (torch.Tensor): (bs, *, 1, h, w)
+ torch.Tensor: (bs, *, num_gaussians, out_channels, h, w)
+ torch.Tensor: (bs, *, 1, 1, 1, 1) or (bs, *, 1, 1, h, w)
+ """
+ gm_means = gm['means']
+ gm_logweights = gm['logweights']
+
+ if 'covs' in gm:
+ batch_shapes = gm_means.shape[:-4]
+ num_gaussians, h, w, out_channels = gm_means.shape[-4:]
+ batch_numel = batch_shapes.numel()
+
+ gm_means = gm_means.reshape(
+ batch_numel, num_gaussians, h, w, out_channels).permute(0, 1, 4, 2, 3).reshape(
+ *batch_shapes, num_gaussians, out_channels, h, w)
+ gm_weights = gm_logweights.exp().unsqueeze(-3) # (bs, *, num_gaussians, 1, h, w)
+
+ gm_vars = gm['covs'].diagonal(
+ offset=0, dim1=-1, dim2=-2).mean(dim=-1).unsqueeze(-3) # (bs, *, 1 or num_gaussians, 1, h, w)
+ if gm_vars.size(-4) > 1:
+ gm_vars = (gm_weights * gm_vars).sum(dim=-4, keepdim=True) # (bs, *, 1, 1, h, w)
+
+ else:
+ if 'gm_weights' in gm:
+ gm_weights = gm['gm_weights']
+ else:
+ gm_weights = gm_logweights.exp() # (bs, *, num_gaussians, 1, h, w)
+ gm['gm_weights'] = gm_weights
+
+ gm_logstds = gm['logstds'] # (bs, *, 1, 1, 1, 1)
+ if 'gm_vars' in gm:
+ gm_vars = gm['gm_vars']
+ else:
+ gm_vars = (gm_logstds * 2).exp() # (bs, *, 1, 1, 1, 1)
+ gm['gm_vars'] = gm_vars
+
+ g_mean, g_var, gm_diffs = gm_to_iso_gaussian_jit(gm_weights, gm_means, gm_vars)
+ gaussian = dict(
+ mean=g_mean.squeeze(-4),
+ var=g_var.squeeze(-4))
+
+ return gaussian, gm_diffs, gm_vars
+
+
+@torch.jit.script
+def gm_to_gaussian_jit(
+ gm_means, gm_weights, gm_covs,
+ batch_numel: int, num_gaussians: int, out_channels: int, h: int, w: int):
+ g_mean = (gm_weights * gm_means).sum(-4, keepdim=True) # (bs, *, 1, out_channels, h, w)
+ # (batch_numel, num_gaussians, h, w, out_channels)
+ gm_diffs = (gm_means - g_mean).reshape(batch_numel, num_gaussians, out_channels, h, w).permute(0, 1, 3, 4, 2)
+ # (batch_numel, h, w, out_channels, out_channels) + (batch_numel, 1, 1, out_channels, out_channels)
+ g_cov = (gm_weights.reshape(
+ batch_numel, num_gaussians, h, w, 1, 1
+ ) * gm_diffs.unsqueeze(-1) * gm_diffs.unsqueeze(-2)).sum(1) + gm_covs
+ return g_mean, g_cov, gm_diffs
+
+
+def gm_to_gaussian(gm, cov_scale=1.0):
+ """
+ Args:
+ gm (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ (Optional) gm_vars (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ (Optional) gm_weights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ or
+ means (torch.Tensor): (bs, *, num_gaussians, h, w, out_channels)
+ covs (torch.Tensor): (bs, *, 1 or num_gaussians, h, w, out_channels, out_channels)
+ logweights (torch.Tensor): (bs, *, num_gaussians, h, w)
+ cov_scale (float): scale factor for the covariance matrix
+
+ Returns:
+ tuple[dict, torch.Tensor, torch.Tensor]:
+ dict: Output gaussian
+ mean (torch.Tensor): (bs, *, h, w, out_channels)
+ cov (torch.Tensor): (bs, *, h, w, out_channels, out_channels)
+ torch.Tensor: (bs, *, num_gaussians, h, w, out_channels)
+ torch.Tensor: (bs, *, out_channels, out_channels)
+ """
+ gm_means = gm['means']
+ gm_logweights = gm['logweights']
+
+ if 'covs' in gm:
+ batch_shapes = gm_means.shape[:-4]
+ num_gaussians, h, w, out_channels = gm_means.shape[-4:]
+ batch_numel = batch_shapes.numel()
+
+ gm_means = gm_means.reshape(
+ batch_numel, num_gaussians, h, w, out_channels).permute(0, 1, 4, 2, 3).reshape(
+ *batch_shapes, num_gaussians, out_channels, h, w)
+ gm_weights = gm_logweights.exp().unsqueeze(-3) # (bs, *, num_gaussians, 1, h, w)
+
+ gm_covs_num_gaussians = gm['covs'].size(-5)
+ if gm_covs_num_gaussians == 1:
+ gm_covs = gm['covs'].reshape(
+ batch_numel, h, w, out_channels, out_channels)
+ else:
+ gm_covs = gm['covs'].reshape(
+ batch_numel, gm['covs'].size(-5), h, w, out_channels, out_channels)
+ gm_covs = (gm_weights.reshape(batch_numel, num_gaussians, h, w, 1, 1) * gm_covs).sum(dim=1)
+
+ else:
+ if 'gm_weights' in gm:
+ gm_weights = gm['gm_weights']
+ else:
+ gm_weights = gm_logweights.exp() # (bs, *, num_gaussians, 1, h, w)
+ gm['gm_weights'] = gm_weights
+
+ gm_logstds = gm['logstds'] # (bs, *, 1, 1, 1, 1)
+ if 'gm_vars' in gm:
+ gm_vars = gm['gm_vars']
+ else:
+ gm_vars = (gm_logstds * 2).exp() # (bs, *, 1, 1, 1, 1)
+ gm['gm_vars'] = gm_vars
+
+ batch_shapes = gm_means.shape[:-4]
+ num_gaussians, out_channels, h, w = gm_means.shape[-4:]
+ batch_numel = batch_shapes.numel()
+ dtype = gm_means.dtype
+ device = gm_means.device
+
+ gm_covs = (torch.eye(out_channels, dtype=dtype, device=device) * gm_vars.reshape(
+ batch_numel, 1, 1, 1, 1)) # (batch_numel, 1, 1, out_channels, out_channels)
+
+ g_mean, g_cov, gm_diffs = gm_to_gaussian_jit(
+ gm_means, gm_weights, gm_covs, batch_numel, num_gaussians, out_channels, h, w)
+ gaussian = dict(
+ mean=g_mean.reshape(batch_numel, out_channels, h, w).permute(
+ 0, 2, 3, 1).reshape(*batch_shapes, h, w, out_channels),
+ cov=g_cov.reshape(*batch_shapes, h, w, out_channels, out_channels) * cov_scale)
+
+ return (gaussian,
+ gm_diffs.reshape(*batch_shapes, num_gaussians, h, w, out_channels),
+ gm_covs.reshape(*batch_shapes, 1, 1, 1, out_channels, out_channels))
+
+
+def gm_mul_gaussian(gm, gaussian, gm_power, gaussian_power, gm_diffs=None, gm_covs=None):
+ """
+ Args:
+ gm (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ (Optional) gm_vars (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ gaussian (dict):
+ mean (torch.Tensor): (bs, *, h, w, out_channels)
+ cov (torch.Tensor): (bs, *, h, w, out_channels, out_channels)
+ gm_power (float): power for the Gaussian mixture
+ gaussian_power (float): power for the Gaussian
+ gm_diffs (torch.Tensor | None): (bs, *, num_gaussians, h, w, out_channels)
+ gm_covs (torch.Tensor | None): (bs, *, 1, 1, 1, out_channels, out_channels)
+
+ Returns:
+ tuple[dict, float]:
+ dict: Output gausian mixture
+ means (torch.Tensor): (bs, *, num_gaussians, h, w, out_channels)
+ covs (torch.Tensor): (bs, *, num_gaussians, h, w, out_channels, out_channels)
+ logweights (torch.Tensor): (bs, *, num_gaussians, h, w)
+ float: Output gaussian mixture power
+ """
+ gm_means = gm['means'] # (bs, *, num_gaussians, out_channels, h, w)
+ gm_logstds = gm['logstds'] # (bs, *, 1, 1, 1, 1)
+ gm_logweights = gm['logweights'] # (bs, *, num_gaussians, 1, h, w)
+ if 'gm_vars' in gm:
+ gm_vars = gm['gm_vars']
+ else:
+ gm_vars = (gm_logstds * 2).exp() # (bs, *, 1, 1, 1, 1)
+ gm['gm_vars'] = gm_vars
+
+ batch_shapes = gm_means.shape[:-4]
+ num_gaussians, out_channels, h, w = gm_means.shape[-4:]
+ batch_numel = batch_shapes.numel()
+ dtype = gm_means.dtype
+ device = gm_means.device
+ eye = torch.eye(out_channels, dtype=dtype, device=device)
+
+ gm_means = gm_means.reshape(batch_numel, num_gaussians, out_channels, h, w).permute(
+ 0, 1, 3, 4, 2) # (batch_numel, num_gaussians, h, w, out_channels)
+ gm_vars = gm_vars.reshape(batch_numel, 1, 1, 1, 1)
+ g_mean = gaussian['mean'].reshape(batch_numel, h, w, out_channels)
+ g_cov = gaussian['cov'].reshape(batch_numel, h, w, out_channels, out_channels)
+
+ if gm_diffs is None:
+ gm_diffs = (gm_means - g_mean.unsqueeze(1)) # (batch_numel, num_gaussians, h, w, out_channels)
+ else:
+ gm_diffs = gm_diffs.reshape(batch_numel, num_gaussians, h, w, out_channels)
+ if gm_covs is None:
+ gm_covs = eye * gm_vars.unsqueeze(-1) # (batch_numel, 1, 1, 1, out_channels, out_channels)
+ else:
+ gm_covs = gm_covs.reshape(batch_numel, 1, 1, 1, out_channels, out_channels)
+
+ gm_weight = eye / gm_vars # (batch_numel, 1, 1, out_channels, out_channels)
+ g_weight = (gaussian_power / gm_power) * psd_inverse(g_cov) # (batch_numel, h, w, out_channels, out_channels)
+ out_covs = psd_inverse(gm_weight + g_weight) # (batch_numel, h, w, out_channels, out_channels)
+ out_means = (out_covs.unsqueeze(1) @ (
+ # (batch_numel, num_gaussians, h, w, out_channels, 1)
+ (gm_weight[..., :1, :1] * gm_means).unsqueeze(-1)
+ # (batch_numel, 1, h, w, out_channels, 1)
+ + (g_weight @ g_mean.unsqueeze(-1)).unsqueeze(1)
+ )).squeeze(-1) # (batch_numel, num_gaussians, h, w, out_channels)
+ # (batch_numel, num_gaussians, h, w, 1, out_channels) @ (batch_numel, 1, h, w, out_channels, out_channels)
+ # @ (batch_numel, num_gaussians, h, w, out_channels, 1) -> (batch_numel, num_gaussians, h, w, 1)
+ logweights_delta = (
+ gm_diffs.unsqueeze(-2)
+ @ psd_inverse(gm_covs * gaussian_power + g_cov.unsqueeze(1) * gm_power)
+ @ gm_diffs.unsqueeze(-1)
+ ).squeeze(-1) * (-0.5 * gaussian_power)
+
+ # (bs, *, num_gaussians, h, w)
+ out_logweights = gm_logweights.squeeze(-3) + logweights_delta.reshape(*batch_shapes, num_gaussians, h, w)
+ out_logweights = out_logweights.log_softmax(dim=-3)
+ out_means = out_means.reshape(*batch_shapes, num_gaussians, h, w, out_channels)
+ out_covs = out_covs.reshape(*batch_shapes, 1, h, w, out_channels, out_channels)
+
+ return dict(means=out_means, covs=out_covs, logweights=out_logweights), gm_power
+
+
+@torch.jit.script
+def gm_mul_iso_gaussian_jit(
+ gm_means, gm_vars, gm_logstds, gm_logweights, g_mean, g_var, g_logstd,
+ gm_power: float, gaussian_power: float, eps: float):
+ gm_diffs = gm_means - g_mean # (bs, *, num_gaussians, out_channels, h, w)
+
+ power_ratio = gaussian_power / gm_power
+ norm_factor = (g_var + power_ratio * gm_vars).clamp(min=eps)
+
+ out_means = (g_var * gm_means + power_ratio * gm_vars * g_mean) / norm_factor
+ # (bs, *, num_gaussians, 1, h, w)
+ logweights_delta = gm_diffs.square().sum(
+ dim=-3, keepdim=True) * (-0.5 * power_ratio / norm_factor)
+
+ out_logweights = (gm_logweights + logweights_delta).log_softmax(dim=-4)
+ out_logstds = gm_logstds + g_logstd - 0.5 * torch.log(norm_factor)
+
+ return out_means, out_logstds, out_logweights
+
+
+def gm_mul_iso_gaussian(gm, gaussian, gm_power, gaussian_power, eps=1e-6):
+ """
+ Args:
+ gm (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ (Optional) gm_vars (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ gaussian (dict):
+ mean (torch.Tensor): (bs, *, out_channels, h, w)
+ var (torch.Tensor): (bs, *, 1, 1, 1)
+ (Optional) logstd (torch.Tensor): (bs, *, 1, 1, 1)
+ gm_power (float): power for the Gaussian mixture
+ gaussian_power (float): power for the Gaussian
+
+ Returns:
+ tuple[dict, float]:
+ dict: Output gausian mixture
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ float: Output gaussian mixture power
+ """
+ gm_means = gm['means'] # (bs, *, num_gaussians, out_channels, h, w)
+ gm_logstds = gm['logstds'] # (bs, *, 1, 1, 1, 1)
+ gm_logweights = gm['logweights'] # (bs, *, num_gaussians, 1, h, w)
+ if 'gm_vars' in gm:
+ gm_vars = gm['gm_vars']
+ else:
+ gm_vars = (gm_logstds * 2).exp() # (bs, *, 1, 1, 1, 1)
+ gm['gm_vars'] = gm_vars
+ if 'logstd' not in gaussian:
+ gaussian['logstd'] = torch.log(gaussian['var']) / 2
+ g_mean = gaussian['mean'].unsqueeze(-4)
+ g_var = gaussian['var'].unsqueeze(-4)
+ g_logstd = gaussian['logstd'].unsqueeze(-4)
+
+ out_means, out_logstds, out_logweights = gm_mul_iso_gaussian_jit(
+ gm_means, gm_vars, gm_logstds, gm_logweights, g_mean, g_var, g_logstd, gm_power, gaussian_power, eps)
+ return dict(means=out_means, logstds=out_logstds, logweights=out_logweights), gm_power
+
+
+@torch.jit.script
+def gm_mul_gm_jit(gm1_means, gm1_logstds, gm1_vars, gm1_logweights, gm2_means, gm2_logstds, gm2_vars, gm2_logweights):
+ gm1_means = gm1_means.unsqueeze(-4) # (bs, *, num_gaussians_1, 1, out_channels, h, w)
+ gm1_vars = gm1_vars.unsqueeze(-4) # (bs, *, 1, 1, 1, 1, 1)
+ gm1_logweights = gm1_logweights.unsqueeze(-4) # (bs, *, num_gaussians_1, 1, 1, h, w)
+
+ gm2_means = gm2_means.unsqueeze(-5) # (bs, *, 1, num_gaussians_2, out_channels, h, w)
+ gm2_vars = gm2_vars.unsqueeze(-5) # (bs, *, 1, 1, 1, 1, 1)
+ gm2_logweights = gm2_logweights.unsqueeze(-5) # (bs, *, 1, num_gaussians_2, 1, h, w)
+
+ # (bs, *, num_gaussians_1, num_gaussians_2, out_channels, h, w)
+ gm_diffs = gm1_means - gm2_means
+ norm_factor = gm1_vars + gm2_vars # (bs, *, 1, 1, 1, 1, 1)
+
+ # (bs, *, num_gaussians_1, num_gaussians_2, out_channels, h, w)
+ out_means = (gm2_vars * gm1_means + gm1_vars * gm2_means) / norm_factor
+ out_means = out_means.flatten(-5, -4) # (bs, *, num_gaussians_1 * num_gaussians_2, out_channels, h, w)
+
+ # (bs, *, num_gaussians_1, num_gaussians_2, 1, h, w)
+ logweights_delta = gm_diffs.square().sum(dim=-3, keepdim=True) * (-0.5 / norm_factor)
+ # (bs, *, num_gaussians_1, num_gaussians_2, 1, h, w)
+ out_logweights = (gm1_logweights + gm2_logweights + logweights_delta)
+ # (bs, *, num_gaussians_1 * num_gaussians_2, 1, h, w)
+ out_logweights = out_logweights.flatten(-5, -4).log_softmax(dim=-4)
+
+ out_logstds = gm1_logstds + gm2_logstds - 0.5 * torch.logaddexp(gm1_logstds * 2, gm2_logstds * 2)
+
+ return out_means, out_logstds, out_logweights
+
+
+def gm_mul_gm(gm1, gm2):
+ """
+ Args:
+ gm1 (dict):
+ means (torch.Tensor): (bs, *, num_gaussians_1, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians_1, 1, h, w)
+ (Optional) gm_vars (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ gm2 (dict):
+ means (torch.Tensor): (bs, *, num_gaussians_2, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians_2, 1, h, w)
+ (Optional) gm_vars (torch.Tensor): (bs, *, 1, 1, 1, 1)
+
+ Returns:
+ dict: Output gausian mixture
+ means (torch.Tensor): (bs, *, num_gaussians_1 * num_gaussians_2, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians_1 * num_gaussians_2, 1, h, w)
+ """
+ gm1_means = gm1['means'] # (bs, *, num_gaussians_1, out_channels, h, w)
+ gm1_logstds = gm1['logstds'] # (bs, *, 1, 1, 1, 1)
+ gm1_logweights = gm1['logweights'] # (bs, *, num_gaussians_1, 1, h, w)
+ if 'gm_vars' in gm1:
+ gm1_vars = gm1['gm_vars']
+ else:
+ gm1_vars = (gm1_logstds * 2).exp() # (bs, *, 1, 1, 1, 1)
+ gm1['gm_vars'] = gm1_vars
+
+ gm2_means = gm2['means'] # (bs, *, num_gaussians_2, out_channels, h, w)
+ gm2_logstds = gm2['logstds'] # (bs, *, 1, 1, 1, 1)
+ gm2_logweights = gm2['logweights'] # (bs, *, num_gaussians_2, 1, h, w)
+ if 'gm_vars' in gm2:
+ gm2_vars = gm2['gm_vars']
+ else:
+ gm2_vars = (gm2_logstds * 2).exp() # (bs, *, 1, 1, 1, 1)
+ gm2['gm_vars'] = gm2_vars
+
+ out_means, out_logstds, out_logweights = gm_mul_gm_jit(
+ gm1_means, gm1_logstds, gm1_vars, gm1_logweights, gm2_means, gm2_logstds, gm2_vars, gm2_logweights)
+ return dict(means=out_means, logstds=out_logstds, logweights=out_logweights)
+
+
+@torch.jit.script
+def gm_to_mean_jit(gm_means, gm_logweights, gm_power: float):
+ mean = ((gm_logweights * gm_power).softmax(dim=-4) * gm_means).sum(dim=-4)
+ return mean
+
+
+def gm_to_mean(gm, gm_power=1):
+ """
+ Args:
+ gm (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ or
+ means (torch.Tensor): (bs, *, num_gaussians, h, w, out_channels)
+ covs (torch.Tensor): (bs, *, 1 or num_gaussians, h, w, out_channels, out_channels)
+ logweights (torch.Tensor): (bs, *, num_gaussians, h, w)
+ gm_power (float)
+
+ Returns:
+ torch.Tensor: (bs, *, out_channels, h, w)
+ """
+ gm_means = gm['means']
+ gm_logweights = gm['logweights']
+ if 'covs' in gm:
+ batch_shapes = gm_means.shape[:-4]
+ num_gaussians, h, w, out_channels = gm_means.shape[-4:]
+ batch_numel = batch_shapes.numel()
+ gm_means = gm_means.reshape(
+ batch_numel, num_gaussians, h, w, out_channels
+ ).permute(0, 1, 4, 2, 3).reshape(*batch_shapes, num_gaussians, out_channels, h, w)
+ gm_logweights = gm_logweights.unsqueeze(-3)
+ return gm_to_mean_jit(gm_means, gm_logweights, gm_power)
+
+
+def gm_to_sample(
+ gm,
+ gm_power=1,
+ n_samples=1,
+ cov_sharpen=False):
+ """
+ Args:
+ gm (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1) or (bs, *, num_gaussians, 1, h, w)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ or
+ means (torch.Tensor): (bs, *, num_gaussians, h, w, out_channels)
+ covs (torch.Tensor): (bs, *, 1 or num_gaussians, h, w, out_channels, out_channels)
+ logweights (torch.Tensor): (bs, *, num_gaussians, h, w)
+ gm_power (float): power for the Gaussian mixture, samples are approximated if power is not 1
+ n_samples (int): number of samples
+ cov_sharpen (bool): whether to sharpen the covariance matrix when power is greater than 1
+
+ Returns:
+ torch.Tensor: (bs, *, n_samples, out_channels, h, w)
+ """
+ gm_means = gm['means']
+ gm_logweights = gm['logweights']
+
+ if 'covs' in gm:
+ gm_covs = gm['covs']
+
+ batch_shapes = gm_means.shape[:-4]
+ num_gaussians, h, w, out_channels = gm_means.shape[-4:]
+ batch_numel = batch_shapes.numel()
+
+ inds = torch.multinomial(
+ (gm_logweights.reshape(batch_numel, num_gaussians, h, w).permute(0, 2, 3, 1).reshape(
+ batch_numel * h * w, num_gaussians) * gm_power).softmax(dim=-1),
+ n_samples, replacement=True
+ ).reshape(batch_numel, h, w, n_samples).permute(0, 3, 1, 2).reshape(*batch_shapes, n_samples, 1, h, w)
+
+ means = gm_means.gather( # (bs, *, n_samples, h, w, out_channels)
+ dim=-4,
+ index=inds.reshape(*batch_shapes, n_samples, h, w, 1).expand(
+ *batch_shapes, n_samples, h, w, out_channels))
+ if gm_covs.size(-5) == 1:
+ tril = torch.linalg.cholesky(gm_covs) # (bs, *, 1, h, w, out_channels, out_channels)
+ elif n_samples < num_gaussians:
+ covs = gm_covs.gather( # (bs, *, n_samples, h, w, out_channels, out_channels)
+ dim=-5,
+ index=inds.reshape(*batch_shapes, n_samples, h, w, 1, 1).expand(
+ *batch_shapes, n_samples, h, w, out_channels, out_channels))
+ if cov_sharpen:
+ covs = covs / gm_power
+ tril = torch.linalg.cholesky(covs)
+ else:
+ tril = torch.linalg.cholesky(gm_covs)
+ if cov_sharpen:
+ tril = tril / math.sqrt(gm_power)
+ tril = tril.gather( # (bs, *, n_samples, h, w, out_channels, out_channels)
+ dim=-5,
+ index=inds.reshape(*batch_shapes, n_samples, h, w, 1, 1).expand(
+ *batch_shapes, n_samples, h, w, out_channels, out_channels))
+
+ # (bs, *, n_samples, h, w, out_channels)
+ samples = (tril @ torch.randn(
+ (*batch_shapes, n_samples, h, w, out_channels, 1),
+ dtype=means.dtype, device=means.device)).squeeze(-1) + means
+ samples = samples.reshape(batch_numel, n_samples, h, w, out_channels).permute(0, 1, 4, 2, 3).reshape(
+ *batch_shapes, n_samples, out_channels, h, w)
+
+ else:
+ batch_shapes = gm_means.shape[:-4]
+ num_gaussians, out_channels, h, w = gm_means.shape[-4:]
+ batch_numel = batch_shapes.numel()
+
+ inds = torch.multinomial(
+ (gm_logweights.reshape(batch_numel, num_gaussians, h, w).permute(0, 2, 3, 1).reshape(
+ batch_numel * h * w, num_gaussians) * gm_power).softmax(dim=-1),
+ n_samples, replacement=True
+ ).reshape(batch_numel, h, w, n_samples).permute(0, 3, 1, 2).reshape(*batch_shapes, n_samples, 1, h, w)
+
+ means = gm_means.gather( # (bs, *, n_samples, out_channels, h, w)
+ dim=-4,
+ index=inds.expand(*batch_shapes, n_samples, out_channels, h, w))
+ stds = gm['logstds'].exp() # (bs, *, 1, 1, 1, 1) or (bs, *, num_gaussians, 1, h, w)
+ if cov_sharpen:
+ stds = stds / math.sqrt(gm_power)
+ if stds.size(-4) == num_gaussians and num_gaussians > 1:
+ stds = stds.gather(dim=-4, index=inds) # (bs, *, n_samples, 1, h, w)
+
+ # (bs, *, n_samples, out_channels, h, w)
+ samples = stds * torch.randn(
+ (*batch_shapes, n_samples, out_channels, h, w),
+ dtype=means.dtype, device=means.device) + means
+
+ return samples
+
+
+def gaussian_mul_gaussian(gaussian_1, gaussian_2, gaussian_1_power, gaussian_2_power):
+ """
+ Args:
+ gaussian_1 (dict):
+ mean (torch.Tensor): (bs, *, h, w, out_channels)
+ cov (torch.Tensor): (bs, *, h, w, out_channels, out_channels)
+ gaussian_2 (dict):
+ mean (torch.Tensor): (bs, *, h, w, out_channels)
+ cov (torch.Tensor): (bs, *, h, w, out_channels, out_channels)
+
+ Returns:
+ dict:
+ mean (torch.Tensor): (bs, *, h, w, out_channels)
+ cov (torch.Tensor): (bs, *, h, w, out_channels, out_channels)
+ """
+ g1_mean = gaussian_1['mean']
+ g1_cov = gaussian_1['cov']
+ g2_mean = gaussian_2['mean']
+ g2_cov = gaussian_2['cov']
+
+ g1_cov_inv = gaussian_1_power * psd_inverse(g1_cov)
+ g2_cov_inv = gaussian_2_power * psd_inverse(g2_cov)
+ out_cov = psd_inverse(g1_cov_inv + g2_cov_inv)
+ out_mean = out_cov @ (g1_cov_inv @ g1_mean.unsqueeze(-1) + g2_cov_inv @ g2_mean.unsqueeze(-1))
+ out_mean = out_mean.squeeze(-1)
+ return dict(mean=out_mean, cov=out_cov)
+
+
+@torch.jit.script
+def iso_gaussian_mul_iso_gaussian_jit(
+ g1_mean, g1_var, g2_mean, g2_var,
+ gaussian_1_power: float, gaussian_2_power: float, eps: float):
+ norm_factor = (gaussian_1_power * g2_var + gaussian_2_power * g1_var).clamp(min=eps)
+ out_var = g2_var * g1_var / norm_factor
+ out_mean = (gaussian_1_power * g2_var * g1_mean + gaussian_2_power * g1_var * g2_mean) / norm_factor
+ return out_mean, out_var
+
+
+def iso_gaussian_mul_iso_gaussian(gaussian_1, gaussian_2, gaussian_1_power, gaussian_2_power, eps=1e-6):
+ """
+ Args:
+ gaussian_1 (dict):
+ mean (torch.Tensor): (bs, *, out_channels, h, w)
+ var (torch.Tensor): (bs, *, 1, h, w)
+ gaussian_2 (dict):
+ mean (torch.Tensor): (bs, *, out_channels, h, w)
+ var (torch.Tensor): (bs, *, 1, h, w)
+
+ Returns:
+ dict:
+ mean (torch.Tensor): (bs, *, out_channels, h, w)
+ var (torch.Tensor): (bs, *, 1, h, w)
+ """
+ out_mean, out_var = iso_gaussian_mul_iso_gaussian_jit(
+ gaussian_1['mean'],
+ gaussian_1['var'],
+ gaussian_2['mean'],
+ gaussian_2['var'],
+ gaussian_1_power, gaussian_2_power, eps)
+ return dict(mean=out_mean, var=out_var)
+
+
+def iso_gaussian_logprob(iso_gaussian, samples):
+ """
+ Args:
+ iso_gaussian (dict):
+ mean (torch.Tensor): (bs, *, out_channels, h, w)
+ var (torch.Tensor): (bs, *, 1, h, w)
+ samples (torch.Tensor): (bs, *, n_samples, out_channels, h, w)
+
+ Returns:
+ torch.Tensor: (bs, *, n_samples, h, w)
+ """
+ mean = iso_gaussian['mean'].unsqueeze(-4) # (bs, *, 1, out_channels, h, w)
+ var = iso_gaussian['var'] # (bs, *, 1, h, w)
+ out_channels = mean.size(-3)
+ const = -0.5 * out_channels * math.log(2 * math.pi)
+ logprob = -0.5 * (samples - mean).square().sum(dim=-3) / var - 0.5 * out_channels * var.log() + const
+ return logprob
+
+
+@torch.jit.script
+def gm_iso_gaussian_logprobs_jit(
+ means, logstds, samples, out_channels: int, const: float):
+ logstds = logstds.unsqueeze(-5)
+ inverse_std = torch.exp(-logstds)
+ # (bs, *, n_samples, num_gaussians, out_channels, h, w)
+ diff_weighted = (samples.unsqueeze(-4) - means.unsqueeze(-5)) * inverse_std
+ # (bs, *, n_samples, num_gaussians, h, w)
+ gaussian_logprobs = -0.5 * (diff_weighted.square().sum(dim=-3)) - out_channels * logstds.squeeze(-3) + const
+ return gaussian_logprobs
+
+
+def gm_logprob(gm, samples):
+ """
+ Args:
+ gm (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1) or (bs, *, num_gaussians, 1, h, w)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ or
+ means (torch.Tensor): (bs, *, num_gaussians, h, w, out_channels)
+ covs (torch.Tensor): (bs, *, 1 or num_gaussians, h, w, out_channels, out_channels)
+ logweights (torch.Tensor): (bs, *, num_gaussians, h, w)
+ (Optional) invcov_trils (torch.Tensor): (bs, *, 1 or num_gaussians, h, w, out_channels, out_channels)
+ lower triangular Cholesky decomposition of the inverse covariance matrix
+ (Optional) logdets (torch.Tensor): (bs, *, 1 or num_gaussians, h, w)
+ log-determinant of the covariance matrix
+ samples (torch.Tensor): (bs, *, n_samples, out_channels, h, w)
+
+ Returns:
+ tuple[torch.Tensor, torch.Tensor]:
+ torch.Tensor: (bs, *, n_samples, h, w)
+ torch.Tensor: (bs, *, n_samples, num_gaussians, h, w)
+ """
+ n_samples = samples.size(-4)
+
+ if 'covs' in gm:
+ means = gm['means']
+ batch_shapes = means.shape[:-4]
+ num_gaussians, h, w, out_channels = means.shape[-4:]
+ batch_numel = batch_shapes.numel()
+
+ const = -0.5 * out_channels * math.log(2 * math.pi)
+
+ covs = gm['covs']
+ if 'invcov_trils' in gm:
+ invcov_trils = gm['invcov_trils']
+ else:
+ invcov_trils = torch.linalg.cholesky(psd_inverse(covs))
+ gm['invcov_trils'] = invcov_trils
+ if 'logdets' in gm:
+ logdets = gm['logdets']
+ else:
+ logdets = torch.logdet(covs)
+ gm['logdets'] = logdets
+
+ samples = samples.reshape(
+ batch_numel, n_samples, out_channels, h, w
+ ).permute(0, 1, 3, 4, 2).reshape(batch_numel, n_samples, 1, h, w, out_channels)
+ # (batch_numel, n_samples, num_gaussians, h, w, out_channels)
+ diffs = samples - means.reshape(batch_numel, 1, num_gaussians, h, w, out_channels)
+ diff_weighted = (diffs.unsqueeze(-2) @ invcov_trils.reshape(
+ batch_numel, 1, invcov_trils.size(-5), h, w, out_channels, out_channels)).squeeze(-2)
+ # (batch_numel, n_samples, num_gaussians, h, w)
+ gaussian_logprobs = (-0.5 * (diff_weighted.square().sum(dim=-1) + logdets.unsqueeze(-4)) + const).reshape(
+ *batch_shapes, n_samples, num_gaussians, h, w)
+
+ else:
+ batch_shapes = gm['means'].shape[:-4]
+ num_gaussians, out_channels, h, w = gm['means'].shape[-4:]
+ const = -0.5 * out_channels * math.log(2 * math.pi)
+ gaussian_logprobs = gm_iso_gaussian_logprobs_jit(
+ gm['means'], gm['logstds'], samples, out_channels, const)
+
+ # (bs, *, n_samples, h, w)
+ logprob = (gm['logweights'].reshape(
+ *batch_shapes, 1, num_gaussians, h, w) + gaussian_logprobs).logsumexp(dim=-3)
+
+ return logprob, gaussian_logprobs
+
+
+def gm_spectral_logprobs(
+ gm, samples, power_spectrum=None, spectral_samples=None, n_axes=None, eps=1e-6, axis_aligned=True):
+ """
+ Args:
+ gm (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1) or (bs, *, num_gaussians, 1, h, w)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ or
+ means (torch.Tensor): (bs, *, num_gaussians, h, w, out_channels)
+ covs (torch.Tensor): (bs, *, 1 or num_gaussians, h, w, out_channels, out_channels)
+ logweights (torch.Tensor): (bs, *, num_gaussians, h, w)
+ (Optional) invcov_trils (torch.Tensor): (bs, *, 1 or num_gaussians, h, w, out_channels, out_channels)
+ lower triangular Cholesky decomposition of the inverse covariance matrix
+ (Optional) logdets (torch.Tensor): (bs, *, 1 or num_gaussians, h, w)
+ log-determinant of the covariance matrix
+ samples (torch.Tensor): (bs, *, n_samples, out_channels, h, w)
+ power_spectrum (torch.Tensor | None): (bs, *, out_channels, h, w)
+
+ Returns:
+ torch.Tensor: (bs, *, n_samples)
+ """
+ logprobs = gm_logprob(gm, samples)[0].sum(dim=(-2, -1))
+ if power_spectrum is not None:
+ if spectral_samples is None:
+ z_kr = gm_samples_to_gaussian_samples(
+ gm, samples, n_axes=n_axes, eps=eps, axis_aligned=axis_aligned)
+ z_kr_fft = torch.fft.fft2(z_kr, norm='ortho')
+ spectral_samples = z_kr_fft.real + z_kr_fft.imag
+ out_channels = spectral_samples.size(-3)
+ spectral_logprob_diff = -0.5 * spectral_samples.square().sum(
+ dim=-3) * (torch.exp(-power_spectrum) - 1) - 0.5 * out_channels * power_spectrum
+ logprobs = logprobs + spectral_logprob_diff.sum(dim=(-2, -1))
+ return logprobs
+
+
+def gm_kl_div(gm_p, gm_q, n_samples=32, use_kr=False, kr_backward_steps=1):
+ """
+ Args:
+ gm_p (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ (Optional) gm_weights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ gm_q (dict):
+ Same as gm_p
+ n_samples (int): number of samples
+
+ Returns:
+ torch.Tensor: (bs, *, 1, h, w)
+ """
+ if use_kr:
+ sample_size = list(gm_p['means'].shape)
+ sample_size[-4] = n_samples
+ gaussian_samples = torch.randn(sample_size, device=gm_p['means'].device, dtype=gm_p['means'].dtype)
+ samples = gaussian_samples_to_gm_samples(
+ gm_p, gaussian_samples, axis_aligned=True, backward_steps=kr_backward_steps)
+ else:
+ samples = gm_to_sample(gm_p, 1.0, n_samples=n_samples)
+ return (gm_logprob(gm_p, samples)[0] - gm_logprob(gm_q, samples)[0]).mean(dim=-3, keepdim=True)
+
+
+def gm_entropy(gm, n_samples=32):
+ samples = gm_to_sample(gm, 1.0, n_samples=n_samples)
+ return -gm_logprob(gm, samples)[0].mean(dim=-3, keepdim=True)
+
+
+def gm_samples_to_gaussian_samples(gm, gm_samples, n_axes=None, eps=1e-6, axis_aligned=True):
+ """
+ Knothe-Rosenblatt transport for transforming samples from a Gaussian mixture to samples from a standard Gaussian.
+
+ Args:
+ gm (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ (Optional) gm_weights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ gm_samples (torch.Tensor): (bs, *, n_samples, out_channels, h, w)
+
+ Returns:
+ torch.Tensor: (bs, *, n_samples, out_channels, h, w)
+ """
+ assert 'covs' not in gm
+ dtype = gm_samples.dtype
+ gm_means = gm['means']
+ batch_shapes = gm_means.shape[:-4]
+ num_gaussians, out_channels, h, w = gm_means.shape[-4:]
+ batch_numel = batch_shapes.numel()
+ n_samples = gm_samples.size(-4)
+ if n_axes is None:
+ n_axes = out_channels
+
+ with torch.no_grad():
+ gaussians = gm_to_gaussian(gm)[0]
+ covs = gaussians['cov']
+ if axis_aligned:
+ covs = covs.mean(dim=(-4, -3), keepdim=True)
+ # (bs, *, h, w, out_channels, out_channels) or (bs, *, 1, 1, out_channels, out_channels)
+ _, eigvecs = torch.linalg.eigh(covs.float())
+ # eigenvalues in descending order
+ eigvecs = eigvecs.flip(-1).to(dtype=dtype)
+
+ eigvecs = eigvecs.reshape(
+ (batch_numel, 1, 1, out_channels, out_channels) if axis_aligned
+ else (batch_numel, h, w, out_channels, out_channels))
+ gm_means = gm_means.reshape(
+ batch_numel, num_gaussians, out_channels, h, w
+ ).permute(0, 3, 4, 1, 2) # (batch_numel, h, w, num_gaussians, out_channels)
+ gm_logstds = gm['logstds'].reshape(batch_numel, 1, 1, 1, 1)
+ gm_logweights = gm['logweights'].reshape(
+ batch_numel, num_gaussians, 1, h, w
+ ).permute(0, 3, 4, 1, 2) # (batch_numel, h, w, num_gaussians, 1)
+ gm_stds = gm_logstds.exp() # (batch_numel, 1, 1, 1, 1)
+ gm_samples = gm_samples.reshape(
+ batch_numel, n_samples, out_channels, h, w
+ ).permute(0, 3, 4, 1, 2) # (batch_numel, h, w, n_samples, out_channels)
+
+ eigvecs_ = eigvecs[..., :n_axes]
+ gm_means_rot = gm_means @ eigvecs_ # (batch_numel, h, w, num_gaussians, n_axes)
+ gm_samples_rot = gm_samples @ eigvecs_ # (batch_numel, h, w, n_samples, n_axes)
+
+ # (batch_numel, h, w, n_samples, num_gaussians, n_axes)
+ gm_norm_diffs = (gm_samples_rot.unsqueeze(-2)
+ - gm_means_rot.unsqueeze(-3)
+ ) / gm_stds.unsqueeze(-3)
+ gm_norm_diffs_sq = gm_norm_diffs.square()
+
+ # (batch_numel, h, w, n_samples, num_gaussians, n_axes - 1)
+ gm_norm_diffs_sq_cumsumprev = torch.cumsum(gm_norm_diffs_sq[..., :-1], dim=-1)
+ gm_slice_logweights = gm_logweights.unsqueeze(-3) - 0.5 * gm_norm_diffs_sq_cumsumprev
+ gm_slice_weights = gm_slice_logweights.softmax(dim=-2)
+
+ if 'gm_weights' in gm:
+ gm_weights = gm['gm_weights'].reshape(
+ batch_numel, num_gaussians, 1, h, w
+ ).permute(0, 3, 4, 1, 2) # (batch_numel, h, w, num_gaussians, 1)
+ else:
+ gm_weights = gm_logweights.exp()
+ gm_slice_weights = torch.cat(
+ [gm_weights.unsqueeze(-3).expand(batch_numel, h, w, n_samples, num_gaussians, 1),
+ gm_slice_weights],
+ dim=-1) # (batch_numel, h, w, n_samples, num_gaussians, n_axes)
+
+ sqrt_2 = math.sqrt(2)
+ # (batch_numel, h, w, n_samples, n_axes)
+ gaussian_samples_rot = (torch.erfinv(
+ (gm_slice_weights * torch.erf(gm_norm_diffs / sqrt_2)).sum(dim=-2).float().clamp(min=-1 + eps, max=1 - eps)
+ ) * sqrt_2).to(dtype=dtype)
+ if n_axes < out_channels:
+ gaussian_samples_rot = torch.cat(
+ [gaussian_samples_rot,
+ torch.randn(
+ (batch_numel, h, w, n_samples, out_channels - n_axes),
+ dtype=gm_samples.dtype, device=gm_samples.device)],
+ dim=-1) # (batch_numel, h, w, n_samples, out_channels)
+
+ if axis_aligned:
+ gaussian_samples = gaussian_samples_rot
+ else:
+ gaussian_samples = gaussian_samples_rot @ eigvecs.transpose(-1, -2)
+
+ return gaussian_samples.permute(0, 3, 4, 1, 2).reshape(*batch_shapes, n_samples, out_channels, h, w)
+
+
+@torch.jit.script
+def gm_kr_slice_jit(last_samples, gm_means_1d, gm_logweights, gm_stds, gm_norm_diffs_sq_cumsumprev):
+ gm_norm_diffs_prev = (last_samples # (batch_numel, n_samples, 1, h, w)
+ - gm_means_1d # (batch_numel, 1, num_gaussians, h, w)
+ ) / gm_stds # (batch_numel, n_samples, num_gaussians, h, w)
+ gm_norm_diffs_sq_cumsumprev = gm_norm_diffs_sq_cumsumprev + gm_norm_diffs_prev * gm_norm_diffs_prev
+ # (batch_numel, n_samples, num_gaussians, h, w)
+ gm_slice_logweights = (gm_logweights - 0.5 * gm_norm_diffs_sq_cumsumprev).log_softmax(dim=-3)
+ return gm_slice_logweights, gm_norm_diffs_sq_cumsumprev
+
+
+def gaussian_samples_to_gm_samples(
+ gm, gaussian_samples, n_axes=None, n_steps=16, backward_steps=0, eps=1e-6,
+ force_fp32=True, axis_aligned=True):
+ """
+ Knothe-Rosenblatt transport for transforming samples from a standard Gaussian to samples from a Gaussian mixture.
+
+ Args:
+ gm (dict):
+ means (torch.Tensor): (bs, *, num_gaussians, out_channels, h, w)
+ logstds (torch.Tensor): (bs, *, 1, 1, 1, 1)
+ logweights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ (Optional) gm_weights (torch.Tensor): (bs, *, num_gaussians, 1, h, w)
+ gaussian_samples (torch.Tensor): (bs, *, n_samples, out_channels, h, w)
+
+ Returns:
+ torch.Tensor: (bs, *, n_samples, out_channels, h, w)
+ """
+ assert 'covs' not in gm
+ ori_dtype = gaussian_samples.dtype
+ if force_fp32:
+ gm = {k: v.float() for k, v in gm.items()}
+ gaussian_samples = gaussian_samples.float()
+
+ dtype = gaussian_samples.dtype
+ gm_means = gm['means']
+ batch_shapes = gm_means.shape[:-4]
+ num_gaussians, out_channels, h, w = gm_means.shape[-4:]
+ batch_numel = batch_shapes.numel()
+ n_samples = gaussian_samples.size(-4)
+ if n_axes is None:
+ n_axes = out_channels
+
+ with torch.no_grad():
+ gaussians = gm_to_gaussian(gm)[0]
+ covs = gaussians['cov']
+ if axis_aligned:
+ covs = covs.mean(dim=(-4, -3), keepdim=True)
+ # (bs, *, h, w, out_channels, out_channels) or (bs, *, 1, 1, out_channels, out_channels)
+ _, eigvecs = torch.linalg.eigh(covs.float())
+ # eigenvalues in descending order
+ eigvecs = eigvecs.flip(-1).to(dtype=dtype)
+
+ eigvecs = eigvecs.reshape(
+ (batch_numel, 1, 1, out_channels, out_channels) if axis_aligned
+ else (batch_numel, h, w, out_channels, out_channels))
+ gm_means = gm_means.reshape(
+ batch_numel, num_gaussians, out_channels, h, w
+ ).permute(0, 3, 4, 1, 2) # (batch_numel, h, w, num_gaussians, out_channels)
+ gaussian_samples = gaussian_samples.reshape(
+ batch_numel, n_samples, out_channels, h, w
+ ).permute(0, 3, 4, 1, 2) # (batch_numel, h, w, n_samples, out_channels)
+
+ eigvecs_ = eigvecs[..., :n_axes]
+ gm_means_rot = gm_means @ eigvecs # (batch_numel, h, w, num_gaussians, out_channels)
+ if axis_aligned:
+ gaussian_samples_rot = gaussian_samples
+ else:
+ gaussian_samples_rot = gaussian_samples @ eigvecs_ # (batch_numel, h, w, n_samples, n_axes)
+
+ gm_logstds = gm['logstds'].reshape(batch_numel, 1, 1, 1, 1)
+ gm_logweights = gm['logweights'].reshape(batch_numel, 1, num_gaussians, h, w)
+ gm_stds = gm_logstds.exp() # (batch_numel, 1, 1, 1, 1)
+ gm_means_rot = gm_means_rot.permute(0, 4, 3, 1, 2) # (batch_numel, n_axes, num_gaussians, h, w)
+ gaussian_samples_rot = gaussian_samples_rot.permute(0, 3, 4, 1, 2) # (batch_numel, n_samples, n_axes, h, w)
+
+ uniform_samples = torch.erf(gaussian_samples_rot / math.sqrt(2)) # in range [-1, 1]
+
+ gm_samples_rot_list = []
+ gm_norm_diffs_sq_cumsumprev = gaussian_samples.new_tensor([0.0])
+ gm_means_1d = gm_means_rot[:, :1] # (batch_numel, 1, num_gaussians, h, w) for axis_id = 0
+ for axis_id in range(n_axes):
+ if axis_id == 0:
+ gm1d = dict(
+ means=gm_means_1d.squeeze(-4), # (batch_numel, num_gaussians, h, w)
+ logstds=gm_logstds.squeeze(-4), # (batch_numel, 1, 1, 1)
+ logweights=gm_logweights.squeeze(-4)) # (batch_numel, num_gaussians, h, w)
+ if 'gm_weights' in gm:
+ gm1d.update(gm_weights=gm['gm_weights'].reshape(batch_numel, num_gaussians, h, w))
+ gm_samples_rot_list.append(
+ gm1d_inverse_cdf(
+ gm1d,
+ uniform_samples[:, :, axis_id], # (batch_numel, n_samples, h, w)
+ n_steps=n_steps,
+ eps=eps,
+ max_step_size=1.5,
+ gaussian_samples=gaussian_samples_rot[:, :, axis_id], # (batch_numel, n_samples, h, w)
+ backward_steps=backward_steps
+ ).unsqueeze(-3) # (batch_numel, n_samples, 1, h, w)
+ )
+ else:
+ gm_slice_logweights, gm_norm_diffs_sq_cumsumprev = gm_kr_slice_jit(
+ gm_samples_rot_list[-1], gm_means_1d, gm_logweights, gm_stds, gm_norm_diffs_sq_cumsumprev)
+ gm_means_1d = gm_means_rot[:, axis_id:axis_id + 1] # (batch_numel, 1, num_gaussians, h, w)
+ gm1d = dict(
+ means=gm_means_1d, # (batch_numel, 1, num_gaussians, h, w)
+ logstds=gm_logstds, # (batch_numel, 1, 1, 1, 1)
+ logweights=gm_slice_logweights) # (batch_numel, n_samples, num_gaussians, h, w)
+ gm_samples_rot_list.append(
+ gm1d_inverse_cdf(
+ gm1d,
+ uniform_samples[:, :, axis_id:axis_id + 1], # (batch_numel, n_samples, 1, h, w)
+ n_steps=n_steps,
+ eps=eps,
+ max_step_size=1.5,
+ gaussian_samples=gaussian_samples_rot[:, :, axis_id:axis_id + 1], # (batch_numel, n_samples, 1, h, w)
+ backward_steps=backward_steps
+ ) # (batch_numel, n_samples, 1, h, w)
+ )
+
+ gm_samples_rot = torch.cat(gm_samples_rot_list, dim=-3) # (batch_numel, n_samples, n_axes, h, w)
+
+ if n_axes < out_channels:
+ gm_slice_logweights, gm_norm_diffs_sq_cumsumprev = gm_kr_slice_jit(
+ gm_samples_rot_list[-1], gm_means_1d, gm_logweights, gm_stds, gm_norm_diffs_sq_cumsumprev)
+ gm_means = gm_means_rot[:, n_axes:] # (batch_numel, num_channels - n_axes, num_gaussians, h, w)
+ gm_slice = dict(
+ means=gm_means.transpose(-3, -4).unsqueeze(-5).expand(
+ batch_numel, n_samples, num_gaussians, out_channels - n_axes, h, w),
+ logstds=gm_logstds.unsqueeze(-3), # (batch_numel, 1, 1, 1, 1, 1)
+ logweights=gm_slice_logweights.unsqueeze(-3)) # (batch_numel, n_samples, num_gaussians, 1, h, w)
+ gm_samples_rot = torch.cat(
+ [gm_samples_rot,
+ gm_to_sample(gm_slice, 1).squeeze(-4)],
+ dim=-3) # (batch_numel, n_samples, out_channels, h, w)
+
+ gm_samples = gm_samples_rot.permute(0, 3, 4, 1, 2) @ eigvecs.transpose(-1, -2)
+
+ return gm_samples.permute(0, 3, 4, 1, 2).reshape(*batch_shapes, n_samples, out_channels, h, w).to(ori_dtype)
+
+
+def gm_transpose_t_first(gm):
+ # (bs, num_gaussians, out_channels, t, h, w) -> (bs, t, num_gaussians, out_channels, h, w)
+ return dict(
+ means=gm['means'].permute(0, 3, 1, 2, 4, 5),
+ logweights=gm['logweights'].permute(0, 3, 1, 2, 4, 5),
+ logstds=gm['logstds'].permute(0, 3, 1, 2, 4, 5)
+ )
+
+
+def gm_temperature(gm, temperature, gm_dim=-4, eps=1e-6):
+ gm = gm.copy()
+ temperature = max(temperature, eps)
+ gm['logweights'] = (gm['logweights'] / temperature).log_softmax(dim=gm_dim)
+ if 'logstds' in gm:
+ gm['logstds'] = gm['logstds'] + (0.5 * math.log(temperature))
+ if 'gm_vars' in gm:
+ gm['gm_vars'] = gm['gm_vars'] * temperature
+ return gm
diff --git a/lakonlab/ops/gmflow_ops/setup.py b/lakonlab/ops/gmflow_ops/setup.py
new file mode 100644
index 0000000000000000000000000000000000000000..8ab5acd698b332c5e7610b8fe576d5fa2d7dd057
--- /dev/null
+++ b/lakonlab/ops/gmflow_ops/setup.py
@@ -0,0 +1,65 @@
+import os
+from setuptools import setup
+from torch.utils.cpp_extension import BuildExtension, CUDAExtension
+
+_src_path = os.path.dirname(os.path.abspath(__file__))
+
+nvcc_flags = [
+ '-O3', '-std=c++17',
+ '-U__CUDA_NO_HALF_OPERATORS__', '-U__CUDA_NO_HALF_CONVERSIONS__', '-U__CUDA_NO_HALF2_OPERATORS__',
+]
+
+if os.name == "posix":
+ c_flags = ['-O3', '-std=c++17']
+elif os.name == "nt":
+ c_flags = ['/O2', '/std:c++17']
+
+ # find cl.exe
+ def find_cl_path():
+ import glob
+ for program_files in [r"C:\\Program Files (x86)", r"C:\\Program Files"]:
+ for edition in ["Enterprise", "Professional", "BuildTools", "Community"]:
+ paths = sorted(glob.glob(r"%s\\Microsoft Visual Studio\\*\\%s\\VC\\Tools\\MSVC\\*\\bin\\Hostx64\\x64" % (program_files, edition)), reverse=True)
+ if paths:
+ return paths[0]
+
+ # If cl.exe is not on path, try to find it.
+ if os.system("where cl.exe >nul 2>nul") != 0:
+ cl_path = find_cl_path()
+ if cl_path is None:
+ raise RuntimeError("Could not locate a supported Microsoft Visual C++ installation")
+ os.environ["PATH"] += ";" + cl_path
+
+'''
+Usage:
+
+python setup.py build_ext --inplace # build extensions locally, do not install (only can be used from the parent directory)
+
+python setup.py install # build extensions and install (copy) to PATH.
+pip install . # ditto but better (e.g., dependency & metadata handling)
+
+python setup.py develop # build extensions and install (symbolic) to PATH.
+pip install -e . # ditto but better (e.g., dependency & metadata handling)
+
+'''
+setup(
+ name='gmflow_ops', # package name, import this to use python API
+ author='Hansheng Chen',
+ author_email='hanshengchen97@gmail.com',
+ ext_modules=[
+ CUDAExtension(
+ name='_gmflow_ops', # extension name, import this to use CUDA API
+ sources=[os.path.join(_src_path, 'src', f) for f in [
+ 'gmflow_ops.cu',
+ 'bindings.cpp',
+ ]],
+ extra_compile_args={
+ 'cxx': c_flags,
+ 'nvcc': nvcc_flags,
+ }
+ ),
+ ],
+ cmdclass={
+ 'build_ext': BuildExtension,
+ }
+)
diff --git a/lakonlab/parallel/__init__.py b/lakonlab/parallel/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..8c4062ceae95f16d27684f9932317fe43ee4635e
--- /dev/null
+++ b/lakonlab/parallel/__init__.py
@@ -0,0 +1,9 @@
+from .distributed import MMDistributedDataParallel
+from .ddp_wrapper import DistributedDataParallelWrapper
+from .fsdp_wrapper import FSDPWrapper
+from .fsdp2_wrapper import FSDP2Wrapper
+from .utils import apply_module_wrapper
+
+__all__ = [
+ 'MMDistributedDataParallel', 'DistributedDataParallelWrapper', 'FSDPWrapper', 'FSDP2Wrapper',
+ 'apply_module_wrapper']
diff --git a/lakonlab/parallel/__pycache__/__init__.cpython-310.pyc b/lakonlab/parallel/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..f27099a02ef993fb6b02a45d6b7f0d50fcb09217
Binary files /dev/null and b/lakonlab/parallel/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/parallel/__pycache__/ddp_wrapper.cpython-310.pyc b/lakonlab/parallel/__pycache__/ddp_wrapper.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..56dfc289b81368c18a1d6351abf3432bcc7f5cec
Binary files /dev/null and b/lakonlab/parallel/__pycache__/ddp_wrapper.cpython-310.pyc differ
diff --git a/lakonlab/parallel/__pycache__/distributed.cpython-310.pyc b/lakonlab/parallel/__pycache__/distributed.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..55bd8d128d6cb508f99be7ab31e9811c6a82bcc7
Binary files /dev/null and b/lakonlab/parallel/__pycache__/distributed.cpython-310.pyc differ
diff --git a/lakonlab/parallel/__pycache__/fsdp2_wrapper.cpython-310.pyc b/lakonlab/parallel/__pycache__/fsdp2_wrapper.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..5bde966637c0be1f1c6ef1886c27aed8328300ed
Binary files /dev/null and b/lakonlab/parallel/__pycache__/fsdp2_wrapper.cpython-310.pyc differ
diff --git a/lakonlab/parallel/__pycache__/fsdp_wrapper.cpython-310.pyc b/lakonlab/parallel/__pycache__/fsdp_wrapper.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..578f17b7bb72d1f89b38938a2044eadbd0ce8cb8
Binary files /dev/null and b/lakonlab/parallel/__pycache__/fsdp_wrapper.cpython-310.pyc differ
diff --git a/lakonlab/parallel/__pycache__/utils.cpython-310.pyc b/lakonlab/parallel/__pycache__/utils.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..0661460f776a6299307d927e1c34718658e58afd
Binary files /dev/null and b/lakonlab/parallel/__pycache__/utils.cpython-310.pyc differ
diff --git a/lakonlab/parallel/ddp_wrapper.py b/lakonlab/parallel/ddp_wrapper.py
new file mode 100644
index 0000000000000000000000000000000000000000..d78cfcdfcd50f9d282d3fe3b42229077e728f313
--- /dev/null
+++ b/lakonlab/parallel/ddp_wrapper.py
@@ -0,0 +1,26 @@
+from mmcv.parallel import MODULE_WRAPPERS
+from mmgen.core.ddp_wrapper import DistributedDataParallelWrapper as _DistributedDataParallelWrapper
+from lakonlab.parallel import MMDistributedDataParallel
+
+
+@MODULE_WRAPPERS.register_module('mmgen.DDPWrapper', force=True)
+class DistributedDataParallelWrapper(_DistributedDataParallelWrapper):
+
+ def to_ddp(self, device_ids, dim, broadcast_buffers,
+ find_unused_parameters, **kwargs):
+ for name, module in self.module._modules.items():
+ if next(module.parameters(), None) is None:
+ module = module.cuda()
+ elif all(not p.requires_grad for p in module.parameters()):
+ module = module.cuda()
+ elif name.endswith('_ema') or name.endswith('_ema2'):
+ module = module.cuda()
+ else:
+ module = MMDistributedDataParallel(
+ module.cuda(),
+ device_ids=device_ids,
+ dim=dim,
+ broadcast_buffers=broadcast_buffers,
+ find_unused_parameters=find_unused_parameters,
+ **kwargs)
+ self.module._modules[name] = module
diff --git a/lakonlab/parallel/distributed.py b/lakonlab/parallel/distributed.py
new file mode 100644
index 0000000000000000000000000000000000000000..abd82dcc0a0741009542671021c57988a808c801
--- /dev/null
+++ b/lakonlab/parallel/distributed.py
@@ -0,0 +1,19 @@
+from typing import Any
+
+from mmcv.parallel.distributed import MMDistributedDataParallel as _MMDistributedDataParallel
+
+
+class MMDistributedDataParallel(_MMDistributedDataParallel):
+
+ def _run_ddp_forward(self, *inputs, **kwargs) -> Any:
+ if hasattr(self, '_use_replicated_tensor_module') and self._use_replicated_tensor_module:
+ module_to_run = self._replicated_tensor_module
+ else:
+ module_to_run = self.module
+
+ if self.device_ids:
+ inputs, kwargs = self.to_kwargs( # type: ignore
+ inputs, kwargs, self.device_ids[0])
+ return module_to_run(*inputs[0], **kwargs[0]) # type: ignore
+ else:
+ return module_to_run(*inputs, **kwargs)
diff --git a/lakonlab/parallel/fsdp2_wrapper.py b/lakonlab/parallel/fsdp2_wrapper.py
new file mode 100644
index 0000000000000000000000000000000000000000..5407acc8f4e05fcab709fa23fd193b6722b90724
--- /dev/null
+++ b/lakonlab/parallel/fsdp2_wrapper.py
@@ -0,0 +1,160 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import importlib
+import torch
+import torch.nn as nn
+import torch.distributed as dist
+
+try:
+ from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy
+except:
+ pass
+from mmcv.parallel.scatter_gather import scatter_kwargs
+from mmcv.parallel import MODULE_WRAPPERS
+
+
+def get_module_object(path):
+ module_path, attribute = path.rsplit('.', 1)
+ module = importlib.import_module(module_path)
+ return getattr(module, attribute)
+
+
+@MODULE_WRAPPERS.register_module()
+class FSDP2Wrapper(nn.Module):
+
+ def __init__(
+ self,
+ module,
+ wrap_frozen_modules=False,
+ ignore_frozen_parameters=False,
+ param_dtype='bfloat16',
+ reduce_dtype='float32',
+ fsdp_modules=None,
+ exclude_keys=(),
+ hybrid_sharding=True,
+ **kwargs):
+ super().__init__()
+ self.module = module
+ fsdp_kwargs = kwargs
+ self.param_dtype = getattr(torch, param_dtype)
+ self.reduce_dtype = getattr(torch, reduce_dtype)
+ if hybrid_sharding:
+ global_world_size = dist.get_world_size()
+ num_devices_per_node = torch.cuda.device_count()
+ mesh = dist.init_device_mesh(
+ 'cuda',
+ (global_world_size // num_devices_per_node, num_devices_per_node),
+ mesh_dim_names=('replicate', 'shard'))
+ fsdp_kwargs.update(mesh=mesh)
+ if fsdp_modules is not None:
+ assert isinstance(fsdp_modules, (list, tuple))
+ fsdp_modules = tuple([get_module_object(m) for m in fsdp_modules])
+ self.to_fsdp(
+ wrap_frozen_modules,
+ ignore_frozen_parameters,
+ exclude_keys,
+ fsdp_modules=fsdp_modules,
+ **fsdp_kwargs)
+
+ def to_fsdp(self,
+ wrap_frozen_modules=False,
+ ignore_frozen_parameters=False,
+ exclude_keys=(),
+ fsdp_modules=(),
+ **kwargs):
+ for name, module in self.module._modules.items():
+ if name in exclude_keys or next(module.parameters(), None) is None:
+ module = module.cuda()
+ elif all(not p.requires_grad for p in module.parameters()):
+ if wrap_frozen_modules:
+ for submodule in module.modules():
+ if isinstance(submodule, fsdp_modules):
+ fsdp_kwargs = kwargs.copy()
+ fsdp_kwargs.update(
+ mp_policy=MixedPrecisionPolicy(
+ param_dtype=self.param_dtype,
+ reduce_dtype=self.reduce_dtype))
+ fully_shard(submodule, **fsdp_kwargs)
+ fsdp_kwargs = kwargs.copy()
+ fsdp_kwargs.update(
+ mp_policy=MixedPrecisionPolicy(
+ param_dtype=self.param_dtype,
+ reduce_dtype=self.reduce_dtype,
+ cast_forward_inputs=False))
+ fully_shard(module, **fsdp_kwargs)
+ else:
+ module = module.cuda()
+ else:
+ if ignore_frozen_parameters:
+ ignored_params = []
+ for p in module.parameters():
+ if not p.requires_grad:
+ p.data = p.data.cuda()
+ ignored_params.append(p)
+ else:
+ ignored_params = None
+ for submodule in module.modules():
+ if isinstance(submodule, fsdp_modules):
+ fsdp_kwargs = kwargs.copy()
+ fsdp_kwargs.update(
+ mp_policy=MixedPrecisionPolicy(
+ param_dtype=self.param_dtype,
+ reduce_dtype=self.reduce_dtype))
+ if ignored_params is not None: # requires torch >= 2.7
+ fsdp_kwargs.update(ignored_params=ignored_params)
+ fully_shard(submodule, **fsdp_kwargs)
+ fsdp_kwargs = kwargs.copy()
+ fsdp_kwargs.update(
+ mp_policy=MixedPrecisionPolicy(
+ param_dtype=self.param_dtype,
+ reduce_dtype=self.reduce_dtype,
+ cast_forward_inputs=False))
+ if ignored_params is not None: # requires torch >= 2.7
+ fsdp_kwargs.update(ignored_params=ignored_params)
+ fully_shard(module, **fsdp_kwargs)
+ self.module._modules[name] = module
+
+ def scatter(self, inputs, kwargs, device_ids):
+ """Scatter function.
+
+ Args:
+ inputs (Tensor): Input Tensor.
+ kwargs (dict): Args for
+ ``mmcv.parallel.scatter_gather.scatter_kwargs``.
+ device_ids (int): Device id.
+ """
+ return scatter_kwargs(inputs, kwargs, device_ids)
+
+ def forward(self, *inputs, **kwargs):
+ """Forward function.
+
+ Args:
+ inputs (tuple): Input data.
+ kwargs (dict): Args for
+ ``mmcv.parallel.scatter_gather.scatter_kwargs``.
+ """
+ inputs, kwargs = self.scatter(inputs, kwargs, [torch.cuda.current_device()])
+ return self.module(*inputs[0], **kwargs[0])
+
+ def train_step(self, *inputs, **kwargs):
+ """Train step function.
+
+ Args:
+ inputs (Tensor): Input Tensor.
+ kwargs (dict): Args for
+ ``mmcv.parallel.scatter_gather.scatter_kwargs``.
+ """
+ inputs, kwargs = self.scatter(inputs, kwargs, [torch.cuda.current_device()])
+ output = self.module.train_step(*inputs[0], **kwargs[0])
+ return output
+
+ def val_step(self, *inputs, **kwargs):
+ """Validation step function.
+
+ Args:
+ inputs (tuple): Input data.
+ kwargs (dict): Args for ``scatter_kwargs``.
+ """
+ inputs, kwargs = self.scatter(inputs, kwargs, [torch.cuda.current_device()])
+ output = self.module.val_step(*inputs[0], **kwargs[0])
+ return output
diff --git a/lakonlab/parallel/fsdp_wrapper.py b/lakonlab/parallel/fsdp_wrapper.py
new file mode 100644
index 0000000000000000000000000000000000000000..e4c356a4bcd662ad5114d67035f847a238d4641f
--- /dev/null
+++ b/lakonlab/parallel/fsdp_wrapper.py
@@ -0,0 +1,288 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import importlib
+import torch
+import torch.nn as nn
+
+from collections import abc
+from typing import Any
+from torch.distributed.fsdp import MixedPrecision, ShardingStrategy, FullyShardedDataParallel
+from torch.distributed.fsdp.wrap import ModuleWrapPolicy
+from torch.distributed.utils import _p_assert
+from torch.distributed.fsdp._runtime_utils import (
+ _pre_forward,
+ _pre_forward_unshard,
+ _root_pre_forward,
+ _post_forward,
+ _post_forward_reshard
+)
+from mmcv.parallel.scatter_gather import scatter_kwargs
+from mmcv.parallel import MODULE_WRAPPERS
+
+
+MODULE_WRAPPERS.register_module(
+ name='FSDP', module=FullyShardedDataParallel)
+
+
+def _clone_if_reqgrad(x, memo):
+ # Avoid cloning the same tensor multiple times across nested structures
+ if isinstance(x, torch.Tensor) and x.requires_grad:
+ xid = id(x)
+ y = memo.get(xid)
+ if y is None:
+ # clone keeps graph connectivity but breaks storage identity
+ # (good for breaking shared-input identity across blocks)
+ y = x.clone()
+ memo[xid] = y
+ return y
+ return x
+
+
+def _map_structure(obj, memo):
+ # Fast path for tensors/leaf types
+ y = _clone_if_reqgrad(obj, memo)
+ if y is not obj:
+ return y
+
+ # Containers
+ if isinstance(obj, (list, tuple)):
+ seq = [_map_structure(v, memo) for v in obj]
+ return type(obj)(seq) if isinstance(obj, tuple) else seq
+ if isinstance(obj, dict):
+ return {k: _map_structure(v, memo) for k, v in obj.items()}
+ if isinstance(obj, abc.Mapping): # other mappings
+ return type(obj)((k, _map_structure(v, memo)) for k, v in obj.items())
+ if isinstance(obj, abc.Sequence) and not isinstance(obj, (str, bytes)):
+ return type(obj)(_map_structure(v, memo) for v in obj)
+
+ # Namedtuple support
+ if hasattr(obj, "_fields") and hasattr(obj, "_asdict"):
+ return type(obj)(**{k: _map_structure(v, memo) for k, v in obj._asdict().items()})
+
+ # Everything else untouched
+ return obj
+
+
+def clone_grad_inputs(*args, **kwargs):
+ """
+ Args: *args, **kwargs (any nested structure)
+ Returns:
+ new_args, new_kwargs with every tensor that has requires_grad=True cloned.
+ Cloning preserves dtype/device/grad requirement and graph connectivity,
+ but breaks storage identity so FSDP's multi-grad hook won't wait on
+ a shared input used by other blocks.
+ """
+ memo = {}
+ new_args = tuple(_map_structure(a, memo) for a in args)
+ new_kwargs = {k: _map_structure(v, memo) for k, v in kwargs.items()}
+ return new_args, new_kwargs
+
+
+def get_module_object(path):
+ module_path, attribute = path.rsplit('.', 1)
+ module = importlib.import_module(module_path)
+ return getattr(module, attribute)
+
+
+class FullyShardedDataParallelFix(FullyShardedDataParallel):
+
+ def forward(self, *args: Any, **kwargs: Any) -> Any:
+ handle = self._handle
+ with torch.autograd.profiler.record_function(
+ "FullyShardedDataParallel.forward"
+ ):
+ args, kwargs = _root_pre_forward(self, self, args, kwargs)
+ unused = None
+ # =================================================
+ # clone the input tensors that require grad if the module is frozen
+ flat_param = handle.flat_param
+ already_registered = hasattr(flat_param, "_post_backward_hook_handle")
+ if not already_registered and not flat_param.requires_grad:
+ args, kwargs = clone_grad_inputs(*args, **kwargs)
+ # =================================================
+ args, kwargs = _pre_forward(
+ self,
+ handle,
+ _pre_forward_unshard,
+ self._fsdp_wrapped_module,
+ args,
+ kwargs,
+ )
+ if handle:
+ _p_assert(
+ handle.flat_param.device == self.compute_device,
+ "Expected `FlatParameter` to be on the compute device "
+ f"{self.compute_device} but got {handle.flat_param.device}",
+ )
+ output = self._fsdp_wrapped_module(*args, **kwargs)
+ return _post_forward(
+ self, handle, _post_forward_reshard, self, unused, output
+ )
+
+
+def tie_fsdp_modules(tgt_module, src_module, recursive=True):
+ if isinstance(src_module, FullyShardedDataParallel) and isinstance(tgt_module, FullyShardedDataParallel):
+
+ old_forward = tgt_module.forward
+
+ def new_forward(*args, **kwargs):
+ handle = src_module._handle
+ args, kwargs = _root_pre_forward(src_module, src_module, args, kwargs)
+ unused = None
+ # =================================================
+ # clone the input tensors that require grad if the module is frozen
+ flat_param = handle.flat_param
+ already_registered = hasattr(flat_param, "_post_backward_hook_handle")
+ if not already_registered and not flat_param.requires_grad:
+ args, kwargs = clone_grad_inputs(*args, **kwargs)
+ # =================================================
+ args, kwargs = _pre_forward(
+ src_module,
+ handle,
+ _pre_forward_unshard,
+ src_module._fsdp_wrapped_module,
+ args,
+ kwargs,
+ )
+ if handle:
+ _p_assert(
+ handle.flat_param.device == src_module.compute_device,
+ "Expected `FlatParameter` to be on the compute device "
+ f"{src_module.compute_device} but got {handle.flat_param.device}",
+ )
+ output = old_forward(*args, **kwargs)
+ return _post_forward(
+ src_module, handle, _post_forward_reshard, src_module, unused, output
+ )
+
+ tgt_module.forward = new_forward
+
+ if recursive:
+ for key, val in src_module._modules.items():
+ if key in tgt_module._modules:
+ tie_fsdp_modules(tgt_module._modules[key], val, recursive)
+
+
+@MODULE_WRAPPERS.register_module()
+class FSDPWrapper(nn.Module):
+
+ def __init__(
+ self,
+ module,
+ device_id,
+ wrap_frozen_modules=False,
+ ignore_frozen_parameters=False,
+ param_dtype='bfloat16',
+ reduce_dtype='float32',
+ buffer_dtype='bfloat16',
+ fsdp_modules=None,
+ exclude_keys=(),
+ tie_key_mappings=None,
+ use_orig_params=True,
+ sharding_strategy='HYBRID_SHARD',
+ **kwargs):
+ super().__init__()
+ self.module = module
+ if fsdp_modules is not None:
+ assert isinstance(fsdp_modules, (list, tuple))
+ fsdp_modules = [get_module_object(m) for m in fsdp_modules]
+ else:
+ fsdp_modules = []
+ fsdp_kwargs = kwargs
+ fsdp_kwargs.update(
+ use_orig_params=use_orig_params,
+ mixed_precision=MixedPrecision(
+ param_dtype=getattr(torch, param_dtype),
+ reduce_dtype=getattr(torch, reduce_dtype),
+ buffer_dtype=getattr(torch, buffer_dtype),
+ cast_root_forward_inputs=False),
+ sharding_strategy=getattr(ShardingStrategy, sharding_strategy.upper()),
+ auto_wrap_policy=ModuleWrapPolicy(fsdp_modules))
+ self.to_fsdp(
+ device_id, wrap_frozen_modules, ignore_frozen_parameters, exclude_keys, tie_key_mappings, **fsdp_kwargs)
+
+ def to_fsdp(
+ self, device_id, wrap_frozen_modules=False, ignore_frozen_parameters=False,
+ exclude_keys=(), tie_key_mappings=None, **kwargs):
+ for name, module in self.module._modules.items():
+ if name in exclude_keys or next(module.parameters(), None) is None:
+ module = module.cuda()
+ elif all(not p.requires_grad for p in module.parameters()):
+ if wrap_frozen_modules:
+ fsdp_kwargs = kwargs.copy()
+ fsdp_kwargs.update(use_orig_params=False)
+ module = FullyShardedDataParallelFix(
+ module,
+ device_id=device_id,
+ **fsdp_kwargs)
+ else:
+ module = module.cuda()
+ else:
+ fsdp_kwargs = kwargs.copy()
+ if ignore_frozen_parameters:
+ ignored_states = []
+ for p in module.parameters():
+ if not p.requires_grad:
+ p.data = p.data.cuda()
+ ignored_states.append(p)
+ fsdp_kwargs.update(ignored_states=ignored_states)
+ module = FullyShardedDataParallelFix(
+ module,
+ device_id=device_id,
+ **fsdp_kwargs)
+ self.module._modules[name] = module
+
+ if tie_key_mappings is not None:
+ # parse tie_key_mappings in the format ('teacher->diffusion', 'teacher->diffusion_ema')
+ for mapping in tie_key_mappings:
+ src_key, tgt_key = mapping.split('->')
+ tie_fsdp_modules(
+ self.module._modules[tgt_key], self.module._modules[src_key])
+
+ def scatter(self, inputs, kwargs, device_ids):
+ """Scatter function.
+
+ Args:
+ inputs (Tensor): Input Tensor.
+ kwargs (dict): Args for
+ ``mmcv.parallel.scatter_gather.scatter_kwargs``.
+ device_ids (int): Device id.
+ """
+ return scatter_kwargs(inputs, kwargs, device_ids)
+
+ def forward(self, *inputs, **kwargs):
+ """Forward function.
+
+ Args:
+ inputs (tuple): Input data.
+ kwargs (dict): Args for
+ ``mmcv.parallel.scatter_gather.scatter_kwargs``.
+ """
+ inputs, kwargs = self.scatter(inputs, kwargs,
+ [torch.cuda.current_device()])
+ return self.module(*inputs[0], **kwargs[0])
+
+ def train_step(self, *inputs, **kwargs):
+ """Train step function.
+
+ Args:
+ inputs (Tensor): Input Tensor.
+ kwargs (dict): Args for
+ ``mmcv.parallel.scatter_gather.scatter_kwargs``.
+ """
+ inputs, kwargs = self.scatter(inputs, kwargs,
+ [torch.cuda.current_device()])
+ output = self.module.train_step(*inputs[0], **kwargs[0])
+ return output
+
+ def val_step(self, *inputs, **kwargs):
+ """Validation step function.
+
+ Args:
+ inputs (tuple): Input data.
+ kwargs (dict): Args for ``scatter_kwargs``.
+ """
+ inputs, kwargs = self.scatter(inputs, kwargs,
+ [torch.cuda.current_device()])
+ output = self.module.val_step(*inputs[0], **kwargs[0])
+ return output
diff --git a/lakonlab/parallel/utils.py b/lakonlab/parallel/utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..4e32ca8f339bd3467b46a48b569ea813aabafb55
--- /dev/null
+++ b/lakonlab/parallel/utils.py
@@ -0,0 +1,36 @@
+import torch
+import mmcv
+
+from . import MMDistributedDataParallel, DistributedDataParallelWrapper, FSDPWrapper, FSDP2Wrapper
+
+
+def apply_module_wrapper(model, module_wrapper, cfg):
+ if module_wrapper is None:
+ model = MMDistributedDataParallel(
+ model.cuda(),
+ device_ids=[torch.cuda.current_device()],
+ broadcast_buffers=False,
+ find_unused_parameters=cfg.get('find_unused_parameters', False))
+ elif module_wrapper.lower() == 'ddp':
+ mmcv.print_log('Use DDP Wrapper.', 'mmgen')
+ model = DistributedDataParallelWrapper(
+ model,
+ device_ids=[torch.cuda.current_device()],
+ broadcast_buffers=False,
+ find_unused_parameters=cfg.get('find_unused_parameters', False))
+ elif module_wrapper.lower() == 'fsdp':
+ mmcv.print_log('Use FSDP Wrapper.', 'mmgen')
+ fsdp_kwargs = cfg.get('fsdp_kwargs', {})
+ model = FSDPWrapper(
+ model,
+ device_id=torch.cuda.current_device(),
+ **fsdp_kwargs)
+ elif module_wrapper.lower() == 'fsdp2':
+ mmcv.print_log('Use FSDP2 Wrapper.', 'mmgen')
+ fsdp_kwargs = cfg.get('fsdp_kwargs', {})
+ model = FSDP2Wrapper(
+ model,
+ **fsdp_kwargs)
+ else:
+ raise ValueError(f'Unsupported module wrapper: {module_wrapper}.')
+ return model
diff --git a/lakonlab/pipelines/__init__.py b/lakonlab/pipelines/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/lakonlab/pipelines/gmdit_pipeline.py b/lakonlab/pipelines/gmdit_pipeline.py
new file mode 100644
index 0000000000000000000000000000000000000000..482a35ee3115cae83cabd35d585ebaedc1aad16c
--- /dev/null
+++ b/lakonlab/pipelines/gmdit_pipeline.py
@@ -0,0 +1,152 @@
+# Copyright (c) 2025 Hansheng Chen
+
+from typing import Dict, List, Optional, Tuple, Union
+
+import torch
+
+from diffusers.models import AutoencoderKL
+from diffusers.pipelines import DiTPipeline
+from diffusers.utils.torch_utils import randn_tensor
+from diffusers.pipelines.pipeline_utils import ImagePipelineOutput
+from lakonlab.models.architecture.gmflow.gmdit import _GMDiTTransformer2DModel as GMDiTTransformer2DModel
+from lakonlab.models.architecture.gmflow.spectrum_mlp import _SpectrumMLP as SpectrumMLP
+from lakonlab.models.diffusions.schedulers import FlowSDEScheduler, FlowEulerODEScheduler
+from lakonlab.models.diffusions.gmflow import probabilistic_guidance_jit, GMFlowMixin
+from lakonlab.ops.gmflow_ops.gmflow_ops import (
+ gm_to_mean, iso_gaussian_mul_iso_gaussian, gm_mul_iso_gaussian, gm_to_iso_gaussian)
+
+
+class GMDiTPipeline(DiTPipeline, GMFlowMixin):
+
+ def __init__(
+ self,
+ transformer: GMDiTTransformer2DModel,
+ spectrum_net: SpectrumMLP,
+ vae: AutoencoderKL,
+ scheduler: FlowSDEScheduler | FlowEulerODEScheduler,
+ id2label: Optional[Dict[int, str]] = None):
+ super(DiTPipeline, self).__init__()
+ self.register_modules(transformer=transformer, spectrum_net=spectrum_net, vae=vae, scheduler=scheduler)
+
+ self.labels = {}
+ if id2label is not None:
+ for key, value in id2label.items():
+ for label in value.split(","):
+ self.labels[label.lstrip().rstrip()] = int(key)
+ self.labels = dict(sorted(self.labels.items()))
+
+ @torch.inference_mode()
+ def __call__(
+ self,
+ class_labels: List[int],
+ guidance_scale: float = 0.45,
+ generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
+ num_inference_steps: int = 32,
+ num_inference_substeps: int = 4,
+ output_mode: str = "mean",
+ order=2,
+ orthogonal_guidance: float = 1.0,
+ gm2_coefs=[0.005, 1.0],
+ gm2_correction_steps=0,
+ output_type: Optional[str] = "pil",
+ return_dict: bool = True,
+ ) -> Union[ImagePipelineOutput, Tuple]:
+ assert 0 <= guidance_scale < 1, "guidance_scale must be in [0, 1)"
+
+ batch_size = len(class_labels)
+ latent_size = self.transformer.config.sample_size
+ latent_channels = self.transformer.config.in_channels
+
+ use_guidance = guidance_scale > 0.0
+
+ x_t = randn_tensor(
+ shape=(batch_size, latent_channels, latent_size, latent_size),
+ generator=generator,
+ device=self._execution_device,
+ )
+
+ class_labels = torch.tensor(class_labels, device=self._execution_device).reshape(-1)
+ class_null = torch.tensor([1000] * batch_size, device=self._execution_device)
+ class_labels_input = torch.cat([class_labels, class_null], 0) if use_guidance else class_labels
+
+ # set step values
+ self.scheduler.set_timesteps(num_inference_steps * num_inference_substeps, device=self._execution_device)
+
+ self.init_gm_cache()
+
+ for timestep_id in self.progress_bar(range(num_inference_steps)):
+ t = self.scheduler.timesteps[timestep_id * num_inference_substeps]
+
+ x_t_input = x_t
+ if use_guidance:
+ x_t_input = torch.cat([x_t_input, x_t], dim=0)
+
+ gm_output = self.transformer(
+ x_t_input.to(dtype=self.transformer.dtype),
+ timestep=t.expand(x_t_input.size(0)),
+ class_labels=class_labels_input)
+ gm_output = {k: v.to(torch.float32) for k, v in gm_output.items()}
+ gm_output = self.u_to_x_0(gm_output, x_t_input, t)
+
+ # ========== Probabilistic CFG ==========
+ if use_guidance:
+ gm_cond = {k: v[:batch_size] for k, v in gm_output.items()}
+ gm_uncond = {k: v[batch_size:] for k, v in gm_output.items()}
+ uncond_mean = gm_to_mean(gm_uncond)
+ gaussian_cond = gm_to_iso_gaussian(gm_cond)[0]
+ gaussian_cond['var'] = gaussian_cond['var'].mean(dim=(-1, -2), keepdim=True)
+ gaussian_output, cfg_bias, avg_var = probabilistic_guidance_jit(
+ gaussian_cond['mean'], gaussian_cond['var'], uncond_mean, guidance_scale,
+ orthogonal=orthogonal_guidance)
+ gm_output = gm_mul_iso_gaussian(
+ gm_cond, iso_gaussian_mul_iso_gaussian(gaussian_output, gaussian_cond, 1, -1),
+ 1, 1)[0]
+ else:
+ gaussian_output = gm_to_iso_gaussian(gm_output)[0]
+ gm_cond = gaussian_cond = avg_var = cfg_bias = None
+
+ # ========== 2nd order GM ==========
+ if order == 2:
+ if timestep_id < num_inference_steps - 1:
+ h = t - self.scheduler.timesteps[(timestep_id + 1) * num_inference_substeps]
+ else:
+ h = t
+ gm_output, gaussian_output = self.gm_2nd_order(
+ gm_output, gaussian_output, x_t, t, h,
+ guidance_scale, gm_cond, gaussian_cond, avg_var, cfg_bias,
+ ca=gm2_coefs[0], cb=gm2_coefs[1], gm2_correction_steps=gm2_correction_steps)
+
+ # ========== GM SDE step or GM ODE substeps ==========
+ x_t_base = x_t
+ t_base = t
+ for substep_id in range(num_inference_substeps):
+ if substep_id == 0:
+ if output_mode == 'sample':
+ power_spectrum = self.spectrum_net(gaussian_output)
+ else:
+ power_spectrum = None
+ model_output = self.gm_to_model_output(
+ gm_output, output_mode, power_spectrum=power_spectrum)
+ else:
+ assert output_mode == 'mean'
+ t = self.scheduler.timesteps[timestep_id * num_inference_substeps + substep_id]
+ model_output = self.gmflow_posterior_mean(
+ gm_output, x_t, x_t_base, t, t_base, prediction_type='x0')
+ x_t = self.scheduler.step(model_output, t, x_t, return_dict=False, prediction_type='x0')[0]
+
+ x_t = x_t / self.vae.config.scaling_factor
+ samples = self.vae.decode(x_t.to(self.vae.dtype)).sample
+
+ samples = (samples / 2 + 0.5).clamp(0, 1)
+ samples = samples.cpu().permute(0, 2, 3, 1).float().numpy()
+
+ if output_type == "pil":
+ samples = self.numpy_to_pil(samples)
+
+ # Offload all models
+ self.maybe_free_model_hooks()
+
+ if not return_dict:
+ return (samples,)
+
+ return ImagePipelineOutput(images=samples)
diff --git a/lakonlab/pipelines/piflow_loader.py b/lakonlab/pipelines/piflow_loader.py
new file mode 100644
index 0000000000000000000000000000000000000000..a82152a7efea7b5957d1d1dbfc8068585f200f76
--- /dev/null
+++ b/lakonlab/pipelines/piflow_loader.py
@@ -0,0 +1,275 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import os
+from typing import Union, Optional
+
+import torch
+import accelerate
+import diffusers
+from diffusers.models import AutoModel
+from diffusers.models.modeling_utils import (
+ load_state_dict,
+ _LOW_CPU_MEM_USAGE_DEFAULT,
+ no_init_weights,
+ ContextManagers
+)
+from diffusers.utils import (
+ SAFETENSORS_WEIGHTS_NAME,
+ WEIGHTS_NAME,
+ _add_variant,
+ _get_model_file,
+ is_accelerate_available,
+ is_torch_version,
+ logging,
+)
+from diffusers.loaders.peft import _SET_ADAPTER_SCALE_FN_MAPPING
+from lakonlab.models.architecture.gmflow.gmflux import _GMFluxTransformer2DModel
+from lakonlab.models.architecture.gmflow.gmqwen import _GMQwenImageTransformer2DModel
+
+
+LOCAL_CLASS_MAPPING = {
+ "GMFluxTransformer2DModel": _GMFluxTransformer2DModel,
+ "GMQwenImageTransformer2DModel": _GMQwenImageTransformer2DModel,
+}
+
+_SET_ADAPTER_SCALE_FN_MAPPING.update(
+ _GMFluxTransformer2DModel=lambda model_cls, weights: weights,
+ _GMQwenImageTransformer2DModel=lambda model_cls, weights: weights,
+)
+
+logger = logging.get_logger(__name__)
+
+
+class PiFlowLoaderMixin:
+
+ def load_piflow_adapter(
+ self,
+ pretrained_model_name_or_path: Union[str, os.PathLike],
+ target_module_name: str = "transformer",
+ adapter_name: Optional[str] = None,
+ **kwargs
+ ):
+ r"""
+ Load a PiFlow adapter from a pretrained model repository into the target module.
+
+ Args:
+ pretrained_model_name_or_path (`str` or `os.PathLike`):
+ Can be either:
+
+ - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on
+ the Hub.
+ - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved
+ with [`~ModelMixin.save_pretrained`].
+
+ target_module_name (`str`, *optional*, defaults to `"transformer"`):
+ The module name in the model to load the PiFlow adapter into.
+ adapter_name (`str`, *optional*):
+ The name to assign to the loaded adapter. If not provided, it defaults to
+ `"{target_module_name}_piflow"`.
+ cache_dir (`Union[str, os.PathLike]`, *optional*):
+ Path to a directory where a downloaded pretrained model configuration is cached if the standard cache
+ is not used.
+ force_download (`bool`, *optional*, defaults to `False`):
+ Whether or not to force the (re-)download of the model weights and configuration files, overriding the
+ cached versions if they exist.
+ proxies (`Dict[str, str]`, *optional*):
+ A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128',
+ 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request.
+ local_files_only(`bool`, *optional*, defaults to `False`):
+ Whether to only load local model weights and configuration files or not. If set to `True`, the model
+ won't be downloaded from the Hub.
+ token (`str` or *bool*, *optional*):
+ The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from
+ `diffusers-cli login` (stored in `~/.huggingface`) is used.
+ revision (`str`, *optional*, defaults to `"main"`):
+ The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier
+ allowed by Git.
+ subfolder (`str`, *optional*, defaults to `""`):
+ The subfolder location of a model file within a larger model repository on the Hub or locally.
+ low_cpu_mem_usage (`bool`, *optional*, defaults to `True` if torch version >= 1.9.0 else `False`):
+ Speed up model loading only loading the pretrained weights and not initializing the weights. This also
+ tries to not use more than 1x model size in CPU memory (including peak memory) while loading the model.
+ Only supported for PyTorch >= 1.9.0. If you are using an older version of PyTorch, setting this
+ argument to `True` will raise an error.
+ variant (`str`, *optional*):
+ Load weights from a specified `variant` filename such as `"fp16"` or `"ema"`. This is ignored when
+ loading `from_flax`.
+ use_safetensors (`bool`, *optional*, defaults to `None`):
+ If set to `None`, the `safetensors` weights are downloaded if they're available **and** if the
+ `safetensors` library is installed. If set to `True`, the model is forcibly loaded from `safetensors`
+ weights. If set to `False`, `safetensors` weights are not loaded.
+ disable_mmap ('bool', *optional*, defaults to 'False'):
+ Whether to disable mmap when loading a Safetensors model. This option can perform better when the model
+ is on a network mount or hard drive, which may not handle the seeky-ness of mmap very well.
+
+ Returns:
+ `str` or `None`: The name assigned to the loaded adapter, or `None` if no LoRA weights were found.
+ """
+ cache_dir = kwargs.pop("cache_dir", None)
+ force_download = kwargs.pop("force_download", False)
+ proxies = kwargs.pop("proxies", None)
+ token = kwargs.pop("token", None)
+ local_files_only = kwargs.pop("local_files_only", False)
+ revision = kwargs.pop("revision", None)
+ subfolder = kwargs.pop("subfolder", None)
+ low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT)
+ variant = kwargs.pop("variant", None)
+ use_safetensors = kwargs.pop("use_safetensors", None)
+ disable_mmap = kwargs.pop("disable_mmap", False)
+
+ allow_pickle = False
+ if use_safetensors is None:
+ use_safetensors = True
+ allow_pickle = True
+
+ if low_cpu_mem_usage and not is_accelerate_available():
+ low_cpu_mem_usage = False
+ logger.warning(
+ "Cannot initialize model with low cpu memory usage because `accelerate` was not found in the"
+ " environment. Defaulting to `low_cpu_mem_usage=False`. It is strongly recommended to install"
+ " `accelerate` for faster and less memory-intense model loading. You can do so with: \n```\npip"
+ " install accelerate\n```\n."
+ )
+
+ if low_cpu_mem_usage is True and not is_torch_version(">=", "1.9.0"):
+ raise NotImplementedError(
+ "Low memory initialization requires torch >= 1.9.0. Please either update your PyTorch version or set"
+ " `low_cpu_mem_usage=False`."
+ )
+
+ user_agent = {
+ "diffusers": diffusers.__version__,
+ "file_type": "model",
+ "framework": "pytorch",
+ }
+
+ # 1. Determine model class from config
+
+ load_config_kwargs = {
+ "cache_dir": cache_dir,
+ "force_download": force_download,
+ "proxies": proxies,
+ "token": token,
+ "local_files_only": local_files_only,
+ "revision": revision,
+ }
+
+ config = AutoModel.load_config(pretrained_model_name_or_path, subfolder=subfolder, **load_config_kwargs)
+
+ orig_class_name = config["_class_name"]
+
+ if orig_class_name in LOCAL_CLASS_MAPPING:
+ model_cls = LOCAL_CLASS_MAPPING[orig_class_name]
+
+ else:
+ load_config_kwargs.update({"subfolder": subfolder})
+
+ from diffusers.pipelines.pipeline_loading_utils import ALL_IMPORTABLE_CLASSES, get_class_obj_and_candidates
+
+ model_cls, _ = get_class_obj_and_candidates(
+ library_name="diffusers",
+ class_name=orig_class_name,
+ importable_classes=ALL_IMPORTABLE_CLASSES,
+ pipelines=None,
+ is_pipeline_module=False,
+ )
+
+ if model_cls is None:
+ raise ValueError(f"Can't find a model linked to {orig_class_name}.")
+
+ # 2. Get model file
+
+ model_file = None
+
+ if use_safetensors:
+ try:
+ model_file = _get_model_file(
+ pretrained_model_name_or_path,
+ weights_name=_add_variant(SAFETENSORS_WEIGHTS_NAME, variant),
+ cache_dir=cache_dir,
+ force_download=force_download,
+ proxies=proxies,
+ local_files_only=local_files_only,
+ token=token,
+ revision=revision,
+ subfolder=subfolder,
+ user_agent=user_agent,
+ )
+
+ except IOError as e:
+ logger.error(f"An error occurred while trying to fetch {pretrained_model_name_or_path}: {e}")
+ if not allow_pickle:
+ raise
+ logger.warning(
+ "Defaulting to unsafe serialization. Pass `allow_pickle=False` to raise an error instead."
+ )
+
+ if model_file is None:
+ model_file = _get_model_file(
+ pretrained_model_name_or_path,
+ weights_name=_add_variant(WEIGHTS_NAME, variant),
+ cache_dir=cache_dir,
+ force_download=force_download,
+ proxies=proxies,
+ local_files_only=local_files_only,
+ token=token,
+ revision=revision,
+ subfolder=subfolder,
+ user_agent=user_agent,
+ )
+
+ # 3. Initialize model
+
+ base_module = getattr(self, target_module_name)
+
+ torch_dtype = base_module.dtype
+ device = base_module.device
+ dtype_orig = model_cls._set_default_torch_dtype(torch_dtype)
+
+ init_contexts = [no_init_weights()]
+
+ if low_cpu_mem_usage:
+ init_contexts.append(accelerate.init_empty_weights())
+
+ with ContextManagers(init_contexts):
+ piflow_module = model_cls.from_config(config).eval()
+
+ torch.set_default_dtype(dtype_orig)
+
+ # 4. Load model weights
+
+ if model_file is not None:
+ base_state_dict = base_module.state_dict()
+ lora_state_dict = dict()
+
+ adapter_state_dict = load_state_dict(model_file, disable_mmap=disable_mmap)
+ for k in adapter_state_dict.keys():
+ adapter_state_dict[k] = adapter_state_dict[k].to(dtype=torch_dtype, device=device)
+ if "lora" in k:
+ lora_state_dict[k.removeprefix(f"{target_module_name}.")] = adapter_state_dict[k]
+ else:
+ base_state_dict[k.removeprefix(f"{target_module_name}.")] = adapter_state_dict[k]
+
+ if len(lora_state_dict) == 0:
+ adapter_name = None
+
+ else:
+ if adapter_name is None:
+ adapter_name = f"{target_module_name}_piflow"
+
+ piflow_module.load_state_dict(
+ base_state_dict, strict=False, assign=True)
+ piflow_module.load_lora_adapter(
+ lora_state_dict, prefix=None, adapter_name=adapter_name)
+
+ setattr(self, target_module_name, piflow_module)
+
+ else:
+ adapter_name = None
+
+ if adapter_name is None:
+ logger.warning(
+ f"No LoRA weights were found in {pretrained_model_name_or_path}."
+ )
+
+ return adapter_name
diff --git a/lakonlab/pipelines/piflux_pipeline.py b/lakonlab/pipelines/piflux_pipeline.py
new file mode 100644
index 0000000000000000000000000000000000000000..c6aead37962b0034ebbbe13399f84855ec6fa175
--- /dev/null
+++ b/lakonlab/pipelines/piflux_pipeline.py
@@ -0,0 +1,491 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import numpy as np
+import torch
+
+from typing import Dict, List, Optional, Union, Any, Callable
+from functools import partial
+from transformers import (
+ CLIPImageProcessor,
+ CLIPTextModel,
+ CLIPTokenizer,
+ CLIPVisionModelWithProjection,
+ T5EncoderModel,
+ T5TokenizerFast,
+)
+from diffusers.utils import is_torch_xla_available
+from diffusers.image_processor import PipelineImageInput
+from diffusers.models import AutoencoderKL, FluxTransformer2DModel
+from diffusers.pipelines.flux.pipeline_flux import (
+ FluxPipeline, calculate_shift, FluxPipelineOutput, retrieve_timesteps)
+from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
+from lakonlab.models.diffusions.piflow_policies import POLICY_CLASSES
+from .piflow_loader import PiFlowLoaderMixin
+
+
+if is_torch_xla_available():
+ import torch_xla.core.xla_model as xm
+
+ XLA_AVAILABLE = True
+else:
+ XLA_AVAILABLE = False
+
+
+def retrieve_raw_timesteps(
+ num_inference_steps: int,
+ total_substeps: int,
+ final_step_size_scale: float
+):
+ r"""
+ Retrieve the raw times and the number of substeps for each inference step.
+
+ Args:
+ num_inference_steps (`int`):
+ Number of inference steps.
+ total_substeps (`int`):
+ Total number of substeps (e.g., 128).
+ final_step_size_scale (`float`):
+ Scale for the final step size (e.g., 0.5).
+
+ Returns:
+ `Tuple[List[float], List[int], int]`: A tuple where the first element is the raw timestep schedule, the second
+ element is the number of substeps for each inference step, and the third element is the rounded total number of
+ substeps.
+ """
+ base_segment_size = 1 / (num_inference_steps - 1 + final_step_size_scale)
+ raw_timesteps = []
+ num_inference_substeps = []
+ _raw_t = 1.0
+ for i in range(num_inference_steps):
+ if i < num_inference_steps - 1:
+ segment_size = base_segment_size
+ else:
+ segment_size = base_segment_size * final_step_size_scale
+ _num_inference_substeps = max(round(segment_size * total_substeps), 1)
+ num_inference_substeps.append(_num_inference_substeps)
+ raw_timesteps.extend(np.linspace(
+ _raw_t, _raw_t - segment_size, _num_inference_substeps, endpoint=False).clip(min=0.0).tolist())
+ _raw_t = _raw_t - segment_size
+ total_substeps = sum(num_inference_substeps)
+ return raw_timesteps, num_inference_substeps, total_substeps
+
+
+class PiFluxPipeline(FluxPipeline, PiFlowLoaderMixin):
+ r"""
+ The policy-based Flux pipeline for text-to-image generation.
+
+ Reference: https://arxiv.org/abs/2510.14974
+
+ Args:
+ transformer ([`FluxTransformer2DModel`]):
+ Conditional Transformer (MMDiT) architecture to denoise the encoded image latents.
+ scheduler ([`FlowMatchEulerDiscreteScheduler`]):
+ A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
+ vae ([`AutoencoderKL`]):
+ Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
+ text_encoder ([`CLIPTextModel`]):
+ [CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), specifically
+ the [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14) variant.
+ text_encoder_2 ([`T5EncoderModel`]):
+ [T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
+ the [google/t5-v1_1-xxl](https://huggingface.co/google/t5-v1_1-xxl) variant.
+ tokenizer (`CLIPTokenizer`):
+ Tokenizer of class
+ [CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer).
+ tokenizer_2 (`T5TokenizerFast`):
+ Second Tokenizer of class
+ [T5TokenizerFast](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5TokenizerFast).
+ policy_type (`str`, *optional*, defaults to `"GMFlow"`):
+ The type of flow policy to use. Currently supports `"GMFlow"` and `"DX"`.
+ policy_kwargs (`Dict`, *optional*):
+ Additional keyword arguments to pass to the policy class.
+ """
+
+ def __init__(
+ self,
+ scheduler: FlowMatchEulerDiscreteScheduler,
+ vae: AutoencoderKL,
+ text_encoder: CLIPTextModel,
+ tokenizer: CLIPTokenizer,
+ text_encoder_2: T5EncoderModel,
+ tokenizer_2: T5TokenizerFast,
+ transformer: FluxTransformer2DModel,
+ image_encoder: CLIPVisionModelWithProjection = None,
+ feature_extractor: CLIPImageProcessor = None,
+ policy_type: str = 'GMFlow',
+ policy_kwargs: Optional[Dict[str, Any]] = None,
+ ):
+ super().__init__(
+ scheduler,
+ vae,
+ text_encoder,
+ tokenizer,
+ text_encoder_2,
+ tokenizer_2,
+ transformer,
+ image_encoder,
+ feature_extractor
+ )
+ assert policy_type in POLICY_CLASSES, f'Invalid policy: {policy_type}. Supported policies are {list(POLICY_CLASSES.keys())}.'
+ self.policy_type = policy_type
+ self.policy_class = partial(
+ POLICY_CLASSES[policy_type], **policy_kwargs
+ ) if policy_kwargs else POLICY_CLASSES[policy_type]
+
+ def _unpack_gm(self, gm, height, width, num_channels_latents, patch_size=2, gm_patch_size=1):
+ c = num_channels_latents * patch_size * patch_size
+ h = (int(height) // (self.vae_scale_factor * patch_size))
+ w = (int(width) // (self.vae_scale_factor * patch_size))
+ bs = gm['means'].size(0)
+ k = self.transformer.num_gaussians
+ scale = patch_size // gm_patch_size
+ gm['means'] = gm['means'].reshape(
+ bs, h, w, k, c // (scale * scale), scale, scale
+ ).permute(
+ 0, 3, 4, 1, 5, 2, 6
+ ).reshape(
+ bs, k, c // (scale * scale), h * scale, w * scale)
+ gm['logweights'] = gm['logweights'].reshape(
+ bs, h, w, k, 1, scale, scale
+ ).permute(
+ 0, 3, 4, 1, 5, 2, 6
+ ).reshape(
+ bs, k, 1, h * scale, w * scale)
+ gm['logstds'] = gm['logstds'].reshape(bs, 1, 1, 1, 1)
+ return gm
+
+ @staticmethod
+ def _pack_latents(latents, batch_size, num_channels_latents, height, width, patch_size=1, target_patch_size=2):
+ scale = target_patch_size // patch_size
+ latents = latents.view(
+ batch_size,
+ num_channels_latents * patch_size * patch_size,
+ height // target_patch_size, scale, width // target_patch_size, scale)
+ latents = latents.permute(0, 2, 4, 1, 3, 5)
+ latents = latents.reshape(
+ batch_size,
+ (height // target_patch_size) * (width // target_patch_size),
+ num_channels_latents * target_patch_size * target_patch_size)
+
+ return latents
+
+ @staticmethod
+ def _unpack_latents(latents, height, width, vae_scale_factor, patch_size=2, target_patch_size=1):
+ batch_size, num_patches, channels = latents.shape
+ scale = patch_size // target_patch_size
+
+ # VAE applies 8x compression on images but we must also account for packing which requires
+ # latent height and width to be divisible by 2.
+ height = (int(height) // (vae_scale_factor * patch_size))
+ width = (int(width) // (vae_scale_factor * patch_size))
+
+ latents = latents.view(
+ batch_size, height, width, channels // (scale * scale), scale, scale)
+ latents = latents.permute(0, 3, 1, 4, 2, 5)
+
+ latents = latents.reshape(batch_size, channels // (scale * scale), height * scale, width * scale)
+
+ return latents
+
+ @torch.inference_mode()
+ def __call__(
+ self,
+ prompt: Union[str, List[str]] = None,
+ prompt_2: Optional[Union[str, List[str]]] = None,
+ height: Optional[int] = None,
+ width: Optional[int] = None,
+ num_inference_steps: int = 4,
+ total_substeps: int = 128,
+ final_step_size_scale: float = 0.5,
+ temperature: Union[float, str] = 'auto',
+ guidance_scale: float = 3.5,
+ num_images_per_prompt: Optional[int] = 1,
+ generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
+ latents: Optional[torch.FloatTensor] = None,
+ prompt_embeds: Optional[torch.FloatTensor] = None,
+ pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
+ ip_adapter_image: Optional[PipelineImageInput] = None,
+ ip_adapter_image_embeds: Optional[List[torch.Tensor]] = None,
+ output_type: Optional[str] = "pil",
+ return_dict: bool = True,
+ joint_attention_kwargs: Optional[Dict[str, Any]] = None,
+ callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
+ callback_on_step_end_tensor_inputs: List[str] = ["latents"],
+ max_sequence_length: int = 512,
+ ):
+ r"""
+ Function invoked when calling the pipeline for generation.
+
+ Args:
+ prompt (`str` or `List[str]`, *optional*):
+ The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
+ instead.
+ prompt_2 (`str` or `List[str]`, *optional*):
+ The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
+ will be used instead.
+ height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
+ The height in pixels of the generated image. This is set to 1024 by default for the best results.
+ width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
+ The width in pixels of the generated image. This is set to 1024 by default for the best results.
+ num_inference_steps (`int`, *optional*, defaults to 50):
+ The number of denoising steps.
+ total_substeps (`int`, *optional*, defaults to 128):
+ The total number of substeps for policy-based flow integration.
+ final_step_size_scale (`float`, *optional*, defaults to 0.5):
+ The scale for the final step size.
+ temperature (`float` or `"auto"`, *optional*, defaults to `"auto"`):
+ The tmperature parameter for the flow policy.
+ guidance_scale (`float`, *optional*, defaults to 3.5):
+ Embedded guiddance scale is enabled by setting `guidance_scale` > 1. Higher `guidance_scale` encourages
+ a model to generate images more aligned with `prompt` at the expense of lower image quality.
+
+ Guidance-distilled models approximates true classifer-free guidance for `guidance_scale` > 1. Refer to
+ the [paper](https://huggingface.co/papers/2210.03142) to learn more.
+ num_images_per_prompt (`int`, *optional*, defaults to 1):
+ The number of images to generate per prompt.
+ generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
+ One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
+ to make generation deterministic.
+ latents (`torch.FloatTensor`, *optional*):
+ Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
+ generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
+ tensor will be generated by sampling using the supplied random `generator`.
+ prompt_embeds (`torch.FloatTensor`, *optional*):
+ Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
+ provided, text embeddings will be generated from `prompt` input argument.
+ pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
+ Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
+ If not provided, pooled text embeddings will be generated from `prompt` input argument.
+ ip_adapter_image: (`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters.
+ ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*):
+ Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of
+ IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not
+ provided, embeddings are computed from the `ip_adapter_image` input argument.
+ output_type (`str`, *optional*, defaults to `"pil"`):
+ The output format of the generate image. Choose between
+ [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
+ return_dict (`bool`, *optional*, defaults to `True`):
+ Whether or not to return a [`~pipelines.flux.FluxPipelineOutput`] instead of a plain tuple.
+ joint_attention_kwargs (`dict`, *optional*):
+ A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
+ `self.processor` in
+ [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
+ callback_on_step_end (`Callable`, *optional*):
+ A function that calls at the end of each denoising steps during the inference. The function is called
+ with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
+ callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
+ `callback_on_step_end_tensor_inputs`.
+ callback_on_step_end_tensor_inputs (`List`, *optional*):
+ The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
+ will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
+ `._callback_tensor_inputs` attribute of your pipeline class.
+ max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`.
+
+ Returns:
+ [`~pipelines.flux.FluxPipelineOutput`] or `tuple`: [`~pipelines.flux.FluxPipelineOutput`] if `return_dict`
+ is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated
+ images.
+ """
+
+ height = height or self.default_sample_size * self.vae_scale_factor
+ width = width or self.default_sample_size * self.vae_scale_factor
+
+ # 1. Check inputs. Raise error if not correct
+ self.check_inputs(
+ prompt,
+ prompt_2,
+ height,
+ width,
+ prompt_embeds=prompt_embeds,
+ pooled_prompt_embeds=pooled_prompt_embeds,
+ callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
+ max_sequence_length=max_sequence_length,
+ )
+
+ self._guidance_scale = guidance_scale
+ self._joint_attention_kwargs = joint_attention_kwargs
+ self._current_timestep = None
+ self._interrupt = False
+
+ # 2. Define call parameters
+ if prompt is not None and isinstance(prompt, str):
+ batch_size = 1
+ elif prompt is not None and isinstance(prompt, list):
+ batch_size = len(prompt)
+ else:
+ batch_size = prompt_embeds.shape[0]
+
+ device = self._execution_device
+
+ # 3. Prepare prompt embeddings
+ lora_scale = (
+ self.joint_attention_kwargs.get("scale", None) if self.joint_attention_kwargs is not None else None
+ )
+ (
+ prompt_embeds,
+ pooled_prompt_embeds,
+ text_ids,
+ ) = self.encode_prompt(
+ prompt=prompt,
+ prompt_2=prompt_2,
+ prompt_embeds=prompt_embeds,
+ pooled_prompt_embeds=pooled_prompt_embeds,
+ device=device,
+ num_images_per_prompt=num_images_per_prompt,
+ max_sequence_length=max_sequence_length,
+ lora_scale=lora_scale,
+ )
+
+ # 4. Prepare latent variables
+ num_channels_latents = self.transformer.config.in_channels // 4
+ latents, latent_image_ids = self.prepare_latents(
+ batch_size * num_images_per_prompt,
+ num_channels_latents,
+ height,
+ width,
+ torch.float32,
+ device,
+ generator,
+ latents,
+ )
+
+ # 5. Prepare timesteps
+ raw_timesteps, num_inference_substeps, total_substeps = retrieve_raw_timesteps(
+ num_inference_steps, total_substeps, final_step_size_scale)
+ image_seq_len = latents.shape[1]
+ mu = calculate_shift(
+ image_seq_len,
+ self.scheduler.config.get("base_image_seq_len", 256),
+ self.scheduler.config.get("max_image_seq_len", 4096),
+ self.scheduler.config.get("base_shift", 0.5),
+ self.scheduler.config.get("max_shift", 1.15),
+ )
+ timesteps, _ = retrieve_timesteps(
+ self.scheduler,
+ num_inference_steps,
+ device,
+ sigmas=raw_timesteps,
+ mu=mu,
+ )
+ assert len(timesteps) == total_substeps
+ self._num_timesteps = total_substeps
+
+ # handle guidance
+ if self.transformer.config.guidance_embeds:
+ guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32)
+ guidance = guidance.expand(latents.shape[0])
+ else:
+ guidance = None
+
+ if self.joint_attention_kwargs is None:
+ self._joint_attention_kwargs = {}
+
+ image_embeds = None
+ if ip_adapter_image is not None or ip_adapter_image_embeds is not None:
+ image_embeds = self.prepare_ip_adapter_image_embeds(
+ ip_adapter_image,
+ ip_adapter_image_embeds,
+ device,
+ batch_size * num_images_per_prompt,
+ )
+
+ # 6. Denoising loop
+ self.scheduler.set_begin_index(0)
+ timestep_id = 0
+ with self.progress_bar(total=num_inference_steps) as progress_bar:
+ for i in range(num_inference_steps):
+ if self.interrupt:
+ continue
+
+ t_src = timesteps[timestep_id]
+ sigma_t_src = t_src / self.scheduler.config.num_train_timesteps
+ is_final_step = i == (num_inference_steps - 1)
+
+ self._current_timestep = t_src
+ if image_embeds is not None:
+ self._joint_attention_kwargs["ip_adapter_image_embeds"] = image_embeds
+
+ with self.transformer.cache_context("cond"):
+ denoising_output = self.transformer(
+ hidden_states=latents.to(dtype=self.transformer.dtype),
+ timestep=t_src.expand(latents.shape[0]) / 1000,
+ guidance=guidance,
+ pooled_projections=pooled_prompt_embeds,
+ encoder_hidden_states=prompt_embeds,
+ txt_ids=text_ids,
+ img_ids=latent_image_ids,
+ joint_attention_kwargs=self.joint_attention_kwargs,
+ )
+
+ # unpack and create policy
+ latents = self._unpack_latents(
+ latents, height, width, self.vae_scale_factor, target_patch_size=1)
+ if self.policy_type == 'GMFlow':
+ denoising_output = self._unpack_gm(
+ denoising_output, height, width, num_channels_latents, gm_patch_size=1)
+ denoising_output = {k: v.to(torch.float32) for k, v in denoising_output.items()}
+ policy = self.policy_class(
+ denoising_output, latents, sigma_t_src)
+ if not is_final_step:
+ if temperature == 'auto':
+ temperature = min(max(0.1 * (num_inference_steps - 1), 0), 1)
+ else:
+ assert isinstance(temperature, (float, int))
+ policy.temperature_(temperature)
+ elif self.policy_type == 'DX':
+ denoising_output = denoising_output[0]
+ denoising_output = self._unpack_latents(
+ denoising_output, height, width, self.vae_scale_factor, target_patch_size=1)
+ denoising_output = denoising_output.reshape(latents.size(0), -1, *latents.shape[1:])
+ denoising_output = denoising_output.to(torch.float32)
+ policy = self.policy_class(
+ denoising_output, latents, sigma_t_src)
+ else:
+ raise ValueError(f'Unknown policy type: {self.policy_type}.')
+
+ # compute the previous noisy sample x_t -> x_t-1
+ for _ in range(num_inference_substeps[i]):
+ t = timesteps[timestep_id]
+ sigma_t = t / self.scheduler.config.num_train_timesteps
+ u = policy.pi(latents, sigma_t)
+ latents = self.scheduler.step(u, t, latents, return_dict=False)[0]
+ timestep_id += 1
+
+ # repack
+ latents = self._pack_latents(
+ latents, latents.size(0), num_channels_latents,
+ 2 * (int(height) // (self.vae_scale_factor * 2)),
+ 2 * (int(width) // (self.vae_scale_factor * 2)),
+ patch_size=1)
+
+ if callback_on_step_end is not None:
+ callback_kwargs = {}
+ for k in callback_on_step_end_tensor_inputs:
+ callback_kwargs[k] = locals()[k]
+ callback_outputs = callback_on_step_end(self, i, t_src, callback_kwargs)
+
+ latents = callback_outputs.pop("latents", latents)
+ prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
+
+ progress_bar.update()
+
+ if XLA_AVAILABLE:
+ xm.mark_step()
+
+ self._current_timestep = None
+
+ if output_type == "latent":
+ image = latents
+ else:
+ latents = self._unpack_latents(latents, height, width, self.vae_scale_factor)
+ latents = (latents / self.vae.config.scaling_factor) + self.vae.config.shift_factor
+ image = self.vae.decode(latents.to(self.vae.dtype), return_dict=False)[0]
+ image = self.image_processor.postprocess(image, output_type=output_type)
+
+ # Offload all models
+ self.maybe_free_model_hooks()
+
+ if not return_dict:
+ return (image,)
+
+ return FluxPipelineOutput(images=image)
diff --git a/lakonlab/pipelines/piqwen_pipeline.py b/lakonlab/pipelines/piqwen_pipeline.py
new file mode 100644
index 0000000000000000000000000000000000000000..e5a24dc7f982d3a5c202cca1e1f9ca7e3867d82d
--- /dev/null
+++ b/lakonlab/pipelines/piqwen_pipeline.py
@@ -0,0 +1,429 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import numpy as np
+import torch
+
+from typing import Dict, List, Optional, Union, Any, Callable
+from functools import partial
+from transformers import Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer
+from diffusers.utils import is_torch_xla_available
+from diffusers.models import AutoencoderKLQwenImage, QwenImageTransformer2DModel
+from diffusers.pipelines.qwenimage.pipeline_qwenimage import (
+ QwenImagePipeline, calculate_shift, retrieve_timesteps, QwenImagePipelineOutput)
+from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
+from lakonlab.models.diffusions.piflow_policies import POLICY_CLASSES
+from .piflow_loader import PiFlowLoaderMixin
+
+
+if is_torch_xla_available():
+ import torch_xla.core.xla_model as xm
+
+ XLA_AVAILABLE = True
+else:
+ XLA_AVAILABLE = False
+
+
+def retrieve_raw_timesteps(
+ num_inference_steps: int,
+ total_substeps: int,
+ final_step_size_scale: float
+):
+ r"""
+ Retrieve the raw times and the number of substeps for each inference step.
+
+ Args:
+ num_inference_steps (`int`):
+ Number of inference steps.
+ total_substeps (`int`):
+ Total number of substeps (e.g., 128).
+ final_step_size_scale (`float`):
+ Scale for the final step size (e.g., 0.5).
+
+ Returns:
+ `Tuple[List[float], List[int], int]`: A tuple where the first element is the raw timestep schedule, the second
+ element is the number of substeps for each inference step, and the third element is the rounded total number of
+ substeps.
+ """
+ base_segment_size = 1 / (num_inference_steps - 1 + final_step_size_scale)
+ raw_timesteps = []
+ num_inference_substeps = []
+ _raw_t = 1.0
+ for i in range(num_inference_steps):
+ if i < num_inference_steps - 1:
+ segment_size = base_segment_size
+ else:
+ segment_size = base_segment_size * final_step_size_scale
+ _num_inference_substeps = max(round(segment_size * total_substeps), 1)
+ num_inference_substeps.append(_num_inference_substeps)
+ raw_timesteps.extend(np.linspace(
+ _raw_t, _raw_t - segment_size, _num_inference_substeps, endpoint=False).clip(min=0.0).tolist())
+ _raw_t = _raw_t - segment_size
+ total_substeps = sum(num_inference_substeps)
+ return raw_timesteps, num_inference_substeps, total_substeps
+
+
+class PiQwenImagePipeline(QwenImagePipeline, PiFlowLoaderMixin):
+ r"""
+ The policy-based QwenImage pipeline for text-to-image generation.
+
+ Reference: https://arxiv.org/abs/2510.14974
+
+ Args:
+ transformer ([`QwenImageTransformer2DModel`]):
+ Conditional Transformer (MMDiT) architecture to denoise the encoded image latents.
+ scheduler ([`FlowMatchEulerDiscreteScheduler`]):
+ A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
+ vae ([`AutoencoderKL`]):
+ Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
+ text_encoder ([`Qwen2.5-VL-7B-Instruct`]):
+ [Qwen2.5-VL-7B-Instruct](https://huggingface.co/Qwen/Qwen2.5-VL-7B-Instruct), specifically the
+ [Qwen2.5-VL-7B-Instruct](https://huggingface.co/Qwen/Qwen2.5-VL-7B-Instruct) variant.
+ tokenizer (`QwenTokenizer`):
+ Tokenizer of class
+ [CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer).
+ policy_type (`str`, *optional*, defaults to `"GMFlow"`):
+ The type of flow policy to use. Currently supports `"GMFlow"` and `"DX"`.
+ policy_kwargs (`Dict`, *optional*):
+ Additional keyword arguments to pass to the policy class.
+ """
+
+ def __init__(
+ self,
+ scheduler: FlowMatchEulerDiscreteScheduler,
+ vae: AutoencoderKLQwenImage,
+ text_encoder: Qwen2_5_VLForConditionalGeneration,
+ tokenizer: Qwen2Tokenizer,
+ transformer: QwenImageTransformer2DModel,
+ policy_type: str = 'GMFlow',
+ policy_kwargs: Optional[Dict[str, Any]] = None,
+ ):
+ super().__init__(
+ scheduler,
+ vae,
+ text_encoder,
+ tokenizer,
+ transformer,
+ )
+ assert policy_type in POLICY_CLASSES, f'Invalid policy: {policy_type}. Supported policies are {list(POLICY_CLASSES.keys())}.'
+ self.policy_type = policy_type
+ self.policy_class = partial(
+ POLICY_CLASSES[policy_type], **policy_kwargs
+ ) if policy_kwargs else POLICY_CLASSES[policy_type]
+
+ def _unpack_gm(self, gm, height, width, num_channels_latents, patch_size=2, gm_patch_size=1):
+ c = num_channels_latents * patch_size * patch_size
+ h = (int(height) // (self.vae_scale_factor * patch_size))
+ w = (int(width) // (self.vae_scale_factor * patch_size))
+ bs = gm['means'].size(0)
+ k = self.transformer.num_gaussians
+ scale = patch_size // gm_patch_size
+ gm['means'] = gm['means'].reshape(
+ bs, h, w, k, c // (scale * scale), scale, scale
+ ).permute(
+ 0, 3, 4, 1, 5, 2, 6
+ ).reshape(
+ bs, k, c // (scale * scale), h * scale, w * scale)
+ gm['logweights'] = gm['logweights'].reshape(
+ bs, h, w, k, 1, scale, scale
+ ).permute(
+ 0, 3, 4, 1, 5, 2, 6
+ ).reshape(
+ bs, k, 1, h * scale, w * scale)
+ gm['logstds'] = gm['logstds'].reshape(bs, 1, 1, 1, 1)
+ return gm
+
+ @staticmethod
+ def _pack_latents(latents, batch_size, num_channels_latents, height, width, patch_size=1, target_patch_size=2):
+ scale = target_patch_size // patch_size
+ latents = latents.view(
+ batch_size,
+ num_channels_latents * patch_size * patch_size,
+ height // target_patch_size, scale, width // target_patch_size, scale)
+ latents = latents.permute(0, 2, 4, 1, 3, 5)
+ latents = latents.reshape(
+ batch_size,
+ (height // target_patch_size) * (width // target_patch_size),
+ num_channels_latents * target_patch_size * target_patch_size)
+
+ return latents
+
+ @staticmethod
+ def _unpack_latents(latents, height, width, vae_scale_factor, patch_size=2, target_patch_size=1):
+ batch_size, num_patches, channels = latents.shape
+ scale = patch_size // target_patch_size
+
+ # VAE applies 8x compression on images but we must also account for packing which requires
+ # latent height and width to be divisible by 2.
+ height = (int(height) // (vae_scale_factor * patch_size))
+ width = (int(width) // (vae_scale_factor * patch_size))
+
+ latents = latents.view(
+ batch_size, height, width, channels // (scale * scale), scale, scale)
+ latents = latents.permute(0, 3, 1, 4, 2, 5)
+
+ latents = latents.reshape(batch_size, channels // (scale * scale), height * scale, width * scale)
+
+ return latents
+
+ @torch.inference_mode()
+ def __call__(
+ self,
+ prompt: Union[str, List[str]] = None,
+ height: Optional[int] = None,
+ width: Optional[int] = None,
+ num_inference_steps: int = 4,
+ total_substeps: int = 128,
+ final_step_size_scale: float = 0.5,
+ temperature: Union[float, str] = 'auto',
+ num_images_per_prompt: Optional[int] = 1,
+ generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
+ latents: Optional[torch.FloatTensor] = None,
+ prompt_embeds: Optional[torch.FloatTensor] = None,
+ prompt_embeds_mask: Optional[torch.Tensor] = None,
+ output_type: Optional[str] = "pil",
+ return_dict: bool = True,
+ attention_kwargs: Optional[Dict[str, Any]] = None,
+ callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
+ callback_on_step_end_tensor_inputs: List[str] = ["latents"],
+ max_sequence_length: int = 512,
+ ):
+ r"""
+ Function invoked when calling the pipeline for generation.
+
+ Args:
+ prompt (`str` or `List[str]`, *optional*):
+ The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
+ instead.
+ height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
+ The height in pixels of the generated image. This is set to 1024 by default for the best results.
+ width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
+ The width in pixels of the generated image. This is set to 1024 by default for the best results.
+ num_inference_steps (`int`, *optional*, defaults to 50):
+ The number of denoising steps.
+ total_substeps (`int`, *optional*, defaults to 128):
+ The total number of substeps for policy-based flow integration.
+ final_step_size_scale (`float`, *optional*, defaults to 0.5):
+ The scale for the final step size.
+ temperature (`float` or `"auto"`, *optional*, defaults to `"auto"`):
+ The tmperature parameter for the flow policy.
+ num_images_per_prompt (`int`, *optional*, defaults to 1):
+ The number of images to generate per prompt.
+ generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
+ One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
+ to make generation deterministic.
+ latents (`torch.FloatTensor`, *optional*):
+ Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
+ generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
+ tensor will be generated by sampling using the supplied random `generator`.
+ prompt_embeds (`torch.FloatTensor`, *optional*):
+ Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
+ provided, text embeddings will be generated from `prompt` input argument.
+ output_type (`str`, *optional*, defaults to `"pil"`):
+ The output format of the generate image. Choose between
+ [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
+ return_dict (`bool`, *optional*, defaults to `True`):
+ Whether or not to return a [`~pipelines.flux.FluxPipelineOutput`] instead of a plain tuple.
+ attention_kwargs (`dict`, *optional*):
+ A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
+ `self.processor` in
+ [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
+ callback_on_step_end (`Callable`, *optional*):
+ A function that calls at the end of each denoising steps during the inference. The function is called
+ with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
+ callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
+ `callback_on_step_end_tensor_inputs`.
+ callback_on_step_end_tensor_inputs (`List`, *optional*):
+ The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
+ will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
+ `._callback_tensor_inputs` attribute of your pipeline class.
+ max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`.
+
+ Returns:
+ [`~pipelines.flux.FluxPipelineOutput`] or `tuple`: [`~pipelines.flux.FluxPipelineOutput`] if `return_dict`
+ is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated
+ images.
+ """
+
+ height = height or self.default_sample_size * self.vae_scale_factor
+ width = width or self.default_sample_size * self.vae_scale_factor
+
+ # 1. Check inputs. Raise error if not correct
+ self.check_inputs(
+ prompt,
+ height,
+ width,
+ prompt_embeds=prompt_embeds,
+ prompt_embeds_mask=prompt_embeds_mask,
+ callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
+ max_sequence_length=max_sequence_length,
+ )
+
+ self._attention_kwargs = attention_kwargs
+ self._current_timestep = None
+ self._interrupt = False
+
+ # 2. Define call parameters
+ if prompt is not None and isinstance(prompt, str):
+ batch_size = 1
+ elif prompt is not None and isinstance(prompt, list):
+ batch_size = len(prompt)
+ else:
+ batch_size = prompt_embeds.shape[0]
+
+ device = self._execution_device
+
+ # 3. Prepare prompt embeddings
+ prompt_embeds, prompt_embeds_mask = self.encode_prompt(
+ prompt=prompt,
+ prompt_embeds=prompt_embeds,
+ prompt_embeds_mask=prompt_embeds_mask,
+ device=device,
+ num_images_per_prompt=num_images_per_prompt,
+ max_sequence_length=max_sequence_length,
+ )
+
+ # 4. Prepare latent variables
+ num_channels_latents = self.transformer.config.in_channels // 4
+ latents = self.prepare_latents(
+ batch_size * num_images_per_prompt,
+ num_channels_latents,
+ height,
+ width,
+ torch.float32,
+ device,
+ generator,
+ latents,
+ )
+ img_shapes = [[(1, height // self.vae_scale_factor // 2, width // self.vae_scale_factor // 2)]] * batch_size
+
+ # 5. Prepare timesteps
+ raw_timesteps, num_inference_substeps, total_substeps = retrieve_raw_timesteps(
+ num_inference_steps, total_substeps, final_step_size_scale)
+ image_seq_len = latents.shape[1]
+ mu = calculate_shift(
+ image_seq_len,
+ self.scheduler.config.get("base_image_seq_len", 256),
+ self.scheduler.config.get("max_image_seq_len", 4096),
+ self.scheduler.config.get("base_shift", 0.5),
+ self.scheduler.config.get("max_shift", 1.15),
+ )
+ timesteps, _ = retrieve_timesteps(
+ self.scheduler,
+ num_inference_steps,
+ device,
+ sigmas=raw_timesteps,
+ mu=mu,
+ )
+ assert len(timesteps) == total_substeps
+ self._num_timesteps = total_substeps
+
+ if self.attention_kwargs is None:
+ self._attention_kwargs = {}
+
+ txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() if prompt_embeds_mask is not None else None
+
+ # 6. Denoising loop
+ self.scheduler.set_begin_index(0)
+ timestep_id = 0
+ with self.progress_bar(total=num_inference_steps) as progress_bar:
+ for i in range(num_inference_steps):
+ if self.interrupt:
+ continue
+
+ t_src = timesteps[timestep_id]
+ sigma_t_src = t_src / self.scheduler.config.num_train_timesteps
+ is_final_step = i == (num_inference_steps - 1)
+
+ self._current_timestep = t_src
+
+ with self.transformer.cache_context("cond"):
+ denoising_output = self.transformer(
+ hidden_states=latents.to(dtype=self.transformer.dtype),
+ timestep=t_src.expand(latents.shape[0]) / 1000,
+ encoder_hidden_states_mask=prompt_embeds_mask,
+ encoder_hidden_states=prompt_embeds,
+ img_shapes=img_shapes,
+ txt_seq_lens=txt_seq_lens,
+ attention_kwargs=self.attention_kwargs,
+ )
+
+ # unpack and create policy
+ latents = self._unpack_latents(
+ latents, height, width, self.vae_scale_factor, target_patch_size=1)
+ if self.policy_type == 'GMFlow':
+ denoising_output = self._unpack_gm(
+ denoising_output, height, width, num_channels_latents, gm_patch_size=1)
+ denoising_output = {k: v.to(torch.float32) for k, v in denoising_output.items()}
+ policy = self.policy_class(
+ denoising_output, latents, sigma_t_src)
+ if not is_final_step:
+ if temperature == 'auto':
+ temperature = min(max(0.1 * (num_inference_steps - 1), 0), 1)
+ else:
+ assert isinstance(temperature, (float, int))
+ policy.temperature_(temperature)
+ elif self.policy_type == 'DX':
+ denoising_output = denoising_output[0]
+ denoising_output = self._unpack_latents(
+ denoising_output, height, width, self.vae_scale_factor, target_patch_size=1)
+ denoising_output = denoising_output.reshape(latents.size(0), -1, *latents.shape[1:])
+ denoising_output = denoising_output.to(torch.float32)
+ policy = self.policy_class(
+ denoising_output, latents, sigma_t_src)
+ else:
+ raise ValueError(f'Unknown policy type: {self.policy_type}.')
+
+ # compute the previous noisy sample x_t -> x_t-1
+ for _ in range(num_inference_substeps[i]):
+ t = timesteps[timestep_id]
+ sigma_t = t / self.scheduler.config.num_train_timesteps
+ u = policy.pi(latents, sigma_t)
+ latents = self.scheduler.step(u, t, latents, return_dict=False)[0]
+ timestep_id += 1
+
+ # repack
+ latents = self._pack_latents(
+ latents, latents.size(0), num_channels_latents,
+ 2 * (int(height) // (self.vae_scale_factor * 2)),
+ 2 * (int(width) // (self.vae_scale_factor * 2)),
+ patch_size=1)
+
+ if callback_on_step_end is not None:
+ callback_kwargs = {}
+ for k in callback_on_step_end_tensor_inputs:
+ callback_kwargs[k] = locals()[k]
+ callback_outputs = callback_on_step_end(self, i, t_src, callback_kwargs)
+
+ latents = callback_outputs.pop("latents", latents)
+ prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
+
+ progress_bar.update()
+
+ if XLA_AVAILABLE:
+ xm.mark_step()
+
+ self._current_timestep = None
+
+ if output_type == "latent":
+ image = latents
+ else:
+ latents = self._unpack_latents(latents, height, width, self.vae_scale_factor)[:, :, None]
+ latents_mean = (
+ torch.tensor(self.vae.config.latents_mean)
+ .view(1, self.vae.config.z_dim, 1, 1, 1)
+ .to(latents.device, latents.dtype)
+ )
+ latents_std = torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to(
+ latents.device, latents.dtype
+ )
+ latents = latents * latents_std + latents_mean
+ image = self.vae.decode(latents.to(self.vae.dtype), return_dict=False)[0][:, :, 0]
+ image = self.image_processor.postprocess(image, output_type=output_type)
+
+ # Offload all models
+ self.maybe_free_model_hooks()
+
+ if not return_dict:
+ return (image,)
+
+ return QwenImagePipelineOutput(images=image)
diff --git a/lakonlab/runner/__init__.py b/lakonlab/runner/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..f0fc9e04da8c5bd1ebf9ef92ed519733ce7a4c07
--- /dev/null
+++ b/lakonlab/runner/__init__.py
@@ -0,0 +1,4 @@
+from .hooks import *
+from .optimizer import *
+from .checkpoint import load_from_huggingface, load_from_tmp
+from .dynamic_iter_based_runner import DynamicIterBasedRunnerMod
diff --git a/lakonlab/runner/__pycache__/__init__.cpython-310.pyc b/lakonlab/runner/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..e4beb14182e51f4ecab4925f36cd582aa1946dce
Binary files /dev/null and b/lakonlab/runner/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/runner/__pycache__/checkpoint.cpython-310.pyc b/lakonlab/runner/__pycache__/checkpoint.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..67cf3ce6edf82f8606ffc74939dfca1d958a1ddd
Binary files /dev/null and b/lakonlab/runner/__pycache__/checkpoint.cpython-310.pyc differ
diff --git a/lakonlab/runner/__pycache__/dynamic_iter_based_runner.cpython-310.pyc b/lakonlab/runner/__pycache__/dynamic_iter_based_runner.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..94ac29cb6d224ff63da410fb8da390853b075c0e
Binary files /dev/null and b/lakonlab/runner/__pycache__/dynamic_iter_based_runner.cpython-310.pyc differ
diff --git a/lakonlab/runner/__pycache__/timer.cpython-310.pyc b/lakonlab/runner/__pycache__/timer.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..649f6b8370e895150a12105026d93267404040bb
Binary files /dev/null and b/lakonlab/runner/__pycache__/timer.cpython-310.pyc differ
diff --git a/lakonlab/runner/checkpoint.py b/lakonlab/runner/checkpoint.py
new file mode 100644
index 0000000000000000000000000000000000000000..ed390df0f874e0a667a70c8e8f4f1347867555b5
--- /dev/null
+++ b/lakonlab/runner/checkpoint.py
@@ -0,0 +1,534 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import os
+import os.path as osp
+import logging
+import time
+import tempfile
+import subprocess
+import uuid
+import shutil
+import re
+import torch
+import torch.nn as nn
+import torch.distributed as dist
+import mmcv
+from typing import Union, Callable, Optional, List
+from collections import OrderedDict
+from tempfile import TemporaryDirectory
+from torch.optim import Optimizer
+from torch.distributed.tensor import DTensor
+from torch.distributed.checkpoint.state_dict import set_model_state_dict, StateDictOptions, get_optimizer_state_dict
+from torch.distributed.fsdp import StateDictType, FullStateDictConfig, FullyShardedDataParallel as FSDP
+from safetensors.torch import load_file, load
+from diffusers.utils.hub_utils import _get_checkpoint_shard_files
+from mmcv.runner import CheckpointLoader, get_dist_info, _load_checkpoint
+from mmcv.parallel import is_module_wrapper
+from lakonlab.utils import download_from_huggingface, rgetattr
+from lakonlab.utils.io_utils import S3Backend, TMP_DIR
+from lakonlab.parallel import FSDP2Wrapper
+
+
+def fsdp_load_full_state_dict(module: FSDP,
+ full_sd: dict,
+ prefix: str,
+ fsdp_missing_keys: List[str],
+ err_msg: List[str]) -> None:
+ flat_p = module._flat_param
+ if flat_p is not None:
+ for param_info, shape, shard_info in zip(flat_p._param_infos, flat_p._shapes, flat_p._shard_param_infos):
+ k = f'{prefix}{param_info.module_name}.{param_info.param_name}'
+ if k not in full_sd:
+ fsdp_missing_keys.append(k)
+ continue
+ t = full_sd[k]
+ if t.shape != shape:
+ err_msg.append(
+ f'size mismatch for {k}: copying a param with shape {t.shape} from checkpoint, '
+ f'the shape in current model is {shape}.')
+ elif shard_info.in_shard:
+ flat_p[shard_info.offset_in_shard:shard_info.offset_in_shard + shard_info.numel_in_shard].data.copy_(
+ t.view(-1)[shard_info.intra_param_start_idx:shard_info.intra_param_end_idx + 1])
+ del full_sd[k]
+
+
+def load_full_state_dict(module: nn.Module,
+ state_dict: Union[dict, OrderedDict],
+ strict: bool = False,
+ logger: Optional[logging.Logger] = None,
+ assign: bool = False) -> None:
+ unexpected_keys: List[str] = []
+ all_missing_keys: List[str] = []
+ err_msg: List[str] = []
+ fsdp_missing_keys: List[str] = []
+
+ metadata = getattr(state_dict, '_metadata', None)
+ state_dict = state_dict.copy() # type: ignore
+ if metadata is not None:
+ state_dict._metadata = metadata # type: ignore
+
+ # use _load_from_state_dict to enable checkpoint version control
+ def load(module, prefix=''):
+ # Load full states to sharded FSDP1
+ if isinstance(module, FSDP):
+ fsdp_load_full_state_dict(
+ module, state_dict, prefix, fsdp_missing_keys, err_msg)
+ # recursively check parallel module in case that the model has a
+ # complicated structure, e.g., nn.Module(nn.Module(DDP))
+ if is_module_wrapper(module):
+ module = module.module
+ local_metadata = {} if metadata is None else metadata.get(prefix[:-1], {})
+ if assign:
+ local_metadata['assign_to_params_buffers'] = assign
+ module._load_from_state_dict(
+ state_dict,
+ prefix,
+ local_metadata,
+ True,
+ all_missing_keys,
+ unexpected_keys,
+ err_msg)
+ for name, child in module._modules.items():
+ if child is not None:
+ load(child, prefix + name + '.')
+
+ load(module)
+ # break load->load reference cycle
+ load = None # type: ignore
+
+ # ignore "num_batches_tracked" of BN layers
+ missing_keys = fsdp_missing_keys.copy()
+ for key in all_missing_keys:
+ if not rgetattr(module, key + '._fsdp_flattened', False):
+ missing_keys.append(key)
+
+ missing_keys = [
+ key for key in missing_keys if 'num_batches_tracked' not in key
+ ]
+
+ if unexpected_keys:
+ err_msg.append('unexpected key in source '
+ f'state_dict: {", ".join(unexpected_keys)}\n')
+ if missing_keys:
+ err_msg.append(
+ f'missing keys in source state_dict: {", ".join(missing_keys)}\n')
+
+ rank, _ = get_dist_info()
+ if len(err_msg) > 0 and rank == 0:
+ err_msg.insert(
+ 0, 'The model and loaded state dict do not match exactly\n')
+ err_msg = '\n'.join(err_msg) # type: ignore
+ if strict:
+ raise RuntimeError(err_msg)
+ elif logger is not None:
+ logger.warning(err_msg)
+ else:
+ print(err_msg)
+
+
+def exists_ckpt(filename):
+ if not filename:
+ return False
+ loader_name = CheckpointLoader._get_checkpoint_loader(filename).__name__[10:]
+ if loader_name == 'local':
+ return os.path.exists(filename)
+ elif loader_name == 'tmp':
+ src_file = filename[4:]
+ return os.path.exists(src_file)
+ elif loader_name == 's3':
+ return S3Backend().exists(filename)
+ else:
+ raise NotImplementedError()
+
+
+@CheckpointLoader.register_scheme(prefixes='s3://', force=True)
+def load_from_s3(filename, map_location=None):
+ ext = os.path.splitext(filename)[-1].lower()
+
+ rank, ws = get_dist_info()
+ if ws > 1:
+ local_rank = dist.get_node_local_rank()
+ else:
+ local_rank = 0
+
+ # get the temporary file path
+ if rank == 0:
+ tmp_file = os.path.join(TMP_DIR, str(uuid.uuid4()) + ext)
+ else:
+ tmp_file = None
+ if ws > 1:
+ object_list = [tmp_file]
+ dist.broadcast_object_list(object_list, src=0)
+ tmp_file = object_list[0]
+
+ # download the file to temp dir
+ if local_rank == 0:
+ cmd = ['aws', 's3', 'cp', filename, tmp_file]
+ subprocess.run(cmd, check=True, stdout=subprocess.DEVNULL)
+ if ws > 1:
+ dist.barrier()
+
+ # load the temporary file
+ if ext == '.txt':
+ # if the file is a text file, it contains the name to the actual checkpoint file
+ with open(tmp_file, 'r') as f:
+ _filename = f.read().strip()
+ filename = os.path.join(
+ os.path.dirname(filename), _filename) # get the actual checkpoint file path
+ # remove the temporary file
+ if ws > 1:
+ dist.barrier()
+ if local_rank == 0:
+ os.remove(tmp_file)
+ return load_from_s3(filename, map_location=map_location)
+
+ elif ext == '.safetensors':
+ ckpt = load_file(tmp_file, device=map_location)
+ else:
+ ckpt = torch.load(tmp_file, map_location=map_location)
+
+ # remove the temporary file
+ if ws > 1:
+ dist.barrier()
+ if local_rank == 0:
+ os.remove(tmp_file)
+
+ return ckpt
+
+
+@CheckpointLoader.register_scheme(prefixes='tmp:')
+def load_from_tmp(filename, map_location=None):
+ src_file = filename[4:]
+ assert os.path.exists(src_file)
+ ext = os.path.splitext(src_file)[-1].lower()
+ rank, ws = get_dist_info()
+ if ws > 1:
+ local_rank = dist.get_node_local_rank()
+ else:
+ local_rank = 0
+
+ # get the temporary file path
+ if rank == 0:
+ tmp_file = os.path.join(TMP_DIR, str(uuid.uuid4()) + ext)
+ else:
+ tmp_file = None
+ if ws > 1:
+ object_list = [tmp_file]
+ dist.broadcast_object_list(object_list, src=0)
+ tmp_file = object_list[0]
+
+ # copy the file to temp dir
+ if local_rank == 0:
+ shutil.copy(src_file, tmp_file)
+ if ws > 1:
+ dist.barrier()
+
+ # load the temporary file
+ if ext == '.safetensors':
+ ckpt = load_file(tmp_file, device=map_location)
+ else:
+ ckpt = torch.load(tmp_file, map_location=map_location)
+
+ # remove the temporary file
+ if ws > 1:
+ dist.barrier()
+ if local_rank == 0:
+ os.remove(tmp_file)
+
+ return ckpt
+
+
+@CheckpointLoader.register_scheme(prefixes='huggingface://')
+def load_from_huggingface(filename, map_location=None):
+ cached_file = download_from_huggingface(filename)
+ if cached_file.endswith('.index.json'): # sharded checkpoint
+ filename = filename.replace('huggingface://', '').split('/')
+ repo_id = '/'.join(filename[:2])
+ repo_subfolder = '/'.join(filename[2:-1])
+ is_dist = dist.is_available() and dist.is_initialized()
+ if is_dist:
+ local_rank = dist.get_node_local_rank()
+ else:
+ local_rank = 0
+ if local_rank == 0:
+ sharded_cached_files = _get_checkpoint_shard_files(
+ repo_id,
+ cached_file,
+ subfolder=repo_subfolder)[0]
+ if is_dist:
+ dist.barrier()
+ if local_rank > 0:
+ sharded_cached_files = _get_checkpoint_shard_files(
+ repo_id,
+ cached_file,
+ subfolder=repo_subfolder)[0]
+ ckpt = OrderedDict()
+ for sharded_cached_file in sharded_cached_files:
+ ext = os.path.splitext(sharded_cached_file)[-1].lower()
+ if ext == '.safetensors':
+ ckpt.update(load_file(sharded_cached_file, device=map_location))
+ else:
+ ckpt.update(torch.load(sharded_cached_file, map_location=map_location))
+ return ckpt
+ else:
+ ext = os.path.splitext(cached_file)[-1].lower()
+ if ext == '.safetensors':
+ return load_file(cached_file, device=map_location)
+ else:
+ return torch.load(cached_file, map_location=map_location)
+
+
+@CheckpointLoader.register_scheme(prefixes='', force=True)
+def load_from_local(filename, map_location=None):
+ filename = osp.expanduser(filename)
+ if not osp.isfile(filename):
+ raise FileNotFoundError(f'{filename} can not be found.')
+ ext = os.path.splitext(filename)[-1].lower()
+ if ext == '.safetensors':
+ with open(filename, "rb") as f: # load_file may fail with FUSE/NFS mmap
+ ckpt = load(f.read())
+ if map_location is not None:
+ for k in ckpt:
+ ckpt[k] = ckpt[k].to(map_location)
+ else:
+ ckpt = torch.load(filename, map_location=map_location)
+ return ckpt
+
+
+def load_checkpoint(model: torch.nn.Module,
+ filename: str,
+ map_location: Union[str, Callable, None] = None,
+ strict: bool = False,
+ logger: Optional[logging.Logger] = None,
+ revise_keys: list = [(r'^module\.', '')],
+ assign: bool = False) -> Union[dict, OrderedDict]:
+ checkpoint = _load_checkpoint(filename, map_location, logger)
+ # OrderedDict is a subclass of dict
+ if not isinstance(checkpoint, dict):
+ raise RuntimeError(
+ f'No state_dict found in checkpoint file {filename}')
+ # get state_dict from checkpoint
+ if 'state_dict' in checkpoint:
+ state_dict = checkpoint['state_dict']
+ else:
+ state_dict = checkpoint
+
+ # strip prefix of state_dict
+ metadata = getattr(state_dict, '_metadata', OrderedDict())
+ for p, r in revise_keys:
+ state_dict = OrderedDict(
+ {re.sub(p, r, k): v
+ for k, v in state_dict.items()})
+ # Keep metadata in state_dict
+ state_dict._metadata = metadata
+
+ # load state_dict
+ if isinstance(model, FSDP2Wrapper): # FSDP2
+ for name, submodule in model.module._modules.items():
+ submodule_state_dict = {
+ k[len(name) + 1:]: v for k, v in state_dict.items() if k.startswith(name)}
+ set_model_state_dict(
+ model=submodule,
+ model_state_dict=submodule_state_dict,
+ options=StateDictOptions(
+ full_state_dict=True,
+ broadcast_from_rank0=False,
+ strict=strict))
+ else: # FSDP1, DDP, or non-distributed model
+ load_full_state_dict(model, state_dict, strict, logger, assign)
+ return checkpoint
+
+
+def _save_to_state_dict(module, destination, prefix, keep_vars, trainable_only=False, cpu_offload=False):
+ for name, param in module._parameters.items():
+ if param is not None and (not trainable_only or param.requires_grad):
+ if not keep_vars:
+ param = param.detach()
+ if isinstance(param, DTensor):
+ param = param.full_tensor()
+ if torch.distributed.get_rank() == 0: # only save the full tensor on rank 0
+ if cpu_offload:
+ param = param.cpu()
+ destination[prefix + name] = param
+ else:
+ if cpu_offload:
+ param = param.cpu()
+ destination[prefix + name] = param
+ for name, buf in module._buffers.items():
+ if buf is not None:
+ if not keep_vars:
+ buf = buf.detach()
+ if cpu_offload:
+ buf = buf.cpu()
+ destination[prefix + name] = buf
+
+
+def get_state_dict(module,
+ destination=None,
+ prefix='',
+ keep_vars=False,
+ trainable_only=False,
+ cpu_offload=True):
+ if isinstance(module, FSDP): # FSDP1
+ if trainable_only and len(module.params) == 1 and not module.params[0].requires_grad:
+ return destination # skip frozen module
+
+ with FSDP.state_dict_type(
+ module,
+ StateDictType.FULL_STATE_DICT,
+ FullStateDictConfig(offload_to_cpu=cpu_offload, rank0_only=True)):
+ module.state_dict(destination=destination, prefix=prefix, keep_vars=keep_vars)
+
+ else: # FSDP2, DDP, or non-distributed model
+ # recursively check parallel module in case that the model has a
+ # complicated structure, e.g., nn.Module(nn.Module(DDP))
+ if is_module_wrapper(module):
+ module = module.module
+
+ # below is the same as torch.nn.Module.state_dict() except for the trainable_only argument
+ if destination is None:
+ destination = OrderedDict()
+ destination._metadata = OrderedDict() # type: ignore
+
+ local_metadata = dict(version=module._version)
+ if hasattr(destination, '_metadata'):
+ destination._metadata[prefix[:-1]] = local_metadata
+
+ for hook in module._state_dict_pre_hooks.values():
+ hook(module, prefix, keep_vars)
+ _save_to_state_dict(
+ module, destination, prefix, keep_vars, trainable_only=trainable_only, cpu_offload=cpu_offload)
+ for name, child in module._modules.items():
+ if child is not None:
+ get_state_dict(
+ child, destination, prefix + name + '.',
+ keep_vars=keep_vars, trainable_only=trainable_only, cpu_offload=cpu_offload)
+ for hook in module._state_dict_hooks.values():
+ hook_result = hook(module, destination, prefix, local_metadata)
+ if not getattr(hook, '_from_public_api', False):
+ if hook_result is not None:
+ destination = hook_result
+ else:
+ if hook_result is not None:
+ raise RuntimeError('state_dict post-hook must return None')
+
+ return destination
+
+
+def get_optim_state_dict(model, optimizer, bf16=False):
+ optim_state_dict = get_optimizer_state_dict(
+ model=model,
+ optimizers=optimizer,
+ options=StateDictOptions(
+ full_state_dict=True,
+ cpu_offload=True))
+ if 'state' in optim_state_dict:
+ for state_name, state in optim_state_dict['state'].items():
+ new_state = dict()
+ for k, v in state.items():
+ if bf16 and isinstance(v, torch.Tensor) and v.dtype == torch.float32 and v.numel() > 1:
+ v = v.to(dtype=torch.bfloat16)
+ new_state[k] = v
+ optim_state_dict['state'][state_name] = new_state
+ return optim_state_dict
+
+
+def write_checkpoint_to_file(checkpoint, filepath, create_symlink=False, after_save_hook=None):
+ if filepath.startswith('pavi://'):
+ try:
+ from pavi import modelcloud
+ from pavi.exception import NodeNotFoundError
+ except ImportError:
+ raise ImportError(
+ 'Please install pavi to load checkpoint from modelcloud.')
+ model_path = filepath[7:]
+ root = modelcloud.Folder()
+ model_dir, model_name = osp.split(model_path)
+ try:
+ model = modelcloud.get(model_dir)
+ except NodeNotFoundError:
+ model = root.create_training_model(model_dir)
+ with TemporaryDirectory() as tmp_dir:
+ checkpoint_file = osp.join(tmp_dir, model_name)
+ with open(checkpoint_file, 'wb') as f:
+ torch.save(checkpoint, f)
+ f.flush()
+ model.create_file(checkpoint_file, name=model_name)
+
+ elif filepath.startswith('s3://'):
+ ext = os.path.splitext(filepath)[-1].lower()
+ with tempfile.NamedTemporaryFile(dir=TMP_DIR, suffix=ext, delete=False) as tmp:
+ cached_file = tmp.name
+ torch.save(checkpoint, tmp)
+ tmp.flush()
+ try:
+ cmd = ['aws', 's3', 'cp', cached_file, filepath]
+ subprocess.run(cmd, check=True, stdout=subprocess.DEVNULL)
+ finally:
+ os.remove(cached_file)
+
+ if create_symlink:
+ # S3 does not support real symlinks, so we create a 'latest.txt'
+ # containing the relative path of the latest checkpoint
+ dst_file = osp.join(osp.dirname(filepath), 'latest.txt')
+ S3Backend().put_text(osp.basename(filepath), dst_file)
+
+ else:
+ mmcv.mkdir_or_exist(osp.dirname(filepath))
+ # immediately flush buffer
+ with open(filepath, 'wb') as f:
+ torch.save(checkpoint, f)
+ f.flush()
+
+ if create_symlink:
+ dst_file = osp.join(osp.dirname(filepath), 'latest.pth')
+ mmcv.symlink(osp.basename(filepath), dst_file)
+
+ if after_save_hook is not None:
+ after_save_hook()
+
+
+def get_checkpoint(model,
+ optimizer=None,
+ loss_scaler=None,
+ meta=None,
+ trainable_only=False,
+ fp16=False,
+ fp16_ema=False,
+ bf16_optim=False):
+ if meta is None:
+ meta = {}
+ elif not isinstance(meta, dict):
+ raise TypeError(f'meta must be a dict or None, but got {type(meta)}')
+ meta.update(mmcv_version=mmcv.__version__, time=time.asctime())
+
+ if is_module_wrapper(model):
+ model = model.module
+
+ if hasattr(model, 'CLASSES') and model.CLASSES is not None:
+ # save class name to the meta
+ meta.update(CLASSES=model.CLASSES)
+
+ checkpoint = {
+ 'meta': meta,
+ 'state_dict': get_state_dict(model, trainable_only=trainable_only, cpu_offload=True)}
+ if fp16 or fp16_ema:
+ for k, v in checkpoint['state_dict'].items():
+ if ((fp16 and '_ema.' not in k and '_ema2.' not in k) or (fp16_ema and ('_ema.' in k or '_ema2.' in k))) \
+ and v.dtype == torch.float32:
+ checkpoint['state_dict'][k] = v.half()
+
+ # save optimizer state dict in the checkpoint
+ if isinstance(optimizer, Optimizer):
+ checkpoint['optimizer'] = get_optim_state_dict(model, optimizer, bf16_optim)
+ elif isinstance(optimizer, dict):
+ checkpoint['optimizer'] = {}
+ for name, optim in optimizer.items():
+ submodule = getattr(model, name)
+ checkpoint['optimizer'][name] = get_optim_state_dict(submodule, optim, bf16_optim)
+
+ # save loss scaler for mixed-precision (FP16) training
+ if loss_scaler is not None:
+ checkpoint['loss_scaler'] = loss_scaler.state_dict()
+
+ return checkpoint
diff --git a/lakonlab/runner/dynamic_iter_based_runner.py b/lakonlab/runner/dynamic_iter_based_runner.py
new file mode 100644
index 0000000000000000000000000000000000000000..e03849207608ecb6dc11b9ed47054f2959be1e81
--- /dev/null
+++ b/lakonlab/runner/dynamic_iter_based_runner.py
@@ -0,0 +1,219 @@
+import gc
+import copy
+import warnings
+import time
+import os.path as osp
+import threading
+import torch
+import mmcv
+from typing import Any, Dict
+from torch.optim import Optimizer
+from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
+from torch.distributed.checkpoint.state_dict import set_optimizer_state_dict, StateDictOptions
+from mmcv.runner import RUNNERS, get_dist_info, get_host_info
+from mmgen.core.runners.dynamic_iterbased_runner import DynamicIterBasedRunner, IterLoader
+
+from .checkpoint import get_checkpoint, load_checkpoint, write_checkpoint_to_file
+from lakonlab.utils import rgetattr, gc_context
+
+
+_last_save_thread: threading.Thread | None = None
+_last_save_lock = threading.Lock() # protects _last_save_thread
+
+
+def strip_initial_lr(opt_state: Dict[str, Any]) -> Dict[str, Any]:
+
+ def _strip_in_single_optimizer(sd: Dict[str, Any]) -> None:
+ for group in sd.get("param_groups", []):
+ group.pop("initial_lr", None)
+
+ # case 1: top‑level looks like a regular optimizer state‑dict
+ if "param_groups" in opt_state:
+ _strip_in_single_optimizer(opt_state)
+ return opt_state
+
+ # case 2: multi‑optimizer checkpoint {name -> state‑dict}
+ for key, sub_sd in opt_state.items():
+ if isinstance(sub_sd, dict) and "param_groups" in sub_sd:
+ _strip_in_single_optimizer(sub_sd)
+
+ return opt_state
+
+
+@RUNNERS.register_module()
+class DynamicIterBasedRunnerMod(DynamicIterBasedRunner):
+
+ def __init__(self,
+ *args,
+ ckpt_trainable_only=False,
+ ckpt_fp16=False,
+ ckpt_fp16_ema=False,
+ ckpt_bf16_optim=False,
+ gc_interval=-1,
+ **kwargs):
+ super(DynamicIterBasedRunnerMod, self).__init__(*args, **kwargs)
+ self.ckpt_trainable_only = ckpt_trainable_only
+ self.ckpt_fp16 = ckpt_fp16
+ self.ckpt_fp16_ema = ckpt_fp16_ema
+ self.ckpt_bf16_optim = ckpt_bf16_optim
+ self.gc_interval = gc_interval
+ self.manual_gc = isinstance(gc_interval, int) and gc_interval > 0
+
+ def run(self, data_loaders, workflow, max_iters=None, **kwargs):
+ assert isinstance(data_loaders, list)
+ assert mmcv.is_list_of(workflow, tuple)
+ assert len(data_loaders) == len(workflow)
+ if max_iters is not None:
+ warnings.warn(
+ 'setting max_iters in run is deprecated, '
+ 'please set max_iters in runner_config', DeprecationWarning)
+ self._max_iters = max_iters
+ assert self._max_iters is not None, (
+ 'max_iters must be specified during instantiation')
+
+ work_dir = self.work_dir if self.work_dir is not None else 'NONE'
+ self.logger.info('Start running, host: %s, work_dir: %s',
+ get_host_info(), work_dir)
+ self.logger.info('workflow: %s, max: %d iters', workflow,
+ self._max_iters)
+ self.call_hook('before_run')
+
+ iter_loaders = [IterLoader(x, self) for x in data_loaders]
+
+ self.call_hook('before_epoch')
+
+ while self.iter < self._max_iters:
+ for i, flow in enumerate(workflow):
+ with gc_context(enable=not self.manual_gc):
+ self._inner_iter = 0
+ mode, iters = flow
+ if not isinstance(mode, str) or not hasattr(self, mode):
+ raise ValueError(
+ 'runner has no method named "{}" to run a workflow'.
+ format(mode))
+ iter_runner = getattr(self, mode)
+ for _ in range(iters):
+ if mode == 'train' and self.iter >= self._max_iters:
+ break
+ if self.manual_gc and self._inner_iter % self.gc_interval == 0:
+ gc.collect()
+ iter_runner(iter_loaders[i], **kwargs)
+
+ time.sleep(1) # wait for some hooks like loggers to finish
+ self.call_hook('after_epoch')
+ self.call_hook('after_run')
+
+ def save_checkpoint(self,
+ out_dir,
+ filename_tmpl='iter_{}.pth',
+ meta=None,
+ save_optimizer=True,
+ create_symlink=True,
+ after_save_hook=None,
+ asynchronous=False):
+ if meta is None:
+ meta = dict(iter=self.iter + 1, epoch=self.epoch + 1)
+ elif isinstance(meta, dict):
+ meta.update(iter=self.iter + 1, epoch=self.epoch + 1)
+ else:
+ raise TypeError(
+ f'meta should be a dict or None, but got {type(meta)}')
+ if self.meta is not None:
+ meta.update(self.meta)
+
+ filename = filename_tmpl.format(self.iter + 1)
+ filepath = osp.join(out_dir, filename)
+ optimizer = self.optimizer if save_optimizer else None
+ _loss_scaler = self.loss_scaler if self.with_fp16_grad_scaler else None
+ checkpoint = get_checkpoint(
+ self.model,
+ optimizer=optimizer,
+ loss_scaler=_loss_scaler,
+ meta=meta,
+ trainable_only=self.ckpt_trainable_only,
+ fp16=self.ckpt_fp16,
+ fp16_ema=self.ckpt_fp16_ema,
+ bf16_optim=self.ckpt_bf16_optim)
+
+ rank, _ = get_dist_info()
+ if rank == 0:
+ global _last_save_thread
+ with _last_save_lock:
+ if _last_save_thread is not None and _last_save_thread.is_alive():
+ print('Waiting for the previous write to finish...')
+ _last_save_thread.join() # wait for the previous write
+
+ if asynchronous:
+ _last_save_thread = threading.Thread(
+ target=write_checkpoint_to_file,
+ args=(copy.deepcopy(checkpoint), filepath, create_symlink, after_save_hook),
+ daemon=True)
+ _last_save_thread.start()
+
+ else:
+ write_checkpoint_to_file(
+ checkpoint,
+ filepath,
+ create_symlink=create_symlink,
+ after_save_hook=after_save_hook)
+
+ def load_checkpoint(self,
+ filename,
+ map_location='cpu',
+ strict=False,
+ revise_keys=[(r'^module.', '')]):
+ return load_checkpoint(
+ self.model,
+ filename,
+ map_location,
+ strict,
+ self.logger,
+ revise_keys=revise_keys)
+
+ def resume(self,
+ checkpoint,
+ resume_optimizer=True,
+ resume_loss_scaler=True,
+ map_location='default'):
+ if map_location == 'default':
+ device_id = torch.cuda.current_device()
+ checkpoint = self.load_checkpoint(
+ checkpoint,
+ map_location=lambda storage, loc: storage.cuda(device_id))
+ else:
+ checkpoint = self.load_checkpoint(
+ checkpoint, map_location=map_location)
+
+ self._epoch = checkpoint['meta']['epoch']
+ self._iter = checkpoint['meta']['iter']
+ self._inner_iter = checkpoint['meta']['iter']
+
+ if 'optimizer' in checkpoint and resume_optimizer:
+ optimizer_sd = strip_initial_lr(checkpoint['optimizer'])
+ if isinstance(self.optimizer, Optimizer):
+ set_optimizer_state_dict(
+ model=self.model,
+ optimizers=self.optimizer,
+ optim_state_dict=optimizer_sd,
+ options=StateDictOptions(
+ full_state_dict=isinstance(self.model, FSDP),
+ broadcast_from_rank0=False))
+ elif isinstance(self.optimizer, dict):
+ for k in self.optimizer.keys():
+ m = rgetattr(self.model, k)
+ set_optimizer_state_dict(
+ model=m,
+ optimizers=self.optimizer[k],
+ optim_state_dict=optimizer_sd[k],
+ options=StateDictOptions(
+ full_state_dict=isinstance(m, FSDP),
+ broadcast_from_rank0=False))
+ else:
+ raise TypeError(
+ 'Optimizer should be dict or torch.optim.Optimizer '
+ f'but got {type(self.optimizer)}')
+
+ if 'loss_scaler' in checkpoint and resume_loss_scaler and hasattr(self, 'loss_scaler'):
+ self.loss_scaler.load_state_dict(checkpoint['loss_scaler'])
+
+ self.logger.info(f'resumed from epoch: {self.epoch}, iter {self.iter}')
diff --git a/lakonlab/runner/hooks/__init__.py b/lakonlab/runner/hooks/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..ab8cd471a72b8b24000f59d4e6bba6d7c6a135eb
--- /dev/null
+++ b/lakonlab/runner/hooks/__init__.py
@@ -0,0 +1,5 @@
+from .ema_hook import ExponentialMovingAverageHookMod
+from .checkpoint import CheckpointHook
+from .logger import *
+
+__all__ = ['CheckpointHook', 'ExponentialMovingAverageHookMod']
diff --git a/lakonlab/runner/hooks/__pycache__/__init__.cpython-310.pyc b/lakonlab/runner/hooks/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..cafe571cd0f42845a9b1cdd70b6573b6e8f73f8a
Binary files /dev/null and b/lakonlab/runner/hooks/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/runner/hooks/__pycache__/checkpoint.cpython-310.pyc b/lakonlab/runner/hooks/__pycache__/checkpoint.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..8ca79e707dd2dd1f2c155617410d1a5acdd15f71
Binary files /dev/null and b/lakonlab/runner/hooks/__pycache__/checkpoint.cpython-310.pyc differ
diff --git a/lakonlab/runner/hooks/__pycache__/ema_hook.cpython-310.pyc b/lakonlab/runner/hooks/__pycache__/ema_hook.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..9592a2098bba0b28db93dc074b5a1dcde3f50b1e
Binary files /dev/null and b/lakonlab/runner/hooks/__pycache__/ema_hook.cpython-310.pyc differ
diff --git a/lakonlab/runner/hooks/checkpoint.py b/lakonlab/runner/hooks/checkpoint.py
new file mode 100644
index 0000000000000000000000000000000000000000..ead25f33aa924d990f8b1a45733415849613df52
--- /dev/null
+++ b/lakonlab/runner/hooks/checkpoint.py
@@ -0,0 +1,91 @@
+import copy
+from mmcv.runner.hooks import CheckpointHook as _CheckpointHook
+from mmcv.runner import HOOKS
+from mmcv.runner.dist_utils import allreduce_params, get_dist_info
+from lakonlab.parallel import FSDPWrapper
+
+
+@HOOKS.register_module(force=True)
+class CheckpointHook(_CheckpointHook):
+
+ def __init__(self,
+ interval: int = -1,
+ must_save_interval: int = 1000000000,
+ **kwargs):
+ super().__init__(interval=interval, **kwargs)
+ self.must_save_interval = must_save_interval
+
+ def after_train_epoch(self, runner):
+ if not self.by_epoch:
+ return
+
+ if self.every_n_epochs(runner, self.interval) \
+ or self.every_n_epochs(runner, self.must_save_interval) \
+ or (self.save_last and self.is_last_epoch(runner)):
+ runner.logger.info(
+ f'Saving checkpoint at {runner.epoch + 1} epochs')
+ if self.sync_buffer:
+ allreduce_params(runner.model.buffers())
+ self._save_checkpoint(runner, asynchronous=not self.is_last_epoch(runner))
+
+ def after_train_iter(self, runner):
+ if self.by_epoch:
+ return
+
+ if self.every_n_iters(runner, self.interval) \
+ or self.every_n_iters(runner, self.must_save_interval) \
+ or (self.save_last and self.is_last_iter(runner)):
+ runner.logger.info(
+ f'Saving checkpoint at {runner.iter + 1} iterations')
+ if self.sync_buffer:
+ allreduce_params(runner.model.buffers())
+ self._save_checkpoint(runner, asynchronous=not self.is_last_iter(runner))
+
+ def _save_checkpoint(self, runner, asynchronous=False):
+ rank, _ = get_dist_info()
+ # in FSDP we need to call save_checkpoint on all ranks
+ if rank == 0 or isinstance(runner.model, FSDPWrapper):
+ if self.max_keep_ckpts > 0:
+ # remove other checkpoints
+ cur_progress = copy.deepcopy(runner.epoch) if self.by_epoch else copy.deepcopy(runner.iter)
+
+ def after_save_hook():
+ if self.by_epoch:
+ name = 'epoch_{}.pth'
+ current_ckpt = cur_progress + 1
+ else:
+ name = 'iter_{}.pth'
+ current_ckpt = cur_progress + 1
+ redundant_ckpts = [i for i in range(
+ current_ckpt - self.max_keep_ckpts * self.interval, 0,
+ -self.interval) if i % self.must_save_interval != 0]
+ filename_tmpl = self.args.get('filename_tmpl', name)
+ for _step in redundant_ckpts:
+ ckpt_path = self.file_client.join_path(
+ self.out_dir, filename_tmpl.format(_step))
+ if self.file_client.isfile(ckpt_path):
+ self.file_client.remove(ckpt_path)
+ else:
+ break
+
+ else:
+ after_save_hook = None
+
+ runner.save_checkpoint(
+ self.out_dir,
+ save_optimizer=self.save_optimizer,
+ after_save_hook=after_save_hook,
+ asynchronous=asynchronous,
+ **self.args)
+
+ if rank == 0:
+ if runner.meta is not None:
+ if self.by_epoch:
+ cur_ckpt_filename = self.args.get(
+ 'filename_tmpl', 'epoch_{}.pth').format(runner.epoch + 1)
+ else:
+ cur_ckpt_filename = self.args.get(
+ 'filename_tmpl', 'iter_{}.pth').format(runner.iter + 1)
+ runner.meta.setdefault('hook_msgs', dict())
+ runner.meta['hook_msgs']['last_ckpt'] = self.file_client.join_path(
+ self.out_dir, cur_ckpt_filename)
diff --git a/lakonlab/runner/hooks/ema_hook.py b/lakonlab/runner/hooks/ema_hook.py
new file mode 100644
index 0000000000000000000000000000000000000000..e86210028ac31992b197910c7755c1ca4840a915
--- /dev/null
+++ b/lakonlab/runner/hooks/ema_hook.py
@@ -0,0 +1,133 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import mmcv
+import torch
+
+from copy import deepcopy
+try:
+ from torch.distributed.fsdp import FSDPModule
+except:
+ FSDPModule = None
+from mmcv.parallel import is_module_wrapper
+from mmcv.runner import HOOKS
+from mmgen.core import ExponentialMovingAverageHook
+from lakonlab.utils import rgetattr, rhasattr
+
+
+def get_ori_key(key):
+ ori_key = key.split('.')
+ if ori_key[0].endswith('_ema'):
+ ori_key[0] = ori_key[0][:-4]
+ elif ori_key[0].endswith('_ema2'):
+ ori_key[0] = ori_key[0][:-5]
+ else:
+ raise ValueError(
+ f'Invalid module key {key}, it should be in the format of '
+ '_ema or _ema2, but got {ori_key[0]}')
+ ori_key = '.'.join(ori_key)
+ return ori_key
+
+
+@HOOKS.register_module()
+class ExponentialMovingAverageHookMod(ExponentialMovingAverageHook):
+
+ _registered_momentum_updaters = ['rampup', 'fixed', 'karras']
+
+ def __init__(self,
+ module_keys,
+ trainable_only=True,
+ interp_mode='lerp',
+ interp_cfg=None,
+ interval=-1,
+ start_iter=0,
+ momentum_policy='fixed',
+ momentum_cfg=None):
+ super(ExponentialMovingAverageHook, self).__init__()
+ self.trainable_only = trainable_only
+ # check args
+ assert interp_mode in self._registered_interp_funcs, (
+ 'Supported '
+ f'interpolation functions are {self._registered_interp_funcs}, '
+ f'but got {interp_mode}')
+
+ assert momentum_policy in self._registered_momentum_updaters, (
+ 'Supported momentum policy are'
+ f'{self._registered_momentum_updaters},'
+ f' but got {momentum_policy}')
+
+ assert isinstance(module_keys, str) or mmcv.is_tuple_of(
+ module_keys, str)
+ self.module_keys = (module_keys, ) if isinstance(module_keys,
+ str) else module_keys
+ # sanity check for the format of module keys
+ for k in self.module_keys:
+ module_name = k.split('.')[0]
+ assert module_name.endswith('_ema') or module_name.endswith('_ema2')
+ self.interp_mode = interp_mode
+ self.interp_cfg = dict() if interp_cfg is None else deepcopy(
+ interp_cfg)
+ self.interval = interval
+ self.start_iter = start_iter
+
+ assert hasattr(
+ self, interp_mode
+ ), f'Currently, we do not support {self.interp_mode} for EMA.'
+ self.interp_func = getattr(self, interp_mode)
+
+ self.momentum_cfg = dict() if momentum_cfg is None else deepcopy(
+ momentum_cfg)
+ self.momentum_policy = momentum_policy
+ if momentum_policy != 'fixed':
+ assert hasattr(
+ self, momentum_policy
+ ), f'Currently, we do not support {self.momentum_policy} for EMA.'
+ self.momentum_updater = getattr(self, momentum_policy)
+
+ def karras(self, runner, gamma=7.0, max_momentum=1.0):
+ t = max(runner.iter + 1 - self.start_iter, 1)
+ ema_beta = min((1 - 1 / t) ** (gamma + 1), max_momentum)
+ return dict(momentum=ema_beta)
+
+ def after_train_iter(self, runner):
+ if not self.every_n_iters(runner, self.interval):
+ return
+
+ with torch.no_grad():
+ model = runner.model.module if is_module_wrapper(
+ runner.model) else runner.model
+
+ # update momentum
+ _interp_cfg = deepcopy(self.interp_cfg)
+ if self.momentum_policy != 'fixed':
+ _updated_args = self.momentum_updater(runner, **self.momentum_cfg)
+ _interp_cfg.update(_updated_args)
+
+ for key in self.module_keys:
+ net = rgetattr(model, get_ori_key(key))
+ ema = rgetattr(model, key)
+ if FSDPModule is not None and isinstance(net, FSDPModule): # Root parameters in EMA are unsharded after inference
+ net_is_sharded = net._get_fsdp_state()._fsdp_param_group.is_sharded
+ ema_is_sharded = ema._get_fsdp_state()._fsdp_param_group.is_sharded
+ if net_is_sharded and not ema_is_sharded:
+ ema.reshard()
+
+ for p_net, p_ema in zip(net.parameters(), ema.parameters()):
+ if self.trainable_only and not p_net.requires_grad:
+ continue
+ if runner.iter < self.start_iter:
+ p_ema.data.copy_(p_net.data)
+ else:
+ p_ema.data.copy_(self.interp_func(
+ p_net, p_ema, trainable=p_net.requires_grad, **_interp_cfg))
+
+ for b_net, b_ema in zip(net.buffers(), ema.buffers()):
+ b_ema.data.copy_(b_net.data)
+
+ def before_run(self, runner):
+ model = runner.model.module if is_module_wrapper(
+ runner.model) else runner.model
+ # sanity check for ema model
+ for k in self.module_keys:
+ if not rhasattr(model, k):
+ raise RuntimeError(
+ f'Cannot find {k} network for EMA hook.')
diff --git a/lakonlab/runner/hooks/logger/__init__.py b/lakonlab/runner/hooks/logger/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..2c03780dd4d27d05f90c86208d8d79fd49395c2a
--- /dev/null
+++ b/lakonlab/runner/hooks/logger/__init__.py
@@ -0,0 +1,3 @@
+from .text import TextLoggerHook
+
+__all__ = ['TextLoggerHook']
diff --git a/lakonlab/runner/hooks/logger/__pycache__/__init__.cpython-310.pyc b/lakonlab/runner/hooks/logger/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..4afa8889806499c2c7150ba10ecba54910373db1
Binary files /dev/null and b/lakonlab/runner/hooks/logger/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/runner/hooks/logger/__pycache__/text.cpython-310.pyc b/lakonlab/runner/hooks/logger/__pycache__/text.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..6168ab52980e42a7cd709e65955512a2e147ec5f
Binary files /dev/null and b/lakonlab/runner/hooks/logger/__pycache__/text.cpython-310.pyc differ
diff --git a/lakonlab/runner/hooks/logger/text.py b/lakonlab/runner/hooks/logger/text.py
new file mode 100644
index 0000000000000000000000000000000000000000..37043c33edf60e09221a31301580cd0d9e6eaaa2
--- /dev/null
+++ b/lakonlab/runner/hooks/logger/text.py
@@ -0,0 +1,24 @@
+import torch
+import torch.distributed as dist
+
+from mmcv.runner import HOOKS
+from mmcv.runner import TextLoggerHook as _TextLoggerHook
+
+from lakonlab.parallel import FSDPWrapper
+
+
+@HOOKS.register_module(force=True)
+class TextLoggerHook(_TextLoggerHook):
+
+ def _get_max_memory(self, runner) -> int:
+ if isinstance(runner.model, FSDPWrapper):
+ return 0
+ else:
+ device = getattr(runner.model, 'output_device', None)
+ mem = torch.cuda.max_memory_allocated(device=device)
+ mem_mb = torch.tensor([int(mem) // (1024 * 1024)],
+ dtype=torch.int,
+ device=device)
+ if runner.world_size > 1:
+ dist.reduce(mem_mb, 0, op=dist.ReduceOp.MAX)
+ return mem_mb.item()
diff --git a/lakonlab/runner/optimizer/__init__.py b/lakonlab/runner/optimizer/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..c0b412f37c1e405b9810676183cb0252e26eb4d2
--- /dev/null
+++ b/lakonlab/runner/optimizer/__init__.py
@@ -0,0 +1,3 @@
+from .builder import OPTIMIZERS, build_optimizers
+
+__all__ = ['OPTIMIZERS', 'build_optimizers']
diff --git a/lakonlab/runner/optimizer/__pycache__/__init__.cpython-310.pyc b/lakonlab/runner/optimizer/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..0cd8733d16da0664d3c42bac0007cf262f4ec595
Binary files /dev/null and b/lakonlab/runner/optimizer/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/runner/optimizer/__pycache__/builder.cpython-310.pyc b/lakonlab/runner/optimizer/__pycache__/builder.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..93486b1e271998a0542ff98430fd57a2557e409e
Binary files /dev/null and b/lakonlab/runner/optimizer/__pycache__/builder.cpython-310.pyc differ
diff --git a/lakonlab/runner/optimizer/builder.py b/lakonlab/runner/optimizer/builder.py
new file mode 100644
index 0000000000000000000000000000000000000000..f5e1e6650f4307f1c16edb233521e3a0c716e5d1
--- /dev/null
+++ b/lakonlab/runner/optimizer/builder.py
@@ -0,0 +1,45 @@
+import inspect
+import bitsandbytes
+
+from typing import List
+from mmcv.runner import build_optimizer
+from lakonlab.utils import rgetattr
+
+from mmcv.runner.optimizer.builder import OPTIMIZERS
+
+
+def register_bitsandbytes_optimizers() -> List:
+ bitsandbytes_optimizers = []
+ for module_name in dir(bitsandbytes.optim):
+ if module_name.startswith('__'):
+ continue
+ _optim = getattr(bitsandbytes.optim, module_name)
+ if inspect.isclass(_optim) and issubclass(_optim, bitsandbytes.optim.optimizer.Optimizer2State) \
+ and module_name not in OPTIMIZERS.module_dict:
+ OPTIMIZERS.register_module(module=_optim)
+ bitsandbytes_optimizers.append(module_name)
+ return bitsandbytes_optimizers
+
+
+BNB_OPTIMIZERS = register_bitsandbytes_optimizers()
+
+
+def build_optimizers(model, cfgs):
+ """Modified from MMGeneration
+ """
+ optimizers = {}
+ if hasattr(model, 'module'):
+ model = model.module
+ # determine whether 'cfgs' has several dicts for optimizer
+ is_dict_of_dict = True
+ for key, cfg in cfgs.items():
+ if not isinstance(cfg, dict):
+ is_dict_of_dict = False
+ if is_dict_of_dict:
+ for key, cfg in cfgs.items():
+ cfg_ = cfg.copy()
+ module = rgetattr(model, key)
+ optimizers[key] = build_optimizer(module, cfg_)
+ return optimizers
+
+ return build_optimizer(model, cfgs)
diff --git a/lakonlab/runner/timer.py b/lakonlab/runner/timer.py
new file mode 100644
index 0000000000000000000000000000000000000000..0a12365b80338e6aa4ab17ffc1913f0f29ca9ac4
--- /dev/null
+++ b/lakonlab/runner/timer.py
@@ -0,0 +1,72 @@
+# modified from
+# https://github.com/tjiiv-cprg/EPro-PnP/blob/42412220b641aef9e8943ceba516b3175631d370/EPro-PnP-Det/epropnp_det/utils/timer.py
+
+"""
+Copyright (C) 2010-2022 Alibaba Group Holding Limited.
+"""
+
+import numpy as np
+import torch
+import mmcv
+from mmcv import Timer
+from mmgen.utils import get_root_logger
+
+
+class IterTimer:
+ def __init__(self, name='time', sync=True, enabled=True):
+ self.name = name
+ self.times = []
+ self.timer = Timer(start=False)
+ self.sync = sync
+ self.enabled = enabled
+
+ def __enter__(self):
+ if not self.enabled:
+ return
+ if self.sync:
+ torch.cuda.synchronize()
+ self.timer.start()
+ return self
+
+ def __exit__(self, type, value, traceback):
+ if not self.enabled:
+ return
+ if self.sync:
+ torch.cuda.synchronize()
+ self.timer_record()
+ self.timer._is_running = False
+
+ def timer_start(self):
+ self.timer.start()
+
+ def timer_record(self):
+ self.times.append(self.timer.since_last_check())
+
+ def print_time(self):
+ if not self.enabled:
+ return
+ logger = get_root_logger()
+ mmcv.print_log(f'Average {self.name} = {np.average(self.times):.4f}', logger=logger)
+
+ def reset(self):
+ self.times = []
+
+
+class IterTimers(dict):
+ def __init__(self, *args, **kwargs):
+ super(IterTimers, self).__init__(*args, **kwargs)
+
+ def disable_all(self):
+ for timer in self.values():
+ timer.enabled = False
+
+ def enable_all(self):
+ for timer in self.values():
+ timer.enabled = True
+
+ def add_timer(self, name='time', sync=True, enabled=False):
+ self[name] = IterTimer(
+ name, sync=sync, enabled=enabled)
+
+
+default_timers = IterTimers()
diff --git a/lakonlab/ui/__init__.py b/lakonlab/ui/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/lakonlab/ui/__pycache__/__init__.cpython-310.pyc b/lakonlab/ui/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..0573bed05d142618a0726afd33cd6e13a04c6707
Binary files /dev/null and b/lakonlab/ui/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/ui/gradio/__init__.py b/lakonlab/ui/gradio/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/lakonlab/ui/gradio/create_text_to_img.py b/lakonlab/ui/gradio/create_text_to_img.py
new file mode 100644
index 0000000000000000000000000000000000000000..9d9c836c335320b6874672106f886ace98aae4c1
--- /dev/null
+++ b/lakonlab/ui/gradio/create_text_to_img.py
@@ -0,0 +1,53 @@
+import gradio as gr
+from .shared_opts import create_base_opts, create_generate_bar, set_seed, create_prompt_opts
+
+
+def create_interface_text_to_img(
+ api, prompt='', seed=42, steps=32, min_steps=4, max_steps=50, steps_slider_step=1,
+ height=768, width=1360, hw_slider_step=16,
+ guidance_scale=None, temperature=None, api_name='text_to_img',
+ create_negative_prompt=False, args=['last_seed', 'prompt', 'width', 'height', 'steps', 'guidance_scale']):
+ var_dict = dict()
+ with gr.Blocks(analytics_enabled=False) as interface:
+ var_dict['output_image'] = gr.Image(
+ type='pil', image_mode='RGB', label='Output image', interactive=False, elem_classes=['vh-img'])
+ create_prompt_opts(var_dict, create_negative_prompt=create_negative_prompt, prompt=prompt)
+ with gr.Column(variant='compact', elem_classes=['custom-spacing']):
+ with gr.Row(variant='compact', elem_classes=['force-hide-container']):
+ var_dict['width'] = gr.Slider(
+ label='Width', minimum=64, maximum=2048, step=hw_slider_step, value=width,
+ elem_classes=['force-hide-container'])
+ var_dict['switch_hw'] = gr.Button('\U000021C6', elem_classes=['tool'])
+ var_dict['height'] = gr.Slider(
+ label='Height', minimum=64, maximum=2048, step=hw_slider_step, value=height,
+ elem_classes=['force-hide-container'])
+ var_dict['switch_hw'].click(
+ fn=lambda w, h: (h, w),
+ inputs=[var_dict['width'], var_dict['height']],
+ outputs=[var_dict['width'], var_dict['height']],
+ show_progress=False,
+ api_name=False)
+ create_generate_bar(var_dict, text='Generate', seed=seed)
+ create_base_opts(
+ var_dict,
+ steps=steps,
+ min_steps=min_steps,
+ max_steps=max_steps,
+ steps_slider_step=steps_slider_step,
+ guidance_scale=guidance_scale,
+ temperature=temperature)
+
+ var_dict['run_btn'].click(
+ fn=set_seed,
+ inputs=var_dict['seed'],
+ outputs=var_dict['last_seed'],
+ show_progress=False,
+ api_name=False
+ ).success(
+ fn=api,
+ inputs=[var_dict[arg] for arg in args],
+ outputs=var_dict['output_image'],
+ concurrency_id='default_group', api_name=api_name
+ )
+
+ return interface, var_dict
diff --git a/lakonlab/ui/gradio/shared_opts.py b/lakonlab/ui/gradio/shared_opts.py
new file mode 100644
index 0000000000000000000000000000000000000000..9d6594ae87440d5dc7f02b50abc9646dfa4f1fbd
--- /dev/null
+++ b/lakonlab/ui/gradio/shared_opts.py
@@ -0,0 +1,64 @@
+import random
+import gradio as gr
+
+
+def create_prompt_opts(var_dict, create_negative_prompt=True, prompt='', negatove_prompt=''):
+ var_dict['prompt'] = gr.Textbox(
+ prompt, label='Prompt', show_label=False, lines=2, placeholder='Prompt', container=False, interactive=True)
+ if create_negative_prompt:
+ var_dict['negative_prompt'] = gr.Textbox(
+ negatove_prompt, label='Negative prompt', show_label=False, lines=2,
+ placeholder='Negative prompt', container=False, interactive=True)
+
+
+def create_generate_bar(var_dict, text='Generate', variant='primary', seed=-1):
+ with gr.Row(equal_height=False):
+ var_dict['run_btn'] = gr.Button(text, variant=variant, scale=2)
+ var_dict['seed'] = gr.Number(
+ label='Seed', value=seed, min_width=100, precision=0, minimum=-1, maximum=2 ** 31,
+ elem_classes=['force-hide-container'])
+ var_dict['random_seed'] = gr.Button('\U0001f3b2\ufe0f', elem_classes=['tool'])
+ var_dict['reuse_seed'] = gr.Button('\u267b\ufe0f', elem_classes=['tool'])
+ with gr.Column(visible=False):
+ var_dict['last_seed'] = gr.Number(value=seed, label='Last seed')
+ var_dict['reuse_seed'].click(
+ fn=lambda x: x,
+ inputs=var_dict['last_seed'],
+ outputs=var_dict['seed'],
+ show_progress=False,
+ api_name=False)
+ var_dict['random_seed'].click(
+ fn=lambda: -1,
+ outputs=var_dict['seed'],
+ show_progress=False,
+ api_name=False)
+
+
+def create_base_opts(var_dict,
+ steps=24,
+ min_steps=4,
+ max_steps=50,
+ steps_slider_step=1,
+ guidance_scale=None,
+ temperature=None,
+ render=True):
+ with gr.Column(variant='compact', elem_classes=['custom-spacing'], render=render) as base_opts:
+ with gr.Row(variant='compact', elem_classes=['force-hide-container']):
+ var_dict['steps'] = gr.Slider(
+ min_steps, max_steps, value=steps, step=steps_slider_step, label='Sampling steps',
+ elem_classes=['force-hide-container'])
+ with gr.Row(variant='compact', elem_classes=['force-hide-container']):
+ if guidance_scale is not None:
+ var_dict['guidance_scale'] = gr.Slider(
+ 0.0, 30.0, value=guidance_scale, step=0.5, label='Guidance scale',
+ elem_classes=['force-hide-container'])
+ if temperature is not None:
+ var_dict['temperature'] = gr.Slider(
+ 0.0, 1.0, value=temperature, step=0.01, label='Temperature',
+ elem_classes=['force-hide-container'])
+ return base_opts
+
+
+def set_seed(seed):
+ seed = random.randint(0, 2**31) if seed == -1 else seed
+ return seed
diff --git a/lakonlab/ui/gradio/style.css b/lakonlab/ui/gradio/style.css
new file mode 100644
index 0000000000000000000000000000000000000000..387a5d732892278c5a918e624e53434949e2b2bb
--- /dev/null
+++ b/lakonlab/ui/gradio/style.css
@@ -0,0 +1,59 @@
+.force-hide-container {
+ margin: 0;
+ box-shadow: none;
+ --block-border-width: 0;
+ background: transparent;
+ padding: 0;
+ overflow: visible;
+}
+
+.svelte-sfqy0y {
+ display: flex;
+ flex-direction: inherit;
+ flex-wrap: wrap;
+ gap: 0;
+ box-shadow: none;
+ border: 0;
+ border-radius: 0;
+ background: transparent;
+ overflow-y: hidden;
+}
+
+.custom-spacing {
+ padding: 10px;
+ gap: 20px;
+ flex-grow: 0 !important;
+}
+
+.unequal-height {
+ align-items: flex-end;
+}
+
+.tool{
+ max-width: 40px;
+ min-width: 40px !important;
+}
+
+/* Center the component and allow it to use the full row width */
+.vh-img {
+ display: grid;
+ justify-items: center;
+}
+
+/* Container should size to the image, but never exceed the row width */
+.vh-img .image-container {
+ inline-size: fit-content !important; /* prefers image’s natural width */
+ max-inline-size: 100% !important; /* ...but clamps to available width */
+ margin-inline: auto;
+ overflow: hidden; /* avoid odd overflow on iOS */
+}
+
+/* Image scales by BOTH constraints: height cap and row width */
+.vh-img .image-container img {
+ max-block-size: 700px !important; /* fixed max height cap */
+ max-inline-size: 100%; /* never wider than container */
+ inline-size: auto; /* keep aspect ratio */
+ block-size: auto;
+ object-fit: contain;
+ display: block;
+}
diff --git a/lakonlab/ui/media_viewer/__init__.py b/lakonlab/ui/media_viewer/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..093e0725630b49fbab9f2e945353430535a453df
--- /dev/null
+++ b/lakonlab/ui/media_viewer/__init__.py
@@ -0,0 +1 @@
+from .grid_tools import write_html
diff --git a/lakonlab/ui/media_viewer/__pycache__/__init__.cpython-310.pyc b/lakonlab/ui/media_viewer/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..1f4d4b6af83b6bcf6adf42b4d239cd7c3552b219
Binary files /dev/null and b/lakonlab/ui/media_viewer/__pycache__/__init__.cpython-310.pyc differ
diff --git a/lakonlab/ui/media_viewer/__pycache__/grid_tools.cpython-310.pyc b/lakonlab/ui/media_viewer/__pycache__/grid_tools.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..ca3042dec15d6e62f96842b685dd0c60b2637e92
Binary files /dev/null and b/lakonlab/ui/media_viewer/__pycache__/grid_tools.cpython-310.pyc differ
diff --git a/lakonlab/ui/media_viewer/grid_tools.py b/lakonlab/ui/media_viewer/grid_tools.py
new file mode 100644
index 0000000000000000000000000000000000000000..00cc221679f0886aabee72bba5ffe5f1319c6a86
--- /dev/null
+++ b/lakonlab/ui/media_viewer/grid_tools.py
@@ -0,0 +1,53 @@
+# Copyright (c) 2025 Hansheng Chen
+
+import html
+import json
+import os
+import pathlib
+from urllib.parse import quote
+
+
+ASSET_DIR = pathlib.Path(__file__).parent
+TEMPLATE = (ASSET_DIR / "viewer.html").read_text(encoding="utf-8")
+
+
+def build_thumbnails(entries, n_cols):
+ """return markup, sources list, captions list"""
+ sources, caps, blocks = [], [], []
+ for i, (d, src, cap) in enumerate(entries):
+ if not (src.startswith("http://") or src.startswith("https://")):
+ src = quote(src, safe="/%")
+ esc_src = html.escape(src)
+ sources.append(esc_src)
+ caps.append(cap)
+
+ ext = os.path.splitext(esc_src)[-1].lower()
+ thumb = (f''
+ if ext == '.mp4' else f'
')
+ blocks.append(f'{thumb}'
+ f'
')
+ grid = (f'\n'
+ if n_cols else '
\n')
+ return grid + '\n'.join(blocks) + '\n
', sources, caps
+
+
+def grid_html(entries, n_cols=None, *, inline_assets=False):
+ grid_markup, srcs, caps = build_thumbnails(entries, n_cols)
+ blob = f''
+ page = TEMPLATE.replace("{{GRID_MARKUP}}", grid_markup + blob)
+ if inline_assets:
+ css = (ASSET_DIR / "viewer.css").read_text(encoding="utf-8")
+ js = (ASSET_DIR / "viewer.js").read_text(encoding="utf-8")
+ page = page.replace(
+ '
', f''
+ ).replace(
+ '', f'')
+ return page
+
+
+def write_html(html_path, entries, file_client):
+ if not entries:
+ return
+ formatted = [(d, img, f'[{d}] {name}') for d, img, name in entries]
+ file_client.put_text(grid_html(formatted, inline_assets=True), html_path)
diff --git a/lakonlab/ui/media_viewer/viewer.css b/lakonlab/ui/media_viewer/viewer.css
new file mode 100644
index 0000000000000000000000000000000000000000..b9c457b50ef91a0b538c34d223ed48dadb999941
--- /dev/null
+++ b/lakonlab/ui/media_viewer/viewer.css
@@ -0,0 +1,76 @@
+/* ---------- layout / thumbnails ---------- */
+body{
+ margin:0;
+ font-family:Arial,Helvetica,sans-serif;
+ padding:1rem;
+}
+
+/* will be “auto-fill” by default or 2-col when forced */
+.grid{
+ display:grid;
+ grid-template-columns:repeat(auto-fill,minmax(320px,1fr));
+ gap:1rem;
+}
+
+.item img,
+.item video{
+ width:100%;
+ height:auto;
+ border:1px solid #ccc;
+ box-sizing:border-box;
+ cursor:pointer;
+}
+
+.prompt{
+ width:100%;
+ height:120px;
+ margin-top:6px;
+ overflow:auto;
+ resize:vertical;
+ box-sizing:border-box;
+}
+
+/* ---------- light-box ---------- */
+#overlay{
+ display:none;
+ position:fixed;
+ inset:0;
+ background:rgba(0,0,0,.9);
+ z-index:9999;
+ flex-direction:column;
+ align-items:center;
+ justify-content:center;
+ overflow:auto;
+ overscroll-behavior:contain; /* stop background scrolling */
+}
+
+#overlay.scrollY{ /* vertical overflow when zoomed */
+ justify-content:flex-start;
+}
+
+#overlay img{
+ max-width:90vw;
+ max-height:90vh;
+ cursor:zoom-in;
+}
+
+#overlay img.zoomed{
+ max-width:none;
+ max-height:none;
+ cursor:zoom-out;
+}
+
+#overlay video{
+ max-width:90vw;
+ max-height:90vh;
+ cursor:default; /* videos don’t zoom */
+}
+
+#caption{
+ color:#fff;
+ margin:1rem auto 0 auto;
+ text-align:center;
+ font-size:1rem;
+ padding:0 1rem;
+ max-width:90vw;
+}
diff --git a/lakonlab/ui/media_viewer/viewer.html b/lakonlab/ui/media_viewer/viewer.html
new file mode 100644
index 0000000000000000000000000000000000000000..9b810fd36f4888401924c7dc818742a57e339cf1
--- /dev/null
+++ b/lakonlab/ui/media_viewer/viewer.html
@@ -0,0 +1,12 @@
+
+
+
+
+
+
+
+
+ {{GRID_MARKUP}}
+
+
+
diff --git a/lakonlab/ui/media_viewer/viewer.js b/lakonlab/ui/media_viewer/viewer.js
new file mode 100644
index 0000000000000000000000000000000000000000..1cd6c8a078b4c1a4de7bc0e5272147aaa322f3c4
--- /dev/null
+++ b/lakonlab/ui/media_viewer/viewer.js
@@ -0,0 +1,170 @@
+/* viewer.js
+ =========
+ Light-box / zoom / keyboard logic for the result grid.
+ Expects:
+ • the page to contain thumbnails inside elements with class="item"
+ and data-idx attributes that match the order of
+ window.GRID_DATA.sources & window.GRID_DATA.captions
+ • the CSS rules from viewer.css (same selectors as before)
+*/
+
+const INITIAL_SCALE_WHEN_FITS = 1.5;
+
+/* ----------------------------------------------------------
+ 1. Data from Python-generated