From 554c8fbf2f32274c99132e84bafa8c8c68f67575 Mon Sep 17 00:00:00 2001 From: vladmandic Date: Tue, 24 Feb 2026 18:08:30 +0100 Subject: [PATCH] remove model: hdm Signed-off-by: vladmandic --- modules/sd_models.py | 4 - pipelines/hdm/__init__.py | 0 pipelines/hdm/hdm/__init__.py | 3 - pipelines/hdm/hdm/data/__init__.py | 0 pipelines/hdm/hdm/data/base.py | 251 --------- pipelines/hdm/hdm/data/kohya.py | 227 -------- pipelines/hdm/hdm/loader.py | 123 ----- pipelines/hdm/hdm/modules/base.py | 417 -------------- pipelines/hdm/hdm/modules/rope.py | 107 ---- pipelines/hdm/hdm/modules/text_encoders.py | 330 ----------- pipelines/hdm/hdm/modules/unet_patch.py | 610 --------------------- pipelines/hdm/hdm/modules/xut.py | 98 ---- pipelines/hdm/hdm/pipeline.py | 165 ------ pipelines/hdm/hdm/trainer/__init__.py | 1 - pipelines/hdm/hdm/trainer/callbacks.py | 66 --- pipelines/hdm/hdm/trainer/diffusion.py | 68 --- pipelines/hdm/hdm/trainer/trainer.py | 473 ---------------- pipelines/hdm/hdm/utils/__init__.py | 74 --- pipelines/hdm/hdm/utils/config.py | 38 -- pipelines/hdm/xut/__init__.py | 0 pipelines/hdm/xut/env.py | 9 - pipelines/hdm/xut/modules/__init__.py | 0 pipelines/hdm/xut/modules/adaln.py | 28 - pipelines/hdm/xut/modules/attention.py | 332 ----------- pipelines/hdm/xut/modules/axial_rope.py | 179 ------ pipelines/hdm/xut/modules/layers.py | 56 -- pipelines/hdm/xut/modules/norm.py | 97 ---- pipelines/hdm/xut/modules/patch.py | 74 --- pipelines/hdm/xut/modules/time_emb.py | 34 -- pipelines/hdm/xut/modules/transformer.py | 79 --- pipelines/hdm/xut/utils/__init__.py | 23 - pipelines/hdm/xut/xut.py | 556 ------------------- pipelines/model_hdm.py | 34 -- 33 files changed, 4556 deletions(-) delete mode 100644 pipelines/hdm/__init__.py delete mode 100644 pipelines/hdm/hdm/__init__.py delete mode 100644 pipelines/hdm/hdm/data/__init__.py delete mode 100644 pipelines/hdm/hdm/data/base.py delete mode 100644 pipelines/hdm/hdm/data/kohya.py delete mode 100644 pipelines/hdm/hdm/loader.py delete mode 100644 pipelines/hdm/hdm/modules/base.py delete mode 100644 pipelines/hdm/hdm/modules/rope.py delete mode 100644 pipelines/hdm/hdm/modules/text_encoders.py delete mode 100644 pipelines/hdm/hdm/modules/unet_patch.py delete mode 100644 pipelines/hdm/hdm/modules/xut.py delete mode 100644 pipelines/hdm/hdm/pipeline.py delete mode 100644 pipelines/hdm/hdm/trainer/__init__.py delete mode 100644 pipelines/hdm/hdm/trainer/callbacks.py delete mode 100644 pipelines/hdm/hdm/trainer/diffusion.py delete mode 100644 pipelines/hdm/hdm/trainer/trainer.py delete mode 100644 pipelines/hdm/hdm/utils/__init__.py delete mode 100644 pipelines/hdm/hdm/utils/config.py delete mode 100644 pipelines/hdm/xut/__init__.py delete mode 100644 pipelines/hdm/xut/env.py delete mode 100644 pipelines/hdm/xut/modules/__init__.py delete mode 100644 pipelines/hdm/xut/modules/adaln.py delete mode 100644 pipelines/hdm/xut/modules/attention.py delete mode 100644 pipelines/hdm/xut/modules/axial_rope.py delete mode 100644 pipelines/hdm/xut/modules/layers.py delete mode 100644 pipelines/hdm/xut/modules/norm.py delete mode 100644 pipelines/hdm/xut/modules/patch.py delete mode 100644 pipelines/hdm/xut/modules/time_emb.py delete mode 100644 pipelines/hdm/xut/modules/transformer.py delete mode 100644 pipelines/hdm/xut/utils/__init__.py delete mode 100644 pipelines/hdm/xut/xut.py delete mode 100644 pipelines/model_hdm.py diff --git a/modules/sd_models.py b/modules/sd_models.py index 80bd03d41..ba72964f6 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -456,10 +456,6 @@ def load_diffuser_force(detected_model_type, checkpoint_info, diffusers_load_con from pipelines.model_nextstep import load_nextstep sd_model = load_nextstep(checkpoint_info, diffusers_load_config) # pylint: disable=assignment-from-none allow_post_quant = False - elif model_type in ['hdm']: - from pipelines.model_hdm import load_hdm - sd_model = load_hdm(checkpoint_info, diffusers_load_config) - allow_post_quant = False elif model_type in ['HunyuanImage']: from pipelines.model_hyimage import load_hyimage sd_model = load_hyimage(checkpoint_info, diffusers_load_config) # pylint: disable=assignment-from-none diff --git a/pipelines/hdm/__init__.py b/pipelines/hdm/__init__.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/pipelines/hdm/hdm/__init__.py b/pipelines/hdm/hdm/__init__.py deleted file mode 100644 index 7d1d1c490..000000000 --- a/pipelines/hdm/hdm/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from diffusers.models.modeling_utils import ModelMixin -from .modules.xut import XUDiTConditionModel -from .modules.unet_patch import HDUNet2DConditionModel, RoPEUNet2DConditionModel diff --git a/pipelines/hdm/hdm/data/__init__.py b/pipelines/hdm/hdm/data/__init__.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/pipelines/hdm/hdm/data/base.py b/pipelines/hdm/hdm/data/base.py deleted file mode 100644 index f72315495..000000000 --- a/pipelines/hdm/hdm/data/base.py +++ /dev/null @@ -1,251 +0,0 @@ -import random - -import numpy as np -import torch -import torch.utils.data as Data -from transformers import PreTrainedTokenizer -from tqdm import tqdm, trange - - -class BaseDataset(Data.Dataset): - def collate(self, batch): - samples = torch.stack([x["sample"] for x in batch]) - caption = [x["caption"] for x in batch] - tokenizer_outs = [x["tokenizer_out"] for x in batch] - add_time_ids = [x["add_time_ids"] for x in batch] - tokenizer_outputs = [] - for tokenizer_out in zip(*tokenizer_outs): - input_ids = torch.concat([x["input_ids"] for x in tokenizer_out]) - attention_mask = torch.concat([x["attention_mask"] for x in tokenizer_out]) - tokenizer_outputs.append( - {"input_ids": input_ids, "attention_mask": attention_mask} - ) - return ( - samples, - caption, - tokenizer_outputs, - {"time_ids": torch.concat(add_time_ids).float()}, - ) - - -class DummyDataset(BaseDataset): - def __init__( - self, - # (4, 128, 128) for latent - sample_size: tuple[int] = (3, 1024, 1024), - n_samples: int = 100, - tokenizers: list[PreTrainedTokenizer] = [], - **kwargs, - ): - if not isinstance(sample_size, tuple): - sample_size = tuple(sample_size) - self.samples = [torch.randn(sample_size) for _ in range(n_samples)] - if isinstance(tokenizers, list): - self.tokenizers = tokenizers - else: - self.tokenizers = [tokenizers] - - def __len__(self): - return len(self.samples) - - def __getitem__(self, index): - sample = self.samples[index] - caption = "DUMMY TEST" - return { - "sample": sample, - "caption": caption, - "tokenizer_out": [ - tokenizer( - caption, - padding="max_length", - truncation=True, - return_tensors="pt", - ) - for tokenizer in self.tokenizers - ], - # org_h, org_w, crop_top, crop_left, target_h, target_w - "add_time_ids": torch.tensor([[1024, 1024, 0, 0, 1024, 1024]]), - } - - -class CombineDataset(Data.Dataset): - def __init__( - self, - datasets: list[Data.Dataset], - latent_scale: float = 1.0, - latent_shift: float = 0.0, - tokenizers: list[PreTrainedTokenizer] = [], - shuffle=True, - arb_mode=False, - ): - self.shuffle = shuffle - self.datasets_ref = datasets - self.datasets = datasets - self.shard_string = sum( - ([chr(i).encode()] * len(dataset) for i, dataset in enumerate(datasets)), - [], - ) - self.tokenizers = tokenizers - self.latent_scale = latent_scale - self.latent_shift = latent_shift - - if shuffle: - random.shuffle(self.shard_string) - self.arb_mode = arb_mode - - dataset_ids = [0] * len(datasets) - for i, data in tqdm( - enumerate(self.shard_string), - total=len(self.shard_string), - smoothing=0.01, - desc="Dataset Indexing...", - ): - index = dataset_ids[data[0]] - dataset_ids[data[0]] += 1 - self.shard_string[i] = data + np.base_repr(index, 36).encode() - self.shard_string = np.array(self.shard_string) - - @torch.no_grad() - def collate(self, batch): - if self.arb_mode: - assert len(batch) == 1 - latents = batch[0]["latent"] - caption = batch[0]["caption"] - pos_map = batch[0]["pos_map"] - tokenizer_outs = batch[0]["tokenizer_out"] - if "aspect_ratio" in batch[0]: - return ( - latents, - caption, - tokenizer_outs, - pos_map, - {"addon_info": batch[0]["aspect_ratio"]}, - ) - return latents, caption, tokenizer_outs, pos_map - latents = torch.stack([x["latent"] for x in batch]) - caption = [x["caption"] for x in batch] - pos_map = torch.stack([x["pos_map"] for x in batch]) - tokenizer_outs = [x["tokenizer_out"] for x in batch] - - tokenizer_outputs = [] - for tokenizer_out in zip(*tokenizer_outs): - input_ids = torch.concat([x["input_ids"] for x in tokenizer_out]) - attention_mask = torch.concat([x["attention_mask"] for x in tokenizer_out]) - tokenizer_outputs.append( - {"input_ids": input_ids, "attention_mask": attention_mask} - ) - if "aspect_ratio" in batch[0]: - aspect_ratio = torch.tensor([x["aspect_ratio"] for x in batch]) - return ( - latents, - caption, - tokenizer_outputs, - pos_map, - {"addon_info": aspect_ratio}, - ) - return latents, caption, tokenizer_outputs, pos_map - - def __len__(self): - return sum(len(dataset) for dataset in self.datasets) - - @torch.no_grad() - def __getitem__(self, index): - choosed = self.shard_string[index] - dataset = self.datasets[choosed[0]] - index = int(choosed[1:], 36) - latent, caption, pos_map, *ar = dataset[index] - if self.arb_mode: - tokenizer_out = [ - [ - tokenizer( - c, - padding="max_length", - truncation=True, - return_tensors="pt", - ) - for tokenizer in self.tokenizers - ] - for c in caption - ] - data = { - "latent": (latent * self.latent_scale + self.latent_shift), - "caption": caption, - "pos_map": pos_map, - "tokenizer_out": tokenizer_out, - } - if len(ar) > 0: - aspect_ratio = ar[0] - data["aspect_ratio"] = aspect_ratio - return data - tokenizer_out = [ - tokenizer( - caption, - padding="max_length", - truncation=True, - return_tensors="pt", - ) - for tokenizer in self.tokenizers - ] - data = { - "latent": (latent * self.latent_scale + self.latent_shift), - "caption": caption, - "pos_map": pos_map, - "tokenizer_out": tokenizer_out, - } - if len(ar) > 0: - aspect_ratio = ar[0] - data["aspect_ratio"] = aspect_ratio - return data - - -if __name__ == "__main__": - from transformers import Qwen2Tokenizer - from .kohya import * - - tokenizer = Qwen2Tokenizer.from_pretrained("Qwen/Qwen3-0.6B") - dataset = KohyaDataset( - dataset_folder="/mp34-1/danbooru2023", - keep_token_seperator="|||", - tag_seperator="$$", - seperator=", ", - group_seperator="%%", - tag_shuffle=True, - group_shuffle=True, - tag_dropout_rate=0.0, - group_dropout_rate=0.0, - use_cached_meta=True, - transform=transforms.Compose( - [ - transforms.ToTensor(), - transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5), inplace=True), - ] - ), - use_arb=True, - arb_config={ - "batch_size": 32, - "target_res": 1024, - "res_step": 16, - "seed": 0, - }, - meta_postfix="_filtered", - ) - combine = CombineDataset( - [dataset], tokenizers=[tokenizer], shuffle=True, arb_mode=True - ) - - print(len(combine)) - - dataloader = Data.DataLoader( - combine, batch_size=1, num_workers=0, shuffle=True, collate_fn=combine.collate - ) - for batch in tqdm(dataloader): - latent, caption, tokenizer_out, pos_map, *ar = batch - print(latent.size(), pos_map.size()) - print(len(caption), caption[0]) - print( - len(tokenizer_out), len(tokenizer_out[0]), tokenizer_out[0][0]["input_ids"] - ) - print(len(ar)) - if len(ar) > 0: - print(ar[0]) - break diff --git a/pipelines/hdm/hdm/data/kohya.py b/pipelines/hdm/hdm/data/kohya.py deleted file mode 100644 index 3d2dd9c55..000000000 --- a/pipelines/hdm/hdm/data/kohya.py +++ /dev/null @@ -1,227 +0,0 @@ -import os -import sys -import io -import math -import random -import pickle -import tempfile -from collections import defaultdict - -import torch -import torchvision.transforms as transforms -import torch.utils.data as Data -import numpy as np -import imagesize -from tqdm import tqdm -from PIL import Image - -from xut.modules.axial_rope import make_cropped_pos, make_axial_pos_no_cache - - -def get_files(folder): - if os.path.isdir(folder): - return [ - os.path.join(folder, f) - for f in os.listdir(folder) - if any(f.endswith(ext) for ext in [".jpg", ".png", ".jpeg", ".webp"]) - ] - else: - return None - - -def load_npy(path): - with open(path, "rb") as f: - raw_data = f.read() - if sys.platform == "win32": - data = np.load(io.BytesIO(raw_data)) - else: - with tempfile.NamedTemporaryFile() as tmp: - tmp.write(raw_data) - tmp.flush() - data = np.load(tmp.name, mmap_mode="r") - return data - - -def load_pickle(path): - with open(path, "rb") as f: - raw_data = f.read() - data = pickle.loads(raw_data) - return data - - -def conver_rgb(x): - return x.convert("RGB") - - -class KohyaDataset(Data.Dataset): - def __init__( - self, - size=1024, - dataset_folder="/mp34-1/danbooru2023", - transform=None, - keep_token_seperator="|||", - tag_seperator="$$", - seperator=", ", - group_seperator="%%", - tag_shuffle=True, - group_shuffle=True, - tag_dropout_rate=0.25, - group_dropout_rate=0.3, - use_cached_meta=True, - meta_postfix="_filtered", - ): - self.dataset_folder = dataset_folder - if ( - os.path.isfile(os.path.join(dataset_folder, f"metadata{meta_postfix}.npy")) - and use_cached_meta - ): - self.files = load_npy( - os.path.join(dataset_folder, f"metadata{meta_postfix}.npy") - ) - else: - print("Cached metadata not found, generating...") - files = [] - for entry in os.listdir(dataset_folder): - if os.path.isdir(os.path.join(dataset_folder, entry)): - files.extend(get_files(os.path.join(dataset_folder, entry))) - elif any( - entry.endswith(ext) for ext in [".jpg", ".png", ".jpeg", ".webp"] - ): - files.append(entry) - files = [(i, os.path.splitext(i)[0] + ".txt") for i in files] - self.files = np.array(files) - np.save(os.path.join(dataset_folder, f"metadata{meta_postfix}.npy"), files) - print("Cached metadata generated and saved") - - self.keep_token_seperator = keep_token_seperator - self.tag_seperator = tag_seperator - self.seperator = seperator - self.group_seperator = group_seperator - self.tag_shuffle = tag_shuffle - self.group_shuffle = group_shuffle - self.tag_dropout_rate = tag_dropout_rate - self.group_dropout_rate = group_dropout_rate - - self.size = size - self.transform = transform or transforms.Compose( - [ - transforms.Lambda(conver_rgb), - transforms.Resize( - size, interpolation=transforms.InterpolationMode.BICUBIC - ), - transforms.ToTensor(), - transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5), inplace=True), - ] - ) - - def __len__(self): - return len(self.files) - - def get_caption(self, txt_file): - if not os.path.isfile(txt_file): - return "" - with open(txt_file, "r", encoding="utf-8") as f: - caption = f.read() - - if self.keep_token_seperator in caption: - keep_tokens, rest = caption.split(self.keep_token_seperator) - keep_tokens = [ - i.strip() for i in keep_tokens.split(self.tag_seperator) if i.strip() - ] - else: - keep_tokens = [] - rest = caption - - groups = [i.strip() for i in rest.split(self.group_seperator) if i.strip()] - if self.group_shuffle: - random.shuffle(groups) - - for group in groups: - tags = [ - i.strip() - for i in group.split(self.tag_seperator) - if i.strip() and random.random() > self.tag_dropout_rate - ] - if self.tag_shuffle: - random.shuffle(tags) - if random.random() > self.group_dropout_rate: - keep_tokens.extend(tags) - - return self.seperator.join(keep_tokens) - - def get_data_from_files(self, img_file, txt_file, resize=None): - img_path = os.path.join(self.dataset_folder, img_file) - txt_path = os.path.join(self.dataset_folder, txt_file) - caption = self.get_caption(txt_path) - - with Image.open(img_path) as img: - if resize: - img = img.resize(resize, Image.Resampling.BICUBIC) - img_t = self.transform(img) - return img_t, caption - - def make_cropped_pos(self, img_t, target_h, target_w): - aspect_ratio = target_w / target_h - aspect_ratio = math.log( - aspect_ratio - ) # so we have a:b and b:a have same abs value - crop_h, crop_w = 0, 0 - if target_h > target_w: - crop_h = torch.randint(0, target_h - target_w, (1,)).item() - img = img_t[:, crop_h : crop_h + target_w, :] - elif target_h < target_w: - crop_w = torch.randint(0, target_w - target_h, (1,)).item() - img = img_t[:, :, crop_w : crop_w + target_h] - else: - img = img_t - - return img, make_cropped_pos(crop_h, crop_w, target_h, target_w) - - def _getitem(self, img_file, txt_file): - img_t, caption = self.get_data_from_files(img_file, txt_file) - target_h, target_w = img_t.shape[1:3] - aspect_ratio = target_w / target_h - img, pos_map = self.make_cropped_pos(img_t, target_h, target_w) - - return img, caption, pos_map, aspect_ratio - - def __getitem__(self, index): - img_file, txt_file = self.files[index] - return self._getitem(img_file, txt_file) - - -if __name__ == "__main__": - import random - - dataset = KohyaDataset( - dataset_folder="/mp34-1/danbooru2023", - keep_token_seperator="|||", - tag_seperator="$$", - seperator=", ", - group_seperator="%%", - tag_shuffle=True, - group_shuffle=True, - tag_dropout_rate=0.0, - group_dropout_rate=0.0, - use_cached_meta=True, - transform=transforms.Compose( - [ - transforms.ToTensor(), - transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5), inplace=True), - ] - ), - use_arb=True, - arb_config={ - "batch_size": 32, - "target_res": 1024, - "res_step": 16, - "seed": 0, - }, - meta_postfix="_filtered", - ) - - print(len(dataset.batches)) - k, values = dataset.batches[0] - print(k) - for v in values: - print(v) diff --git a/pipelines/hdm/hdm/loader.py b/pipelines/hdm/hdm/loader.py deleted file mode 100644 index ef60708d3..000000000 --- a/pipelines/hdm/hdm/loader.py +++ /dev/null @@ -1,123 +0,0 @@ -import torch -from diffusers import ( - EulerDiscreteScheduler, - UNet2DConditionModel, - AutoencoderKL, -) - -from .modules.text_encoders import BaseTextEncoder, SimpleTextEncoder -from .trainer import DMTrainer, FlowTrainer -from .utils import instantiate - - -def model_loader( - unet: UNet2DConditionModel | None = None, - unet_class=UNet2DConditionModel, - unet_config=None, - te: BaseTextEncoder | None = None, - te_class=SimpleTextEncoder, - te_config={ - "te_name": "apple/DFN5B-CLIP-ViT-H-14-378", - "device": "cpu", - "dtype": torch.float32, - "zero_for_padding": True, - }, - te_name="", - tokenizers: list[dict] | None = None, - vae: AutoencoderKL | None = None, - vae_class=AutoencoderKL, - vae_config=None, - vae_name="", - scheduler: EulerDiscreteScheduler | None = None, - scheduler_class=EulerDiscreteScheduler, - scheduler_config=None, - scheduler_name=None, - type=None, -): - if unet is None: - unet = instantiate(unet_class)(**unet_config) - else: - unet = instantiate(unet) - if te is None: - if te_name is not None and te_name != "": - te = instantiate(te_class).from_pretrained(te_name) - else: - te = instantiate(te_class)(**te_config) - else: - te = instantiate(te) - if vae is not None: - vae = instantiate(vae) - elif vae_class is not None: - if vae_name is not None and vae_name != "": - vae = instantiate(vae_class).from_pretrained(vae_name) - elif vae_config is not None: - vae = instantiate(vae_class)(**vae_config) - if scheduler is None: - if scheduler_name is not None and scheduler_name != "": - scheduler = instantiate(scheduler_class).from_pretrained(scheduler_name) - elif scheduler_class is not None and scheduler_config is not None: - scheduler = instantiate(scheduler_class)(**scheduler_config) - else: - scheduler = None - else: - scheduler = instantiate(scheduler) - - if hasattr(te, "tokenizers"): - tokenizers = te.tokenizers - elif hasattr(te, "tokenizer") and te.tokenizer is not None: - tokenizers = [te.tokenizer] - elif isinstance(tokenizers, str) and tokenizers != "": - tokenizers = [instantiate(tokenizers)] - elif isinstance(tokenizers, list): - tokenizers = [instantiate(tokenizer) for tokenizer in tokenizers] - else: - tokenizers = None - - return unet, te, tokenizers, vae, scheduler - - -def load_trainer(conf: dict, unet=None, te=None, vae=None, scheduler=None, type=None): - conf = dict(**conf) - if unet is not None: - conf["unet"] = unet - if te is not None: - conf["te"] = te - if vae is not None: - conf["vae"] = vae - if scheduler is not None: - conf["scheduler"] = scheduler - type = type or conf.pop("type", "dm") - if type == "dm": - trainer = DMTrainer(**conf) - elif type == "flow": - conf.pop("scheduler") - trainer = FlowTrainer(**conf) - else: - raise NotImplementedError - return trainer - - -def load_model(conf: dict): - """ - return unet(dit)/te/vae/scheduler - """ - if "model" in conf: - return model_loader(**conf["model"]) - return model_loader(**conf) - - -def load_dataset(conf: dict): - dataset = instantiate(conf) - return dataset - - -def load_all(conf: dict): - dataset_conf = conf.pop("dataset") - dataset = load_dataset(dataset_conf) - model_conf = conf.pop("model") - unet, te, tokenizers, vae, scheduler = load_model(model_conf) - trainer = load_trainer( - conf.pop("trainer"), unet=unet, te=te, vae=vae, scheduler=scheduler - ) - dataset.tokenizers = tokenizers - return dataset, trainer, (unet, te, tokenizers, vae, scheduler) diff --git a/pipelines/hdm/hdm/modules/base.py b/pipelines/hdm/hdm/modules/base.py deleted file mode 100644 index 51fdab237..000000000 --- a/pipelines/hdm/hdm/modules/base.py +++ /dev/null @@ -1,417 +0,0 @@ -from typing import Any, Dict, Optional, Tuple, Union - -import torch -import torch.nn as nn - -from diffusers import UNet2DConditionModel -from diffusers.configuration_utils import ConfigMixin, register_to_config -from diffusers.models.modeling_utils import ModelMixin -from diffusers.models.unets.unet_2d_condition import UNet2DConditionOutput - - -class BasicUNet(ModelMixin, ConfigMixin): - def enable_gradient_checkpointing(self): - raise NotImplementedError - - def disable_gradient_checkpointing(self): - raise NotImplementedError - - def forward( - self, - sample: torch.Tensor, - timestep: Union[torch.Tensor, float, int], - encoder_hidden_states: torch.Tensor, - class_labels: Optional[torch.Tensor] = None, - timestep_cond: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - cross_attention_kwargs: Optional[Dict[str, Any]] = None, - added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None, - down_block_additional_residuals: Optional[Tuple[torch.Tensor]] = None, - mid_block_additional_residual: Optional[torch.Tensor] = None, - down_intrablock_additional_residuals: Optional[Tuple[torch.Tensor]] = None, - encoder_attention_mask: Optional[torch.Tensor] = None, - return_dict: bool = True, - ) -> Union[UNet2DConditionOutput, Tuple]: - raise NotImplementedError - - -class UNetWithPos(UNet2DConditionModel): - @register_to_config - def __init__( - self, - sample_size: Optional[Union[int, Tuple[int, int]]] = None, - in_channels: int = 4, - out_channels: int = 4, - center_input_sample: bool = False, - flip_sin_to_cos: bool = True, - freq_shift: int = 0, - down_block_types: Tuple[str] = ( - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "CrossAttnDownBlock2D", - "DownBlock2D", - ), - mid_block_type: Optional[str] = "UNetMidBlock2DCrossAttn", - up_block_types: Tuple[str] = ( - "UpBlock2D", - "CrossAttnUpBlock2D", - "CrossAttnUpBlock2D", - "CrossAttnUpBlock2D", - ), - only_cross_attention: Union[bool, Tuple[bool]] = False, - block_out_channels: Tuple[int] = (320, 640, 1280, 1280), - layers_per_block: Union[int, Tuple[int]] = 2, - downsample_padding: int = 1, - mid_block_scale_factor: float = 1, - dropout: float = 0.0, - act_fn: str = "silu", - norm_num_groups: Optional[int] = 32, - norm_eps: float = 1e-5, - cross_attention_dim: Union[int, Tuple[int]] = 1280, - transformer_layers_per_block: Union[int, Tuple[int], Tuple[Tuple]] = 1, - reverse_transformer_layers_per_block: Optional[Tuple[Tuple[int]]] = None, - encoder_hid_dim: Optional[int] = None, - encoder_hid_dim_type: Optional[str] = None, - attention_head_dim: Union[int, Tuple[int]] = 8, - num_attention_heads: Optional[Union[int, Tuple[int]]] = None, - dual_cross_attention: bool = False, - use_linear_projection: bool = False, - class_embed_type: Optional[str] = None, - addition_embed_type: Optional[str] = None, - addition_time_embed_dim: Optional[int] = None, - num_class_embeds: Optional[int] = None, - upcast_attention: bool = False, - resnet_time_scale_shift: str = "default", - resnet_skip_time_act: bool = False, - resnet_out_scale_factor: float = 1.0, - time_embedding_type: str = "positional", - time_embedding_dim: Optional[int] = None, - time_embedding_act_fn: Optional[str] = None, - timestep_post_act: Optional[str] = None, - timestep_scale: Optional[float] = 1, - time_cond_proj_dim: Optional[int] = None, - conv_in_kernel: int = 3, - conv_out_kernel: int = 3, - projection_class_embeddings_input_dim: Optional[int] = None, - attention_type: str = "default", - class_embeddings_concat: bool = False, - mid_block_only_cross_attention: Optional[bool] = None, - cross_attention_norm: Optional[str] = None, - addition_embed_type_num_heads: int = 64, - ): - super().__init__( - sample_size=sample_size, - in_channels=in_channels, - out_channels=out_channels, - center_input_sample=center_input_sample, - flip_sin_to_cos=flip_sin_to_cos, - freq_shift=freq_shift, - down_block_types=down_block_types, - mid_block_type=mid_block_type, - up_block_types=up_block_types, - only_cross_attention=only_cross_attention, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - downsample_padding=downsample_padding, - mid_block_scale_factor=mid_block_scale_factor, - dropout=dropout, - act_fn=act_fn, - norm_num_groups=norm_num_groups, - norm_eps=norm_eps, - cross_attention_dim=cross_attention_dim, - transformer_layers_per_block=transformer_layers_per_block, - reverse_transformer_layers_per_block=reverse_transformer_layers_per_block, - encoder_hid_dim=encoder_hid_dim, - encoder_hid_dim_type=encoder_hid_dim_type, - attention_head_dim=attention_head_dim, - num_attention_heads=num_attention_heads, - dual_cross_attention=dual_cross_attention, - use_linear_projection=use_linear_projection, - class_embed_type=class_embed_type, - addition_embed_type=addition_embed_type, - addition_time_embed_dim=addition_time_embed_dim, - num_class_embeds=num_class_embeds, - upcast_attention=upcast_attention, - resnet_time_scale_shift=resnet_time_scale_shift, - resnet_skip_time_act=resnet_skip_time_act, - resnet_out_scale_factor=resnet_out_scale_factor, - time_embedding_type=time_embedding_type, - time_embedding_dim=time_embedding_dim, - time_embedding_act_fn=time_embedding_act_fn, - timestep_post_act=timestep_post_act, - time_cond_proj_dim=time_cond_proj_dim, - conv_in_kernel=conv_in_kernel, - conv_out_kernel=conv_out_kernel, - projection_class_embeddings_input_dim=projection_class_embeddings_input_dim, - attention_type=attention_type, - class_embeddings_concat=class_embeddings_concat, - mid_block_only_cross_attention=mid_block_only_cross_attention - or False, # default to False - cross_attention_norm=cross_attention_norm - or "default", # default to "default" - addition_embed_type_num_heads=addition_embed_type_num_heads, - ) - self.time_proj.scale = timestep_scale - self.pos_enc_conv = nn.Conv2d(2, self.conv_in.out_channels, 1, 1, 0) - nn.init.zeros_(self.pos_enc_conv.weight) - nn.init.zeros_(self.pos_enc_conv.bias) - - def forward( - self, - sample: torch.Tensor, - timestep: Union[torch.Tensor, float, int], - encoder_hidden_states: torch.Tensor, - class_labels: Optional[torch.Tensor] = None, - timestep_cond: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - cross_attention_kwargs: Optional[Dict[str, Any]] = None, - added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None, - down_block_additional_residuals: Optional[Tuple[torch.Tensor]] = None, - mid_block_additional_residual: Optional[torch.Tensor] = None, - down_intrablock_additional_residuals: Optional[Tuple[torch.Tensor]] = None, - encoder_attention_mask: Optional[torch.Tensor] = None, - return_dict: bool = True, - pos_map: Optional[torch.Tensor] = None, - ) -> Union[UNet2DConditionOutput, Tuple]: - B, C, H, W = sample.shape - timestep = timestep.view(-1) if isinstance(timestep, torch.Tensor) else timestep - pos_map = ( - pos_map - if pos_map is not None - else torch.zeros((B, H * W, 2), device=sample.device, dtype=sample.dtype) - ) - pos_map = pos_map.view(B, H, W, 2).permute(0, 3, 1, 2) - - # By default samples have to be AT least a multiple of the overall upsampling factor. - # The overall upsampling factor is equal to 2 ** (# num of upsampling layers). - # However, the upsampling interpolation output size can be forced to fit any upsampling size - # on the fly if necessary. - default_overall_up_factor = 2**self.num_upsamplers - - # upsample size should be forwarded when sample is not a multiple of `default_overall_up_factor` - forward_upsample_size = False - upsample_size = None - - for dim in sample.shape[-2:]: - if dim % default_overall_up_factor != 0: - # Forward upsample size to force interpolation output size. - forward_upsample_size = True - break - - # ensure attention_mask is a bias, and give it a singleton query_tokens dimension - # expects mask of shape: - # [batch, key_tokens] - # adds singleton query_tokens dimension: - # [batch, 1, key_tokens] - # this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes: - # [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn) - # [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn) - if attention_mask is not None: - # assume that mask is expressed as: - # (1 = keep, 0 = discard) - # convert mask into a bias that can be added to attention scores: - # (keep = +0, discard = -10000.0) - attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None: - encoder_attention_mask = ( - 1 - encoder_attention_mask.to(sample.dtype) - ) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - # 0. center input if necessary - if self.config.center_input_sample: - sample = 2 * sample - 1.0 - - # 1. time - t_emb = self.get_time_embed(sample=sample, timestep=timestep) - emb = self.time_embedding(t_emb, timestep_cond) - - class_emb = self.get_class_embed(sample=sample, class_labels=class_labels) - if class_emb is not None: - if self.config.class_embeddings_concat: - emb = torch.cat([emb, class_emb], dim=-1) - else: - emb = emb + class_emb - - aug_emb = self.get_aug_embed( - emb=emb, - encoder_hidden_states=encoder_hidden_states, - added_cond_kwargs=added_cond_kwargs, - ) - if self.config.addition_embed_type == "image_hint": - aug_emb, hint = aug_emb - sample = torch.cat([sample, hint], dim=1) - - emb = emb + aug_emb if aug_emb is not None else emb - - if self.time_embed_act is not None: - emb = self.time_embed_act(emb) - - encoder_hidden_states = self.process_encoder_hidden_states( - encoder_hidden_states=encoder_hidden_states, - added_cond_kwargs=added_cond_kwargs, - ) - - # 2. pre-process - sample = self.conv_in(sample) - pos_enc = self.pos_enc_conv(pos_map) - sample = sample + pos_enc - - # 2.5 GLIGEN position net - if ( - cross_attention_kwargs is not None - and cross_attention_kwargs.get("gligen", None) is not None - ): - cross_attention_kwargs = cross_attention_kwargs.copy() - gligen_args = cross_attention_kwargs.pop("gligen") - cross_attention_kwargs["gligen"] = { - "objs": self.position_net(**gligen_args) - } - - # 3. down - # we're popping the `scale` instead of getting it because otherwise `scale` will be propagated - # to the internal blocks and will raise deprecation warnings. this will be confusing for our users. - if cross_attention_kwargs is not None: - cross_attention_kwargs = cross_attention_kwargs.copy() - lora_scale = cross_attention_kwargs.pop("scale", 1.0) - else: - lora_scale = 1.0 - - is_controlnet = ( - mid_block_additional_residual is not None - and down_block_additional_residuals is not None - ) - # using new arg down_intrablock_additional_residuals for T2I-Adapters, to distinguish from controlnets - is_adapter = down_intrablock_additional_residuals is not None - # maintain backward compatibility for legacy usage, where - # T2I-Adapter and ControlNet both use down_block_additional_residuals arg - # but can only use one or the other - if ( - not is_adapter - and mid_block_additional_residual is None - and down_block_additional_residuals is not None - ): - down_intrablock_additional_residuals = down_block_additional_residuals - is_adapter = True - - down_block_res_samples = (sample,) - for downsample_block in self.down_blocks: - if ( - hasattr(downsample_block, "has_cross_attention") - and downsample_block.has_cross_attention - ): - # For t2i-adapter CrossAttnDownBlock2D - additional_residuals = {} - if is_adapter and len(down_intrablock_additional_residuals) > 0: - additional_residuals["additional_residuals"] = ( - down_intrablock_additional_residuals.pop(0) - ) - - sample, res_samples = downsample_block( - hidden_states=sample, - temb=emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - encoder_attention_mask=encoder_attention_mask, - **additional_residuals, - ) - else: - sample, res_samples = downsample_block(hidden_states=sample, temb=emb) - if is_adapter and len(down_intrablock_additional_residuals) > 0: - sample += down_intrablock_additional_residuals.pop(0) - - down_block_res_samples += res_samples - - if is_controlnet: - new_down_block_res_samples = () - - for down_block_res_sample, down_block_additional_residual in zip( - down_block_res_samples, down_block_additional_residuals - ): - down_block_res_sample = ( - down_block_res_sample + down_block_additional_residual - ) - new_down_block_res_samples = new_down_block_res_samples + ( - down_block_res_sample, - ) - - down_block_res_samples = new_down_block_res_samples - - # 4. mid - if self.mid_block is not None: - if ( - hasattr(self.mid_block, "has_cross_attention") - and self.mid_block.has_cross_attention - ): - sample = self.mid_block( - sample, - emb, - encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - encoder_attention_mask=encoder_attention_mask, - ) - else: - sample = self.mid_block(sample, emb) - - # To support T2I-Adapter-XL - if ( - is_adapter - and len(down_intrablock_additional_residuals) > 0 - and sample.shape == down_intrablock_additional_residuals[0].shape - ): - sample += down_intrablock_additional_residuals.pop(0) - - if is_controlnet: - sample = sample + mid_block_additional_residual - - # 5. up - for i, upsample_block in enumerate(self.up_blocks): - is_final_block = i == len(self.up_blocks) - 1 - - res_samples = down_block_res_samples[-len(upsample_block.resnets) :] - down_block_res_samples = down_block_res_samples[ - : -len(upsample_block.resnets) - ] - - # if we have not reached the final block and need to forward the - # upsample size, we do it here - if not is_final_block and forward_upsample_size: - upsample_size = down_block_res_samples[-1].shape[2:] - - if ( - hasattr(upsample_block, "has_cross_attention") - and upsample_block.has_cross_attention - ): - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - upsample_size=upsample_size, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - ) - else: - sample = upsample_block( - hidden_states=sample, - temb=emb, - res_hidden_states_tuple=res_samples, - upsample_size=upsample_size, - ) - - # 6. post-process - if self.conv_norm_out: - sample = self.conv_norm_out(sample) - sample = self.conv_act(sample) - sample = self.conv_out(sample) - - if not return_dict: - return (sample,) - - return UNet2DConditionOutput(sample=sample) diff --git a/pipelines/hdm/hdm/modules/rope.py b/pipelines/hdm/hdm/modules/rope.py deleted file mode 100644 index 6ab8a6f20..000000000 --- a/pipelines/hdm/hdm/modules/rope.py +++ /dev/null @@ -1,107 +0,0 @@ -import math -from functools import cache - -import torch -import torch.nn as nn -import torch.nn.functional as F - - -@cache -def bounding_box(h, w, pixel_aspect_ratio=1.0): - # Adjusted dimensions - w_adj = w - h_adj = h * pixel_aspect_ratio - - # Adjusted aspect ratio - ar_adj = w_adj / h_adj - - # Determine bounding box based on the adjusted aspect ratio - y_min, y_max, x_min, x_max = -1.0, 1.0, -1.0, 1.0 - if ar_adj > 1: - y_min, y_max = -1 / ar_adj, 1 / ar_adj - elif ar_adj < 1: - x_min, x_max = -ar_adj, ar_adj - - return y_min, y_max, x_min, x_max - - -@cache -def make_grid(h_pos, w_pos): - grid = torch.stack(torch.meshgrid(h_pos, w_pos, indexing="ij"), dim=-1) - h, w, d = grid.shape - return grid.view(h * w, d) - - -@cache -def centers(start, stop, num, dtype=None, device=None): - edges = torch.linspace(start, stop, num + 1, dtype=dtype, device=device) - return (edges[:-1] + edges[1:]) / 2 - - -@cache -def make_axial_pos( - h, w, pixel_aspect_ratio=1.0, align_corners=False, dtype=None, device=None -): - y_min, y_max, x_min, x_max = bounding_box(h, w, pixel_aspect_ratio) - if align_corners: - h_pos = torch.linspace(y_min, y_max, h, dtype=dtype, device=device) - w_pos = torch.linspace(x_min, x_max, w, dtype=dtype, device=device) - else: - h_pos = centers(y_min, y_max, h, dtype=dtype, device=device) - w_pos = centers(x_min, x_max, w, dtype=dtype, device=device) - return make_grid(h_pos, w_pos) - - -def rotate_half(x): - x = torch.stack((-x[..., 0::2], x[..., 1::2]), dim=-1) - return x.flatten(-2, -1) - - -def apply_rotary_emb(freqs, t, start_index=0, scale=1.0): - freqs = freqs.to(t) - rot_dim = freqs.shape[-1] - end_index = start_index + rot_dim - t_left, t, t_right = ( - t[..., :start_index], - t[..., start_index:end_index], - t[..., end_index:], - ) - t = (t * freqs.cos() * scale) + (rotate_half(t) * freqs.sin() * scale) - return torch.cat((t_left, t, t_right), dim=-1) - - -def freqs_pixel_log(max_freq=10.0): - def init(shape): - log_min = math.log(math.pi) - log_max = math.log(max_freq * math.pi / 2) - return torch.linspace(log_min, log_max, shape[-1]).expand(shape) - - return init - - -class AxialRoPE(nn.Module): - def __init__( - self, dim, n_heads, start_index=0, freqs_init=freqs_pixel_log(max_freq=10.0) - ): - super().__init__() - self.n_heads = n_heads - self.start_index = start_index - log_freqs = freqs_init((n_heads, dim // 4)) - self.freqs_h = nn.Parameter(log_freqs.clone()) - self.freqs_w = nn.Parameter(log_freqs.clone()) - - def extra_repr(self): - dim = (self.freqs_h.shape[-1] + self.freqs_w.shape[-1]) * 2 - return f"dim={dim}, n_heads={self.n_heads}, start_index={self.start_index}" - - def get_freqs(self, pos): - if pos.shape[-1] != 2: - raise ValueError("input shape must be (..., 2)") - freqs_h = pos[..., None, None, 0] * self.freqs_h.exp() - freqs_w = pos[..., None, None, 1] * self.freqs_w.exp() - freqs = torch.cat((freqs_h, freqs_w), dim=-1).repeat_interleave(2, dim=-1) - return freqs - - def forward(self, x, pos): - freqs = self.get_freqs(pos) - return apply_rotary_emb(freqs, x, self.start_index) diff --git a/pipelines/hdm/hdm/modules/text_encoders.py b/pipelines/hdm/hdm/modules/text_encoders.py deleted file mode 100644 index 727352750..000000000 --- a/pipelines/hdm/hdm/modules/text_encoders.py +++ /dev/null @@ -1,330 +0,0 @@ -from typing import Any - -import torch -import torch.nn as nn -import torch.nn.functional as F -from transformers import AutoTokenizer, CLIPTextModel, T5EncoderModel, Qwen2Model - -from ..utils import remove_none, instantiate - - -class BaseTextEncoder(nn.Module): - def __init__(self): - super().__init__() - self.tokenizer = None - self.text_model = None - - def tokenize(self, text: str) -> list[int] | list[list[int]] | torch.LongTensor: - raise NotImplementedError - - def encode(self, text: str) -> torch.Tensor: - raise NotImplementedError - - def forward(self, tokenizer_outputs: list[dict[str, torch.Tensor]]): - raise NotImplementedError - - -class SimpleTextEncoder(BaseTextEncoder): - def __init__( - self, - te_name: str = "apple/DFN5B-CLIP-ViT-H-14-378", - te_cls: type = CLIPTextModel, - te_kwargs: dict[str, Any] = {}, - zero_for_padding: bool = True, - max_length: int = 256, - ): - super().__init__() - self.tokenizers = [AutoTokenizer.from_pretrained(te_name, **te_kwargs)] - for tokenizer in self.tokenizers: - if not tokenizer.pad_token: - tokenizer.pad_token = tokenizer.eos_token - if tokenizer.model_max_length > max_length: - tokenizer.model_max_length = max_length - - self.text_model = ( - instantiate(te_cls).from_pretrained(te_name).to(self.device_type) - ) - self.zero_for_padding = zero_for_padding - - def tokenize(self, text, **kwargs): - return [self.tokenizers[0](text, **kwargs)] - - def encode(self, text, **kwargs): - return self.forward(self.tokenize(text, **kwargs)) - - def forward(self, tokenizers_outputs): - tokens = tokenizers_outputs[0] - text_model = self.text_model - - input_ids = tokens["input_ids"].to(self.device_type.device) - attn_mask = tokens["attention_mask"].to(self.device_type.device) - - # In CLIP we have `last_hidden_state = self.final_layer_norm(last_hidden_state)` - # The pooled embedding is also normalized - normed_embedding, pooled_embedding, *embeddings = text_model( - input_ids, - attention_mask=attn_mask, - output_hidden_states=True, - return_dict=False, - ) - if len(embeddings): - embedding = embeddings[-1][-1] - else: - embedding = pooled_embedding[-1] - pooled_embedding = None - if self.zero_for_padding: - while embedding.ndim > attn_mask.ndim: - attn_mask = attn_mask.unsqueeze(-1) - embedding = embedding * attn_mask - normed_embedding = normed_embedding * attn_mask - return embedding, normed_embedding, pooled_embedding, attn_mask - - -class ConcatTextEncoders(BaseTextEncoder): - DEFAULT_SETTINGS = { - "disable_autocast": False, - "concat_buckets": 0, - "use_pooled": False, - "need_mask": False, - "layer_ids": -1, - } - - def __init__( - self, - tokenizers: list[str] = [], - text_models: list[dict] = [], - zero_for_padding: bool = True, - max_length: int = 256, - model_dim: int = -1, - output_dim: int = -1, - pooled_dim: int = -1, - extra_mlp: bool = False, - ): - """ - A text encoder wrapper for multiple tokenizers and text models. - Can support tricky concat config like what SD3 need - - SDXL: - tes: [CLIP-L, openCLIP-G] - concat_buckets: [0, 0] - use_pooled: [True, True] - layer_index: [-1, -2] - SD3: - tes: [CLIP-L, openCLIP-G, T5-xxl] - concat_buckets: [0, 0, 1] - use_pooled: [True, True, False] - """ - super().__init__() - self.tokenizers = [ - AutoTokenizer.from_pretrained(tokenizer) for tokenizer in tokenizers - ] - for tokenizer in self.tokenizers: - if not tokenizer.pad_token: - tokenizer.pad_token = tokenizer.eos_token - if tokenizer.model_max_length > max_length: - tokenizer.model_max_length = max_length - - text_models_configs = [ - (instantiate(config.pop("model")), {**config}) for config in text_models - ] - self.max_bucket = max([i[1]["concat_buckets"] for i in text_models_configs]) - self.register_buffer("_device", torch.tensor(0), persistent=False) - - self.text_models = nn.ModuleList([i[0] for i in text_models_configs]) - self.configs = [i[1] for i in text_models_configs] - self.zero_for_padding = zero_for_padding - - self.emb_mlp = self.pool_mlp = None - if extra_mlp and model_dim != -1: - if output_dim != -1: - self.emb_mlp = nn.Sequential( - nn.LayerNorm(model_dim), - nn.Linear(model_dim, model_dim * 4), - nn.Mish(), - nn.Linear(model_dim * 4, output_dim), - ) - if pooled_dim != -1: - self.pool_mlp = nn.Sequential( - nn.LayerNorm(model_dim), - nn.Linear(model_dim, model_dim * 4), - nn.Mish(), - nn.Linear(model_dim * 4, pooled_dim), - ) - - def trainable_modules(self): - results = [] - if self.emb_mlp is not None: - results.append(self.emb_mlp) - if self.pool_mlp is not None: - results.append(self.pool_mlp) - return results - - def trainable_params(self): - results = [] - if self.emb_mlp is not None: - results.extend(self.emb_mlp.parameters()) - if self.pool_mlp is not None: - results.extend(self.pool_mlp.parameters()) - return results - - @property - def device(self): - return self._device.device - - def tokenize(self, text, **kwargs): - results = [] - for tokenizer in self.tokenizers: - results.append(tokenizer(text, **kwargs, return_tensors="pt")) - return results - - def encode(self, text, **kwargs): - return self.forward(self.tokenize(text, **kwargs)) - - def forward( - self, tokenizers_outputs - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - """ - Returns: - embedding: torch.Tensor - normed_embedding: torch.Tensor - pooled_embedding: torch.Tensor - attn_mask: torch.Tensor - """ - attn_masks = [None for _ in range(self.max_bucket + 1)] - text_embeddings = [[] for _ in range(self.max_bucket + 1)] - normed_text_embeddings = [[] for _ in range(self.max_bucket + 1)] - pooled_text_embeddings = [[] for _ in range(self.max_bucket + 1)] - for idx, (tokens, text_model, config) in enumerate( - zip(tokenizers_outputs, self.text_models, self.configs) - ): - bucket = config["concat_buckets"] - need_mask = config["need_mask"] - use_pooled = config["use_pooled"] - layer_idx = config["layer_idx"] - disable_autocast = config["disable_autocast"] - - input_ids = tokens["input_ids"].to(self.device) - attn_mask = tokens["attention_mask"].to(self.device) - if attn_masks[bucket] is None and need_mask: - attn_masks[bucket] = attn_mask - - with torch.autocast("cuda", enabled=not disable_autocast): - output = text_model( - input_ids, - attention_mask=attn_mask, - output_hidden_states=True, - return_dict=True, - ) - normed_embedding = output.last_hidden_state - # The case of CLIP - if hasattr(output, "pooler_output"): - # embeddings is tuple - embedding = output.hidden_states[layer_idx] - pooled_embedding = output.pooler_output - # The case of T5 or other models - else: - embedding = output.hidden_states[-1] - pooled_embedding = torch.zeros_like(embedding[:, 0, :]) - - if self.zero_for_padding: - while embedding.ndim > attn_mask.ndim: - attn_mask = attn_mask.unsqueeze(-1) - embedding = embedding * attn_mask - normed_embedding = normed_embedding * attn_mask - text_embeddings[bucket].append(embedding) - normed_text_embeddings[bucket].append(normed_embedding) - if use_pooled: - pooled_text_embeddings[bucket].append(pooled_embedding) - - for i in range(len(text_embeddings)): - if text_embeddings[i] == []: - text_embeddings[i] = None - normed_text_embeddings[i] = None - pooled_text_embeddings[i] = None - continue - text_embeddings[i] = torch.cat(text_embeddings[i], dim=-1) - normed_text_embeddings[i] = torch.cat(normed_text_embeddings[i], dim=-1) - if pooled_text_embeddings[i] == []: - pooled_text_embeddings[i] = None - continue - pooled_text_embeddings[i] = torch.cat(pooled_text_embeddings[i], dim=-1) - - max_dim = max( - embedding.size(-1) for embedding in text_embeddings if embedding is not None - ) - for idx, embedding in enumerate(text_embeddings): - if embedding is None: - continue - if embedding.size(-1) < max_dim: - text_embeddings[idx] = torch.nn.functional.pad( - embedding, (0, max_dim - embedding.size(-1)) - ) - for idx, embedding in enumerate(normed_text_embeddings): - if embedding is None: - continue - if embedding.size(-1) < max_dim: - normed_text_embeddings[idx] = torch.nn.functional.pad( - embedding, (0, max_dim - embedding.size(-1)) - ) - if any(mask is not None for mask in attn_masks): - for idx, embedding in enumerate(text_embeddings): - if embedding is None: - continue - elif attn_masks[idx] is None: - attn_masks[idx] = torch.ones( - embedding.size(0), embedding.size(1), device=embedding.device - ).long() - attn_masks = torch.cat(remove_none(attn_masks), dim=1) - else: - attn_masks = None - if any(pooled is not None for pooled in pooled_text_embeddings): - pooled_text_embeddings = torch.cat( - remove_none(pooled_text_embeddings), dim=-1 - ) - else: - pooled_text_embeddings = None - text_embeddings = torch.cat(remove_none(text_embeddings), dim=1) - normed_text_embeddings = torch.cat(remove_none(normed_text_embeddings), dim=1) - - if self.emb_mlp is not None: - text_embeddings = self.emb_mlp(text_embeddings) - normed_text_embeddings = self.emb_mlp(normed_text_embeddings) - if self.pool_mlp is not None and pooled_text_embeddings is not None: - pooled_text_embeddings = self.pool_mlp(pooled_text_embeddings) - - return ( - normed_text_embeddings, - text_embeddings, - pooled_text_embeddings, - attn_masks, - ) - - -if __name__ == "__main__": - te = ConcatTextEncoders( - tokenizers=[ - "openai/clip-vit-large-patch14", - "laion/CLIP-ViT-bigG-14-laion2B-39B-b160k", - "google/t5-v1_1-xxl", - ], - text_models=[ - (CLIPTextModel, "openai/clip-vit-large-patch14", {}), - (CLIPTextModel, "laion/CLIP-ViT-bigG-14-laion2B-39B-b160k", {}), - ( - T5EncoderModel, - "google/t5-v1_1-xxl", - {}, - ), # Need `pip install sentencepiece` - ], - concat_buckets=[0, 0, 1], - use_pooled=[True, True, False], - layer_idx=[-1, -2, -1], - need_mask=[False, False, True], - device="cuda" if torch.cuda.is_available() else "cpu", - ) - with torch.no_grad(): - text_embeddings, normed_text_embeddings, pooled_text_embeddings, attn_masks = ( - te.encode("hello") - ) - print(text_embeddings.shape, normed_text_embeddings.shape) - print(pooled_text_embeddings.shape) diff --git a/pipelines/hdm/hdm/modules/unet_patch.py b/pipelines/hdm/hdm/modules/unet_patch.py deleted file mode 100644 index 9fdc0dd25..000000000 --- a/pipelines/hdm/hdm/modules/unet_patch.py +++ /dev/null @@ -1,610 +0,0 @@ -import json -from typing import Any, Optional, Dict - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from diffusers import UNet2DConditionModel -from diffusers.models.unets.unet_2d_blocks import ( - ResnetBlock2D, -) -from diffusers.models.transformers.transformer_2d import ( - Transformer2DModel, - Transformer2DModelOutput, -) -from diffusers.models.attention import BasicTransformerBlock -from diffusers.models.attention_processor import ( - Attention, - XFormersAttnProcessor, - AttnProcessor2_0, -) - -try: - import xformers - import xformers.ops -except ImportError: - xformers = None - -from .rope import AxialRoPE, make_axial_pos -from ..utils import instantiate - - -class RoPEAttention(Attention): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - - head_dim = self.inner_dim // self.heads - self.axial_rope = AxialRoPE(head_dim, self.heads) - self.set_processor(RoPEAttnProcessor2_0()) - - @classmethod - def apply_to(cls, original: Attention): - original.axial_rope = AxialRoPE( - original.inner_dim // original.heads, original.heads - ) - original.set_processor(RoPEAttnProcessor2_0()) - original.forward = lambda *args, **kwargs: cls.forward( - original, *args, **kwargs - ) - return original - - def forward( - self, - hidden_states: torch.Tensor, - position_map: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - **cross_attention_kwargs, - ) -> torch.Tensor: - return self.processor( - self, - hidden_states, - position_map, - encoder_hidden_states, - attention_mask, - **cross_attention_kwargs, - ) - - -class RoPEAttnProcessor2_0(AttnProcessor2_0): - def __call__( - self, - attn: RoPEAttention, - hidden_states: torch.Tensor, - position_map: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - temb: Optional[torch.Tensor] = None, - *args, - **kwargs, - ) -> torch.Tensor: - residual = hidden_states - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view( - batch_size, channel, height * width - ).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape - if encoder_hidden_states is None - else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask( - attention_mask, sequence_length, batch_size - ) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view( - batch_size, attn.heads, -1, attention_mask.shape[-1] - ) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( - 1, 2 - ) - - query = attn.to_q(hidden_states) - - rotary_k = False - if encoder_hidden_states is None: - rotary_k = True - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states( - encoder_hidden_states - ) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim) - - key = key.view(batch_size, -1, attn.heads, head_dim) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - query = attn.axial_rope(query, position_map).transpose(1, 2) - if rotary_k: - key = attn.axial_rope(key, position_map).transpose(1, 2) - else: - key = key.transpose(1, 2) - - # the output of sdp = (batch, num_heads, seq_len, head_dim) - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape( - batch_size, -1, attn.heads * head_dim - ) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape( - batch_size, channel, height, width - ) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class RoPEXFormersAttnProcessor(XFormersAttnProcessor): - def __call__( - self, - attn: RoPEAttention, - hidden_states: torch.Tensor, - position_map: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - temb: Optional[torch.Tensor] = None, - *args, - **kwargs, - ) -> torch.Tensor: - residual = hidden_states - - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view( - batch_size, channel, height * width - ).transpose(1, 2) - - batch_size, key_tokens, _ = ( - hidden_states.shape - if encoder_hidden_states is None - else encoder_hidden_states.shape - ) - - attention_mask = attn.prepare_attention_mask( - attention_mask, key_tokens, batch_size - ) - if attention_mask is not None: - _, query_tokens, _ = hidden_states.shape - attention_mask = attention_mask.expand(-1, query_tokens, -1) - if attention_mask is not None and attention_mask.ndim == 3: - attention_mask = attention_mask.reshape( - batch_size, -1, *attention_mask.shape[-2:] - ) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( - 1, 2 - ) - - query = attn.to_q(hidden_states) - - rotary_k = False - if encoder_hidden_states is None: - rotary_k = True - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states( - encoder_hidden_states - ) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - query = query.reshape(batch_size, -1, attn.heads, head_dim) - key = key.reshape(batch_size, -1, attn.heads, head_dim) - value = value.reshape(batch_size, -1, attn.heads, head_dim) - - query = attn.axial_rope(query, position_map) - if rotary_k: - key = attn.axial_rope(key, position_map) - - if attention_mask is not None: - attention_mask = attention_mask.to(query) - hidden_states = xformers.ops.memory_efficient_attention( - query, - key, - value, - attn_bias=attention_mask, - op=self.attention_op, - scale=attn.scale, - ) - hidden_states = hidden_states.to(query.dtype) - hidden_states = hidden_states.reshape(batch_size, -1, inner_dim) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape( - batch_size, channel, height, width - ) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - - -class RoPEBasicTransformerBlock(BasicTransformerBlock): - @classmethod - def apply_to(cls, original: BasicTransformerBlock): - original.forward = lambda *args, **kwargs: cls.forward( - original, *args, **kwargs - ) - for module in original.modules(): - if isinstance(module, Attention): - RoPEAttention.apply_to(module) - - def forward( - self, - hidden_states: torch.Tensor, - position_map: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - encoder_hidden_states: Optional[torch.Tensor] = None, - encoder_attention_mask: Optional[torch.Tensor] = None, - timestep: Optional[torch.LongTensor] = None, - cross_attention_kwargs: Dict[str, Any] = None, - class_labels: Optional[torch.LongTensor] = None, - added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None, - ) -> torch.Tensor: - # Notice that normalization is always applied before the real computation in the following blocks. - # 0. Self-Attention - batch_size = hidden_states.shape[0] - - if self.norm_type == "ada_norm": - norm_hidden_states = self.norm1(hidden_states, timestep) - elif self.norm_type == "ada_norm_zero": - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1( - hidden_states, timestep, class_labels, hidden_dtype=hidden_states.dtype - ) - elif self.norm_type in ["layer_norm", "layer_norm_i2vgen"]: - norm_hidden_states = self.norm1(hidden_states) - elif self.norm_type == "ada_norm_continuous": - norm_hidden_states = self.norm1( - hidden_states, added_cond_kwargs["pooled_text_emb"] - ) - elif self.norm_type == "ada_norm_single": - shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( - self.scale_shift_table[None] + timestep.reshape(batch_size, 6, -1) - ).chunk(6, dim=1) - norm_hidden_states = self.norm1(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa - norm_hidden_states = norm_hidden_states.squeeze(1) - else: - raise ValueError("Incorrect norm used") - - if self.pos_embed is not None: - norm_hidden_states = self.pos_embed(norm_hidden_states) - - # 1. Prepare GLIGEN inputs - cross_attention_kwargs = ( - cross_attention_kwargs.copy() if cross_attention_kwargs is not None else {} - ) - gligen_kwargs = cross_attention_kwargs.pop("gligen", None) - - attn_output = self.attn1( - norm_hidden_states, - position_map, - encoder_hidden_states=( - encoder_hidden_states if self.only_cross_attention else None - ), - attention_mask=attention_mask, - **cross_attention_kwargs, - ) - if self.norm_type == "ada_norm_zero": - attn_output = gate_msa.unsqueeze(1) * attn_output - elif self.norm_type == "ada_norm_single": - attn_output = gate_msa * attn_output - - hidden_states = attn_output + hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - # 1.2 GLIGEN Control - if gligen_kwargs is not None: - hidden_states = self.fuser(hidden_states, gligen_kwargs["objs"]) - - # 3. Cross-Attention - if self.attn2 is not None: - if self.norm_type == "ada_norm": - norm_hidden_states = self.norm2(hidden_states, timestep) - elif self.norm_type in ["ada_norm_zero", "layer_norm", "layer_norm_i2vgen"]: - norm_hidden_states = self.norm2(hidden_states) - elif self.norm_type == "ada_norm_single": - # For PixArt norm2 isn't applied here: - # https://github.com/PixArt-alpha/PixArt-alpha/blob/0f55e922376d8b797edd44d25d0e7464b260dcab/diffusion/model/nets/PixArtMS.py#L70C1-L76C103 - norm_hidden_states = hidden_states - elif self.norm_type == "ada_norm_continuous": - norm_hidden_states = self.norm2( - hidden_states, added_cond_kwargs["pooled_text_emb"] - ) - else: - raise ValueError("Incorrect norm") - - if self.pos_embed is not None and self.norm_type != "ada_norm_single": - norm_hidden_states = self.pos_embed(norm_hidden_states) - - attn_output = self.attn2( - norm_hidden_states, - position_map, - encoder_hidden_states=encoder_hidden_states, - attention_mask=encoder_attention_mask, - **cross_attention_kwargs, - ) - hidden_states = attn_output + hidden_states - - # 4. Feed-forward - # i2vgen doesn't have this norm 🤷‍♂️ - if self.norm_type == "ada_norm_continuous": - norm_hidden_states = self.norm3( - hidden_states, added_cond_kwargs["pooled_text_emb"] - ) - elif not self.norm_type == "ada_norm_single": - norm_hidden_states = self.norm3(hidden_states) - - if self.norm_type == "ada_norm_zero": - norm_hidden_states = ( - norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - ) - - if self.norm_type == "ada_norm_single": - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp - - ff_output = self.ff(norm_hidden_states) - - if self.norm_type == "ada_norm_zero": - ff_output = gate_mlp.unsqueeze(1) * ff_output - elif self.norm_type == "ada_norm_single": - ff_output = gate_mlp * ff_output - - hidden_states = ff_output + hidden_states - if hidden_states.ndim == 4: - hidden_states = hidden_states.squeeze(1) - - return hidden_states - - -class RoPETransformer2DModel(Transformer2DModel): - _org_init = Transformer2DModel.__init__ - - def __init__(self, *args, **kwargs): - RoPETransformer2DModel._org_init(self, *args, **kwargs) - for block in self.transformer_blocks: - if isinstance(block, BasicTransformerBlock): - RoPEBasicTransformerBlock.apply_to(block) - - def forward( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: Optional[torch.Tensor] = None, - timestep: Optional[torch.LongTensor] = None, - added_cond_kwargs: Dict[str, torch.Tensor] = None, - class_labels: Optional[torch.LongTensor] = None, - cross_attention_kwargs: Dict[str, Any] = None, - attention_mask: Optional[torch.Tensor] = None, - encoder_attention_mask: Optional[torch.Tensor] = None, - position_map: Optional[torch.Tensor] = None, - return_dict: bool = True, - ): - if attention_mask is not None and attention_mask.ndim == 2: - # assume that mask is expressed as: - # (1 = keep, 0 = discard) - # convert mask into a bias that can be added to attention scores: - # (keep = +0, discard = -10000.0) - attention_mask = (1 - attention_mask.to(hidden_states.dtype)) * -10000.0 - attention_mask = attention_mask.unsqueeze(1) - - # convert encoder_attention_mask to a bias the same way we do for attention_mask - if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2: - encoder_attention_mask = ( - 1 - encoder_attention_mask.to(hidden_states.dtype) - ) * -10000.0 - encoder_attention_mask = encoder_attention_mask.unsqueeze(1) - - # 1. Input - if self.is_input_continuous: - batch_size, _, height, width = hidden_states.shape - residual = hidden_states - hidden_states, inner_dim = self._operate_on_continuous_inputs(hidden_states) - elif self.is_input_vectorized: - height = self.latent_image_embedding.height - width = self.latent_image_embedding.width - hidden_states = self.latent_image_embedding(hidden_states) - elif self.is_input_patches: - height, width = ( - hidden_states.shape[-2] // self.patch_size, - hidden_states.shape[-1] // self.patch_size, - ) - hidden_states, encoder_hidden_states, timestep, embedded_timestep = ( - self._operate_on_patched_inputs( - hidden_states, encoder_hidden_states, timestep, added_cond_kwargs - ) - ) - if position_map is None: - position_map = make_axial_pos( - h=height, - w=width, - device=hidden_states.device, - dtype=hidden_states.dtype, - ) - else: - position_map = position_map.to(hidden_states) - assert position_map.shape[-3:] == (height, width, 2) - - # 2. Blocks - for block in self.transformer_blocks: - if self.training 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 - - ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} - hidden_states = torch.utils.checkpoint.checkpoint( - create_custom_forward(block), - hidden_states, - position_map, - attention_mask, - encoder_hidden_states, - encoder_attention_mask, - timestep, - cross_attention_kwargs, - class_labels, - **ckpt_kwargs, - ) - else: - hidden_states = block( - hidden_states, - position_map, - attention_mask=attention_mask, - encoder_hidden_states=encoder_hidden_states, - encoder_attention_mask=encoder_attention_mask, - timestep=timestep, - cross_attention_kwargs=cross_attention_kwargs, - class_labels=class_labels, - ) - - # 3. Output - if self.is_input_continuous: - output = self._get_output_for_continuous_inputs( - hidden_states=hidden_states, - residual=residual, - batch_size=batch_size, - height=height, - width=width, - inner_dim=inner_dim, - ) - elif self.is_input_vectorized: - output = self._get_output_for_vectorized_inputs(hidden_states) - elif self.is_input_patches: - output = self._get_output_for_patched_inputs( - hidden_states=hidden_states, - timestep=timestep, - class_labels=class_labels, - embedded_timestep=embedded_timestep, - height=height, - width=width, - ) - - if not return_dict: - return (output,) - - return Transformer2DModelOutput(sample=output) - - -org_init = Transformer2DModel.__init__ -org_forward = Transformer2DModel.forward - - -def apply_patch(): - import diffusers.models.transformers.transformer_2d as transformer_2d - - transformer_2d.Transformer2DModel.__init__ = RoPETransformer2DModel.__init__ - transformer_2d.Transformer2DModel.forward = RoPETransformer2DModel.forward - - -def restore(): - import diffusers.models.transformers.transformer_2d as transformer_2d - - transformer_2d.Transformer2DModel.__init__ = org_init - transformer_2d.Transformer2DModel.forward = org_forward - - -class HDUNet2DConditionModel(UNet2DConditionModel): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - for module in self.modules(): - if isinstance(module, BasicTransformerBlock): - nn.init.constant_(module.attn1.to_out[0].weight, 0.0) - if module.attn2 is not None: - nn.init.constant_(module.attn2.to_out[0].weight, 0.0) - if isinstance(module.ff.net[-2], nn.Linear): - nn.init.constant_(module.ff.net[-2].weight, 0.0) - nn.init.constant_(module.ff.net[-2].bias, 0.0) - else: - nn.init.constant_(module.ff.net[-1].weight, 0.0) - nn.init.constant_(module.ff.net[-1].bias, 0.0) - if isinstance(module, ResnetBlock2D): - nn.init.constant_(module.conv2.weight, 0.0) - nn.init.constant_(module.conv2.bias, 0.0) - nn.init.constant_(self.conv_out.weight, 0.0) - - @classmethod - def from_config(cls, arch: dict): - if isinstance(arch, str): - with open(arch, "r") as f: - arch = json.load(f) - return cls(**instantiate(arch)) - - -class RoPEUNet2DConditionModel(HDUNet2DConditionModel): - def __init__(self, *args, **kwargs): - apply_patch() - super().__init__(*args, **kwargs) - restore() - if xformers is not None: - self.set_attn_processor(RoPEXFormersAttnProcessor()) - - @classmethod - def from_config(cls, arch: dict): - if isinstance(arch, str): - with open(arch, "r") as f: - arch = json.load(f) - return cls(**instantiate(arch)) - - def forward(self, *args, **kwargs): - apply_patch() - result = super().forward(*args, **kwargs) - restore() - return result diff --git a/pipelines/hdm/hdm/modules/xut.py b/pipelines/hdm/hdm/modules/xut.py deleted file mode 100644 index 619b79903..000000000 --- a/pipelines/hdm/hdm/modules/xut.py +++ /dev/null @@ -1,98 +0,0 @@ -import json -import torch - -from diffusers.configuration_utils import ConfigMixin, register_to_config -from diffusers.models.modeling_utils import ModelMixin - -from ...xut.xut import XUDiT -from .base import * - - -class XUDiTConditionModel(ModelMixin, ConfigMixin): - _supports_gradient_checkpointing = True - - @register_to_config - def __init__( - self, - patch_size=2, - input_dim=4, - dim=1024, - ctx_dim=1024, - ctx_size=256, - heads=16, - dim_head=64, - mlp_dim=3072, - depth=8, - enc_blocks=1, - dec_blocks=2, - dec_ctx=False, - class_cond=0, - shared_adaln=True, - concat_ctx=True, - use_dyt=False, - double_t=False, - addon_info_embs_dim=None, - tread_config=None, - ): - super().__init__() - self.model = XUDiT( - patch_size=patch_size, - input_dim=input_dim, - dim=dim, - ctx_dim=ctx_dim, - ctx_size=ctx_size, - heads=heads, - dim_head=dim_head, - mlp_dim=mlp_dim, - depth=depth, - enc_blocks=enc_blocks, - dec_blocks=dec_blocks, - dec_ctx=dec_ctx, - class_cond=class_cond, - shared_adaln=shared_adaln, - concat_ctx=concat_ctx, - use_dyt=use_dyt, - double_t=double_t, - addon_info_embs_dim=addon_info_embs_dim, - tread_config=tread_config, - ) - - @classmethod - def from_config(cls, config: Dict[str, Any] | str) -> "XUDiTConditionModel": - if isinstance(config, str): - with open(config, "r") as f: - config = json.load(f) - return cls(**config) - - def enable_gradient_checkpointing(self): - return self.model.set_grad_ckpt(True) - - def disable_gradient_checkpointing(self): - return self.model.set_grad_ckpt(False) - - def forward( - self, - sample: torch.Tensor, - timestep: Union[torch.Tensor, float, int], - encoder_hidden_states: torch.Tensor, - class_labels: Optional[torch.Tensor] = None, - timestep_cond: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - cross_attention_kwargs: Optional[Dict[str, Any]] = None, - added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None, - down_block_additional_residuals: Optional[Tuple[torch.Tensor]] = None, - mid_block_additional_residual: Optional[torch.Tensor] = None, - down_intrablock_additional_residuals: Optional[Tuple[torch.Tensor]] = None, - encoder_attention_mask: Optional[torch.Tensor] = None, - return_dict: bool = True, - pos_map: Optional[torch.Tensor] = None, - ) -> Union[UNet2DConditionOutput, Tuple]: - if added_cond_kwargs is None: - added_cond_kwargs = {} - result = self.model( - sample, timestep, encoder_hidden_states, pos_map, **added_cond_kwargs - ) - if return_dict: - return UNet2DConditionOutput(sample=result) - else: - return (sample,) diff --git a/pipelines/hdm/hdm/pipeline.py b/pipelines/hdm/hdm/pipeline.py deleted file mode 100644 index 1f9eaa4b2..000000000 --- a/pipelines/hdm/hdm/pipeline.py +++ /dev/null @@ -1,165 +0,0 @@ -from typing import Optional, Tuple, Union - -import torch - -from diffusers import DiffusionPipeline, ImagePipelineOutput -from diffusers import AutoencoderKL -from transformers import Qwen3Model, Qwen2Tokenizer - -from .modules.xut import XUDiTConditionModel -from ..xut.modules.axial_rope import make_axial_pos_no_cache - - -class HDMXUTPipeline(DiffusionPipeline): - transformer: XUDiTConditionModel - tokenizer = Qwen2Tokenizer - text_encoder: Qwen3Model - vae: AutoencoderKL - - def __init__( - self, - transformer: XUDiTConditionModel, - text_encoder: Qwen3Model, - tokenizer: Qwen2Tokenizer, - vae: AutoencoderKL, - scheduler, - ): - super().__init__() - self.register_modules( - transformer=transformer, - text_encoder=text_encoder, - tokenizer=tokenizer, - vae=vae, - scheduler=scheduler, - ) - self.vae_mean = torch.tensor(self.vae.config.latents_mean)[None, :, None, None] - self.vae_std = torch.tensor(self.vae.config.latents_std)[None, :, None, None] - - def apply_compile(self, *args, **kwargs): - self.transformer.model.prev_tread_trns = torch.compile( - self.transformer.model.prev_tread_trns, *args, **kwargs - ) - self.transformer.model.backbone = torch.compile( - self.transformer.model.backbone, *args, **kwargs - ) - self.transformer.model.post_tread_trns = torch.compile( - self.transformer.model.post_tread_trns, *args, **kwargs - ) - self.vae.encoder = torch.compile(self.vae.encoder, *args, **kwargs) - self.vae.decoder = torch.compile(self.vae.decoder, *args, **kwargs) - - @torch.no_grad() - def __call__( - self, - prompt: str = "a photo of a dog", - negative_prompt: str = "", - width: int = 1024, - height: int = 1024, - cfg_scale: float = 3.0, - num_inference_steps: int = 16, - camera_param: dict[str, float] = { - "zoom": 1.0, - "x_shift": 0.0, - "y_shift": 0.0, - }, - tread_gamma1: float = 0.0, - tread_gamma2: float = 0.25, - generator: Optional[torch.Generator] = None, - output_type: Optional[str] = "pil", - return_dict: bool = True, - **kwargs, - ) -> Union[ImagePipelineOutput, Tuple]: - if isinstance(prompt, str): - prompt = [prompt] - if isinstance(negative_prompt, str): - negative_prompt = [negative_prompt] - if len(negative_prompt) == 1: - negative_prompt = negative_prompt * len(prompt) - - prompt_tokens = self.tokenizer( - prompt, - padding="longest", - return_tensors="pt", - ) - negative_prompt_tokens = self.tokenizer( - negative_prompt, - padding="longest", - return_tensors="pt", - ) - - prompt_emb = self.text_encoder( - input_ids=prompt_tokens.input_ids.to(self.device), - attention_mask=prompt_tokens.attention_mask.to(self.device), - ).last_hidden_state - negative_prompt_emb = self.text_encoder( - input_ids=negative_prompt_tokens.input_ids.to(self.device), - attention_mask=negative_prompt_tokens.attention_mask.to(self.device), - ).last_hidden_state - - # Sample gaussian noise to begin loop - image = torch.randn( - ( - len(prompt), - self.transformer.config.input_dim, - height // 16 * 2, - width // 16 * 2, - ), - generator=generator[0], - ) - image = image.to(self.device).to(self.dtype) - aspect_ratio = ( - torch.tensor([width / height], device=self.device) - .log() - .repeat(image.size(0)) - ).to(self.dtype) - - latent_h, latent_w = image.shape[-2:] - pos_map = make_axial_pos_no_cache(latent_h, latent_w, device=self.device) - pos_map[..., 0] = pos_map[..., 0] + camera_param.get("y_shift", 0.0) - pos_map[..., 1] = pos_map[..., 1] + camera_param.get("x_shift", 0.0) - pos_map = pos_map / camera_param.get("zoom", 1.0) - pos_map = pos_map[None].expand(image.size(0), -1, -1).to(self.dtype) - - t = torch.tensor([1] * image.size(0), device=self.device).to(self.dtype) - current_t = 1.0 - dt = 1.0 / num_inference_steps - - for _ in (pbar := self.progress_bar(range(num_inference_steps))): - cond = self.transformer( - image.to(self.dtype), - t, - prompt_emb, - added_cond_kwargs={ - "addon_info": aspect_ratio, - "tread_rate": tread_gamma1, - }, - pos_map=pos_map, - ).sample.float() - uncond = self.transformer( - image.to(self.dtype), - t, - negative_prompt_emb, - added_cond_kwargs={ - "addon_info": aspect_ratio, - "tread_rate": tread_gamma2, - }, - pos_map=pos_map, - ).sample.float() - cfg_flow = uncond + cfg_scale * (cond - uncond) - image = image - dt * cfg_flow - t = t - dt - current_t -= dt - - torch.cuda.empty_cache() - image = image * self.vae_std.to(self.device) + self.vae_mean.to(self.device) - image = torch.concat([self.vae.decode(i[None].to(self.dtype)).sample for i in image]) - image = (image.float() / 2 + 0.5).clamp(0, 1) - image = image.cpu().permute(0, 2, 3, 1).numpy() - - if output_type == "pil": - image = self.numpy_to_pil(image) - - if not return_dict: - return (image,) - - return ImagePipelineOutput(images=image) diff --git a/pipelines/hdm/hdm/trainer/__init__.py b/pipelines/hdm/hdm/trainer/__init__.py deleted file mode 100644 index 7bdfa7e7b..000000000 --- a/pipelines/hdm/hdm/trainer/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from .trainer import DMTrainer, FlowTrainer diff --git a/pipelines/hdm/hdm/trainer/callbacks.py b/pipelines/hdm/hdm/trainer/callbacks.py deleted file mode 100644 index 16fb22fdf..000000000 --- a/pipelines/hdm/hdm/trainer/callbacks.py +++ /dev/null @@ -1,66 +0,0 @@ -from operator import is_ -import os - -import torch -import wandb - -from lightning.pytorch import Callback, Trainer -from hdm.trainer import DMTrainer - - -class ImageGenCallback(Callback): - def __init__(self, config, img_gen_func): - self.config = { - "period": 100, - "num": 4, - "preview_num": 4, - "batch_size": 4, - "steps": 24, - } - self.config.update(config) - self.img_gen = img_gen_func - - @torch.no_grad() - def on_train_batch_start( - self, trainer: Trainer, pl_module: DMTrainer, batch, batch_idx - ): - if batch_idx % self.config["period"] == 0: - is_training = pl_module.training - pl_module.eval() - torch.cuda.empty_cache() - captions, images = self.img_gen(pl_module, batch, self.config) - torch.cuda.empty_cache() - - if hasattr(trainer.logger, "id"): - id = trainer.logger.id - elif hasattr(trainer.logger, "experiment"): - id = getattr(trainer.logger.experiment, "id", self.config.get("id", 0)) - else: - id = self.config.get("id", 0) - if not isinstance(id, (str, bytes, int, float)): - id = self.config.get("id", 0) - if "id" in self.config: - id = self.config["id"] - - rank = trainer.local_rank - base_idx = rank * self.config["num"] - - os.makedirs(f"./sample/{id}/{trainer.global_step}", exist_ok=True) - data = [] - for idx, (caption, image) in enumerate(zip(captions, images)): - idx = base_idx + idx - image.save(f"./sample/{id}/{trainer.global_step}/{idx}.png") - data.append( - [ - caption, - wandb.Image(f"./sample/{id}/{trainer.global_step}/{idx}.png"), - ] - ) - if trainer.is_global_zero: - trainer.logger.log_table( - key="sample/images", - columns=["caption", "image"], - data=data[: self.config["preview_num"]], - ) - torch.cuda.empty_cache() - pl_module.train(is_training) diff --git a/pipelines/hdm/hdm/trainer/diffusion.py b/pipelines/hdm/hdm/trainer/diffusion.py deleted file mode 100644 index 7d2209c3b..000000000 --- a/pipelines/hdm/hdm/trainer/diffusion.py +++ /dev/null @@ -1,68 +0,0 @@ -import torch - -# import torch.nn as nn -# import torch.nn.functional as F -# import torch.optim as optim - -from diffusers import EulerDiscreteScheduler - - -def get_noise_noisy_latents_and_timesteps( - noise_scheduler: EulerDiscreteScheduler, latents -): - noise = torch.randn_like(latents, device=latents.device) - b_size = latents.shape[0] - min_timestep = 0 - max_timestep = noise_scheduler.config.num_train_timesteps - - timesteps = torch.randint( - min_timestep, max_timestep, (b_size,), device=latents.device - ) - - sigmas = noise_scheduler.sigmas.to(device=latents.device, dtype=latents.dtype) - schedule_timesteps = noise_scheduler.timesteps.to(latents.device) - timesteps = timesteps.to(latents.device) - step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] - sigma = sigmas[step_indices].flatten() - while len(sigma.shape) < len(latents.shape): - sigma = sigma.unsqueeze(-1) - - # Diffusion Forward process - noisy_samples = latents + noise * sigma - scale = 1 / (sigma**2 + 1) ** 0.5 - return noisy_samples * scale, noise, timesteps - - -def apply_snr_weight(loss, timesteps, noise_scheduler, gamma, v_prediction=False): - snr = torch.stack([noise_scheduler.all_snr[t] for t in timesteps]) - min_snr_gamma = torch.minimum(snr, torch.full_like(snr, gamma)) - if v_prediction: - snr_weight = torch.div(min_snr_gamma, snr + 1).float().to(loss.device) - else: - snr_weight = torch.div(min_snr_gamma, snr).float().to(loss.device) - loss = loss * snr_weight - return loss - - -def apply_debiased_estimation(loss, timesteps, noise_scheduler): - snr_t = torch.stack([noise_scheduler.all_snr[t] for t in timesteps]) # batch_size - snr_t = torch.minimum( - snr_t, torch.ones_like(snr_t) * 1000 - ) # if timestep is 0, snr_t is inf, so limit it to 1000 - weight = 1 / torch.sqrt(snr_t) - loss = weight * loss - return loss - - -def prepare_scheduler_for_custom_training(noise_scheduler, device): - if hasattr(noise_scheduler, "all_snr"): - return - - alphas_cumprod = noise_scheduler.alphas_cumprod - sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod) - sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod) - alpha = sqrt_alphas_cumprod - sigma = sqrt_one_minus_alphas_cumprod - all_snr = (alpha / sigma) ** 2 - - noise_scheduler.all_snr = all_snr.to(device) diff --git a/pipelines/hdm/hdm/trainer/trainer.py b/pipelines/hdm/hdm/trainer/trainer.py deleted file mode 100644 index a195d39f3..000000000 --- a/pipelines/hdm/hdm/trainer/trainer.py +++ /dev/null @@ -1,473 +0,0 @@ -import os -from typing import Any, Iterator - -import torch -import torch.nn as nn -import torch.nn.functional as F -import torch.optim as optim -import lightning.pytorch as pl - -from diffusers import ( - EulerDiscreteScheduler, - UNet2DConditionModel, - AutoencoderKL, -) -from anyschedule import AnySchedule - -from ..utils import instantiate -from ..modules.text_encoders import BaseTextEncoder -from .diffusion import ( - get_noise_noisy_latents_and_timesteps, - prepare_scheduler_for_custom_training, -) - - -class BaseTrainer(pl.LightningModule): - def __init__( - self, - *args, - name: str = "", - lr: float = 1e-5, - optimizer: type[optim.Optimizer] = optim.AdamW, - opt_configs: dict[str, Any] = { - "weight_decay": 0.01, - "betas": (0.9, 0.999), - }, - lr_sch_configs: dict[str, Any] = { - "lr": { - "mode": "cosine", - "end": 100000, - "min_value": 0.001, - } - }, - use_warm_up: bool = True, - warm_up_period: int = 1000, - **kwargs, - ): - super().__init__() - self.name = name - self.train_params: Iterator[nn.Parameter] = None - self.optimizer = instantiate(optimizer) - self.opt_configs = opt_configs - self.lr = lr - self.lr_sch_configs = lr_sch_configs - self.use_warm_up = use_warm_up - self.warm_up_period = warm_up_period - - def configure_optimizers(self): - parameters = [] - assert self.train_params is not None - for param in self.train_params: - if param.ndim < 2: # bias, norm, ... - fan_in = param.numel() - elif param.ndim > 2: # Conv layer in patch embedding - # For conv layers, fan_in is channels_in * kernel_size^2 - fan_ins = param.shape[1:] - fan_in = 1 - for fan_in_i in fan_ins: - fan_in *= fan_in_i - else: # Linear layers, including attention and MLP - fan_in = param.shape[1] - parameters.append( - { - "params": param, - "lr": self.lr / fan_in, - } - ) - optimizer = self.optimizer(parameters, lr=self.lr, **self.opt_configs) - - lr_scheduler = None - if bool(self.lr_sch_configs): - lr_scheduler = AnySchedule(optimizer=optimizer, config=self.lr_sch_configs) - - if lr_scheduler is None: - return optimizer - else: - return { - "optimizer": optimizer, - "lr_scheduler": {"scheduler": lr_scheduler, "interval": "step"}, - } - - -class DMTrainer(BaseTrainer): - def __init__( - self, - unet: UNet2DConditionModel, - te: BaseTextEncoder, - vae: AutoencoderKL | None = None, - unet_compile: bool = False, - te_compile: bool = False, - vae_compile: bool = False, - te_use_normed_ctx: bool = False, - te_freeze: bool = True, - vae_std: float = 7.5, - vae_mean: float = 1.125, - scheduler: EulerDiscreteScheduler | None = None, - lycoris_model: nn.Module | None = None, - *args, - name: str = "", - lr: float = 1e-5, - optimizer: type[optim.Optimizer] = optim.AdamW, - opt_configs: dict[str, Any] = { - "weight_decay": 0.01, - "betas": (0.9, 0.999), - }, - lr_sch_configs: dict[str, Any] = {}, - use_warm_up: bool = True, - warm_up_period: int = 1000, - full_config: dict[str, Any] = {}, - **kwargs, - ): - super(DMTrainer, self).__init__( - name=name, - lr=lr, - optimizer=optimizer, - opt_configs=opt_configs, - lr_sch_configs=lr_sch_configs, - use_warm_up=use_warm_up, - warm_up_period=warm_up_period, - ) - self.save_hyperparameters( - ignore=["unet", "scheduler", "te", "vae", "lycoris_model", "args", "kwargs"] - ) - prepare_scheduler_for_custom_training(scheduler, self.device) - - if unet_compile: - unet = torch.compile(unet) - - if te_compile: - te = torch.compile(te) - - if vae_compile and vae is not None: - vae = torch.compile(vae) - - if te_freeze: - te.requires_grad_(False).eval() - - if vae is not None: - vae.requires_grad_(False).eval() - - self.unet = unet - self.te = te - self.vae = vae - self.scheduler = scheduler - - self.te_use_normed_ctx = te_use_normed_ctx - self.vae_std = vae_std - self.vae_mean = vae_mean - - self.lycoris_model = lycoris_model - - self.epoch = 0 - self.opt_step = 0 - self.ema_loss = 0 - self.ema_decay = 0.99 - - if lycoris_model is not None: - self.lycoris_model.train() - self.train_params = self.lycoris_model.parameters() - else: - self.unet.requires_grad_(True).train() - self.train_params = self.unet.parameters() - - def on_train_epoch_end(self) -> None: - self.epoch += 1 - if self.lycoris_model is not None: - dir = "./lycoris_weight" - epoch = self.epoch - if self._trainer is not None: - trainer = self._trainer - epoch = trainer.current_epoch - if len(trainer.loggers) > 0: - if trainer.loggers[0].save_dir is not None: - save_dir = trainer.loggers[0].save_dir - else: - save_dir = trainer.default_root_dir - name = trainer.loggers[0].name - version = trainer.loggers[0].version - version = ( - version if isinstance(version, str) else f"version_{version}" - ) - dir = os.path.join(save_dir, str(name), version, "lycoris_weight") - else: - # if no loggers, use default_root_dir - dir = os.path.join(trainer.default_root_dir, "lycoris_weight") - os.makedirs(dir, exist_ok=True) - model_weight = { - k: v for k, v in self.unet.named_parameters() if v.requires_grad - } - lycoris_weight = self.lycoris_model.state_dict() | model_weight - torch.save(lycoris_weight, os.path.join(dir, f"epoch={epoch}.pt")) - - def training_step(self, batch, idx): - x, captions, tokenizer_outputs = batch - # print(type(x), type(captions), type(tokenizer_outputs), type(added_cond)) - - if self.vae is not None: - with torch.no_grad(): - latent_dist = self.vae.encode(x).latent_dist - x = latent_dist.sample() - x = (x - self.vae_mean) / self.vae_std - - b, c, h, w = x.shape - - noisy_latent, noise, timesteps = get_noise_noisy_latents_and_timesteps( - self.scheduler, x - ) - - if self.scheduler.config.prediction_type == "epsilon": - target = noise - elif self.scheduler.config.prediction_type == "v_prediction": - target = self.scheduler.get_velocity(x, noise, timesteps) - elif self.scheduler.config.prediction_type == "sample": - target = x - else: - raise ValueError( - f"Unknown prediction type {self.scheduler.config.prediction_type}" - ) - - with torch.no_grad(): - if isinstance(self.te, BaseTextEncoder): - normed_embedding, embedding, pooled_embedding, attn_mask = self.te( - tokenizer_outputs - ) - else: - normed_embedding, pooled_embedding, *embeddings = self.te( - **tokenizer_outputs[0], return_dict=False, output_hidden_states=True - ) - embedding = embeddings[-1][-1] - if self.te_use_normed_ctx: - ctx = normed_embedding - else: - ctx = embedding - - model_output = self.unet( - noisy_latent.to(self.dtype), - timesteps, - encoder_hidden_states=ctx.to(self.dtype), - encoder_attention_mask=attn_mask, - )[0] - loss = F.mse_loss(model_output, target) - - ema_decay = min(self.opt_step / (10 + self.opt_step), self.ema_decay) - self.ema_loss = ema_decay * self.ema_loss + (1 - ema_decay) * loss.item() - self.opt_step += 1 - - if self._trainer is not None: - self.log("train/loss", loss.item(), on_step=True, logger=True) - self.log( - "train/ema_loss", - self.ema_loss, - on_step=True, - logger=True, - prog_bar=True, - ) - return loss - - -class FlowTrainer(BaseTrainer): - def __init__( - self, - unet: nn.Module, - te: BaseTextEncoder, - vae: AutoencoderKL | None = None, - unet_compile: bool = False, - te_compile: bool = False, - vae_compile: bool = False, - te_use_normed_ctx: bool = False, - te_freeze: bool = True, - vae_std: float = 7.5, - vae_mean: float = 1.125, - lycoris_model: nn.Module | None = None, - *args, - name: str = "", - lr: float = 1e-5, - optimizer: type[optim.Optimizer] = optim.AdamW, - opt_configs: dict[str, Any] = { - "weight_decay": 0.01, - "betas": (0.9, 0.999), - }, - lr_sch_configs: dict[str, Any] = {}, - use_warm_up: bool = True, - warm_up_period: int = 1000, - full_config: dict[str, Any] = {}, - **kwargs, - ): - super(FlowTrainer, self).__init__( - name=name, - lr=lr, - optimizer=optimizer, - opt_configs=opt_configs, - lr_sch_configs=lr_sch_configs, - use_warm_up=use_warm_up, - warm_up_period=warm_up_period, - ) - self.save_hyperparameters( - ignore=[ - "unet", - "te", - "vae", - "lycoris_model", - "args", - "kwargs", - "full_config", - "opt_configs", - "lr_sch_configs", - ] - ) - - if unet_compile: - unet = torch.compile(unet) - - if te_compile: - te = torch.compile(te) - - if vae_compile and vae is not None: - vae = torch.compile(vae) - - if te_freeze: - te.requires_grad_(False).eval() - - if vae is not None: - vae.requires_grad_(False).eval() - - self.unet = unet - self.te = te - self.vae = vae - - self.te_use_normed_ctx = te_use_normed_ctx - if self.vae is not None: - vae_std = self.vae.config["latents_std"] - vae_mean = self.vae.config["latents_mean"] - self.register_buffer("vae_std", torch.tensor(vae_std).view(1, -1, 1, 1)) - self.register_buffer("vae_mean", torch.tensor(vae_mean).view(1, -1, 1, 1)) - else: - self.vae_std = vae_std - self.vae_mean = vae_mean - - self.lycoris_model = lycoris_model - - self.epoch = 0 - self.opt_step = 0 - self.ema_loss = 0 - self.ema_decay = 0.995 - - if lycoris_model is not None: - self.lycoris_model.train() - self.train_params = self.lycoris_model.parameters() - else: - self.unet.requires_grad_(True).train() - self.train_params = self.unet.parameters() - - def on_train_epoch_end(self) -> None: - self.epoch += 1 - if self.lycoris_model is not None: - dir = "./lycoris_weight" - epoch = self.epoch - if self._trainer is not None: - trainer = self._trainer - epoch = trainer.current_epoch - if len(trainer.loggers) > 0: - if trainer.loggers[0].save_dir is not None: - save_dir = trainer.loggers[0].save_dir - else: - save_dir = trainer.default_root_dir - name = trainer.loggers[0].name - version = trainer.loggers[0].version - version = ( - version if isinstance(version, str) else f"version_{version}" - ) - dir = os.path.join(save_dir, str(name), version, "lycoris_weight") - else: - # if no loggers, use default_root_dir - dir = os.path.join(trainer.default_root_dir, "lycoris_weight") - os.makedirs(dir, exist_ok=True) - model_weight = { - k: v for k, v in self.unet.named_parameters() if v.requires_grad - } - lycoris_weight = self.lycoris_model.state_dict() | model_weight - torch.save(lycoris_weight, os.path.join(dir, f"epoch={epoch}.pt")) - - def training_step(self, batch, idx): - x, captions, tokenizer_outputs, pos_map, *addon_info = batch - - if self.vae is not None: - if pos_map is not None: - pos_map = pos_map.unflatten(1, x.shape[-2:]) # (B, H, W, 2) - with torch.no_grad(): - x = x.to(self.device) - x = torch.concat( - [ - self.vae.encode(x[i : i + 4]).latent_dist.sample() - for i in range(0, x.shape[0], 4) - ] - ) - x = (x - self.vae_mean) / self.vae_std - pos_map = pos_map.permute(0, 3, 1, 2) - pos_map = ( - F.interpolate(pos_map, x.shape[-2:], mode="area") - .permute(0, 2, 3, 1) - .flatten(1, 2) - ) - - b, c, h, w = x.shape - - noise = torch.randn_like(x) - t = torch.sigmoid(torch.randn(b, 1, 1, 1, device=x.device)) - noisy_latent = t * noise + (1 - t) * x - target = noise - x - - with torch.no_grad(): - if isinstance(self.te, BaseTextEncoder): - normed_embedding, embedding, pooled_embedding, attn_mask = self.te( - tokenizer_outputs - ) - else: - normed_embedding, pooled_embedding, *embeddings = self.te( - **tokenizer_outputs[0], return_dict=False, output_hidden_states=True - ) - embedding = embeddings[-1][-1] - if self.te_use_normed_ctx: - ctx = normed_embedding - else: - ctx = embedding - - if pooled_embedding is not None: - added_cond_kwargs = { - "time_ids": torch.tensor([[1024, 1024, 0, 0, 1024, 1024]]).to( - noisy_latent - ), - "text_embeds": pooled_embedding.to(noisy_latent), - } - else: - added_cond_kwargs = {} - - if len(addon_info) > 0: - for addon in addon_info: - added_cond_kwargs.update(addon) - - model_output = self.unet( - noisy_latent.to(self.dtype), - t, - encoder_hidden_states=ctx.to(self.dtype), - encoder_attention_mask=attn_mask, - pos_map=pos_map, - added_cond_kwargs=added_cond_kwargs, - )[0] - loss = F.mse_loss(model_output, target) - if torch.isnan(loss): - raise ValueError("loss is nan") - - ema_decay = min(self.opt_step / (10 + self.opt_step), self.ema_decay) - self.ema_loss = ema_decay * self.ema_loss + (1 - ema_decay) * loss.item() - self.opt_step += 1 - - if self._trainer is not None: - self.log("train/loss", loss.item(), logger=True) - self.log( - "train/ema_loss", - self.ema_loss, - logger=True, - prog_bar=True, - ) - return loss diff --git a/pipelines/hdm/hdm/utils/__init__.py b/pipelines/hdm/hdm/utils/__init__.py deleted file mode 100644 index f6507e2f5..000000000 --- a/pipelines/hdm/hdm/utils/__init__.py +++ /dev/null @@ -1,74 +0,0 @@ -import importlib -from inspect import isfunction -from random import shuffle - -import torch -import torch.nn as nn - - -def get_obj_from_str(string, reload=False): - module, cls = string.rsplit(".", 1) - if reload: - module_imp = importlib.import_module(module) - importlib.reload(module_imp) - return getattr(importlib.import_module(module, package=None), cls) - - -def instantiate(obj): - from installer import install - install('omegaconf') - import omegaconf - if isinstance(obj, omegaconf.DictConfig): - obj = dict(**obj) - if isinstance(obj, dict) and "class" in obj: - obj_factory = instantiate(obj["class"]) - if "factory" in obj: - obj_factory = getattr(obj_factory, obj["factory"]) - return obj_factory(*obj.get("args", []), **obj.get("kwargs", {})) - if isinstance(obj, str): - return get_obj_from_str(obj) - return obj - - -def exists(val): - return val is not None - - -def uniq(arr): - return {el: True for el in arr}.keys() - - -def default(val, d): - if val is not None: - return val - return d() if isfunction(d) else d - - -def zero_module(module: nn.Module): - """ - Zero out the parameters of a module and return it. - """ - for p in module.parameters(): - p.detach().zero_() - return module - - -def random_choice( - x: torch.Tensor, - num: int, -): - rand_x = list(x) - shuffle(rand_x) - - return torch.stack(rand_x[:num]) - - -def count_params(model, verbose=False): - total_params = sum(p.numel() for p in model.parameters()) - if verbose: - print(f"{model.__class__.__name__} has {total_params * 1.e-6:.2f} M params.") - return total_params - - -def remove_none(list_x): - return [i for i in list_x if i is not None] diff --git a/pipelines/hdm/hdm/utils/config.py b/pipelines/hdm/hdm/utils/config.py deleted file mode 100644 index 454b7d532..000000000 --- a/pipelines/hdm/hdm/utils/config.py +++ /dev/null @@ -1,38 +0,0 @@ -import os - - -def load_train_config(file): - from installer import install - install('omegaconf') - install('toml') - import omegaconf - import toml - - config = toml.load(file) - - model = config["model"] - model["config"] = omegaconf.OmegaConf.to_container( - omegaconf.OmegaConf.load(model["config"]), resolve=True - ) - dataset = config["dataset"] - trainer = config["trainer"] - lightning = config["lightning"] - - if "logger" in lightning and not lightning["logger"].get("version", None): - lightning["logger"]["version"] = os.urandom(4).hex() - - if "scaling_factor" in model and "scaling_factor" not in dataset: - dataset["scaling_factor"] = model["scaling_factor"] - if "scaling_factor" in dataset and "scaling_factor" not in model: - model["scaling_factor"] = dataset["scaling_factor"] - if "scaling_factor" not in model and "scaling_factor" not in dataset: - model["scaling_factor"] = dataset["scaling_factor"] = 1.0 - - if "latent_shift" in model and "latent_shift" not in dataset: - dataset["latent_shift"] = model["latent_shift"] - if "latent_shift" in dataset and "latent_shift" not in model: - model["latent_shift"] = dataset["latent_shift"] - if "latent_shift" not in model and "latent_shift" not in dataset: - model["latent_shift"] = dataset["latent_shift"] = 0.0 - - return model, dataset, trainer, lightning diff --git a/pipelines/hdm/xut/__init__.py b/pipelines/hdm/xut/__init__.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/pipelines/hdm/xut/env.py b/pipelines/hdm/xut/env.py deleted file mode 100644 index 17a0874df..000000000 --- a/pipelines/hdm/xut/env.py +++ /dev/null @@ -1,9 +0,0 @@ -TORCH_COMPILE = False -USE_LIGER = True -USE_VANILLA = True -USE_XFORMERS = False -USE_XFORMERS_LAYERS = False -COMPILE_ARGS = { - "mode": "default", - "dynamic": True, -} diff --git a/pipelines/hdm/xut/modules/__init__.py b/pipelines/hdm/xut/modules/__init__.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/pipelines/hdm/xut/modules/adaln.py b/pipelines/hdm/xut/modules/adaln.py deleted file mode 100644 index a4fb1f0e7..000000000 --- a/pipelines/hdm/xut/modules/adaln.py +++ /dev/null @@ -1,28 +0,0 @@ -import torch -import torch.nn as nn -import torch.nn.functional as F -import torch.optim as optim - -from .norm import RMSNorm - - -class AdaLN(nn.Module): - def __init__(self, dim, y_dim, gate=True, norm_layer=RMSNorm, shared=False): - super().__init__() - self.norm = norm_layer(dim) - self.gate = gate - if shared: - self.adaln = None - else: - self.adaln = nn.Linear(y_dim, dim * (2 + bool(gate))) - nn.init.constant_(self.adaln.bias, 0) - nn.init.constant_(self.adaln.weight, 0) - - def forward(self, x, y, shared_adaln=None): - if shared_adaln is None: - scale, shift, *gate = self.adaln(y).chunk(2 + bool(self.gate), dim=-1) - else: - scale, shift, *gate = shared_adaln - normed_x, _ = self.norm(x) - result = normed_x * (scale + 1.0) + shift - return result, (gate[0] + 1) if self.gate else 1 diff --git a/pipelines/hdm/xut/modules/attention.py b/pipelines/hdm/xut/modules/attention.py deleted file mode 100644 index ce39ca9c6..000000000 --- a/pipelines/hdm/xut/modules/attention.py +++ /dev/null @@ -1,332 +0,0 @@ -import math -from functools import cache - -import torch -import torch.nn as nn -import torch.nn.functional as F - -try: - import xformers - - XFORMERS_AVAILABLE = True -except ImportError: - XFORMERS_AVAILABLE = False -if XFORMERS_AVAILABLE: - from xformers.ops import memory_efficient_attention -else: - memory_efficient_attention = None - -from .. import env -from ..utils import compile_wrapper -from .axial_rope import AxialRoPE - - -if not env.USE_XFORMERS: - memory_efficient_attention = None -if env.USE_VANILLA: - - @compile_wrapper - def memory_efficient_attention(query, key, value, attn_bias=None, p=0.0): - scale = 1.0 / query.shape[-1] ** 0.5 - query = query * scale - query = query.transpose(1, 2) - key = key.transpose(1, 2) - value = value.transpose(1, 2) - attn = query @ key.transpose(-2, -1) - if attn_bias is not None: - attn = attn + attn_bias - attn = attn.softmax(-1) - attn = F.dropout(attn, p) - attn = attn @ value - return attn.transpose(1, 2).contiguous() - - -class SelfAttention(nn.Module): - def __init__(self, dim, n_heads=8, head_dim=-1, pos_dim=2): - super().__init__() - self.dim = dim - self.n_heads = n_heads - self.head_dim = head_dim if head_dim > 0 else dim // n_heads - self.n_heads = dim // self.head_dim - assert ( - self.n_heads * self.head_dim == dim - ), "dim must be divisible by n_heads or head_dim" - - self.qkv = nn.Linear(dim, dim * 3, bias=False) - self.out = nn.Linear(dim, dim) - self.rope = AxialRoPE(self.head_dim, self.n_heads, pos_dim) - self.attn = memory_efficient_attention or F.scaled_dot_product_attention - self.xformers = memory_efficient_attention is not None - - def forward(self, x, pos_map=None, mask=None): - b, n, _, h = *x.shape, self.n_heads - q, k, v = self.qkv(x).chunk(3, dim=-1) - - if pos_map is not None: - q = self.rope(q.reshape(b, n, h, -1).transpose(1, 2), pos_map) - k = self.rope(k.reshape(b, n, h, -1).transpose(1, 2), pos_map) - v = v.reshape(b, n, h, -1) - if self.xformers: - q = q.transpose(1, 2) - k = k.transpose(1, 2) - else: - v = v.transpose(1, 2) - else: - q, k, v = map(lambda t: t.reshape(b, n, h, -1), (q, k, v)) - if not self.xformers: - q = q.transpose(1, 2) - k = k.transpose(1, 2) - v = v.transpose(1, 2) - - if mask is not None: - if mask.ndim == 2: - mask = mask[None, None] - elif mask.ndim == 3: - mask = mask[:, None] - if n % 8 and self.xformers: - align_n = math.ceil(n / 8) * 8 - mask_align = torch.empty( - *mask.shape[:3], align_n, device=mask.device, dtype=mask.dtype - ) - mask_align[..., :n] = mask - mask = mask_align.to(q).expand(b, h, n, align_n)[..., :n] - else: - mask = mask.to(q).expand(b, h, n, n) - - attn = self.attn(q, k, v, mask) - if not self.xformers: - attn = attn.transpose(1, 2) - attn = attn.reshape(b, n, h * self.head_dim) - attn = self.out(attn) - return attn - - -class CrossAttention(nn.Module): - def __init__(self, dim, ctx_dim, n_heads=8, head_dim=-1, pos_dim=2): - super().__init__() - self.dim = dim - self.n_heads = n_heads - self.head_dim = head_dim if head_dim > 0 else dim // n_heads - self.n_heads = dim // self.head_dim - assert ( - self.n_heads * self.head_dim == dim - ), "dim must be divisible by n_heads or head_dim" - - self.q = nn.Linear(dim, dim, bias=False) - self.kv = nn.Linear(ctx_dim, dim * 2, bias=False) - self.out = nn.Linear(dim, dim) - self.rope = AxialRoPE(self.head_dim, self.n_heads, pos_dim) - self.attn = memory_efficient_attention or F.scaled_dot_product_attention - self.xformers = memory_efficient_attention is not None - - def forward(self, x, ctx, pos_map=None, ctx_pos_map=None, mask=None): - b, n, _, h = *x.shape, self.n_heads - ctx_n = ctx.shape[1] - q = self.q(x) - k, v = self.kv(ctx).chunk(2, dim=-1) - - if pos_map is not None: - q = self.rope(q.reshape(b, n, h, -1).transpose(1, 2), pos_map) - q = q if not self.xformers else q.transpose(1, 2) - else: - q = q.reshape(b, n, h, -1) - q = q if self.xformers else q.transpose(1, 2) - if ctx_pos_map is not None: - k = self.rope(k.reshape(b, ctx_n, h, -1).transpose(1, 2), ctx_pos_map) - k = k if not self.xformers else k.transpose(1, 2) - else: - k = k.reshape(b, ctx_n, h, -1) - k = k if self.xformers else k.transpose(1, 2) - v = v.reshape(b, ctx_n, h, -1) - v = v if self.xformers else v.transpose(1, 2) - - if mask is not None: - if mask.ndim == 2: - mask = mask[None, None] - elif mask.ndim == 3: - mask = mask[:, None] - if ctx_n % 8 and self.xformers: - align_n = math.ceil(ctx_n / 8) * 8 - mask_align = torch.empty( - *mask.shape[:3], align_n, device=mask.device, dtype=mask.dtype - ) - mask_align[..., :ctx_n] = mask - mask = mask_align.to(q).expand(b, h, n, align_n)[..., :ctx_n] - else: - mask = mask.to(q).expand(b, h, n, ctx_n) - - attn = self.attn(q, k, v, mask) - if not self.xformers: - attn = attn.transpose(1, 2) - attn = attn.reshape(b, n, h * self.head_dim) - attn = self.out(attn) - return attn - - -class AttentionPooling(CrossAttention): - def __init__(self, dim, n_heads=8, head_dim=-1, pos_dim=2): - super().__init__(dim, dim, n_heads, head_dim, pos_dim) - self.query_token = nn.Parameter(torch.randn(1, 1, dim) * 1 / dim**0.5) - - def forward(self, x, pos_map=None, mask=None): - query = self.query_token.expand(x.shape[0], -1, -1) - return super().forward(query, x, None, pos_map, mask).squeeze(1) - - -class AttentiveProbe(CrossAttention): - def __init__(self, dim, out_dim, n_heads=8, head_dim=-1, pos_dim=2, n_probes=1): - super().__init__(dim, dim, n_heads, head_dim, pos_dim) - self.query_token = nn.Parameter(torch.randn(1, n_probes, dim) * 1 / dim**0.5) - self.token_proj = nn.Linear(dim * n_probes, out_dim) - - def forward(self, x, pos_map=None, mask=None): - query = self.query_token.expand(x.shape[0], -1, -1) - output_embedding = super().forward(query, x, None, pos_map, mask) - output_embedding = output_embedding.flatten(-2, -1) - return self.token_proj(output_embedding) - - -@cache -def prefix_causal_attention_mask( - q_len, kv_len, prefix_len=0, is_self_attn=False, dtype=None, device=None -): - """ - **Made by claude 3.7 sonnet without thinking** - Generate attention masks and biases for transformer models. - - Parameters: - ----------- - q_len : int - Length of the query sequence - kv_len : int - Length of the key/value sequence - prefix_len : int, optional - Length of the prefix for which we allow full attention (no causal masking) - Default: 0 (standard causal mask) - is_self_attn : bool, optional - Whether this is for self-attention (q_len == kv_len and they represent the same sequence) - Enables faster mask generation - Default: False - dtype : torch.dtype, optional - Data type for the output tensors - Default: None (will use torch.bool for mask, torch.float for bias) - device : torch.device, optional - Device on which to create the tensors - Default: None (will use the default torch device) - - Returns: - -------- - tuple: (attention_mask, attention_bias) - - attention_mask: Boolean tensor of shape (q_len, kv_len) where True values indicate - positions that should be attended to - - attention_bias: Tensor of same shape with dtype specified (or float), containing - 0.0 for positions to attend to and -float('inf') for positions to mask out - """ - # Fast path for self-attention with no prefix - if is_self_attn and prefix_len == 0: - # Simple lower triangular matrix for standard causal self-attention - attention_mask = torch.tril( - torch.ones(q_len, q_len, dtype=torch.bool, device=device) - ) - - # Fast path for self-attention with prefix - elif is_self_attn and prefix_len > 0: - attention_mask = torch.tril( - torch.ones(q_len, q_len, dtype=torch.bool, device=device) - ) - - # Add the prefix part (allow full attention to the prefix) - if prefix_len < q_len: - # Set the prefix columns to all True (we use indexing which is faster than cat) - attention_mask[:, :prefix_len] = True - - # General case for cross-attention or when fast path is not used - else: - # Create base causal mask (lower triangular) - # Each query position i can attend to key positions j where j <= i - causal_mask = torch.tril( - torch.ones(q_len, kv_len, dtype=torch.bool, device=device) - ) - - # If there's a prefix, allow full attention within that prefix - if prefix_len > 0: - # Combine masks: - # - For the prefix part of kv, use all True - # - For the rest, use causal mask - if prefix_len < kv_len: - attention_mask = torch.cat( - [ - torch.ones(q_len, prefix_len, dtype=torch.bool, device=device), - causal_mask[:, prefix_len:], - ], - dim=1, - ) - else: - # If prefix_len >= kv_len, the entire sequence gets full attention - attention_mask = torch.ones( - q_len, kv_len, dtype=torch.bool, device=device - ) - else: - # Without prefix, just use the causal mask - attention_mask = causal_mask - - # Convert boolean mask to attention bias - # True -> 0.0, False -> -inf - float_dtype = torch.float if dtype is None else dtype - attention_bias = torch.zeros_like(attention_mask, dtype=float_dtype, device=device) - attention_bias = attention_bias.masked_fill(~attention_mask, float("-inf")) - - return attention_mask, attention_bias - - -# Example usage: -if __name__ == "__main__": - # Standard causal mask for sequence length 6 - mask, bias = prefix_causal_attention_mask(q_len=6, kv_len=6) - print("Standard causal mask:") - print(mask) - print("\nStandard causal bias:") - print(bias) - - # Same with self-attention flag - mask_self, bias_self = prefix_causal_attention_mask( - q_len=6, kv_len=6, is_self_attn=True - ) - print("\nSelf-attention causal mask (should be identical):") - print(mask_self) - print("Masks are identical:", torch.all(mask == mask_self).item()) - - # Causal mask with prefix_len=3 (first 2 tokens get full attention) - mask, bias = prefix_causal_attention_mask(q_len=6, kv_len=6, prefix_len=3) - print("\nCausal mask with prefix_len=3:") - print(mask) - print("\nCausal bias with prefix_len=3:") - print(bias) - - # Same with self-attention flag - mask_self, bias_self = prefix_causal_attention_mask( - q_len=6, kv_len=6, prefix_len=3, is_self_attn=True - ) - print("\nSelf-attention mask with prefix_len=3 (should be identical):") - print(mask_self) - print("Masks are identical:", torch.all(mask == mask_self).item()) - - # Handling different q_len and kv_len (for cross-attention) - mask, bias = prefix_causal_attention_mask(q_len=4, kv_len=6, prefix_len=3) - print("\nCross-attention mask with q_len=4, kv_len=6, prefix_len=3:") - print(mask) - print("\nCross-attention bias:") - print(bias) - - self_attn = SelfAttention(64, 8).cuda().half() - x = torch.randn(1, 16, 64).cuda().half() - mask, bias = prefix_causal_attention_mask( - 16, 16, is_self_attn=True, device=x.device, dtype=x.dtype - ) - test_out = self_attn(x, mask=bias) - torch.sum(test_out).backward() - - print(x.shape, mask.shape, bias.shape) - print(test_out.shape) - print(torch.isnan(test_out).any()) - print(torch.norm(next(self_attn.parameters()).grad)) diff --git a/pipelines/hdm/xut/modules/axial_rope.py b/pipelines/hdm/xut/modules/axial_rope.py deleted file mode 100644 index e5cadcd25..000000000 --- a/pipelines/hdm/xut/modules/axial_rope.py +++ /dev/null @@ -1,179 +0,0 @@ -import math -from functools import lru_cache - -import torch -from torch import nn - -from ..utils import compile_wrapper - - -@compile_wrapper -def rotate_half(x): - x1, x2 = x[..., 0::2], x[..., 1::2] - x = torch.stack((-x2, x1), dim=-1) - *shape, d, r = x.shape - return x.view(*shape, d * r) - - -@compile_wrapper -def apply_rotary_emb(freqs, t, start_index=0, scale=1.0): - freqs = freqs.to(t) - rot_dim = freqs.shape[-1] - end_index = start_index + rot_dim - assert ( - rot_dim <= t.shape[-1] - ), f"feature dimension {t.shape[-1]} is not of sufficient size to rotate in all the positions {rot_dim}" - t_left, t, t_right = ( - t[..., :start_index], - t[..., start_index:end_index], - t[..., end_index:], - ) - t = (t * freqs.cos() * scale) + (rotate_half(t) * freqs.sin() * scale) - return torch.cat((t_left, t, t_right), dim=-1) - - -def centers(start, stop, num, dtype=None, device=None): - edges = torch.linspace(start, stop, num + 1, dtype=dtype, device=device) - return (edges[:-1] + edges[1:]) / 2 - - -def make_grid(h_pos, w_pos): - grid = torch.stack(torch.meshgrid(h_pos, w_pos, indexing="ij"), dim=-1) - return grid.flatten(0, 1) - - -def bounding_box(h, w, pixel_aspect_ratio=1.0): - # Adjusted dimensions - w_adj = w - h_adj = h * pixel_aspect_ratio - - # Adjusted aspect ratio - ar_adj = w_adj / h_adj - - # Determine bounding box based on the adjusted aspect ratio - y_min, y_max, x_min, x_max = -1.0, 1.0, -1.0, 1.0 - if ar_adj > 1: - y_min, y_max = -1 / ar_adj, 1 / ar_adj - elif ar_adj < 1: - x_min, x_max = -ar_adj, ar_adj - - return torch.tensor([y_min, y_max, x_min, x_max]) - - -@lru_cache(maxsize=8) -def make_axial_pos( - h, w, pixel_aspect_ratio=1.0, align_corners=False, dtype=None, device=None -): - y_min, y_max, x_min, x_max = bounding_box(h, w, pixel_aspect_ratio) - if align_corners: - h_pos = torch.linspace(y_min, y_max, h, dtype=dtype, device=device) - w_pos = torch.linspace(x_min, x_max, w, dtype=dtype, device=device) - else: - h_pos = centers(y_min, y_max, h, dtype=dtype, device=device) - w_pos = centers(x_min, x_max, w, dtype=dtype, device=device) - return make_grid(h_pos, w_pos) - - -def make_axial_pos_no_cache( - h, w, pixel_aspect_ratio=1.0, align_corners=False, dtype=None, device=None -): - y_min, y_max, x_min, x_max = bounding_box(h, w, pixel_aspect_ratio) - if align_corners: - h_pos = torch.linspace(y_min, y_max, h, dtype=dtype, device=device) - w_pos = torch.linspace(x_min, x_max, w, dtype=dtype, device=device) - else: - h_pos = centers(y_min, y_max, h, dtype=dtype, device=device) - w_pos = centers(x_min, x_max, w, dtype=dtype, device=device) - return make_grid(h_pos, w_pos) - - -def make_cropped_pos(crop_h, crop_w, target_h, target_w): - pos_map = make_axial_pos_no_cache(target_h, target_w).unflatten( - 0, (target_h, target_w) - ) - if target_h > target_w: - pos_map = pos_map[crop_h : crop_h + target_w, :] - elif target_h < target_w: - pos_map = pos_map[:, crop_w : crop_w + target_h] - return pos_map.flatten(0, 1) - - -def freqs_pixel(max_freq=10.0): - def init(shape): - freqs = torch.linspace(1.0, max_freq / 2, shape[-1]) * math.pi - return freqs.log().expand(shape) - - return init - - -def freqs_pixel_log(max_freq=10.0): - def init(shape): - log_min = math.log(math.pi) - log_max = math.log(max_freq * math.pi / 2) - return torch.linspace(log_min, log_max, shape[-1]).expand(shape) - - return init - - -class AxialRoPE(nn.Module): - def __init__( - self, - dim, - n_heads, - pos_dim=2, - start_index=0, - freqs_init=freqs_pixel_log(max_freq=10.0), - ): - super().__init__() - self.n_heads = n_heads - self.start_index = start_index - log_freqs = freqs_init((n_heads, dim // (2 * pos_dim), 1)) - self.freqs = nn.Parameter(log_freqs.clone().repeat(1, 1, pos_dim)) - - def extra_repr(self): - dim = self.freqs.shape[-1] - return f"dim={dim}, n_heads={self.n_heads}, start_index={self.start_index}" - - def get_freqs(self, pos): - if pos.shape[-1] != self.freqs.shape[-1]: - raise ValueError(f"input shape must be (..., {self.freqs.shape[-1]})") - freqs = pos[..., None, None, :] * self.freqs.exp() - freqs = freqs.flatten(-2, -1).repeat_interleave(2, dim=-1) - return freqs.transpose(-2, -3) - - @compile_wrapper - def forward(self, x, pos): - freqs = self.get_freqs(pos) - return apply_rotary_emb(freqs, x, self.start_index) - - -class AdditiveAxialRoPE(AxialRoPE): - """ - https://arxiv.org/abs/2405.10436 - """ - - def __init__( - self, - dim, - n_heads, - pos_dim=2, - start_index=0, - freqs_init=freqs_pixel_log(max_freq=10.0), - ): - super().__init__(dim, n_heads, pos_dim, start_index, freqs_init) - self.emb = nn.Parameter(torch.randn(dim) / dim**0.5) - - def forward(self, x, pos): - pos_emb = torch.zeros_like(x) - pos_emb = pos_emb + self.emb - freqs = self.get_freqs(pos) - if x.ndim == 3: - pos_emb = pos_emb.unsqueeze(1) - return x + apply_rotary_emb(freqs, pos_emb, self.start_index).view(x.shape) - - -if __name__ == "__main__": - x = torch.randn(2, 1, 4 * 4, 16) - pos = torch.randn(2, 16, 1) - model = AxialRoPE(16, 1, 1) - print(model(x, pos).shape) diff --git a/pipelines/hdm/xut/modules/layers.py b/pipelines/hdm/xut/modules/layers.py deleted file mode 100644 index bf8ec4060..000000000 --- a/pipelines/hdm/xut/modules/layers.py +++ /dev/null @@ -1,56 +0,0 @@ -import torch -import torch.nn as nn -import torch.nn.functional as F - -try: - import xformers - - XFORMERS_AVAILABLE = True -except ImportError: - XFORMERS_AVAILABLE = False - -from .. import env -from ..utils import compile_wrapper - - -class SwiGLUTorch(nn.Module): - def __init__( - self, in_features, hidden_features, out_features, bias=True, _pack_weights=True - ): - super().__init__() - self.in_features = in_features - self.hidden_features = hidden_features or in_features - self.out_features = out_features or in_features - if _pack_weights: - self.w12 = torch.nn.Linear(in_features, 2 * hidden_features, bias=bias) - else: - self.w1 = torch.nn.Linear(in_features, hidden_features, bias=bias) - self.w2 = torch.nn.Linear(in_features, hidden_features, bias=bias) - self.w3 = torch.nn.Linear(hidden_features, out_features, bias=bias) - - @compile_wrapper - def forward(self, x): - if self.w12 is not None: - x1, x2 = self.w12(x).chunk(2, dim=-1) - else: - x1 = self.w1(x) - x2 = self.w2(x) - return self.w3(F.silu(x1) * x2) - - -if XFORMERS_AVAILABLE: - from xformers.ops import SwiGLU -else: - SwiGLU = SwiGLUTorch -if not env.USE_XFORMERS_LAYERS: - SwiGLU = SwiGLUTorch - - -if __name__ == "__main__": - x = torch.randn(2, 16, 128) - model1 = SwiGLU(128, 256, 128) - model2 = SwiGLUTorch(128, 256, 128) - - model1.load_state_dict(model2.state_dict()) - - print(F.mse_loss(model1(x), model2(x)), torch.norm(model1(x))) diff --git a/pipelines/hdm/xut/modules/norm.py b/pipelines/hdm/xut/modules/norm.py deleted file mode 100644 index f51b6074a..000000000 --- a/pipelines/hdm/xut/modules/norm.py +++ /dev/null @@ -1,97 +0,0 @@ -import torch -import torch.nn as nn -import torch.nn.functional as F - -try: - from liger_kernel.transformers.rms_norm import LigerRMSNorm -except ImportError: - LigerRMSNorm = None - -from .. import env -from ..utils import compile_wrapper - - -class DyT(nn.Module): - """ - Transformers without Normalization - https://arxiv.org/abs/2503.10622 - """ - - def __init__(self, hidden_size, init_alpha=1.0): - super().__init__() - self.hidden_size = hidden_size - self.in_weight = nn.Parameter(torch.ones(hidden_size) * init_alpha) - - @compile_wrapper - def forward(self, hidden_states): - hidden_states = torch.tanh(self.in_weight * hidden_states) - return hidden_states, 1.0 - - -class RMSNormTorch(nn.RMSNorm): - def __init__(self, hidden_size, *args, eps=1e-6, offset=0.0, **kwargs): - super().__init__((hidden_size,), *args, eps=eps, **kwargs) - self.offset = offset - - @compile_wrapper - def forward(self, hidden_states): - return ( - F.rms_norm( - hidden_states, - self.normalized_shape, - self.weight + self.offset, - self.eps, - ), - 1.0, - ) - - -if LigerRMSNorm is None or not env.USE_LIGER: - RMSNorm = RMSNormTorch - -else: - - class RMSNorm(LigerRMSNorm): - def __init__( - self, - hidden_size, - eps=1e-6, - offset=0.0, - casting_mode="llama", - init_fn="ones", - in_place=True, - ): - super().__init__( - hidden_size, - eps=eps, - offset=offset, - casting_mode=casting_mode, - init_fn=init_fn, - in_place=in_place, - ) - - def forward(self, hidden_states): - return super().forward(hidden_states), 1.0 - - -def Norm(module: nn.Module): - module.org_forward = module.forward - module.forward = lambda *args, **kwargs: module.org_forward(*args, **kwargs)[0] - return module - - -if __name__ == "__main__": - if LigerRMSNorm is None: - print("LigerRMSNorm is available") - exit() - - hidden_size = 512 - hidden_states = torch.randn(2, hidden_size).cuda() - - norm1 = RMSNorm(hidden_size).cuda() - norm2 = RMSNormTorch(hidden_size).cuda() - - nn.init.normal_(norm1.weight) - norm2.load_state_dict(norm1.state_dict()) - - print(F.mse_loss(norm1(hidden_states)[0], norm2(hidden_states)[0])) diff --git a/pipelines/hdm/xut/modules/patch.py b/pipelines/hdm/xut/modules/patch.py deleted file mode 100644 index 127ddcb59..000000000 --- a/pipelines/hdm/xut/modules/patch.py +++ /dev/null @@ -1,74 +0,0 @@ -import torch -import torch.nn as nn -import torch.nn.functional as F - - -class PatchEmbed(nn.Module): - def __init__( - self, - patch_size=4, - in_channels=3, - embed_dim=512, - norm_layer=None, - flatten=True, - bias=True, - ): - super().__init__() - self.patch_size = patch_size - self.flatten = flatten - - self.proj = nn.Conv2d(in_channels, embed_dim, patch_size, patch_size, bias=bias) - self.norm = nn.Identity() if norm_layer is None else norm_layer(embed_dim) - - def forward(self, x, pos_map=None): - b, _, h, w = x.shape - x = self.proj(x) - b, _, new_h, new_w = x.shape - if pos_map is not None: - pos_map = ( - F.interpolate( - pos_map.reshape(b, h, w, -1).permute(0, 3, 1, 2), - (new_h, new_w), - mode="bilinear", - antialias=True, - ) - .permute(0, 2, 3, 1) - .flatten(1, 2) - ) - if self.flatten: - x = x.flatten(2).transpose(1, 2) - x = self.norm(x) - return x, pos_map - - -class UnPatch(nn.Module): - def __init__(self, patch_size=4, input_dim=512, out_channel=3, proj=True): - super().__init__() - self.patch_size = patch_size - self.c = out_channel - - if proj: - self.proj = nn.Linear(input_dim, patch_size**2 * out_channel) - else: - self.proj = nn.Identity() - - def forward(self, x: torch.Tensor, axis1=None, axis2=None, loss_mask=None): - b, n, _ = x.shape - p = q = self.patch_size - if axis1 is None and axis2 is None: - w = h = int(n**0.5) - assert h * w == n - else: - h = axis1 // p if axis1 else n // (axis2 // p) - w = axis2 // p if axis2 else n // h - assert h * w == n - - x = self.proj(x) - if loss_mask is not None: - x = torch.where(loss_mask[..., None], x, x.detach()) - x = ( - x.reshape(b, h, w, p, q, self.c) - .permute(0, 5, 1, 3, 2, 4) - .reshape(b, self.c, h * p, w * q) - ) - return x diff --git a/pipelines/hdm/xut/modules/time_emb.py b/pipelines/hdm/xut/modules/time_emb.py deleted file mode 100644 index 290e75149..000000000 --- a/pipelines/hdm/xut/modules/time_emb.py +++ /dev/null @@ -1,34 +0,0 @@ -import math - -import torch -import torch.nn as nn - -from ..utils import compile_wrapper - - -class TimestepEmbedding(nn.Module): - def __init__(self, dim, max_period=10000, time_factor: float = 1000.0): - super().__init__() - self.dim = dim - self.max_period = max_period - self.time_factor = time_factor - self.register_buffer( - "freqs", - torch.exp( - -math.log(max_period) - * torch.arange(start=0, end=dim // 2, dtype=torch.float32) - / (dim // 2) - )[None], - ) - self.proj = nn.Sequential(nn.Linear(dim, dim), nn.Mish()) - - @compile_wrapper - def forward(self, t): - t = self.time_factor * t - args = t[:, None] * self.freqs - embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) - if self.dim % 2: - embedding = torch.cat( - [embedding, torch.zeros_like(embedding[:, :1])], dim=-1 - ) - return self.proj(embedding) diff --git a/pipelines/hdm/xut/modules/transformer.py b/pipelines/hdm/xut/modules/transformer.py deleted file mode 100644 index 4172ee648..000000000 --- a/pipelines/hdm/xut/modules/transformer.py +++ /dev/null @@ -1,79 +0,0 @@ -import torch.nn as nn - -from .layers import SwiGLU -from .attention import SelfAttention, CrossAttention -from .norm import RMSNorm -from .adaln import AdaLN - - -class TransformerBlock(nn.Module): - def __init__( - self, - dim, - ctx_dim, - heads, - dim_head, - mlp_dim, - pos_dim, - use_adaln=False, - use_shared_adaln=False, - ctx_from_self=False, - norm_layer=RMSNorm, - ): - super().__init__() - self.use_adaln = use_adaln - self.attn = SelfAttention(dim, heads, dim_head, pos_dim) - if ctx_dim is None: - self.xattn_pre_norm = None - self.xattn = None - else: - self.ctx_from_self = ctx_from_self - self.xattn = CrossAttention(dim, ctx_dim, heads, dim_head, pos_dim) - self.mlp = SwiGLU(dim, mlp_dim, dim) - - if self.use_adaln: - self.attn_pre_norm = AdaLN( - dim, dim, norm_layer=norm_layer, shared=use_shared_adaln - ) - self.mlp_pre_norm = AdaLN( - dim, dim, norm_layer=norm_layer, shared=use_shared_adaln - ) - if self.xattn is not None: - self.xattn_pre_norm = AdaLN( - dim, dim, norm_layer=norm_layer, shared=use_shared_adaln - ) - else: - self.attn_pre_norm = norm_layer(dim) - self.mlp_pre_norm = norm_layer(dim) - if self.xattn is not None: - self.xattn_pre_norm = norm_layer(dim) - - def forward( - self, - x, - ctx, - pos_map=None, - ctx_pos_map=None, - y=None, - x_mask=None, - ctx_mask=None, - shared_adaln=None, - ): - y = [y] if y is not None else [] - y = y if shared_adaln is None else [y[0], shared_adaln[0]] - x, gate = self.attn_pre_norm(x, *y) - x = x + self.attn(x, pos_map, mask=x_mask) * gate - - if self.xattn is not None: - if shared_adaln is not None: - y[1] = shared_adaln[1] - x, gate = self.xattn_pre_norm(x, *y) - if self.ctx_from_self: - ctx_mask = x_mask - x = x + self.xattn(x, ctx, pos_map, ctx_pos_map, mask=ctx_mask) * gate - - if shared_adaln is not None: - y[1] = shared_adaln[-1] - x, gate = self.mlp_pre_norm(x, *y) - x = x + self.mlp(x) * gate - return x diff --git a/pipelines/hdm/xut/utils/__init__.py b/pipelines/hdm/xut/utils/__init__.py deleted file mode 100644 index 2d6a9cb72..000000000 --- a/pipelines/hdm/xut/utils/__init__.py +++ /dev/null @@ -1,23 +0,0 @@ -import torch -from .. import env - - -def isiterable(obj): - try: - iter(obj) - except TypeError: - return False - return True - - -def compile_wrapper(func, **kwargs): - kwargs.update(env.COMPILE_ARGS) - compiled = torch.compile(func, **kwargs) - - def runner(*args, **kwargs): - if env.TORCH_COMPILE: - return compiled(*args, **kwargs) - else: - return func(*args, **kwargs) - - return runner diff --git a/pipelines/hdm/xut/xut.py b/pipelines/hdm/xut/xut.py deleted file mode 100644 index 2dfdab9dc..000000000 --- a/pipelines/hdm/xut/xut.py +++ /dev/null @@ -1,556 +0,0 @@ -import torch -import torch.nn as nn -import torch.nn.functional as F -from torch.utils.checkpoint import checkpoint - -from .modules.norm import RMSNorm -from .modules.transformer import TransformerBlock -from .modules.patch import PatchEmbed, UnPatch -from .modules.axial_rope import make_axial_pos -from .modules.time_emb import TimestepEmbedding -from .modules.norm import RMSNorm, DyT -from .utils import isiterable - - -class TBackBone(nn.Module): - """ - Basic backbone of transformer - """ - - def __init__( - self, - dim=1024, - ctx_dim=1024, - heads=16, - dim_head=64, - mlp_dim=3072, - pos_dim=2, - depth=8, - use_adaln=False, - use_shared_adaln=False, - use_dyt=False, - ): - super().__init__() - self.blocks = nn.ModuleList( - [ - TransformerBlock( - dim, - ctx_dim, - heads, - dim_head, - mlp_dim, - pos_dim, - use_adaln, - use_shared_adaln, - norm_layer=DyT if use_dyt else RMSNorm, - ) - for _ in range(depth) - ] - ) - self.grad_ckpt = False - - def init_weight(self): - for param in self.parameters(): - if param.ndim == 1: - nn.init.normal_(param, mean=0.0, std=(1 / param.size(0)) ** 0.5) - elif param.ndim == 2: - fan_in = param.size(1) - nn.init.normal_(param, mean=0.0, std=(1 / fan_in) ** 0.5) - elif param.ndim >= 3: - fan_out, *fan_ins = param.shape - # cumprod - fan_in = 1 - for f in fan_ins: - fan_in *= f - nn.init.normal_(param, mean=0.0, std=(1 / fan_in) ** 0.5) - - def forward( - self, - x, - ctx=None, - x_mask=None, - ctx_mask=None, - pos_map=None, - y=None, - shared_adaln=None, - ): - if pos_map is not None: - assert pos_map.size(1) == x.size(1) - - for block in self.blocks: - if self.grad_ckpt: - x = checkpoint( - block, - x, - ctx, - pos_map, - None, - y, - x_mask, - ctx_mask, - shared_adaln, - use_reentrant=False, - ) - else: - x = block(x, ctx, pos_map, None, y, x_mask, ctx_mask, shared_adaln) - - return x - - -class XUTBackBone(nn.Module): - """ - Basic backbone of cross-U-transformer. - """ - - def __init__( - self, - dim=1024, - ctx_dim=None, - heads=16, - dim_head=64, - mlp_dim=3072, - pos_dim=2, - depth=8, - enc_blocks=1, - dec_blocks=2, - dec_ctx=False, - use_adaln=False, - use_shared_adaln=False, - use_dyt=False, - ): - super().__init__() - if isiterable(enc_blocks): - enc_blocks = list(enc_blocks) - assert len(enc_blocks) == depth - else: - enc_blocks = [int(enc_blocks)] * depth - if isiterable(dec_blocks): - dec_blocks = list(dec_blocks) - assert len(dec_blocks) == depth - else: - dec_blocks = [int(dec_blocks)] * depth - - self.enc_blocks = nn.ModuleList() - for i in range(depth): - blocks = [ - TransformerBlock( - dim, - ctx_dim, - heads, - dim_head, - mlp_dim, - pos_dim, - use_adaln, - use_shared_adaln, - norm_layer=DyT if use_dyt else RMSNorm, - ) - for _ in range(enc_blocks[i]) - ] - self.enc_blocks.append(nn.ModuleList(blocks)) - - self.dec_ctx = dec_ctx - self.dec_blocks = nn.ModuleList() - for i in range(depth): - blocks = [ - TransformerBlock( - dim, - dim if bid == 0 else ctx_dim if dec_ctx else None, - heads, - dim_head, - mlp_dim, - pos_dim, - use_adaln, - use_shared_adaln, - ctx_from_self=bid == 0, - norm_layer=DyT if use_dyt else RMSNorm, - ) - for bid in range(dec_blocks[i]) - ] - self.dec_blocks.append(nn.ModuleList(blocks)) - - self.grad_ckpt = False - - def init_weight(self): - for param in self.parameters(): - if param.ndim == 1: - nn.init.normal_(param, mean=0.0, std=(1 / param.size(0)) ** 0.5) - elif param.ndim == 2: - fan_in = param.size(1) - nn.init.normal_(param, mean=0.0, std=(1 / fan_in) ** 0.5) - elif param.ndim >= 3: - fan_out, *fan_ins = param.shape - # cumprod - fan_in = 1 - for f in fan_ins: - fan_in *= f - nn.init.normal_(param, mean=0.0, std=(1 / fan_in) ** 0.5) - - def forward( - self, - x, - ctx=None, - x_mask=None, - ctx_mask=None, - pos_map=None, - y=None, - shared_adaln=None, - return_enc_out=False, - ): - if pos_map is not None: - assert pos_map.size(1) == x.size(1) - - self_ctx = [] - for blocks in self.enc_blocks: - for block in blocks: - if self.grad_ckpt: - x = checkpoint( - block, - x, - ctx, - pos_map, - None, - y, - x_mask, - ctx_mask, - shared_adaln, - use_reentrant=False, - ) - else: - x = block(x, ctx, pos_map, None, y, x_mask, ctx_mask, shared_adaln) - self_ctx.append(x) - enc_out = x - - for blocks in self.dec_blocks: - first_block = blocks[0] - if self.grad_ckpt: - x = checkpoint( - first_block, - x, - self_ctx[-1], - pos_map, - pos_map, - y, - x_mask, - ctx_mask, - shared_adaln, - use_reentrant=False, - ) - else: - x = first_block( - x, self_ctx[-1], pos_map, pos_map, y, x_mask, ctx_mask, shared_adaln - ) - - for block in blocks[1:]: - if self.grad_ckpt: - x = checkpoint( - block, - x, - ctx if self.dec_ctx else None, - pos_map, - None, - y, - x_mask, - ctx_mask, - shared_adaln, - use_reentrant=False, - ) - else: - x = block( - x, - ctx if self.dec_ctx else None, - pos_map, - None, - y, - x_mask, - ctx_mask, - shared_adaln, - ) - - if return_enc_out: - return x, enc_out - return x - - -class XUDiT(nn.Module): - """ - Xross-U-Transformer for Image Gen (XUDiT). - """ - - def __init__( - self, - patch_size=2, - input_dim=4, - dim=1024, - ctx_dim=1024, - ctx_size=256, - heads=16, - dim_head=64, - mlp_dim=3072, - depth=8, - enc_blocks=1, - dec_blocks=2, - dec_ctx=False, - class_cond=0, - shared_adaln=True, - concat_ctx=True, - use_dyt=False, - double_t=False, - addon_info_embs_dim=None, - tread_config=None, - ): - super().__init__() - self.backbone = XUTBackBone( - dim, - None if concat_ctx else ctx_dim, - heads, - dim_head, - mlp_dim, - 2, - depth, - enc_blocks, - dec_blocks, - use_adaln=True, - use_shared_adaln=shared_adaln, - dec_ctx=dec_ctx, - use_dyt=use_dyt, - ) - - self.use_tread = False - if tread_config is not None: - self.use_tread = True - self.dropout_ratio = tread_config["dropout_ratio"] - self.prev_tread_trns = TBackBone( - dim, - None if concat_ctx else ctx_dim, - heads, - dim_head, - mlp_dim, - 2, - tread_config["prev_trns_depth"], - use_adaln=True, - use_shared_adaln=shared_adaln, - use_dyt=use_dyt, - ) - self.post_tread_trns = TBackBone( - dim, - None if concat_ctx else ctx_dim, - heads, - dim_head, - mlp_dim, - 2, - tread_config["post_trns_depth"], - use_adaln=True, - use_shared_adaln=shared_adaln, - use_dyt=use_dyt, - ) - - self.patch_size = patch_size - self.in_patch = PatchEmbed(patch_size, input_dim, dim) - self.out_patch = UnPatch(patch_size, dim, input_dim) - self.time_emb = TimestepEmbedding(dim) - if double_t: - self.r_emb = TimestepEmbedding(dim) - if shared_adaln: - self.shared_adaln_attn = nn.Sequential( - nn.LayerNorm(dim), - nn.Linear(dim, dim * 4), - nn.Mish(), - nn.Linear(dim * 4, dim * 3), - ) - nn.init.constant_(self.shared_adaln_attn[-1].bias, 0) - nn.init.constant_(self.shared_adaln_attn[-1].weight, 0) - self.shared_adaln_xattn = nn.Sequential( - nn.LayerNorm(dim), - nn.Linear(dim, dim * 4), - nn.Mish(), - nn.Linear(dim * 4, dim * 3), - ) - nn.init.constant_(self.shared_adaln_xattn[-1].bias, 0) - nn.init.constant_(self.shared_adaln_xattn[-1].weight, 0) - self.shared_adaln_ffw = nn.Sequential( - nn.LayerNorm(dim), - nn.Linear(dim, dim * 4), - nn.Mish(), - nn.Linear(dim * 4, dim * 3), - ) - nn.init.constant_(self.shared_adaln_ffw[-1].bias, 0) - nn.init.constant_(self.shared_adaln_ffw[-1].weight, 0) - if class_cond > 0: - self.class_token = nn.Embedding(class_cond, dim) - else: - self.class_token = None - if concat_ctx and ctx_dim is not None: - self.ctx_proj = nn.Linear(ctx_dim, dim) - else: - self.ctx_proj = None - if addon_info_embs_dim is not None: - self.addon_info_embs_proj = nn.Sequential( - nn.Linear(addon_info_embs_dim, dim), nn.Mish(), nn.Linear(dim, dim) - ) - nn.init.constant_(self.addon_info_embs_proj[-1].bias, 0) - nn.init.constant_(self.addon_info_embs_proj[-1].weight, 0) - - self.concat_ctx = concat_ctx - self.shared_adaln = shared_adaln - self.need_ctx = ctx_dim is not None - self.ctx_dim = ctx_dim - self.ctx_size = ctx_size - self.grad_ckpt = False - self.init_weight() - - def init_weight(self): - if isinstance(self.out_patch.proj, nn.Linear): - nn.init.normal_( - self.out_patch.proj.weight, - mean=0.0, - std=1 / self.out_patch.proj.in_features**2, - ) - - def set_grad_ckpt(self, grad_ckpt): - self.backbone.grad_ckpt = grad_ckpt - self.grad_ckpt = grad_ckpt - if self.use_tread: - self.prev_tread_trns.grad_ckpt = grad_ckpt - self.post_tread_trns.grad_ckpt = grad_ckpt - - def forward( - self, - x, - t, - ctx=None, - pos_map=None, - r=None, - addon_info=None, - tread_rate=None, - return_enc_out=False, - ): - n, c, h, w = x.size() - t = t.reshape(n, -1) - x, pos_map = self.in_patch(x, pos_map) - x = x.contiguous() - if pos_map is None: - pos_map = ( - make_axial_pos( - h // self.patch_size, - w // self.patch_size, - dtype=x.dtype, - device=x.device, - ) - .unsqueeze(0) - .expand(n, -1, -1) - ) - t_emb = self.time_emb(t) - if r is not None: - t_emb = t_emb + self.r_emb((t - r.reshape(n, -1))) - if self.class_token is not None and ctx is not None: - if ctx.ndim == 1: - ctx = ctx[:, None] - t_emb = t_emb + self.class_token(ctx) - ctx = None - if addon_info is not None: - if addon_info.ndim == 1: - # [B] -> [B, 1] for single value info - addon_info = addon_info[:, None] - # [B, D] -> [B, 1, D] for t_emb shape - addon_embs = self.addon_info_embs_proj(addon_info)[:, None] - t_emb = t_emb + addon_embs - if ctx == None and self.need_ctx: - ctx = torch.zeros(n, self.ctx_size, self.ctx_dim, device=x.device) - - if self.shared_adaln: - shared_adaln_state = [ - self.shared_adaln_attn(t_emb).chunk(3, dim=-1), - self.shared_adaln_xattn(t_emb).chunk(3, dim=-1), - self.shared_adaln_ffw(t_emb).chunk(3, dim=-1), - ] - else: - shared_adaln_state = None - - length = x.size(1) - if self.ctx_proj is not None: - ctx = self.ctx_proj(ctx) - x = torch.cat([x, ctx], dim=1) - if pos_map is not None: - pos_map = torch.cat( - [ - pos_map, - torch.zeros(n, ctx.size(1), pos_map.size(2), device=x.device), - ], - dim=1, - ) - ctx = None - - if self.use_tread: - x = self.prev_tread_trns( - x, - ctx=ctx, - pos_map=pos_map, - y=t_emb, - shared_adaln=shared_adaln_state, - ) - if self.training or tread_rate is not None: - xt_selection_length = selection_length = length - int( - length * (tread_rate or self.dropout_ratio) - ) - selection = torch.stack( - [ - torch.randperm(length, device=x.device) < selection_length - for _ in range(n) - ] - ) - if self.ctx_proj is not None: - ctx_length = x.size(1) - length - selection = torch.concat( - [ - selection, - torch.ones( - n, ctx_length, device=x.device, dtype=torch.bool - ), - ], - dim=1, - ) - selection_length += ctx_length - full_length = x.size(1) - not_masked_part = x[~selection, :] - masked_part = x[selection, :].unflatten(0, (n, selection_length)) - x = masked_part - raw_pos_map = pos_map - pos_map = pos_map[selection, :].unflatten(0, (n, selection_length)) - backbone_out = self.backbone( - x, - ctx=ctx, - pos_map=pos_map, - y=t_emb, - shared_adaln=shared_adaln_state, - return_enc_out=return_enc_out, - ) - if return_enc_out: - backbone_out, enc_out = backbone_out - if self.use_tread: - if self.training or tread_rate is not None: - out = torch.empty( - n, full_length, x.size(2), device=x.device, dtype=x.dtype - ) - out[~selection, :] = not_masked_part - out[selection, :] = backbone_out.flatten(0, 1) - pos_map = raw_pos_map - else: - out = backbone_out - out = self.post_tread_trns( - out, - ctx=ctx, - pos_map=pos_map, - y=t_emb, - shared_adaln=shared_adaln_state, - ) - else: - out = backbone_out - out = out[:, :length] - out = self.out_patch(out, h, w) - - if return_enc_out: - length = ( - xt_selection_length if self.use_tread and self.training else full_length - ) - return out, enc_out[:, :length] - return out diff --git a/pipelines/model_hdm.py b/pipelines/model_hdm.py deleted file mode 100644 index 03c49406f..000000000 --- a/pipelines/model_hdm.py +++ /dev/null @@ -1,34 +0,0 @@ -import sys -import torch -import diffusers -from modules import shared, devices, sd_models, errors -from modules.logger import log - - -def load_hdm(checkpoint_info, diffusers_load_config=None): # pylint: disable=unused-argument - if diffusers_load_config is None: - diffusers_load_config = {} - repo_id = sd_models.path_to_repo(checkpoint_info) - sd_models.hf_auth_check(checkpoint_info) - - try: - devices.dtype = torch.float16 - diffusers_load_config['torch_dtype'] = torch.float16 - torch.set_float32_matmul_precision("high") - from pipelines.hdm import hdm - sys.modules['hdm'] = hdm - from pipelines.hdm.hdm.pipeline import HDMXUTPipeline - diffusers.HDMXUTPipeline = HDMXUTPipeline - pipe = diffusers.HDMXUTPipeline.from_pretrained( - repo_id, - cache_dir=shared.opts.diffusers_dir, - trust_remote_code=True, - **diffusers_load_config, - ).to(devices.device) - except Exception as e: - log.error(f'Load HDM-XUT: path="{checkpoint_info.path}" {e}') - errors.display(e, 'hdm') - return None - - devices.torch_gc(force=True, reason='load') - return pipe