diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index d7237ee1e..69f7574c4 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -14,7 +14,7 @@ repos: - id: check-yaml - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.16.2 + rev: v0.16.1 hooks: # Run the Ruff linter, but not the formatter. - id: ruff diff --git a/modules/dataLoader/LTXBaseDataLoader.py b/modules/dataLoader/LTXBaseDataLoader.py new file mode 100644 index 000000000..a3f7aafd1 --- /dev/null +++ b/modules/dataLoader/LTXBaseDataLoader.py @@ -0,0 +1,145 @@ +import os + +from modules.dataLoader.BaseDataLoader import BaseDataLoader +from modules.dataLoader.mixin.DataLoaderText2ImageMixin import DataLoaderText2ImageMixin +from modules.model.LTXModel import PROMPT_MAX_LENGTH, PROMPT_PADDING_SIDE, LTXModel +from modules.modelSetup.BaseLTXSetup import BaseLTXSetup +from modules.util import factory +from modules.util.config.TrainConfig import TrainConfig +from modules.util.enum.ModelType import ModelType +from modules.util.TrainProgress import TrainProgress + +from mgds.pipelineModules.DecodeTokens import DecodeTokens +from mgds.pipelineModules.DecodeVAE import DecodeVAE +from mgds.pipelineModules.EncodeGemma3Text import EncodeGemma3Text +from mgds.pipelineModules.EncodeLTX2Connectors import EncodeLTX2Connectors +from mgds.pipelineModules.EncodeVAE import EncodeVAE +from mgds.pipelineModules.RescaleImageChannels import RescaleImageChannels +from mgds.pipelineModules.SampleVAEDistribution import SampleVAEDistribution +from mgds.pipelineModules.SaveImage import SaveImage +from mgds.pipelineModules.SaveText import SaveText +from mgds.pipelineModules.SaveVideo import SaveVideo +from mgds.pipelineModules.ScaleImage import ScaleImage +from mgds.pipelineModules.Tokenize import Tokenize + + +@factory.register(BaseDataLoader, ModelType.LTX_2) +class LTXBaseDataLoader( + BaseDataLoader, + DataLoaderText2ImageMixin, +): + def _preparation_modules(self, config: TrainConfig, model: LTXModel): + rescale_image = RescaleImageChannels(image_in_name='image', image_out_name='image', in_range_min=0, in_range_max=1, out_range_min=-1, out_range_max=1) + encode_image = EncodeVAE(in_name='image', out_name='latent_image_distribution', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) + image_sample = SampleVAEDistribution(in_name='latent_image_distribution', out_name='latent_image', mode='mean') + # LTX VAE is 32x spatial (not the usual 8x), so the mask latent scales by 1/32, not 1/8 + downscale_mask = ScaleImage(in_name='mask', out_name='latent_mask', factor=1/32) + tokenize_prompt = Tokenize(in_name='prompt', tokens_out_name='tokens', mask_out_name='tokens_mask', tokenizer=model.tokenizer, max_token_length=PROMPT_MAX_LENGTH, padding_side=PROMPT_PADDING_SIDE) + encode_prompt = EncodeGemma3Text(tokens_name='tokens', tokens_attention_mask_in_name='tokens_mask', hidden_state_out_name='text_encoder_hidden_state', tokens_attention_mask_out_name='tokens_mask', + text_encoder=model.text_encoder, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) + # cache the connectors' small per-modality output instead of the huge stacked TE output: ~30x smaller + # text cache and no connector forward at train time. No Prune/Pad - the connector output is dense + # fixed-length, since learnable registers replace padding. + encode_connectors = EncodeLTX2Connectors( + hidden_state_in_name='text_encoder_hidden_state', tokens_attention_mask_in_name='tokens_mask', + video_embeds_out_name='connector_video_embeds', audio_embeds_out_name='connector_audio_embeds', + connectors=model.connectors, padding_side=PROMPT_PADDING_SIDE, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) + + modules = [rescale_image, encode_image, image_sample] + if config.masked_training or config.model_type.has_mask_input(): + modules.append(downscale_mask) + modules += [tokenize_prompt, encode_prompt, encode_connectors] + + return modules + + def _cache_modules(self, config: TrainConfig, model: LTXModel, model_setup: BaseLTXSetup): + image_split_names = ['latent_image', 'original_resolution', 'crop_offset'] + + if config.masked_training or config.model_type.has_mask_input(): + image_split_names.append('latent_mask') + + image_aggregate_names = ['crop_resolution', 'image_path'] + + text_split_names = [] + + sort_names = image_aggregate_names + image_split_names + [ + 'prompt', 'tokens', 'tokens_mask', 'connector_video_embeds', + 'concept' + ] + + text_split_names += ['tokens', 'tokens_mask', 'connector_video_embeds'] + + return self._cache_modules_from_names( + model, model_setup, + image_split_names=image_split_names, + image_aggregate_names=image_aggregate_names, + text_split_names=text_split_names, + sort_names=sort_names, + config=config, + text_caching=True, + ) + + def _output_modules(self, config: TrainConfig, model: LTXModel, model_setup: BaseLTXSetup): + output_names = [ + 'image_path', 'latent_image', + 'prompt', + 'tokens', + 'tokens_mask', + 'original_resolution', 'crop_resolution', 'crop_offset', + 'connector_video_embeds', + ] + + if config.masked_training or config.model_type.has_mask_input(): + output_names.append('latent_mask') + + return self._output_modules_from_out_names( + model, model_setup, + output_names=output_names, + config=config, + use_conditioning_image=False, + vae=model.vae, + autocast_context=[model.autocast_context], + train_dtype=model.train_dtype, + ) + + def _debug_modules(self, config: TrainConfig, model: LTXModel): + debug_dir = os.path.join(config.debug_dir, "dataloader") + + def before_save_fun(): + model.materialize_only("vae") + + decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) + upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=32) + decode_prompt = DecodeTokens(in_name='tokens', out_name='decoded_prompt', tokenizer=model.tokenizer) + + # SaveVideo instead of SaveImage: latents here are 5D (vae_frame_dim), so the decode is [C, F, H, W] + # and SaveImage's ToPILImage can't take it (the #1015 FIXME the other video loaders still carry). + save_video = SaveVideo(video_in_name='decoded_image', original_path_in_name='image_path', path=debug_dir, in_range_min=-1, in_range_max=1, fps=model.NATIVE_FPS, before_save_fun=before_save_fun) + + save_mask = SaveImage(image_in_name='decoded_mask', original_path_in_name='image_path', path=debug_dir, in_range_min=0, in_range_max=1, before_save_fun=before_save_fun) + save_prompt = SaveText(text_in_name='decoded_prompt', original_path_in_name='image_path', path=debug_dir, before_save_fun=before_save_fun) + + modules = [decode_image, save_video] + + if config.masked_training or config.model_type.has_mask_input(): + modules += [upscale_mask, save_mask] + + modules += [decode_prompt, save_prompt] + + return modules + + def _create_dataset( + self, + config: TrainConfig, + model: LTXModel, + model_setup: BaseLTXSetup, + train_progress: TrainProgress, + is_validation: bool = False, + ): + return DataLoaderText2ImageMixin._create_dataset(self, + config, model, model_setup, train_progress, is_validation, + aspect_bucketing_quantization=64, + frame_dim_enabled=True, + allow_video_files=True, + vae_frame_dim=True, + ) diff --git a/modules/model/AnimaModel.py b/modules/model/AnimaModel.py index 9c365273d..51fdd1a93 100644 --- a/modules/model/AnimaModel.py +++ b/modules/model/AnimaModel.py @@ -7,7 +7,6 @@ from modules.util.convert_util import add_prefix from modules.util.enum.DataType import DataType from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -39,8 +38,6 @@ class AnimaModel(BaseModel): text_encoder_train_dtype: DataType - text_encoder_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent lora training data transformer_lora: LoRAModuleWrapper | None @@ -66,8 +63,6 @@ def __init__( self.text_encoder_train_dtype = DataType.FLOAT_32 - self.text_encoder_offload_conductor = None - self.transformer_offload_conductor = None self.transformer_lora = None self.lora_state_dict = None diff --git a/modules/model/BaseModel.py b/modules/model/BaseModel.py index 20d7ae0ea..cf75827f3 100644 --- a/modules/model/BaseModel.py +++ b/modules/model/BaseModel.py @@ -1,4 +1,5 @@ from abc import ABCMeta +from collections.abc import Callable from contextlib import nullcontext from uuid import uuid4 @@ -6,12 +7,14 @@ from modules.module.LoRAModule import LoRAModuleWrapper from modules.util.config.TrainConfig import TrainConfig from modules.util.convert_util import qkv_fusion +from modules.util.disk_stream import stream_module_to from modules.util.enum.DataType import DataType from modules.util.enum.ModelFormat import ModelFormat from modules.util.enum.ModelType import ModelType +from modules.util.LayerOffloadConductor import LayerOffloadConductor from modules.util.modelSpec.ModelSpec import ModelSpec from modules.util.NamedParameterGroup import NamedParameterGroupCollection -from modules.util.torch_util import device_equals, torch_gc +from modules.util.torch_util import create_mem_pool, device_equals, mem_pool_context, supports_mem_pool, torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -21,6 +24,21 @@ from transformers import PreTrainedTokenizer +def _module_is_on(module: torch.nn.Module | LoRAModuleWrapper, device: torch.device) -> bool: + # Read a module's residency off its first tensor (parameters first, buffers for a parameterless module). + # module.to() is already a per-tensor no-op when the device matches, but whether it moved anything has to be + # known before the call to decide whether the collection afterwards is worth running. One tensor is enough: + # the modules asked here are moved as a unit by the branch below, so they are never split across devices. + # Both callers are duck-typed on .to(): a LoRAModuleWrapper is not an nn.Module, returns a plain list from + # parameters() (hence iter(), not next() on it directly) and has no buffers() at all. + tensor = next(iter(module.parameters()), None) + if tensor is None and hasattr(module, "buffers"): + tensor = next(iter(module.buffers()), None) + if tensor is None: + return True # nothing to move, so it is trivially where it was asked to be + return device_equals(tensor.device, device) + + class BaseModelEmbedding: def __init__( self, @@ -79,6 +97,9 @@ class BaseModel(metaclass=ABCMeta): embedding_state_dicts: dict[str, dict[str, Tensor]] | None autocast_context: torch.autocast | nullcontext train_dtype: DataType + cache_in_ram: dict[str, bool] + offload_conductor: dict[str, LayerOffloadConductor] + materialize_fn: dict[str, Callable] def __init__( self, @@ -86,6 +107,9 @@ def __init__( ): self.model_type = model_type self.parameters = None + self.cache_in_ram = {} + self.offload_conductor = {} + self.materialize_fn = {} self.optimizer = None self.optimizer_state_dict = None self.param_group_mapping = None @@ -97,6 +121,8 @@ def __init__( self.autocast_context = nullcontext() self.train_dtype = DataType.FLOAT_32 + self._mem_pools = {} + @property def train_device(self) -> torch.device: return torch.device(self.train_config.train_device) @@ -112,9 +138,15 @@ def materialize(self, *parts: str): def evict(self, *parts: str): # Move `parts` onto temp_device. No parts given -> every component in ModelType.model_parts(). + # Collect once at the end, and only if something was actually freed: materialize_only() evicts every + # part it doesn't want on every call, so most evictions here are of parts that are already evicted and + # have nothing to reclaim. Each part reports whether it moved rather than the model tracking residency, + # so the answer always comes from the component that did (or didn't) do the move. + moved = False for part in parts or self.model_type.model_parts(): - self._move_part(part, self.temp_device) - torch_gc() + moved |= self._move_part(part, self.temp_device) + if moved: + torch_gc() def materialize_only(self, *parts: str): # Materialize exactly `parts` on train_device; evict every other component in ModelType.model_parts() @@ -131,28 +163,68 @@ def materialize_only_text_encoders(self): # call this before encode_text, which reads every text encoder the model has. self.materialize_only(*self.model_type.text_encoder_parts()) - def _move_part(self, part: str, device: torch.device): - # The generic per-component move: `part` (or `part_1` for the first of several split text encoders), - # its LoRA (`{part}_lora`), and its layer-offload conductor (`{part}_offload_conductor`), if present. + def _move_part(self, part: str, device: torch.device) -> bool: + # Move a component (`part`, or `part_1` for the first of several split text encoders) and its LoRA. The + # dispatch below routes through an offload conductor and/or a disk-stream materialize closure if present. + # Returns whether any weight actually moved: each of the three paths is idempotent and a repeat call is + # common (materialize_only states a set, not a delta), so the caller needs to know whether this one did + # anything before paying for a collection. stem = f"{part}_1" if hasattr(self, f"{part}_1") else part - conductor = getattr(self, f"{stem}_offload_conductor", None) + conductor = self.offload_conductor.get(stem) + materialize_fn = self.materialize_fn.get(stem) + cache_in_ram = self.cache_in_ram.get(stem, True) + + moved = False if conductor is not None: if device_equals(device, self.temp_device): - conductor.evict() + to_meta = materialize_fn is not None and not cache_in_ram + moved = conductor.evict(to_meta=to_meta) else: assert device_equals(device, self.train_device), f"unexpected device {device} for part {part}" - conductor.materialize() + train_dtype = getattr(self, f"{stem}_train_dtype", self.train_dtype) + moved = conductor.materialize( + train_dtype, name=part, materialize_fn=materialize_fn, + cache_in_ram=cache_in_ram) + elif materialize_fn is not None: + streamed_component = getattr(self, stem) + train_dtype = getattr(self, f"{stem}_train_dtype", self.train_dtype) + moved = stream_module_to( + streamed_component, device, materialize_fn, train_dtype, + cache_in_ram=cache_in_ram, name=part, temp_device=self.temp_device) + + # move into the shared stem pool: the base component itself (unless a conductor or stream owns its move) plus + # the LoRA. getattr(self, stem) is None for a part in model_parts() that was never populated (e.g. an omitted + # text encoder), so it drops out below. + to_move = [] + if conductor is None and materialize_fn is None: + to_move.append(getattr(self, stem)) + lora = getattr(self, f"{stem}_lora", None) + to_move.append(lora) + to_move = [module for module in to_move if module is not None] + if not to_move: + return moved + + moved |= any(not _module_is_on(module, device) for module in to_move) + + if supports_mem_pool(device): + # The component (when not conductor/stream-managed) and its LoRA share a per-stem MemPool so both release + # together on evict, keeping the LoRA's small tensors from pinning freed default-pool segments across the + # part's evict/reload cycle. A conductor keeps its own pool, so the stem pool then holds only the LoRA. + pool = self._mem_pools.get(stem) + if pool is None: + pool = self._mem_pools[stem] = create_mem_pool(device) + with mem_pool_context(pool): + for module in to_move: + module.to(device=device) else: - component = getattr(self, stem) # raises if `part` doesn't name a real attribute - # None when the part is excluded from training (e.g. a text encoder with include_text_encoder off): - # it stays in model_parts() but the loader never populated it, so there is nothing to move. - if component is not None: - component.to(device=device) + # the target has no MemPool (CPU): move normally and drop this stem's pool from the earlier GPU move, + # so evict()'s torch_gc can release its segments + for module in to_move: + module.to(device=device) + self._mem_pools.pop(stem, None) - lora = getattr(self, f"{stem}_lora", None) - if lora is not None: - lora.to(device) + return moved def eval(self): # Put every present component on eval(); driven by the same part registry as materialize()/evict(). diff --git a/modules/model/ChromaModel.py b/modules/model/ChromaModel.py index 4fbcf430b..3f15219e2 100644 --- a/modules/model/ChromaModel.py +++ b/modules/model/ChromaModel.py @@ -8,7 +8,6 @@ from modules.util.enum.DataType import DataType from modules.util.enum.ModelFormat import ModelFormat from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -52,8 +51,6 @@ class ChromaModel(BaseModel): text_encoder_train_dtype: DataType - text_encoder_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent embedding training data embedding: ChromaModelEmbedding | None @@ -84,8 +81,6 @@ def __init__( self.text_encoder_train_dtype = DataType.FLOAT_32 - self.text_encoder_offload_conductor = None - self.transformer_offload_conductor = None self.embedding = None self.additional_embeddings = [] diff --git a/modules/model/ErnieModel.py b/modules/model/ErnieModel.py index 42bb1be74..2e920a47a 100644 --- a/modules/model/ErnieModel.py +++ b/modules/model/ErnieModel.py @@ -8,7 +8,6 @@ from modules.model.BaseModel import BaseModel from modules.module.LoRAModule import LoRAModuleWrapper from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -34,8 +33,6 @@ class ErnieModel(BaseModel): # autocast context text_encoder_autocast_context: torch.autocast | nullcontext - text_encoder_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None transformer_lora: LoRAModuleWrapper | None lora_state_dict: dict | None @@ -56,8 +53,6 @@ def __init__( self.text_encoder_autocast_context = nullcontext() - self.text_encoder_offload_conductor = None - self.transformer_offload_conductor = None self.transformer_lora = None self.lora_state_dict = None diff --git a/modules/model/Flux2Model.py b/modules/model/Flux2Model.py index 79eb02006..f7887e3f2 100644 --- a/modules/model/Flux2Model.py +++ b/modules/model/Flux2Model.py @@ -6,7 +6,6 @@ from modules.module.LoRAModule import LoRAModuleWrapper from modules.util.convert_util import chunk_swap from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -43,8 +42,6 @@ class Flux2Model(BaseModel): # autocast context text_encoder_autocast_context: torch.autocast | nullcontext - text_encoder_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None transformer_lora: LoRAModuleWrapper | None lora_state_dict: dict | None @@ -65,8 +62,6 @@ def __init__( self.text_encoder_autocast_context = nullcontext() - self.text_encoder_offload_conductor = None - self.transformer_offload_conductor = None self.transformer_lora = None self.lora_state_dict = None diff --git a/modules/model/FluxModel.py b/modules/model/FluxModel.py index 5236b3498..e61d1a4f3 100644 --- a/modules/model/FluxModel.py +++ b/modules/model/FluxModel.py @@ -11,7 +11,6 @@ from modules.util.enum.DataType import DataType from modules.util.enum.ModelFormat import ModelFormat from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -66,8 +65,6 @@ class FluxModel(BaseModel): text_encoder_2_train_dtype: DataType - text_encoder_2_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent embedding training data embedding: FluxModelEmbedding | None @@ -103,8 +100,6 @@ def __init__( self.text_encoder_2_train_dtype = DataType.FLOAT_32 - self.text_encoder_2_offload_conductor = None - self.transformer_offload_conductor = None self.embedding = None self.additional_embeddings = [] diff --git a/modules/model/HiDreamModel.py b/modules/model/HiDreamModel.py index f99b5116b..8d39d36e8 100644 --- a/modules/model/HiDreamModel.py +++ b/modules/model/HiDreamModel.py @@ -9,7 +9,6 @@ from modules.util.enum.DataType import DataType from modules.util.enum.ModelFormat import ModelFormat from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -98,9 +97,6 @@ class HiDreamModel(BaseModel): text_encoder_3_train_dtype: DataType transformer_train_dtype: DataType - text_encoder_3_offload_conductor: LayerOffloadConductor | None - text_encoder_4_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent embedding training data embedding: HiDreamModelEmbedding | None @@ -149,9 +145,6 @@ def __init__( self.text_encoder_3_train_dtype = DataType.FLOAT_32 self.transformer_train_dtype = DataType.FLOAT_32 - self.text_encoder_3_offload_conductor = None - self.text_encoder_4_offload_conductor = None - self.transformer_offload_conductor = None self.embedding = None self.additional_embeddings = [] diff --git a/modules/model/HunyuanVideoModel.py b/modules/model/HunyuanVideoModel.py index 9ec6b4516..15bf26b7b 100644 --- a/modules/model/HunyuanVideoModel.py +++ b/modules/model/HunyuanVideoModel.py @@ -10,7 +10,6 @@ from modules.util.enum.DataType import DataType from modules.util.enum.ModelFormat import ModelFormat from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -80,8 +79,6 @@ class HunyuanVideoModel(BaseModel): transformer_train_dtype: DataType - text_encoder_1_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent embedding training data embedding: HunyuanVideoModelEmbedding | None @@ -118,8 +115,6 @@ def __init__( self.transformer_train_dtype = DataType.FLOAT_32 - self.text_encoder_1_offload_conductor = None - self.transformer_offload_conductor = None self.embedding = None self.additional_embeddings = [] diff --git a/modules/model/IdeogramModel.py b/modules/model/IdeogramModel.py index 70fae8444..755b6cb3e 100644 --- a/modules/model/IdeogramModel.py +++ b/modules/model/IdeogramModel.py @@ -7,7 +7,6 @@ from modules.model.BaseModel import BaseModel from modules.module.LoRAModule import LoRAModuleWrapper from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -36,9 +35,6 @@ class IdeogramModel(BaseModel): # autocast context text_encoder_autocast_context: torch.autocast | nullcontext - text_encoder_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None - unconditional_transformer_offload_conductor: LayerOffloadConductor | None transformer_lora: LoRAModuleWrapper | None lora_state_dict: dict | None @@ -60,9 +56,6 @@ def __init__( self.text_encoder_autocast_context = nullcontext() - self.text_encoder_offload_conductor = None - self.transformer_offload_conductor = None - self.unconditional_transformer_offload_conductor = None self.transformer_lora = None self.lora_state_dict = None diff --git a/modules/model/Krea2Model.py b/modules/model/Krea2Model.py index cf5ca0d69..c1c7a8976 100644 --- a/modules/model/Krea2Model.py +++ b/modules/model/Krea2Model.py @@ -5,7 +5,6 @@ from modules.model.BaseModel import BaseModel from modules.module.LoRAModule import LoRAModuleWrapper from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch import torch.nn.functional as F @@ -51,8 +50,6 @@ class Krea2Model(BaseModel): # autocast context text_encoder_autocast_context: torch.autocast | nullcontext - text_encoder_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent lora training data transformer_lora: LoRAModuleWrapper | None @@ -74,8 +71,6 @@ def __init__( self.text_encoder_autocast_context = nullcontext() - self.text_encoder_offload_conductor = None - self.transformer_offload_conductor = None self.transformer_lora = None self.lora_state_dict = None diff --git a/modules/model/LTXModel.py b/modules/model/LTXModel.py new file mode 100644 index 000000000..358c9587c --- /dev/null +++ b/modules/model/LTXModel.py @@ -0,0 +1,281 @@ +import math +from contextlib import nullcontext + +from modules.model.BaseModel import BaseModel +from modules.module.LoRAModule import LoRAModuleWrapper +from modules.util.convert_util import add_prefix +from modules.util.enum.DataType import DataType +from modules.util.enum.ModelType import ModelType + +import torch +from torch import Tensor + +from diffusers import ( + AutoencoderKLLTX2Audio, + AutoencoderKLLTX2Video, + DiffusionPipeline, + FlowMatchEulerDiscreteScheduler, + LTX2Pipeline, + LTX2VideoTransformer3DModel, +) +from diffusers.pipelines.ltx2.connectors import LTX2TextConnectors +from diffusers.pipelines.ltx2.vocoder import LTX2Vocoder, LTX2VocoderWithBWE +from transformers import ( + Gemma3ForConditionalGeneration, + Gemma4UnifiedForConditionalGeneration, + GemmaTokenizer, + GemmaTokenizerFast, +) + +PROMPT_MAX_LENGTH = 1024 +# Gemma expects left padding for chat-style prompts, matching LTX2Pipeline._get_gemma_prompt_embeds. The +# checkpoint's tokenizer_config.json doesn't set it (transformers defaults to "right"). +PROMPT_PADDING_SIDE = "left" + + +class LTXModel(BaseModel): + NATIVE_FPS = 24 + + # base model data + # LTX 2.3 names the fast tokenizer and the Gemma 3 encoder, LTX 2.5 the slow tokenizer and the Gemma 4 one + tokenizer: GemmaTokenizer | GemmaTokenizerFast | None + noise_scheduler: FlowMatchEulerDiscreteScheduler | None + text_encoder: Gemma3ForConditionalGeneration | Gemma4UnifiedForConditionalGeneration | None + vae: AutoencoderKLLTX2Video | None + connectors: LTX2TextConnectors | None + transformer: LTX2VideoTransformer3DModel | None + low_noise_transformer: LTX2VideoTransformer3DModel | None + + # audio branch - frozen, never run + audio_vae: AutoencoderKLLTX2Audio | None + vocoder: LTX2Vocoder | LTX2VocoderWithBWE | None + + transformer_autocast_context: torch.autocast | nullcontext + low_noise_transformer_autocast_context: torch.autocast | nullcontext + text_encoder_autocast_context: torch.autocast | nullcontext + connectors_autocast_context: torch.autocast | nullcontext + transformer_train_dtype: DataType + low_noise_transformer_train_dtype: DataType + text_encoder_train_dtype: DataType + connectors_train_dtype: DataType + + transformer_lora: LoRAModuleWrapper | None + lora_state_dict: dict | None + + def __init__( + self, + model_type: ModelType, + ): + super().__init__( + model_type=model_type, + ) + + self.tokenizer = None + self.noise_scheduler = None + self.text_encoder = None + self.vae = None + self.connectors = None + self.transformer = None + self.low_noise_transformer = None + + self.audio_vae = None + self.vocoder = None + + self.transformer_autocast_context = nullcontext() + self.low_noise_transformer_autocast_context = nullcontext() + self.text_encoder_autocast_context = nullcontext() + self.connectors_autocast_context = nullcontext() + self.transformer_train_dtype = DataType.FLOAT_32 + self.low_noise_transformer_train_dtype = DataType.FLOAT_32 + self.text_encoder_train_dtype = DataType.FLOAT_32 + self.connectors_train_dtype = DataType.FLOAT_32 + + self.transformer_lora = None + self.lora_state_dict = None + + @staticmethod + def _attn(name: str) -> tuple: + # one attention module's leaves. Only the qk norms are renamed; the projections are spelled the same in + # both namespaces and are listed so the block rule stays strict -- a leaf with no rule is an error, not a + # silent passthrough. + return (name, name, [ + ("norm_q", "q_norm"), + ("norm_k", "k_norm"), + ("to_q", "to_q"), + ("to_k", "to_k"), + ("to_v", "to_v"), + ("to_out.0", "to_out.0"), + ("to_gate_logits", "to_gate_logits"), + ]) + + def diffusers_to_original(self) -> list | None: + # rename only -- LTX-2's native (Lightricks/ComfyUI) layout differs from diffusers in module names alone. + # The two namespaces name the same modules, so this one body serves both the full checkpoint and a LoRA; + # they differ only in the top prefix each format adds (see checkpoint_diffusers_to_original). + return [ + ("proj_in", "patchify_proj"), + ("audio_proj_in", "audio_patchify_proj"), + ("proj_out", "proj_out"), + ("audio_proj_out", "audio_proj_out"), + # every modulation predictor is an "adaln_single" natively; diffusers names each after what it feeds + ("time_embed", "adaln_single"), + ("audio_time_embed", "audio_adaln_single"), + ("prompt_adaln", "prompt_adaln_single"), + ("audio_prompt_adaln", "audio_prompt_adaln_single"), + ("av_cross_attn_video_scale_shift", "av_ca_video_scale_shift_adaln_single"), + ("av_cross_attn_video_a2v_gate", "av_ca_a2v_gate_adaln_single"), + ("av_cross_attn_audio_scale_shift", "av_ca_audio_scale_shift_adaln_single"), + ("av_cross_attn_audio_v2a_gate", "av_ca_v2a_gate_adaln_single"), + ("scale_shift_table", "scale_shift_table"), + ("audio_scale_shift_table", "audio_scale_shift_table"), + # LTX 2.5 only + ("keyframes_abs_pos_embedding", "keyframes_abs_pos_embedding"), + ("transformer_blocks.{i}", "transformer_blocks.{i}", [ + ("video_a2v_cross_attn_scale_shift_table", "scale_shift_table_a2v_ca_video"), + ("audio_a2v_cross_attn_scale_shift_table", "scale_shift_table_a2v_ca_audio"), + ("scale_shift_table", "scale_shift_table"), + ("audio_scale_shift_table", "audio_scale_shift_table"), + ("prompt_scale_shift_table", "prompt_scale_shift_table"), + ("audio_prompt_scale_shift_table", "audio_prompt_scale_shift_table"), + self._attn("attn1"), + self._attn("attn2"), + self._attn("audio_attn1"), + self._attn("audio_attn2"), + self._attn("audio_to_video_attn"), + self._attn("video_to_audio_attn"), + ("ff", "ff"), + ("audio_ff", "audio_ff"), + ]), + ] + + def checkpoint_diffusers_to_original(self) -> list | None: + # the full checkpoint nests the transformer under model.diffusion_model., beside the vae/text-projection + # branches of the same file (a LoRA carries only diffusion_model., which the saver adds). + return [self.diffusers_to_original(), add_prefix("model.diffusion_model")] + + def create_pipeline(self) -> DiffusionPipeline: + return LTX2Pipeline( + scheduler=self.noise_scheduler, + vae=self.vae, + audio_vae=self.audio_vae, + text_encoder=self.text_encoder, + tokenizer=self.tokenizer, + connectors=self.connectors, + transformer=self.transformer, + vocoder=self.vocoder, + ) + + def encode_text( + self, + train_device: torch.device, + text: str | list[str] | None = None, + connector_video_embeds: Tensor | None = None, + text_encoder_dropout_probability: float | None = None, + ) -> Tensor: + if text_encoder_dropout_probability is not None and text_encoder_dropout_probability > 0.0: + raise NotImplementedError # needs a cached null-caption embedding, not zero-out + + if connector_video_embeds is not None: + return connector_video_embeds + + text_encoder_outputs, tokens_mask = self.encode_text_encoder(text=text) + return self.encode_connectors(text_encoder_outputs, tokens_mask, train_device)[0] + + def encode_text_encoder( + self, + text: str | list[str] | None = None, + tokens: Tensor | None = None, + tokens_mask: Tensor | None = None, + ) -> tuple[list[Tensor], Tensor]: + if tokens is None and text is not None: + if isinstance(text, str): + text = [text] + + tokenizer_output = self.tokenizer( + [t.strip() for t in text], + padding='max_length', + padding_side=PROMPT_PADDING_SIDE, + max_length=PROMPT_MAX_LENGTH, + truncation=True, + return_tensors='pt', + add_special_tokens=True, + ) + tokens = tokenizer_output.input_ids.to(self.text_encoder.device) + tokens_mask = tokenizer_output.attention_mask.to(self.text_encoder.device) + + # every layer's hidden state, incl. the embedding layer, flattened to 3D - the "Pack to 3D" step in + # LTX2Pipeline._get_gemma_prompt_embeds + text_encoder_outputs = [] + with self.text_encoder_autocast_context: + for i in range(tokens.shape[0]): + output = self.text_encoder( + tokens[i:i + 1], + attention_mask=tokens_mask[i:i + 1], + output_hidden_states=True, + use_cache=False, + ) + stacked = torch.stack(output.hidden_states, dim=-1).flatten(2, 3) # [1, T, H*L] + # Gemma's residual stream comes out fp32 under autocast; downcast to the TE dtype, which + # halves every full-width copy downstream + text_encoder_outputs.append(stacked.to(self.text_encoder_train_dtype.torch_dtype())) + del output, stacked # drop the retained per-layer hidden_states tuple before the next prompt + + return text_encoder_outputs, tokens_mask + + def encode_connectors( + self, + text_encoder_outputs: list[Tensor], + tokens_mask: Tensor, + train_device: torch.device, + ) -> tuple[Tensor, Tensor]: + video_embeds, audio_embeds = [], [] + for i, text_encoder_output in enumerate(text_encoder_outputs): + video_embed, audio_embed, mask = self.connectors( + text_encoder_output.to(train_device), tokens_mask[i:i + 1].to(train_device), + padding_side=PROMPT_PADDING_SIDE, + ) + # the connectors replace padded positions with learnable registers, so the mask they return is + # all-attend. _assert_async queues the check as a kernel, so it costs no device sync. + torch._assert_async(mask.all(), "connector attention mask is not all-True") + video_embeds.append(video_embed) + audio_embeds.append(audio_embed) + return torch.cat(video_embeds), torch.cat(audio_embeds) + + # video latents [B, C, F, H, W] <-> patch sequence [B, S, D] + @staticmethod + def pack_latents(latents: Tensor, patch_size: int = 1, patch_size_t: int = 1) -> Tensor: + return LTX2Pipeline._pack_latents(latents, patch_size, patch_size_t) + + @staticmethod + def unpack_latents( + latents: Tensor, num_frames: int, height: int, width: int, patch_size: int = 1, patch_size_t: int = 1, + ) -> Tensor: + return LTX2Pipeline._unpack_latents(latents, num_frames, height, width, patch_size, patch_size_t) + + # audio latents [B, C, L, M] -> patch sequence [B, S, D]. Without a patch size the pipeline packs all mel + # bins into one patch, which is what the checkpoint's audio_patch_size config asks for. + @staticmethod + def pack_audio_latents(latents: Tensor) -> Tensor: + return LTX2Pipeline._pack_audio_latents(latents) + + def scale_latents(self, latents: Tensor) -> Tensor: + return LTX2Pipeline._normalize_latents( + latents, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor) + + def unscale_latents(self, latents: Tensor) -> Tensor: + return LTX2Pipeline._denormalize_latents( + latents, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor) + + def calculate_timestep_shift(self, num_latent_frames: int, latent_height: int, latent_width: int) -> float: + # resolution/length-dependent flow-matching shift, matching Lightricks' get_normal_shift. patch_size + # is 1, so the latent element count is the transformer sequence length. + base_seq_len = self.noise_scheduler.config.base_image_seq_len + max_seq_len = self.noise_scheduler.config.max_image_seq_len + base_shift = self.noise_scheduler.config.base_shift + max_shift = self.noise_scheduler.config.max_shift + + image_seq_len = num_latent_frames * latent_height * latent_width + m = (max_shift - base_shift) / (max_seq_len - base_seq_len) + b = base_shift - m * base_seq_len + mu = image_seq_len * m + b + return math.exp(mu) diff --git a/modules/model/PixArtAlphaModel.py b/modules/model/PixArtAlphaModel.py index bf36d1a21..066413fac 100644 --- a/modules/model/PixArtAlphaModel.py +++ b/modules/model/PixArtAlphaModel.py @@ -8,7 +8,6 @@ from modules.util.enum.DataType import DataType from modules.util.enum.ModelFormat import ModelFormat from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -54,8 +53,6 @@ class PixArtAlphaModel(BaseModel): text_encoder_train_dtype: DataType - text_encoder_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent embedding training data embedding: PixArtAlphaModelEmbedding | None @@ -86,8 +83,6 @@ def __init__( self.text_encoder_train_dtype = DataType.FLOAT_32 - self.text_encoder_offload_conductor = None - self.transformer_offload_conductor = None self.embedding = None self.additional_embeddings = [] diff --git a/modules/model/QwenModel.py b/modules/model/QwenModel.py index 3b69d8f8b..f1b5f06a0 100644 --- a/modules/model/QwenModel.py +++ b/modules/model/QwenModel.py @@ -7,7 +7,6 @@ from modules.util.enum.DataType import DataType from modules.util.enum.ModelFormat import ModelFormat from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -38,8 +37,6 @@ class QwenModel(BaseModel): text_encoder_train_dtype: DataType - text_encoder_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent lora training data text_encoder_lora: LoRAModuleWrapper | None @@ -64,8 +61,6 @@ def __init__( self.text_encoder_train_dtype = DataType.FLOAT_32 - self.text_encoder_offload_conductor = None - self.transformer_offload_conductor = None self.text_encoder_lora = None self.transformer_lora = None diff --git a/modules/model/SanaModel.py b/modules/model/SanaModel.py index f43e98ed5..10b932f03 100644 --- a/modules/model/SanaModel.py +++ b/modules/model/SanaModel.py @@ -8,7 +8,6 @@ from modules.util.enum.DataType import DataType from modules.util.enum.ModelFormat import ModelFormat from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -55,8 +54,6 @@ class SanaModel(BaseModel): text_encoder_train_dtype: DataType vae_train_dtype: DataType - text_encoder_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent embedding training data embedding: SanaModelEmbedding | None @@ -88,8 +85,6 @@ def __init__( self.text_encoder_train_dtype = DataType.FLOAT_32 self.vae_train_dtype = DataType.FLOAT_32 - self.text_encoder_offload_conductor = None - self.transformer_offload_conductor = None self.embedding = None self.additional_embeddings = [] diff --git a/modules/model/StableDiffusion3Model.py b/modules/model/StableDiffusion3Model.py index e9324bb1b..076e58aed 100644 --- a/modules/model/StableDiffusion3Model.py +++ b/modules/model/StableDiffusion3Model.py @@ -10,7 +10,6 @@ from modules.util.enum.DataType import DataType from modules.util.enum.ModelFormat import ModelFormat from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -77,8 +76,6 @@ class StableDiffusion3Model(BaseModel): text_encoder_3_train_dtype: DataType - text_encoder_3_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent embedding training data embedding: StableDiffusion3ModelEmbedding | None @@ -119,8 +116,6 @@ def __init__( self.text_encoder_3_train_dtype = DataType.FLOAT_32 - self.text_encoder_3_offload_conductor = None - self.transformer_offload_conductor = None self.embedding = None self.additional_embeddings = [] diff --git a/modules/model/ZImageModel.py b/modules/model/ZImageModel.py index 8b60f2c28..942a6ce08 100644 --- a/modules/model/ZImageModel.py +++ b/modules/model/ZImageModel.py @@ -7,7 +7,6 @@ from modules.util.convert_util import fuse from modules.util.enum.DataType import DataType from modules.util.enum.ModelType import ModelType -from modules.util.LayerOffloadConductor import LayerOffloadConductor import torch from torch import Tensor @@ -42,8 +41,6 @@ class ZImageModel(BaseModel): text_encoder_train_dtype: DataType - text_encoder_offload_conductor: LayerOffloadConductor | None - transformer_offload_conductor: LayerOffloadConductor | None # persistent lora training data text_encoder_lora: LoRAModuleWrapper | None @@ -68,8 +65,6 @@ def __init__( self.text_encoder_train_dtype = DataType.FLOAT_32 #TODO - self.text_encoder_offload_conductor = None - self.transformer_offload_conductor = None self.transformer_lora = None self.lora_state_dict = None diff --git a/modules/modelLoader/AnimaModelLoader.py b/modules/modelLoader/AnimaModelLoader.py index f4d3662bc..edc38eade 100644 --- a/modules/modelLoader/AnimaModelLoader.py +++ b/modules/modelLoader/AnimaModelLoader.py @@ -18,7 +18,6 @@ AutoencoderKLQwenImage, CosmosTransformer3DModel, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, ) from transformers import Qwen2Tokenizer, Qwen3Model, T5TokenizerFast @@ -38,10 +37,12 @@ def __load_internal( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, quantization, + stream_from_disk, ) else: raise Exception("not an internal model") @@ -55,83 +56,56 @@ def __load_diffusers( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): - tokenizer = Qwen2Tokenizer.from_pretrained( + model.tokenizer = Qwen2Tokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) - t5_tokenizer = T5TokenizerFast.from_pretrained( + model.t5_tokenizer = T5TokenizerFast.from_pretrained( base_model_name, subfolder="t5_tokenizer", ) - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( Qwen3Model, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) # conditioner is always bfloat16 — small adapter, no user dtype control - text_conditioner = AnimaTextConditioner.from_pretrained( + model.text_conditioner = AnimaTextConditioner.from_pretrained( base_model_name, subfolder="text_conditioner", torch_dtype=torch.bfloat16, ) - if vae_model_name: #TODO simplify - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) - - if transformer_model_name: - transformer = CosmosTransformer3DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - #avoid loading the transformer in float32: - torch_dtype=torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - CosmosTransformer3DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + model.vae = self._load_vae( + AutoencoderKLQwenImage, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - model.model_type = model_type - model.tokenizer = tokenizer - model.t5_tokenizer = t5_tokenizer - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.text_conditioner = text_conditioner - model.vae = vae - model.transformer = transformer + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( + CosmosTransformer3DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + stream_from_disk=stream_from_disk, + ) def load( #TODO share code between models self, @@ -140,12 +114,14 @@ def load( #TODO share code between models model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] try: self.__load_internal( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + stream_from_disk, ) return except Exception: @@ -154,6 +130,7 @@ def load( #TODO share code between models try: self.__load_diffusers( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + stream_from_disk, ) return except Exception: diff --git a/modules/modelLoader/BaseModelLoader.py b/modules/modelLoader/BaseModelLoader.py index 4a560c2f1..44d908e22 100644 --- a/modules/modelLoader/BaseModelLoader.py +++ b/modules/modelLoader/BaseModelLoader.py @@ -49,5 +49,7 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, + cache_in_ram: dict[str, bool] | None = None, ) -> BaseModel | None: pass diff --git a/modules/modelLoader/ErnieModelLoader.py b/modules/modelLoader/ErnieModelLoader.py index af268c25c..0c9e3b85b 100644 --- a/modules/modelLoader/ErnieModelLoader.py +++ b/modules/modelLoader/ErnieModelLoader.py @@ -11,13 +11,10 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKLFlux2, ErnieImageTransformer2DModel, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, ) from transformers import AutoTokenizer, Mistral3Model @@ -37,11 +34,12 @@ def __load_internal( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, - quantization, + quantization, stream_from_disk, ) else: raise Exception("not an internal model") @@ -55,68 +53,44 @@ def __load_diffusers( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): - if transformer_model_name: - transformer = ErnieImageTransformer2DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - torch_dtype=torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - ErnieImageTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( + ErnieImageTransformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + stream_from_disk=stream_from_disk, + ) - tokenizer = AutoTokenizer.from_pretrained( + model.tokenizer = AutoTokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) - text_encoder = self._load_transformers_sub_module( + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( Mistral3Model, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKLFlux2, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLFlux2, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) - - model.model_type = model_type - model.tokenizer = tokenizer - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.vae = vae - model.transformer = transformer + model.vae = self._load_vae( + AutoencoderKLFlux2, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) def __load_safetensors( self, @@ -140,13 +114,14 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] try: self.__load_internal( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, - model_names.vae_model, quantization, + model_names.vae_model, quantization, stream_from_disk, ) return except Exception: @@ -155,7 +130,7 @@ def load( try: self.__load_diffusers( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, - model_names.vae_model, quantization, + model_names.vae_model, quantization, stream_from_disk, ) return except Exception: diff --git a/modules/modelLoader/Flux2ModelLoader.py b/modules/modelLoader/Flux2ModelLoader.py index 33f3fe518..0f27c3cbc 100644 --- a/modules/modelLoader/Flux2ModelLoader.py +++ b/modules/modelLoader/Flux2ModelLoader.py @@ -11,13 +11,10 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKLFlux2, FlowMatchEulerDiscreteScheduler, Flux2Transformer2DModel, - GGUFQuantizationConfig, ) from transformers import ( Mistral3ForConditionalGeneration, @@ -42,10 +39,12 @@ def __load_internal( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( - model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, quantization, + model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, + quantization, stream_from_disk, ) else: raise Exception("not an internal model") @@ -59,82 +58,52 @@ def __load_diffusers( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): - if transformer_model_name: - transformer = Flux2Transformer2DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - #avoid loading the transformer in float32: - torch_dtype=torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - Flux2Transformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( + Flux2Transformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + stream_from_disk=stream_from_disk, + ) - if transformer.config.num_attention_heads == 48: #Flux2.Dev - tokenizer = PixtralProcessor.from_pretrained( + if model.transformer.config.num_attention_heads == 48: #Flux2.Dev + model.tokenizer = PixtralProcessor.from_pretrained( base_model_name, subfolder="tokenizer", ).tokenizer - - text_encoder = self._load_transformers_sub_module( - Mistral3ForConditionalGeneration, - weight_dtypes.text_encoder, - weight_dtypes.fallback_train_dtype, - base_model_name, - "text_encoder", - ) + text_encoder_class = Mistral3ForConditionalGeneration else: #Flux2.Klein - tokenizer = Qwen2Tokenizer.from_pretrained( + model.tokenizer = Qwen2Tokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) - text_encoder = self._load_transformers_sub_module( - Qwen3ForCausalLM, - weight_dtypes.text_encoder, - weight_dtypes.fallback_train_dtype, - base_model_name, - "text_encoder", - ) + text_encoder_class = Qwen3ForCausalLM - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( + text_encoder_class, + weight_dtypes.text_encoder, + weight_dtypes.fallback_train_dtype, base_model_name, - subfolder="scheduler", + "text_encoder", + stream_from_disk=stream_from_disk, ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKLFlux2, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLFlux2, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + base_model_name, + subfolder="scheduler", + ) - model.model_type = model_type - model.tokenizer = tokenizer - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.vae = vae - model.transformer = transformer + model.vae = self._load_vae( + AutoencoderKLFlux2, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) def __load_safetensors( self, @@ -156,12 +125,14 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] try: self.__load_internal( - model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, + model_names.vae_model, quantization, stream_from_disk, ) return except Exception: @@ -169,7 +140,8 @@ def load( try: self.__load_diffusers( - model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, + model_names.vae_model, quantization, stream_from_disk, ) return except Exception: diff --git a/modules/modelLoader/GenericEmbeddingModelLoader.py b/modules/modelLoader/GenericEmbeddingModelLoader.py index 019502fb4..7106cd1f0 100644 --- a/modules/modelLoader/GenericEmbeddingModelLoader.py +++ b/modules/modelLoader/GenericEmbeddingModelLoader.py @@ -36,16 +36,26 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, + cache_in_ram: dict[str, bool] | None = None, ) -> model_class | None: base_model_loader = model_loader_class() embedding_loader = embedding_loader_class() model = model_class(model_type=model_type) + cache_in_ram = cache_in_ram or {} + if not stream_from_disk and any(not cache_in_ram.get(part, False) for part in model_type.model_parts()): + print("Warning: 'stream from disk' is off, so every component stays fully in RAM; disabling " + "'cache in ram' cannot free it without streaming.") + model.cache_in_ram = { + part: not stream_from_disk or cache_in_ram.get(part, False) + for part in model_type.model_parts() + } self._load_internal_data(model, model_names.embedding.model_name) model.model_spec = self._load_default_model_spec(model_type) if model_names.base_model is not None: - base_model_loader.load(model, model_type, model_names, weight_dtypes, quantization) + base_model_loader.load(model, model_type, model_names, weight_dtypes, quantization, stream_from_disk) embedding_loader.load(model, model_names.embedding.model_name, model_names) return model diff --git a/modules/modelLoader/GenericFineTuneModelLoader.py b/modules/modelLoader/GenericFineTuneModelLoader.py index 09915388f..410c3efcb 100644 --- a/modules/modelLoader/GenericFineTuneModelLoader.py +++ b/modules/modelLoader/GenericFineTuneModelLoader.py @@ -40,17 +40,27 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, + cache_in_ram: dict[str, bool] | None = None, ) -> model_class | None: base_model_loader = model_loader_class() if embedding_loader_class is not None: embedding_loader = embedding_loader_class() model = model_class(model_type=model_type) + cache_in_ram = cache_in_ram or {} + if not stream_from_disk and any(not cache_in_ram.get(part, False) for part in model_type.model_parts()): + print("Warning: 'stream from disk' is off, so every component stays fully in RAM; disabling " + "'cache in ram' cannot free it without streaming.") + model.cache_in_ram = { + part: not stream_from_disk or cache_in_ram.get(part, False) + for part in model_type.model_parts() + } self._load_internal_data(model, model_names.base_model) model.model_spec = self._load_default_model_spec(model_type) - base_model_loader.load(model, model_type, model_names, weight_dtypes, quantization) + base_model_loader.load(model, model_type, model_names, weight_dtypes, quantization, stream_from_disk) if embedding_loader_class is not None: embedding_loader.load(model, model_names.base_model, model_names) diff --git a/modules/modelLoader/GenericLoRAModelLoader.py b/modules/modelLoader/GenericLoRAModelLoader.py index d120eb008..4f005123b 100644 --- a/modules/modelLoader/GenericLoRAModelLoader.py +++ b/modules/modelLoader/GenericLoRAModelLoader.py @@ -37,6 +37,8 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, + cache_in_ram: dict[str, bool] | None = None, ) -> model_class | None: base_model_loader = model_loader_class() lora_model_loader = lora_loader_class() @@ -44,11 +46,19 @@ def load( embedding_loader = embedding_loader_class() model = model_class(model_type=model_type) + cache_in_ram = cache_in_ram or {} + if not stream_from_disk and any(not cache_in_ram.get(part, False) for part in model_type.model_parts()): + print("Warning: 'stream from disk' is off, so every component stays fully in RAM; disabling " + "'cache in ram' cannot free it without streaming.") + model.cache_in_ram = { + part: not stream_from_disk or cache_in_ram.get(part, False) + for part in model_type.model_parts() + } self._load_internal_data(model, model_names.lora) model.model_spec = self._load_default_model_spec(model_type) if model_names.base_model is not None: - base_model_loader.load(model, model_type, model_names, weight_dtypes, quantization) + base_model_loader.load(model, model_type, model_names, weight_dtypes, quantization, stream_from_disk) lora_model_loader.load(model, model_names) if embedding_loader_class is not None: embedding_loader.load(model, model_names.lora, model_names) diff --git a/modules/modelLoader/IdeogramModelLoader.py b/modules/modelLoader/IdeogramModelLoader.py index aef400f08..eca9d168b 100644 --- a/modules/modelLoader/IdeogramModelLoader.py +++ b/modules/modelLoader/IdeogramModelLoader.py @@ -34,9 +34,12 @@ def __load_internal( vae_model_name: str, include_unconditional_transformer: bool, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): - self.__load_diffusers(model, model_type, weight_dtypes, base_model_name, vae_model_name, include_unconditional_transformer, quantization) + self.__load_diffusers( + model, model_type, weight_dtypes, base_model_name, vae_model_name, include_unconditional_transformer, + quantization, stream_from_disk) else: raise Exception("not an internal model") @@ -49,71 +52,59 @@ def __load_diffusers( vae_model_name: str, include_unconditional_transformer: bool, quantization: QuantizationConfig, + stream_from_disk: bool, ): - transformer = self._load_diffusers_sub_module( + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( Ideogram4Transformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, + weight_dtypes, base_model_name, - "transformer", + "", quantization, + stream_from_disk=stream_from_disk, ) # the unconditional transformer is frozen and only used for the negative branch of the dual-network CFG at - # sampling, so it has its own weight dtype independent of the trainable transformer's. It is optional: if not - # loaded, only cfg_scale<=1 sampling is possible. + # sampling, so it has its own weight dtype independent of the trainable transformer's. It is optional. It uses + # _load_diffusers_sub_module directly (not _load_transformer) because of its own subfolder and dtype. if include_unconditional_transformer: - unconditional_transformer = self._load_diffusers_sub_module( - Ideogram4Transformer2DModel, - weight_dtypes.unconditional_transformer, - weight_dtypes.train_dtype, - base_model_name, - "unconditional_transformer", - quantization, - ) + model.unconditional_transformer, model.materialize_fn["unconditional_transformer"] = \ + self._load_diffusers_sub_module( + Ideogram4Transformer2DModel, + weight_dtypes.unconditional_transformer, + weight_dtypes.train_dtype, + base_model_name, + "unconditional_transformer", + quantization, + stream_from_disk=stream_from_disk, + ) else: - unconditional_transformer = None + model.unconditional_transformer = None - text_encoder = self._load_transformers_sub_module( + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( Qwen3VLModel, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) - tokenizer = AutoTokenizer.from_pretrained( + model.tokenizer = AutoTokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKLFlux2, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLFlux2, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) - - model.model_type = model_type - model.tokenizer = tokenizer - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.vae = vae - model.transformer = transformer - model.unconditional_transformer = unconditional_transformer + model.vae = self._load_vae( + AutoencoderKLFlux2, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) def __load_safetensors( self, @@ -134,13 +125,14 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] try: self.__load_internal( model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, - model_names.include_unconditional_transformer, quantization, + model_names.include_unconditional_transformer, quantization, stream_from_disk, ) return except Exception: @@ -149,7 +141,7 @@ def load( try: self.__load_diffusers( model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, - model_names.include_unconditional_transformer, quantization, + model_names.include_unconditional_transformer, quantization, stream_from_disk, ) return except Exception: diff --git a/modules/modelLoader/LTXModelLoader.py b/modules/modelLoader/LTXModelLoader.py new file mode 100644 index 000000000..3e33b6d79 --- /dev/null +++ b/modules/modelLoader/LTXModelLoader.py @@ -0,0 +1,284 @@ +import os +import traceback + +from modules.model.LTXModel import LTXModel +from modules.modelLoader.GenericFineTuneModelLoader import make_fine_tune_model_loader +from modules.modelLoader.GenericLoRAModelLoader import make_lora_model_loader +from modules.modelLoader.mixin.HFModelLoaderMixin import HFModelLoaderMixin +from modules.modelLoader.mixin.LoRALoaderMixin import LoRALoaderMixin +from modules.util.config.TrainConfig import QuantizationConfig +from modules.util.enum.ModelType import ModelType +from modules.util.ModelNames import ModelNames +from modules.util.ModelWeightDtypes import ModelWeightDtypes + +import torch + +from diffusers import ( + AutoencoderKLLTX2Audio, + AutoencoderKLLTX2Video, + FlowMatchEulerDiscreteScheduler, + GGUFQuantizationConfig, + LTX2VideoTransformer3DModel, +) +from diffusers.pipelines.ltx2.connectors import LTX2TextConnectors +from diffusers.pipelines.ltx2.vocoder import LTX2VocoderWithBWE +from transformers import ( + MODEL_FOR_IMAGE_TEXT_TO_TEXT_MAPPING, + AutoConfig, + AutoTokenizer, +) + +import huggingface_hub + + +class LTXModelLoader( + HFModelLoaderMixin, +): + def __init__(self): + super().__init__() + + def __load_internal( + self, + model: LTXModel, + model_type: ModelType, + weight_dtypes: ModelWeightDtypes, + base_model_name: str, + transformer_model_name: str, + low_noise_transformer_model_name: str, + include_low_noise_transformer: bool, + text_encoder_model_name: str, + vae_model_name: str, + quantization: QuantizationConfig, + stream_from_disk: bool, + ): + if os.path.isfile(os.path.join(base_model_name, "meta.json")): + self.__load_diffusers( + model, model_type, weight_dtypes, base_model_name, transformer_model_name, + low_noise_transformer_model_name, include_low_noise_transformer, text_encoder_model_name, vae_model_name, + quantization, stream_from_disk, + ) + else: + raise Exception("not an internal model") + + def __load_diffusers( + self, + model: LTXModel, + model_type: ModelType, + weight_dtypes: ModelWeightDtypes, + base_model_name: str, + transformer_model_name: str, + low_noise_transformer_model_name: str, + include_low_noise_transformer: bool, + text_encoder_model_name: str, + vae_model_name: str, + quantization: QuantizationConfig, + stream_from_disk: bool, + ): + # LTX 2.5 ships two DiTs in one repo: transformer/ is the distilled one, transformer_full/ the full/SFT + # model training wants. 2.3 has only transformer/ and it is already the full model, so preferring + # transformer_full/ picks the trainable DiT in both generations. + has_transformer_full = os.path.isdir(os.path.join(base_model_name, "transformer_full")) \ + if os.path.isdir(base_model_name) \ + else huggingface_hub.file_exists(base_model_name, "transformer_full/config.json") + transformer_subfolder = "transformer_full" if has_transformer_full else "transformer" + + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( + LTX2VideoTransformer3DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + stream_from_disk=stream_from_disk, + subfolder=transformer_subfolder, + ) + + # The distilled low-noise expert, from the base repo's transformer/ where the repo ships one (2.5), else + # from whatever the user named. Optional: with neither it stays None. _load_diffusers_sub_module rather + # than _load_transformer, which hardcodes weight_dtypes.transformer and would give this part the main + # transformer's dtype. + low_noise_transformer_source = low_noise_transformer_model_name or (base_model_name if has_transformer_full else None) + if include_low_noise_transformer and low_noise_transformer_source: + if os.path.isfile(low_noise_transformer_source): + # single-file checkpoints load whole into RAM: streaming reads a shard map keyed by the module + # tree, which only a diffusers folder provides. The config comes from the base repo's + # transformer/, which the distilled DiT matches field for field. + low_noise_transformer_dtype = weight_dtypes.low_noise_transformer.torch_dtype() + model.low_noise_transformer = LTX2VideoTransformer3DModel.from_single_file( + low_noise_transformer_source, + config=base_model_name, + subfolder="transformer", + # avoid loading the expert in float32: + torch_dtype=torch.bfloat16 if low_noise_transformer_dtype is None else low_noise_transformer_dtype, + quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) + if weight_dtypes.low_noise_transformer.is_gguf() else None, + ) + model.low_noise_transformer = self._convert_diffusers_sub_module_to_dtype( + model.low_noise_transformer, weight_dtypes.low_noise_transformer, weight_dtypes.train_dtype, quantization, + ) + else: + model.low_noise_transformer, model.materialize_fn["low_noise_transformer"] = self._load_diffusers_sub_module( + LTX2VideoTransformer3DModel, + weight_dtypes.low_noise_transformer, + weight_dtypes.train_dtype, + low_noise_transformer_source, + "transformer", + quantization, + stream_from_disk=stream_from_disk, + ) + + model.tokenizer = AutoTokenizer.from_pretrained( + base_model_name, + subfolder="tokenizer", + ) + # padding='max_length' needs one, and every LTX checkpoint's tokenizer config sets it + assert model.tokenizer.pad_token is not None + + if text_encoder_model_name: + # 2.3 bundles the encoder as float32 (48.7GB vs 24.4GB) for identical values, so naming a bf16 repo + # is worth it there; 2.5 already ships bf16. ComfyUI's stock google/gemma-3-12b-it differs from the + # bundled QAT weights by up to ~5% relative - a fidelity choice. A standalone repo has no subfolder. + text_encoder_base, text_encoder_subfolder = text_encoder_model_name, "" + else: + text_encoder_base, text_encoder_subfolder = base_model_name, "text_encoder" + + # LTX 2.3 ships a Gemma 3 encoder, LTX 2.5 a Gemma 4 one, and both load through the same ModelType, so + # the class comes from the checkpoint's config - resolved the way AutoModelForImageTextToText would. + # The auto class itself can't be used: the loader builds a meta skeleton from config_class, which only + # the concrete class has. + text_encoder_config = AutoConfig.from_pretrained( + text_encoder_base, + subfolder=text_encoder_subfolder, + ) + if type(text_encoder_config) not in MODEL_FOR_IMAGE_TEXT_TO_TEXT_MAPPING: + raise NotImplementedError( + f"unsupported LTX text encoder model_type '{text_encoder_config.model_type}'") + + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( + MODEL_FOR_IMAGE_TEXT_TO_TEXT_MAPPING[type(text_encoder_config)], + weight_dtypes.text_encoder, + weight_dtypes.fallback_train_dtype, + text_encoder_base, + text_encoder_subfolder, + stream_from_disk=stream_from_disk, + ) + + model.connectors, model.materialize_fn["connectors"] = self._load_diffusers_sub_module( + LTX2TextConnectors, + weight_dtypes.connectors, + weight_dtypes.fallback_train_dtype, + base_model_name, + "connectors", + stream_from_disk=stream_from_disk, + ) + + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + base_model_name, + subfolder="scheduler", + ) + + if vae_model_name: + model.vae = self._load_diffusers_sub_module( + AutoencoderKLLTX2Video, + weight_dtypes.vae, + weight_dtypes.train_dtype, + vae_model_name, + ) + else: + model.vae = self._load_diffusers_sub_module( + AutoencoderKLLTX2Video, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + "vae", + ) + + # the audio branch is frozen and never trained, but LTX2Pipeline requires both, so they are loaded + # to keep a saved diffusers repo complete. They piggyback on the vae's dtype config. + model.audio_vae = self._load_diffusers_sub_module( + AutoencoderKLLTX2Audio, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + "audio_vae", + ) + + model.vocoder = self._load_diffusers_sub_module( + LTX2VocoderWithBWE, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + "vocoder", + ) + + model.model_type = model_type + + def load( + self, + model: LTXModel, + model_type: ModelType, + model_names: ModelNames, + weight_dtypes: ModelWeightDtypes, + quantization: QuantizationConfig, + stream_from_disk: bool = False, + ): + stacktraces = [] + + try: + self.__load_internal( + model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, + model_names.low_noise_transformer_model, model_names.include_low_noise_transformer, + model_names.text_encoder_model, model_names.vae_model, quantization, stream_from_disk, + ) + return + except Exception: + stacktraces.append(traceback.format_exc()) + + try: + self.__load_diffusers( + model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, + model_names.low_noise_transformer_model, model_names.include_low_noise_transformer, + model_names.text_encoder_model, model_names.vae_model, quantization, stream_from_disk, + ) + return + except Exception: + stacktraces.append(traceback.format_exc()) + + for stacktrace in stacktraces: + print(stacktrace) + raise Exception("could not load model: " + model_names.base_model) + + +class LTXLoRALoader( + LoRALoaderMixin +): + def __init__(self): + super().__init__() + + + def load( + self, + model: LTXModel, + model_names: ModelNames, + ): + return self._load(model, model_names) + + +LTXLoRAModelLoader = make_lora_model_loader( + model_spec_map={ + ModelType.LTX_2: "resources/sd_model_spec/ltx_2-lora.json", + }, + model_class=LTXModel, + model_loader_class=LTXModelLoader, + lora_loader_class=LTXLoRALoader, + embedding_loader_class=None, +) + +LTXFineTuneModelLoader = make_fine_tune_model_loader( + model_spec_map={ + ModelType.LTX_2: "resources/sd_model_spec/ltx_2.json", + }, + model_class=LTXModel, + model_loader_class=LTXModelLoader, + embedding_loader_class=None, +) diff --git a/modules/modelLoader/ZImageModelLoader.py b/modules/modelLoader/ZImageModelLoader.py index 308232823..63df43787 100644 --- a/modules/modelLoader/ZImageModelLoader.py +++ b/modules/modelLoader/ZImageModelLoader.py @@ -11,12 +11,9 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKL, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, ZImageTransformer2DModel, ) from transformers import ( @@ -40,10 +37,12 @@ def __load_internal( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, quantization, + stream_from_disk, ) else: raise Exception("not an internal model") @@ -57,67 +56,43 @@ def __load_diffusers( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): - tokenizer = Qwen2Tokenizer.from_pretrained( + model.tokenizer = Qwen2Tokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( Qwen3ForCausalLM, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) - - if transformer_model_name: - transformer = ZImageTransformer2DModel.from_single_file( - transformer_model_name, - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - ZImageTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + model.vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - model.model_type = model_type - model.tokenizer = tokenizer - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.vae = vae - model.transformer = transformer + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( + ZImageTransformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + stream_from_disk=stream_from_disk, + ) def __load_safetensors( self, @@ -139,12 +114,14 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] try: self.__load_internal( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + stream_from_disk, ) return except Exception: @@ -153,6 +130,7 @@ def load( try: self.__load_diffusers( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + stream_from_disk, ) return except Exception: diff --git a/modules/modelLoader/chroma/ChromaModelLoader.py b/modules/modelLoader/chroma/ChromaModelLoader.py index 7dcbef794..59994818d 100644 --- a/modules/modelLoader/chroma/ChromaModelLoader.py +++ b/modules/modelLoader/chroma/ChromaModelLoader.py @@ -9,13 +9,10 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKL, ChromaTransformer2DModel, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, ) from transformers import T5EncoderModel, T5Tokenizer @@ -35,10 +32,12 @@ def __load_internal( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, quantization, + stream_from_disk, ) else: raise Exception("not an internal model") @@ -52,68 +51,44 @@ def __load_diffusers( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): - tokenizer = T5Tokenizer.from_pretrained( + model.tokenizer = T5Tokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) + model.orig_tokenizer = copy.deepcopy(model.tokenizer) - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( T5EncoderModel, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) - - if transformer_model_name: - transformer = ChromaTransformer2DModel.from_single_file( - transformer_model_name, - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - ChromaTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + model.vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - model.model_type = model_type - model.tokenizer = tokenizer - model.orig_tokenizer = copy.deepcopy(tokenizer) - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.vae = vae - model.transformer = transformer + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( + ChromaTransformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + stream_from_disk=stream_from_disk, + ) def __load_safetensors( self, @@ -135,12 +110,14 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] try: self.__load_internal( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + stream_from_disk, ) return except Exception: @@ -149,6 +126,7 @@ def load( try: self.__load_diffusers( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + stream_from_disk, ) return except Exception: diff --git a/modules/modelLoader/flux/FluxModelLoader.py b/modules/modelLoader/flux/FluxModelLoader.py index d4f21ea2a..0a2cf9546 100644 --- a/modules/modelLoader/flux/FluxModelLoader.py +++ b/modules/modelLoader/flux/FluxModelLoader.py @@ -16,7 +16,6 @@ FlowMatchEulerDiscreteScheduler, FluxPipeline, FluxTransformer2DModel, - GGUFQuantizationConfig, ) from transformers import CLIPTextModel, CLIPTokenizer, T5EncoderModel, T5Tokenizer @@ -38,11 +37,12 @@ def __load_internal( include_text_encoder_1: bool, include_text_encoder_2: bool, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, - include_text_encoder_1, include_text_encoder_2, quantization, + include_text_encoder_1, include_text_encoder_2, quantization, stream_from_disk, ) else: raise Exception("not an internal model") @@ -58,96 +58,71 @@ def __load_diffusers( include_text_encoder_1: bool, include_text_encoder_2: bool, quantization: QuantizationConfig, + stream_from_disk: bool, ): if include_text_encoder_1: - tokenizer_1 = CLIPTokenizer.from_pretrained( + model.tokenizer_1 = CLIPTokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) else: - tokenizer_1 = None + model.tokenizer_1 = None + model.orig_tokenizer_1 = copy.deepcopy(model.tokenizer_1) if include_text_encoder_2: - tokenizer_2 = T5Tokenizer.from_pretrained( + model.tokenizer_2 = T5Tokenizer.from_pretrained( base_model_name, subfolder="tokenizer_2", ) else: - tokenizer_2 = None + model.tokenizer_2 = None + model.orig_tokenizer_2 = copy.deepcopy(model.tokenizer_2) - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) if include_text_encoder_1: - text_encoder_1 = self._load_transformers_sub_module( + model.text_encoder_1, model.materialize_fn["text_encoder_1"] = self._load_text_encoder( CLIPTextModel, weight_dtypes.text_encoder, weight_dtypes.train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) else: - text_encoder_1 = None + model.text_encoder_1 = None if include_text_encoder_2: - text_encoder_2 = self._load_transformers_sub_module( + model.text_encoder_2, model.materialize_fn["text_encoder_2"] = self._load_text_encoder( T5EncoderModel, weight_dtypes.text_encoder_2, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder_2", + stream_from_disk=stream_from_disk, ) else: - text_encoder_2 = None - - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + model.text_encoder_2 = None - if transformer_model_name: - transformer = FluxTransformer2DModel.from_single_file( - transformer_model_name, - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - FluxTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + model.vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - model.model_type = model_type - model.tokenizer_1 = tokenizer_1 - model.orig_tokenizer_1 = copy.deepcopy(tokenizer_1) - model.tokenizer_2 = tokenizer_2 - model.orig_tokenizer_2 = copy.deepcopy(tokenizer_2) - model.noise_scheduler = noise_scheduler - model.text_encoder_1 = text_encoder_1 - model.text_encoder_2 = text_encoder_2 - model.vae = vae - model.transformer = transformer + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( + FluxTransformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + stream_from_disk=stream_from_disk, + ) def __load_safetensors( self, @@ -233,13 +208,14 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] try: self.__load_internal( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, - model_names.include_text_encoder, model_names.include_text_encoder_2, quantization, + model_names.include_text_encoder, model_names.include_text_encoder_2, quantization, stream_from_disk, ) return except Exception: @@ -248,7 +224,7 @@ def load( try: self.__load_diffusers( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, - model_names.include_text_encoder, model_names.include_text_encoder_2, quantization, + model_names.include_text_encoder, model_names.include_text_encoder_2, quantization, stream_from_disk, ) return except Exception: diff --git a/modules/modelLoader/hiDream/HiDreamModelLoader.py b/modules/modelLoader/hiDream/HiDreamModelLoader.py index b3e20f23c..b5ac66425 100644 --- a/modules/modelLoader/hiDream/HiDreamModelLoader.py +++ b/modules/modelLoader/hiDream/HiDreamModelLoader.py @@ -44,11 +44,13 @@ def __load_internal( include_text_encoder_3: bool, include_text_encoder_4: bool, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( model, model_type, weight_dtypes, base_model_name, text_encoder_4_model_name, vae_model_name, - include_text_encoder_1, include_text_encoder_2, include_text_encoder_3, include_text_encoder_4, quantization, + include_text_encoder_1, include_text_encoder_2, include_text_encoder_3, include_text_encoder_4, + quantization, stream_from_disk, ) else: raise Exception("not an internal model") @@ -66,122 +68,109 @@ def __load_diffusers( include_text_encoder_3: bool, include_text_encoder_4: bool, quantization: QuantizationConfig, + stream_from_disk: bool, ): - tokenizer_1 = CLIPTokenizer.from_pretrained( + model.tokenizer_1 = CLIPTokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) if include_text_encoder_1 else None - tokenizer_2 = CLIPTokenizer.from_pretrained( + model.tokenizer_2 = CLIPTokenizer.from_pretrained( base_model_name, subfolder="tokenizer_2", ) if include_text_encoder_2 else None - tokenizer_3 = T5Tokenizer.from_pretrained( + model.tokenizer_3 = T5Tokenizer.from_pretrained( base_model_name, subfolder="tokenizer_3", ) if include_text_encoder_3 else None - tokenizer_4 = LlamaTokenizerFast.from_pretrained( + model.tokenizer_4 = LlamaTokenizerFast.from_pretrained( text_encoder_4_model_name, ) if include_text_encoder_4 else None - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) if include_text_encoder_1: - text_encoder_1 = self._load_transformers_sub_module( + model.text_encoder_1, model.materialize_fn["text_encoder_1"] = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder, weight_dtypes.train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) else: - text_encoder_1 = None + model.text_encoder_1 = None if include_text_encoder_2: - text_encoder_2 = self._load_transformers_sub_module( + model.text_encoder_2, model.materialize_fn["text_encoder_2"] = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder_2, weight_dtypes.train_dtype, base_model_name, "text_encoder_2", + stream_from_disk=stream_from_disk, ) else: - text_encoder_2 = None + model.text_encoder_2 = None if include_text_encoder_3: - text_encoder_3 = self._load_transformers_sub_module( + model.text_encoder_3, model.materialize_fn["text_encoder_3"] = self._load_text_encoder( T5EncoderModel, weight_dtypes.text_encoder_3, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder_3", + stream_from_disk=stream_from_disk, ) else: - text_encoder_3 = None + model.text_encoder_3 = None if include_text_encoder_4: if text_encoder_4_model_name: - text_encoder_4 = self._load_transformers_sub_module( + # override repo holds text_encoder_4 at its root, not in a base-model subfolder, so it bypasses + # _load_text_encoder (which always loads from a base-repo subfolder) and loads directly. + model.text_encoder_4, model.materialize_fn["text_encoder_4"] = self._load_transformers_sub_module( LlamaForCausalLM, weight_dtypes.text_encoder_4, weight_dtypes.train_dtype, text_encoder_4_model_name, + stream_from_disk=stream_from_disk, ) else: - text_encoder_4 = self._load_transformers_sub_module( + model.text_encoder_4, model.materialize_fn["text_encoder_4"] = self._load_text_encoder( LlamaForCausalLM, weight_dtypes.text_encoder_4, weight_dtypes.train_dtype, base_model_name, "text_encoder_4", + stream_from_disk=stream_from_disk, ) else: - text_encoder_4 = None + model.text_encoder_4 = None - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + model.vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - transformer = self._load_diffusers_sub_module( + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( HiDreamImageTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, + weight_dtypes, base_model_name, - "transformer", + "", quantization, + stream_from_disk=stream_from_disk, ) - model.model_type = model_type - model.tokenizer_1 = tokenizer_1 - model.tokenizer_2 = tokenizer_2 - model.tokenizer_3 = tokenizer_3 - model.tokenizer_4 = tokenizer_4 - model.noise_scheduler = noise_scheduler - model.text_encoder_1 = text_encoder_1 - model.text_encoder_2 = text_encoder_2 - model.text_encoder_3 = text_encoder_3 - model.text_encoder_4 = text_encoder_4 - model.vae = vae - model.transformer = transformer - def __load_safetensors( self, model: HiDreamModel, @@ -268,6 +257,7 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] @@ -277,6 +267,7 @@ def load( model_names.text_encoder_4, model_names.vae_model, model_names.include_text_encoder, model_names.include_text_encoder_2, model_names.include_text_encoder_3, model_names.include_text_encoder_4, quantization, + stream_from_disk, ) self.__after_load(model) return @@ -289,12 +280,18 @@ def load( model_names.text_encoder_4, model_names.vae_model, model_names.include_text_encoder, model_names.include_text_encoder_2, model_names.include_text_encoder_3, model_names.include_text_encoder_4, quantization, + stream_from_disk, ) self.__after_load(model) return except Exception: stacktraces.append(traceback.format_exc()) + if stream_from_disk: + # the single-file loader below builds a full pipeline via from_single_file, which can't stream; fall + # back to loading it fully into RAM. + print(f"Warning: 'stream from disk' is not supported for single-file {model_type}; loading fully into RAM.") + try: self.__load_safetensors( model, model_type, weight_dtypes, model_names.base_model, diff --git a/modules/modelLoader/hunyuanVideo/HunyuanVideoModelLoader.py b/modules/modelLoader/hunyuanVideo/HunyuanVideoModelLoader.py index 85c91699b..f8a219d67 100644 --- a/modules/modelLoader/hunyuanVideo/HunyuanVideoModelLoader.py +++ b/modules/modelLoader/hunyuanVideo/HunyuanVideoModelLoader.py @@ -38,11 +38,12 @@ def __load_internal( include_text_encoder_1: bool, include_text_encoder_2: bool, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, - include_text_encoder_1, include_text_encoder_2, quantization, + include_text_encoder_1, include_text_encoder_2, quantization, stream_from_disk, ) else: raise Exception("not an internal model") @@ -58,96 +59,70 @@ def __load_diffusers( include_text_encoder_1: bool, include_text_encoder_2: bool, quantization: QuantizationConfig, + stream_from_disk: bool, ): if include_text_encoder_1: - tokenizer_1 = LlamaTokenizerFast.from_pretrained( + model.tokenizer_1 = LlamaTokenizerFast.from_pretrained( base_model_name, subfolder="tokenizer", ) else: - tokenizer_1 = None + model.tokenizer_1 = None if include_text_encoder_2: - tokenizer_2 = CLIPTokenizer.from_pretrained( + model.tokenizer_2 = CLIPTokenizer.from_pretrained( base_model_name, subfolder="tokenizer_2", ) else: - tokenizer_2 = None + model.tokenizer_2 = None - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) if include_text_encoder_1: - text_encoder_1 = self._load_transformers_sub_module( + model.text_encoder_1, model.materialize_fn["text_encoder_1"] = self._load_text_encoder( LlamaModel, weight_dtypes.text_encoder, weight_dtypes.train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) else: - text_encoder_1 = None + model.text_encoder_1 = None if include_text_encoder_2: - text_encoder_2 = self._load_transformers_sub_module( + model.text_encoder_2, model.materialize_fn["text_encoder_2"] = self._load_text_encoder( CLIPTextModel, weight_dtypes.text_encoder_2, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder_2", + stream_from_disk=stream_from_disk, ) else: - text_encoder_2 = None + model.text_encoder_2 = None - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKLHunyuanVideo, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLHunyuanVideo, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) - - if transformer_model_name: - transformer = HunyuanVideoTransformer3DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization - ) - else: - transformer = self._load_diffusers_sub_module( - HunyuanVideoTransformer3DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + model.vae = self._load_vae( + AutoencoderKLHunyuanVideo, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - model.model_type = model_type - model.tokenizer_1 = tokenizer_1 - model.tokenizer_2 = tokenizer_2 - model.noise_scheduler = noise_scheduler - model.text_encoder_1 = text_encoder_1 - model.text_encoder_2 = text_encoder_2 - model.vae = vae - model.transformer = transformer + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( + HunyuanVideoTransformer3DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + stream_from_disk=stream_from_disk, + ) def __load_safetensors( self, @@ -233,13 +208,14 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] try: self.__load_internal( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, - model_names.include_text_encoder, model_names.include_text_encoder_2, quantization, + model_names.include_text_encoder, model_names.include_text_encoder_2, quantization, stream_from_disk, ) self.__after_load(model) return @@ -249,7 +225,7 @@ def load( try: self.__load_diffusers( model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, - model_names.include_text_encoder, model_names.include_text_encoder_2, quantization, + model_names.include_text_encoder, model_names.include_text_encoder_2, quantization, stream_from_disk, ) self.__after_load(model) return diff --git a/modules/modelLoader/krea2/Krea2ModelLoader.py b/modules/modelLoader/krea2/Krea2ModelLoader.py index c3987e97c..b4789e5d6 100644 --- a/modules/modelLoader/krea2/Krea2ModelLoader.py +++ b/modules/modelLoader/krea2/Krea2ModelLoader.py @@ -8,12 +8,9 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKLQwenImage, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, Krea2Transformer2DModel, ) from transformers import Qwen2Tokenizer, Qwen3VLModel @@ -34,10 +31,12 @@ def __load_internal( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( - model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, quantization, + model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, + quantization, stream_from_disk, ) else: raise Exception("not an internal model") @@ -51,69 +50,44 @@ def __load_diffusers( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): - tokenizer = Qwen2Tokenizer.from_pretrained( + model.tokenizer = Qwen2Tokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( Qwen3VLModel, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) - if vae_model_name: #TODO simplify - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) - - if transformer_model_name: - transformer = Krea2Transformer2DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - Krea2Transformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + model.vae = self._load_vae( + AutoencoderKLQwenImage, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - model.model_type = model_type - model.tokenizer = tokenizer - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.vae = vae - model.transformer = transformer + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( + Krea2Transformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + stream_from_disk=stream_from_disk, + ) def __load_safetensors( self, @@ -134,12 +108,14 @@ def load( #TODO share code between models model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] try: self.__load_internal( - model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, + model_names.vae_model, quantization, stream_from_disk, ) return except Exception: @@ -147,7 +123,8 @@ def load( #TODO share code between models try: self.__load_diffusers( - model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, + model_names.vae_model, quantization, stream_from_disk, ) return except Exception: diff --git a/modules/modelLoader/mixin/HFModelLoaderMixin.py b/modules/modelLoader/mixin/HFModelLoaderMixin.py index f2f196257..908adc81b 100644 --- a/modules/modelLoader/mixin/HFModelLoaderMixin.py +++ b/modules/modelLoader/mixin/HFModelLoaderMixin.py @@ -1,36 +1,282 @@ +import contextlib import json import logging import os +import queue +import threading from abc import ABCMeta from itertools import repeat +from modules.module.quantized.mixin.QuantizedModuleMixin import QuantizedModuleMixin from modules.util.config.TrainConfig import QuantizationConfig from modules.util.enum.DataType import DataType +from modules.util.ModelWeightDtypes import ModelWeightDtypes from modules.util.quantization_util import ( + is_quantized_module, is_quantized_parameter, replace_linear_with_quantized_layers, ) +from modules.util.torch_util import mem_pool_context import torch from torch import nn +from diffusers import GGUFQuantizationConfig from transformers.conversion_mapping import get_checkpoint_conversion_mapping from transformers.core_model_loading import rename_source_key import accelerate import huggingface_hub +from accelerate.utils import set_module_tensor_to_device from huggingface_hub.utils import EntryNotFoundError +from safetensors import safe_open from safetensors.torch import load_file +from tqdm import tqdm # huggingface_hub 1.16+ uses httpx, which logs every HTTP request/response at INFO level. logging.getLogger("httpx").setLevel(logging.WARNING) +# reader threads striping the checkpoint into host RAM while the main thread does H2D + inline quant +STREAM_READER_THREADS = 4 + + +def __stream_reader( + tid: int, + nthreads: int, + work: list[tuple], + key_to_file: dict[str, str], + source_key_map: dict[str, str] | None, + out_queue: queue.Queue, + done, + stop: threading.Event, +): + # prefetch reader thread: reads a stripe of the work list into host RAM and feeds the bounded queue. Each thread + # owns its safe_open handles (a handle is not safe for concurrent get_tensor). stop lets the main thread break the + # stripe early on abort/OOM, so no reader is left executing inside safetensors when the stream unwinds. + thread_handles: dict[str, object] = {} + try: + for i in range(tid, len(work), nthreads): + if stop.is_set(): + break + item = work[i] + path = key_to_file[item[0]] + handle = thread_handles.get(path) + if handle is None: + handle = thread_handles[path] = safe_open(path, framework="pt", device="cpu") + # cache key is the renamed module-layout key; the file stores the original, so read by the original + # (identity when no rename map was built). + read_key = source_key_map.get(item[0], item[0]) if source_key_map else item[0] + # get_tensor returns a lazy mmap view; .clone() forces the read off disk into host RAM + out_queue.put((item, handle.get_tensor(read_key).clone())) + except Exception as e: + out_queue.put(e) + finally: + out_queue.put(done) + + +def _drop_page_cache(paths: set[str]): + # Releases the shards' page cache so it cannot grow large enough to push the host-side offload buffers into swap. + # A later re-read of a shard costs far less than that. Linux only; elsewhere the cache is left alone. + if not hasattr(os, "posix_fadvise"): + return + for path in paths: + with contextlib.suppress(OSError): + fd = os.open(path, os.O_RDONLY) + try: + os.posix_fadvise(fd, 0, 0, os.POSIX_FADV_DONTNEED) + finally: + os.close(fd) + + +def _intended_float_dtype( + module: nn.Module, + module_name: str, + tensor_name: str, + dtype: DataType, + train_dtype: DataType, + keep_in_fp32_modules: list[str], +) -> torch.dtype | None: + # target dtype for a streamed float tensor, or None to leave it unchanged. A param the quantizer will pack keeps + # its dtype (the quantizer converts it); keep-in-fp32 modules and a quantized component's leftover params go to + # train_dtype; everything else to the weight dtype. + if is_quantized_parameter(module, tensor_name): + return None + if dtype.is_quantized() or module_name in keep_in_fp32_modules: + # a caller without a train_dtype yet (budget sizing) gets None -> the budget over-estimates these from the + # fp32 skeleton; the stream-time caller always passes a real train_dtype. + return train_dtype.torch_dtype() if train_dtype is not None else None + return dtype.torch_dtype() + + +def _stamp_skeleton_float_dtypes( + sub_module: nn.Module, + dtype: DataType, + train_dtype: DataType, + keep_in_fp32_modules: list[str], +): + # stamp each meta-skeleton float param with the dtype the stream will give it, so the offload VRAM budget (which + # sizes the still-meta skeleton) measures the real post-load footprint, not init_empty_weights' fp32 default. Free: + # a meta tensor holds no data, so .to() only rewrites its declared dtype. Uses the same _intended_float_dtype helper + # as the stream-time cast so the two agree; quantized weights are left alone (sized via predict_offload_bytes). + # Buffers are not stamped: they never enter the offload budget. + for name, module in sub_module.named_modules(): + module_name = name.split(".")[-1] + for tensor_name, param in module.named_parameters(recurse=False): + if not torch.is_floating_point(param): + continue + target = _intended_float_dtype(module, module_name, tensor_name, dtype, train_dtype, keep_in_fp32_modules) + if target is not None and param.dtype != target: + param.data = param.data.to(dtype=target) + + +def stream_module_from_checkpoint( + module: nn.Module, + device: torch.device, + key_to_file: dict[str, str], + dtype: DataType, + train_dtype: DataType, + keep_in_fp32_modules: list[str], + tied_weights_keys: dict[str, str] | None, + quantize: bool, + key_prefix: str = "", + source_key_map: dict[str, str] | None = None, + part_name: str | None = None, + dest_pool=None, +): + # Fill a meta skeleton by streaming its checkpoint weights one tensor at a time, so the full checkpoint never lands + # in RAM. key_prefix scopes the lookup to one sub-module; keys stay checkpoint-absolute. dest_pool routes + # non-quantized weights straight into a MemPool (quantized modules pack in the default pool). + def dest_pool_for(sub_module): + return dest_pool if (dest_pool is not None and not is_quantized_module(sub_module)) else None + + # flat work list of every checkpoint-backed skeleton tensor, so the reader threads below can drive the reads. + work = [] # (key, sub_module, tensor_name, is_buffer, module_name) + for name, sub_module in module.named_modules(): + module_name = name.split(".")[-1] + # gradient checkpointing in compile mode wraps each block in a CheckpointLayer, inserting a "checkpoint" + # level into the live path; the checkpoint keys have none, so strip it before lookup (as LoRAModule does). + lookup_name = ".".join(p for p in name.split(".") if p != "checkpoint") + for tensor_name, param in list(sub_module.named_parameters(recurse=False)): + key = ".".join(p for p in (key_prefix, lookup_name, tensor_name) if p) + if key in key_to_file and param.is_meta: + work.append((key, sub_module, tensor_name, False, module_name)) + for tensor_name, _buffer in list(sub_module.named_buffers(recurse=False)): + # non-persistent buffers (rotary inv_freq etc.) are config-derived, not stored in the checkpoint + if tensor_name in sub_module._non_persistent_buffers_set: + continue + # no is_meta guard (unlike params): init_empty_weights materializes persistent buffers as REAL init values, + # so is_meta can't mean "not yet filled" -- always stream, else the init value survives (mis-normalizing the VAE). + key = ".".join(p for p in (key_prefix, lookup_name, tensor_name) if p) + if key in key_to_file: + work.append((key, sub_module, tensor_name, True, module_name)) + + # place() lands one tensor: cast floats to their intended dtype (quantizer-packed params keep theirs), move to the + # compute device, quantize inline once a layer's weight arrives so VRAM never holds the whole unquantized module. + # bar: one tick per streamed tensor; only a whole-module stream (part_name set) shows it, per-layer conductor calls stay silent. + bar = tqdm(total=len(work), unit="tensor", desc=f"streaming {part_name}", leave=False, smoothing=0.05) \ + if part_name is not None else None + + def quantize_if_ready(sub_module): + # quantize a module whose weight has landed (no longer meta): quantize() self-guards against a second call, so + # firing it the moment the weight arrives (rather than in the batch pass quantize_layers() does) is always safe. + if isinstance(sub_module, QuantizedModuleMixin) and not sub_module.weight.is_meta: + sub_module.compute_dtype = train_dtype.torch_dtype() + sub_module.quantize(device=device) + + def place(item, value): + _key, sub_module, tensor_name, is_buffer, module_name = item + # tensors that will be quantized stay at their original dtype (the quantizer converts them); everything else is + # cast to its intended dtype here. + if torch.is_floating_point(value): + target = _intended_float_dtype(sub_module, module_name, tensor_name, dtype, train_dtype, keep_in_fp32_modules) + if target is not None: + value = value.to(dtype=target) + with mem_pool_context(dest_pool_for(sub_module)): + set_module_tensor_to_device(sub_module, tensor_name, device, value=value, dtype=value.dtype) + # quantize outside the pool context so a quantized module's dequant scratch stays in the default pool + if quantize: + quantize_if_ready(sub_module) + if bar is not None: + bar.update(1) + + # reader threads stripe the work list into host RAM and feed a bounded queue; the main thread drains it and does + # H2D + inline quantize on the default stream. Both the parallel reads and overlapping them with the GPU work are + # wins. Each reader clones the tensor off its mmap and drops its safetensors handles when it exits (right after its + # stripe), so the file mmaps are released early rather than pinned until first use -- keeps page-cache pressure + # down. Tensors may land out of order -- place() addresses each by name and inline quant is order-free. + nthreads = STREAM_READER_THREADS + out_queue: queue.Queue = queue.Queue(maxsize=2 * nthreads) + done = object() + stop = threading.Event() + + threads = [ + threading.Thread( + target=__stream_reader, + args=(tid, nthreads, work, key_to_file, source_key_map, out_queue, done, stop), + name=f"stream-reader-{tid}", daemon=True, + ) + for tid in range(nthreads) + ] + for t in threads: + t.start() + finished = 0 + try: + while finished < nthreads: + got = out_queue.get() + if got is done: + finished += 1 + elif isinstance(got, Exception): + raise got + else: + place(*got) + finally: + # On the happy path this just joins the already-finished readers. On an exception (place() OOM, a reader + # error) it signals the readers to stop and keeps draining so any reader blocked on a full queue can post its + # done sentinel and exit -- so no daemon reader is ever left executing inside safetensors when the stream + # unwinds, which on Windows would segfault (0xC0000005) when the thread is force-killed at teardown. + stop.set() + while finished < nthreads: + if out_queue.get() is done: + finished += 1 + for t in threads: + t.join() + _drop_page_cache({key_to_file[key] for key, *_ in work}) + + # tied weights (e.g. Qwen3 lm_head <-> embed_tokens) are saved once, so the target stays meta; fill it with an + # independent clone of the source (not an alias -- in-place quantize would corrupt both), then quantize. Both keys + # are module-root-relative, so whole-module streams only (key_prefix == ""). + if not key_prefix: + for target_key, source_key in (tied_weights_keys or {}).items(): + parent_path, _, target_name = target_key.rpartition(".") + target_module = module.get_submodule(parent_path) + if target_module._parameters[target_name].is_meta: + source = module.get_parameter(source_key) + with mem_pool_context(dest_pool_for(target_module)): + set_module_tensor_to_device( + target_module, target_name, device, value=source.detach().clone(), dtype=source.dtype) + if quantize: + quantize_if_ready(target_module) + + # non-persistent buffers (rotary inv_freq etc.) are skipped above but materialized REAL on cpu by init_empty_weights; + # move them to the device so the forward doesn't see cpu buffers vs device activations. Whole-module streams only. + if not key_prefix and device.type != "meta": + for sub_module in module.modules(): + for buffer_name in sub_module._non_persistent_buffers_set: + buffer = sub_module._buffers.get(buffer_name) + if buffer is not None and not buffer.is_meta: + with mem_pool_context(dest_pool_for(sub_module)): + sub_module._buffers[buffer_name] = buffer.to(device) + + if bar is not None: + bar.close() + class HFModelLoaderMixin(metaclass=ABCMeta): def __init__(self): super().__init__() - def __load_sub_module( + # ===== LEGACY (non-streaming) load path -- used only when Stream From Disk is off ===== + def __load_sub_module_legacy( self, sub_module: nn.Module, dtype: DataType, @@ -131,7 +377,8 @@ def __load_sub_module( new_state_dict = {} for k, v in state_dict.items(): new_k, _ = rename_source_key( - k, weight_renamings, [], prefix=sub_module.base_model_prefix, meta_state_dict=meta_state_dict, + k, weight_renamings, [], base_model_prefix=sub_module.base_model_prefix, + meta_state_dict=meta_state_dict, ) new_state_dict[new_k] = v state_dict = new_state_dict @@ -157,7 +404,13 @@ def __load_sub_module( if torch.is_floating_point(old_value): old_type = type(old_value) if not is_quantized_parameter(module, tensor_name): - if dtype.is_quantized() or module_name in keep_in_fp32_modules: + if module_name in keep_in_fp32_modules: + value = value.to(dtype=train_dtype.torch_dtype()) + elif dtype.is_quantized() and type(module) is nn.Linear: + # a plain Linear that the quantization layer filter excluded + fallback_dtype = quantization.fallback_dtype if quantization is not None else DataType.BFLOAT_16 + value = value.to(dtype=fallback_dtype.torch_dtype()) + elif dtype.is_quantized(): value = value.to(dtype=train_dtype.torch_dtype()) else: value = value.to(dtype=dtype.torch_dtype()) @@ -189,6 +442,7 @@ def __load_sub_module( module._parameters[tensor_name] = type(module._parameters[tensor_name])(source) return sub_module + # ===== end LEGACY load path ===== def _load_transformers_sub_module( self, @@ -197,7 +451,12 @@ def _load_transformers_sub_module( train_dtype: DataType, pretrained_model_name_or_path: str, subfolder: str = "", + stream_from_disk: bool | None = None, ): + # stream_from_disk None means the caller never streams this sub-module, and a bare module is returned. + # Passing True or False marks the call site as stream-capable and always returns a + # (sub_module, materialize_fn) pair -- materialize_fn None when not streaming -- so that caller needs no + # branch on the flag. user_agent = { "file_type": "model", "framework": "pytorch", @@ -213,19 +472,116 @@ def _load_transformers_sub_module( with accelerate.init_empty_weights(): sub_module = module_type(config) - return self.__load_sub_module( - sub_module=sub_module, - dtype=dtype, - train_dtype=train_dtype, - keep_in_fp32_modules=module_type._keep_in_fp32_modules, - quantization=None, - pretrained_model_name_or_path=pretrained_model_name_or_path, - subfolder=subfolder, + if not stream_from_disk: + # LEGACY fallback: whole-checkpoint-into-RAM load + sub_module = self.__load_sub_module_legacy( + sub_module=sub_module, + dtype=dtype, + train_dtype=train_dtype, + keep_in_fp32_modules=module_type._keep_in_fp32_modules, + quantization=None, + pretrained_model_name_or_path=pretrained_model_name_or_path, + subfolder=subfolder, + model_filename="model.safetensors", + pytorch_model_filename="pytorch_model.bin", + shard_index_filename="model.safetensors.index.json", + ) + if stream_from_disk is None: + return sub_module + else: + # the weights are already in RAM, so there is nothing left to materialize + return sub_module, None + + keep_in_fp32_modules = module_type._keep_in_fp32_modules or [] + replace_linear_with_quantized_layers(sub_module, dtype, keep_in_fp32_modules, None, copy_parameters=False) + # stamp with train_dtype=None: the per-part train_dtype is only known at materialize, so keep-in-fp32 and + # quantized-leftover params stay fp32 in the skeleton (a safe budget over-estimate); see _intended_float_dtype + _stamp_skeleton_float_dtypes(sub_module, dtype, None, keep_in_fp32_modules) + + key_to_file = self.__resolve_shard_key_to_file( + pretrained_model_name_or_path, subfolder, model_filename="model.safetensors", - pytorch_model_filename="pytorch_model.bin", shard_index_filename="model.safetensors.index.json", ) + # some checkpoints (e.g. Ernie's Mistral3, Qwen's Qwen2_5_VL text encoders) were saved with an older module + # layout than transformers builds from the config now. Reuse transformers' own checkpoint conversion registry + # to rename the checkpoint keys to the module's layout so the streamed lookup finds them. diffusers sub-modules + # have no such registry (plain FrozenDict config, no model_type) and never need this. + weight_renamings = get_checkpoint_conversion_mapping(sub_module.config.model_type) \ + if hasattr(sub_module.config, 'model_type') else None + source_key_map = None + if weight_renamings: + meta_state_dict = sub_module.state_dict() + renamed_key_to_file = {} + # the rename maps each checkpoint key to the module's layout so the streamed lookup and the offload cache + # find it; the file itself still stores the original key, so keep renamed->original to read the tensor. + source_key_map = {} + for key, file in key_to_file.items(): + renamed = rename_source_key( + key, weight_renamings, [], base_model_prefix=sub_module.base_model_prefix, + meta_state_dict=meta_state_dict, + )[0] + renamed_key_to_file[renamed] = file + source_key_map[renamed] = key + key_to_file = renamed_key_to_file + + return self.__finish_sub_module_load( + sub_module, dtype, train_dtype, keep_in_fp32_modules, key_to_file, source_key_map=source_key_map) + + def __resolve_shard_key_to_file( + self, + pretrained_model_name_or_path: str, + subfolder: str, + model_filename: str, + shard_index_filename: str, + ) -> dict[str, str]: + # map every checkpoint tensor key to the local safetensors file that holds it (downloading shards from the + # hub if the source is a repo id), so the streaming fill can read each tensor on demand. + is_local = os.path.isdir(pretrained_model_name_or_path) + + def resolve(filename: str) -> str | None: + # return a local path to `filename` (downloading it from the hub if needed), or None if it is absent + if is_local: + if subfolder: + path = os.path.join(pretrained_model_name_or_path, subfolder, filename) + else: + path = os.path.join(pretrained_model_name_or_path, filename) + return path if os.path.isfile(path) else None + try: + return huggingface_hub.hf_hub_download( + repo_id=pretrained_model_name_or_path, subfolder=subfolder, filename=filename) + except EntryNotFoundError: + return None + + key_to_file = {} + + index_path = resolve(shard_index_filename) + if index_path is not None: + with open(index_path, "r") as f: + weight_map = json.loads(f.read())["weight_map"] + shard_paths = {shard: resolve(shard) for shard in set(weight_map.values())} + for key, shard in weight_map.items(): + key_to_file[key] = shard_paths[shard] + return key_to_file + + # non-sharded: prefer the full-precision safetensors, fall back to the fp16 variant (some older repos, e.g. + # stable-diffusion-inpainting, ship only *.fp16.safetensors next to legacy pickle .bin files). Pickle .bin + # weights are not supported -- safe_open needs safetensors for random per-tensor reads. + fp16_filename = model_filename.replace(".safetensors", ".fp16.safetensors") + full_filename = resolve(model_filename) or resolve(fp16_filename) + if full_filename is None: + location = f"{pretrained_model_name_or_path}/{subfolder}" if subfolder else pretrained_model_name_or_path + raise FileNotFoundError( + f"No safetensors weights found for '{location}' (looked for {model_filename} and {fp16_filename}). " + f"Only pickle .bin checkpoints are present, which are not supported; convert the model to " + f"safetensors.") + with safe_open(full_filename, framework="pt") as f: + for key in f.keys(): # noqa: SIM118 -- safe_open handle, not a dict + key_to_file[key] = full_filename + + return key_to_file + def _load_diffusers_sub_module( self, module_type, @@ -234,7 +590,12 @@ def _load_diffusers_sub_module( pretrained_model_name_or_path: str, subfolder: str | None = None, quantization: QuantizationConfig | None = None, + stream_from_disk: bool | None = None, ): + # stream_from_disk None means the caller never streams this sub-module, and a bare module is returned. + # Passing True or False marks the call site as stream-capable and always returns a + # (sub_module, materialize_fn) pair -- materialize_fn None when not streaming -- so that caller needs no + # branch on the flag. user_agent = { "file_type": "model", "framework": "pytorch", @@ -250,19 +611,72 @@ def _load_diffusers_sub_module( with accelerate.init_empty_weights(): sub_module = module_type.from_config(config) - return self.__load_sub_module( - sub_module=sub_module, - dtype=dtype, - train_dtype=train_dtype, - keep_in_fp32_modules=module_type._keep_in_fp32_modules, - quantization=quantization, - pretrained_model_name_or_path=pretrained_model_name_or_path, - subfolder=subfolder, + if not stream_from_disk: + # LEGACY fallback: whole-checkpoint-into-RAM load + sub_module = self.__load_sub_module_legacy( + sub_module=sub_module, + dtype=dtype, + train_dtype=train_dtype, + keep_in_fp32_modules=module_type._keep_in_fp32_modules, + quantization=quantization, + pretrained_model_name_or_path=pretrained_model_name_or_path, + subfolder=subfolder, + model_filename="diffusion_pytorch_model.safetensors", + pytorch_model_filename="diffusion_pytorch_model.bin", + shard_index_filename="diffusion_pytorch_model.safetensors.index.json", + ) + if stream_from_disk is None: + return sub_module + else: + # the weights are already in RAM, so there is nothing left to materialize + return sub_module, None + + keep_in_fp32_modules = module_type._keep_in_fp32_modules or [] + replace_linear_with_quantized_layers(sub_module, dtype, keep_in_fp32_modules, quantization, copy_parameters=False) + # stamp with train_dtype=None: the per-part train_dtype is only known at materialize, so keep-in-fp32 and + # quantized-leftover params stay fp32 in the skeleton (a safe budget over-estimate); see _intended_float_dtype + _stamp_skeleton_float_dtypes(sub_module, dtype, None, keep_in_fp32_modules) + + key_to_file = self.__resolve_shard_key_to_file( + pretrained_model_name_or_path, subfolder, model_filename="diffusion_pytorch_model.safetensors", - pytorch_model_filename="diffusion_pytorch_model.bin", shard_index_filename="diffusion_pytorch_model.safetensors.index.json", ) + # diffusers renamed deprecated attention-block weights (query->to_q etc.); older single-file checkpoints still + # use the old names. _fix_state_dict_keys_on_load rewrites them to the current layout, and since it only + # renames dict keys, applying it to the key->file map matches applying it to a state_dict. No-op for modern + # architectures. + if hasattr(sub_module, '_fix_state_dict_keys_on_load'): + sub_module._fix_state_dict_keys_on_load(key_to_file) + + return self.__finish_sub_module_load( + sub_module, dtype, train_dtype, keep_in_fp32_modules, key_to_file) + + def __finish_sub_module_load( + self, + sub_module: nn.Module, + dtype: DataType, + train_dtype: DataType, + keep_in_fp32_modules: list[str], + key_to_file: dict[str, str], + source_key_map: dict[str, str] | None = None, + ): + tied_weights_keys = getattr(sub_module, "_tied_weights_keys", None) + + # module/key_prefix let the layer-offload conductor reuse this same closure to stream one layer at a time + # (module=that layer, key_prefix=its path in the checkpoint) as well as the non-layer remainder + # (module=the whole sub-module, key_prefix=""). Whole-module callers pass neither and stream everything. + def materialize_fn( + module: nn.Module, device: torch.device, train_dtype: DataType, key_prefix: str = "", + part_name: str | None = None, dest_pool=None): + stream_module_from_checkpoint( + module, device, key_to_file, dtype, train_dtype, + keep_in_fp32_modules, tied_weights_keys, quantize=True, key_prefix=key_prefix, + source_key_map=source_key_map, part_name=part_name, dest_pool=dest_pool) + + return sub_module, materialize_fn + def __convert_sub_module_to_dtype( self, sub_module: nn.Module, @@ -283,7 +697,13 @@ def __convert_sub_module_to_dtype( if value is not None and torch.is_floating_point(value): old_type = type(value) if not is_quantized_parameter(module, tensor_name): - if dtype.is_quantized() or module_name in keep_in_fp32_modules: + if module_name in keep_in_fp32_modules: + value = value.to(dtype=train_dtype.torch_dtype()) + elif dtype.is_quantized() and type(module) is nn.Linear: + # a plain Linear that the quantization layer filter excluded + fallback_dtype = quantization.fallback_dtype if quantization is not None else DataType.BFLOAT_16 + value = value.to(dtype=fallback_dtype.torch_dtype()) + elif dtype.is_quantized(): value = value.to(dtype=train_dtype.torch_dtype()) else: value = value.to(dtype=dtype.torch_dtype()) @@ -328,3 +748,99 @@ def _convert_diffusers_sub_module_to_dtype( None, quantization, ) + + def _load_transformer( + self, + module_type, + weight_dtypes: ModelWeightDtypes, + base_model_name: str, + transformer_model_name: str, + quantization: QuantizationConfig, + config: str | None = None, + stream_from_disk: bool = False, + subfolder: str = "transformer", + ): + # a single-file (optionally GGUF-quantized) checkpoint is loaded directly, using + # a separate repo to source the model config if the checkpoint doesn't carry one; + # otherwise the transformer is loaded from its subfolder in the base model repo. subfolder is that name: + # "transformer" everywhere except where a repo ships more than one DiT and the trainable one is not the + # repo's default (LTX 2.5's transformer_full/). + # Returns a (transformer, materialize_fn) pair, materialize_fn None when not streamed. + if transformer_model_name: + single_file_kwargs = {} + if config is not None: + single_file_kwargs["config"] = config + single_file_kwargs["subfolder"] = subfolder + + transformer = module_type.from_single_file( + transformer_model_name, + **single_file_kwargs, + #avoid loading the transformer in float32: + torch_dtype=torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), + quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, + ) + transformer = self._convert_diffusers_sub_module_to_dtype( + transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, + ) + return transformer, None + else: + # when streaming, this yields a meta skeleton + materialize closure: weights are streamed and quantized + # to the compute device on use and evicted back to meta afterwards, so the full unquantized module never + # lands in RAM, and train_dtype is applied per-materialize rather than here. + return self._load_diffusers_sub_module( + module_type, + weight_dtypes.transformer, + weight_dtypes.train_dtype, + base_model_name, + subfolder, + quantization, + stream_from_disk=stream_from_disk, + ) + + def _load_text_encoder( + self, + module_type, + dtype: DataType, + train_dtype: DataType, + base_model_name: str, + subfolder: str, + stream_from_disk: bool = False, + ): + # text encoders have no single-file override and always load from their subfolder. Returns a + # (text_encoder, materialize_fn) pair, materialize_fn None when not streamed, mirroring _load_transformer. + # dtype/train_dtype are explicit rather than a weight_dtypes bundle since a model can hold several encoders + # (text_encoder, text_encoder_2, ...) with differing dtypes. + return self._load_transformers_sub_module( + module_type, + dtype, + train_dtype, + base_model_name, + subfolder, + stream_from_disk=stream_from_disk, + ) + + def _load_vae( + self, + module_type, + dtype: DataType, + train_dtype: DataType, + base_model_name: str, + vae_model_name: str, + ): + # a separate vae repo overrides the base model's vae subfolder when given. train_dtype is explicit + # since some models (e.g. SDXL) upgrade the vae to fallback_train_dtype to avoid fp16 overflow + if vae_model_name: + return self._load_diffusers_sub_module( + module_type, + dtype, + train_dtype, + vae_model_name, + ) + else: + return self._load_diffusers_sub_module( + module_type, + dtype, + train_dtype, + base_model_name, + "vae", + ) diff --git a/modules/modelLoader/pixartAlpha/PixArtAlphaModelLoader.py b/modules/modelLoader/pixartAlpha/PixArtAlphaModelLoader.py index 467c29c8a..9c6d647e6 100644 --- a/modules/modelLoader/pixartAlpha/PixArtAlphaModelLoader.py +++ b/modules/modelLoader/pixartAlpha/PixArtAlphaModelLoader.py @@ -27,9 +27,11 @@ def __load_internal( base_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): - self.__load_diffusers(model, model_type, weight_dtypes, base_model_name, vae_model_name, quantization) + self.__load_diffusers( + model, model_type, weight_dtypes, base_model_name, vae_model_name, quantization, stream_from_disk) else: raise Exception("not an internal model") @@ -41,58 +43,45 @@ def __load_diffusers( base_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): - tokenizer = T5Tokenizer.from_pretrained( + model.tokenizer = T5Tokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) + model.orig_tokenizer = copy.deepcopy(model.tokenizer) - noise_scheduler = DDIMScheduler.from_pretrained( + model.noise_scheduler = DDIMScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( T5EncoderModel, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + model.vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - transformer = self._load_diffusers_sub_module( + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( Transformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, + weight_dtypes, base_model_name, - "transformer", + "", quantization, + stream_from_disk=stream_from_disk, ) - model.model_type = model_type - model.tokenizer = tokenizer - model.orig_tokenizer = copy.deepcopy(tokenizer) - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.vae = vae - model.transformer = transformer - def load( self, model: PixArtAlphaModel, @@ -100,19 +89,22 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ) -> PixArtAlphaModel | None: stacktraces = [] base_model_name = model_names.base_model try: - self.__load_internal(model, model_type, weight_dtypes, base_model_name, model_names.vae_model, quantization) + self.__load_internal( + model, model_type, weight_dtypes, base_model_name, model_names.vae_model, quantization, stream_from_disk) return except Exception: stacktraces.append(traceback.format_exc()) try: - self.__load_diffusers(model, model_type, weight_dtypes, base_model_name, model_names.vae_model, quantization) + self.__load_diffusers( + model, model_type, weight_dtypes, base_model_name, model_names.vae_model, quantization, stream_from_disk) return except Exception: stacktraces.append(traceback.format_exc()) diff --git a/modules/modelLoader/qwen/QwenModelLoader.py b/modules/modelLoader/qwen/QwenModelLoader.py index 953f15bfb..21a4e76f5 100644 --- a/modules/modelLoader/qwen/QwenModelLoader.py +++ b/modules/modelLoader/qwen/QwenModelLoader.py @@ -8,12 +8,9 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKLQwenImage, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, QwenImageTransformer2DModel, ) from transformers import Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer @@ -34,10 +31,12 @@ def __load_internal( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( - model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, quantization, + model, model_type, weight_dtypes, base_model_name, transformer_model_name, vae_model_name, + quantization, stream_from_disk, ) else: raise Exception("not an internal model") @@ -51,69 +50,44 @@ def __load_diffusers( transformer_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): - tokenizer = Qwen2Tokenizer.from_pretrained( + model.tokenizer = Qwen2Tokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( Qwen2_5_VLForConditionalGeneration, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) - if vae_model_name: #TODO simplify - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) - - if transformer_model_name: - transformer = QwenImageTransformer2DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - QwenImageTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + model.vae = self._load_vae( + AutoencoderKLQwenImage, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - model.model_type = model_type - model.tokenizer = tokenizer - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.vae = vae - model.transformer = transformer + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( + QwenImageTransformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + stream_from_disk=stream_from_disk, + ) def __load_safetensors( self, @@ -135,12 +109,14 @@ def load( #TODO share code between models model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] try: self.__load_internal( - model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, + model_names.vae_model, quantization, stream_from_disk, ) return except Exception: @@ -148,7 +124,8 @@ def load( #TODO share code between models try: self.__load_diffusers( - model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, model_names.vae_model, quantization, + model, model_type, weight_dtypes, model_names.base_model, model_names.transformer_model, + model_names.vae_model, quantization, stream_from_disk, ) return except Exception: diff --git a/modules/modelLoader/sana/SanaModelLoader.py b/modules/modelLoader/sana/SanaModelLoader.py index a904e3996..74700bf54 100644 --- a/modules/modelLoader/sana/SanaModelLoader.py +++ b/modules/modelLoader/sana/SanaModelLoader.py @@ -27,9 +27,11 @@ def __load_internal( base_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): - self.__load_diffusers(model, model_type, weight_dtypes, base_model_name, vae_model_name, quantization) + self.__load_diffusers( + model, model_type, weight_dtypes, base_model_name, vae_model_name, quantization, stream_from_disk) else: raise Exception("not an internal model") @@ -41,58 +43,45 @@ def __load_diffusers( base_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): - tokenizer = GemmaTokenizer.from_pretrained( + model.tokenizer = GemmaTokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) + model.orig_tokenizer = copy.deepcopy(model.tokenizer) - noise_scheduler = DPMSolverMultistepScheduler.from_pretrained( + model.noise_scheduler = DPMSolverMultistepScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + model.text_encoder, model.materialize_fn["text_encoder"] = self._load_text_encoder( Gemma2Model, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderDC, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderDC, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + model.vae = self._load_vae( + AutoencoderDC, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - transformer = self._load_diffusers_sub_module( + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( SanaTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, + weight_dtypes, base_model_name, - "transformer", + "", quantization, + stream_from_disk=stream_from_disk, ) - model.model_type = model_type - model.tokenizer = tokenizer - model.orig_tokenizer = copy.deepcopy(tokenizer) - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.vae = vae - model.transformer = transformer - def load( self, model: SanaModel, @@ -100,19 +89,22 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ) -> SanaModel | None: stacktraces = [] base_model_name = model_names.base_model try: - self.__load_internal(model, model_type, weight_dtypes, base_model_name, model_names.vae_model, quantization) + self.__load_internal( + model, model_type, weight_dtypes, base_model_name, model_names.vae_model, quantization, stream_from_disk) return except Exception: stacktraces.append(traceback.format_exc()) try: - self.__load_diffusers(model, model_type, weight_dtypes, base_model_name, model_names.vae_model, quantization) + self.__load_diffusers( + model, model_type, weight_dtypes, base_model_name, model_names.vae_model, quantization, stream_from_disk) return except Exception: stacktraces.append(traceback.format_exc()) diff --git a/modules/modelLoader/stableDiffusion/StableDiffusionModelLoader.py b/modules/modelLoader/stableDiffusion/StableDiffusionModelLoader.py index aa610f485..50dd4d12b 100644 --- a/modules/modelLoader/stableDiffusion/StableDiffusionModelLoader.py +++ b/modules/modelLoader/stableDiffusion/StableDiffusionModelLoader.py @@ -73,21 +73,22 @@ def __load_diffusers( vae_model_name: str, quantization: QuantizationConfig, ): - tokenizer = CLIPTokenizer.from_pretrained( + model.tokenizer = CLIPTokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) + model.orig_tokenizer = copy.deepcopy(model.tokenizer) noise_scheduler = DDIMScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - noise_scheduler = create.create_noise_scheduler( + model.noise_scheduler = create.create_noise_scheduler( noise_scheduler=NoiseScheduler.DDIM, original_noise_scheduler=noise_scheduler, ) - text_encoder = self._load_transformers_sub_module( + model.text_encoder, _ = self._load_text_encoder( CLIPTextModel, weight_dtypes.text_encoder, weight_dtypes.train_dtype, @@ -95,23 +96,15 @@ def __load_diffusers( "text_encoder", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + model.vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - unet = self._load_diffusers_sub_module( + model.unet = self._load_diffusers_sub_module( UNet2DConditionModel, weight_dtypes.unet, weight_dtypes.train_dtype, @@ -120,27 +113,17 @@ def __load_diffusers( quantization, ) - image_depth_processor = DPTImageProcessor.from_pretrained( + model.image_depth_processor = DPTImageProcessor.from_pretrained( base_model_name, subfolder="feature_extractor", ) if model_type.has_depth_input() else None - depth_estimator = DPTForDepthEstimation.from_pretrained( + model.depth_estimator = DPTForDepthEstimation.from_pretrained( base_model_name, subfolder="depth_estimator", torch_dtype=weight_dtypes.unet.torch_dtype(), # TODO: use depth estimator dtype ) if model_type.has_depth_input() else None - model.model_type = model_type - model.tokenizer = tokenizer - model.orig_tokenizer = copy.deepcopy(tokenizer) - model.noise_scheduler = noise_scheduler - model.text_encoder = text_encoder - model.vae = vae - model.unet = unet - model.image_depth_processor = image_depth_processor - model.depth_estimator = depth_estimator - def __fix_nai_model(self, state_dict: dict) -> dict: # fix for loading models with an empty state_dict key while 'state_dict' in state_dict and len(state_dict['state_dict']) > 0: @@ -280,9 +263,17 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] + if stream_from_disk: + # SD 1.5 / 2.x checkpoints are almost always single-file (.ckpt/.safetensors loaded via + # download_from_original_stable_diffusion_ckpt), which builds a full pipeline and can't stream from a meta + # skeleton. The diffusers-subfolder path could stream its unet/text encoder like SDXL does, but wasn't + # wired up, as this is legacy. So the toggle is ignored here. + print(f"Warning: 'stream from disk' is not supported for {model_type}; loading the model fully into RAM.") + model.sd_config = self._load_sd_config(model_type, model_names.base_model) model.sd_config_filename = self._get_sd_config_name(model_type, model_names.base_model) diff --git a/modules/modelLoader/stableDiffusion3/StableDiffusion3ModelLoader.py b/modules/modelLoader/stableDiffusion3/StableDiffusion3ModelLoader.py index 47d87da74..633e0a72d 100644 --- a/modules/modelLoader/stableDiffusion3/StableDiffusion3ModelLoader.py +++ b/modules/modelLoader/stableDiffusion3/StableDiffusion3ModelLoader.py @@ -30,11 +30,13 @@ def __load_internal( include_text_encoder_2: bool, include_text_encoder_3: bool, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): self.__load_diffusers( model, model_type, weight_dtypes, base_model_name, vae_model_name, include_text_encoder_1, include_text_encoder_2, include_text_encoder_3, quantization, + stream_from_disk, ) else: raise Exception("not an internal model") @@ -50,108 +52,93 @@ def __load_diffusers( include_text_encoder_2: bool, include_text_encoder_3: bool, quantization: QuantizationConfig, + stream_from_disk: bool, ): if include_text_encoder_1: - tokenizer_1 = CLIPTokenizer.from_pretrained( + model.tokenizer_1 = CLIPTokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) else: - tokenizer_1 = None + model.tokenizer_1 = None + model.orig_tokenizer_1 = copy.deepcopy(model.tokenizer_1) if include_text_encoder_2: - tokenizer_2 = CLIPTokenizer.from_pretrained( + model.tokenizer_2 = CLIPTokenizer.from_pretrained( base_model_name, subfolder="tokenizer_2", ) else: - tokenizer_2 = None + model.tokenizer_2 = None + model.orig_tokenizer_2 = copy.deepcopy(model.tokenizer_2) if include_text_encoder_3: - tokenizer_3 = T5Tokenizer.from_pretrained( + model.tokenizer_3 = T5Tokenizer.from_pretrained( base_model_name, subfolder="tokenizer_3", ) else: - tokenizer_3 = None + model.tokenizer_3 = None + model.orig_tokenizer_3 = copy.deepcopy(model.tokenizer_3) - noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + model.noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) if include_text_encoder_1: - text_encoder_1 = self._load_transformers_sub_module( + model.text_encoder_1, model.materialize_fn["text_encoder_1"] = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder, weight_dtypes.train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) else: - text_encoder_1 = None + model.text_encoder_1 = None if include_text_encoder_2: - text_encoder_2 = self._load_transformers_sub_module( + model.text_encoder_2, model.materialize_fn["text_encoder_2"] = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder_2, weight_dtypes.train_dtype, base_model_name, "text_encoder_2", + stream_from_disk=stream_from_disk, ) else: - text_encoder_2 = None + model.text_encoder_2 = None if include_text_encoder_3: - text_encoder_3 = self._load_transformers_sub_module( + model.text_encoder_3, model.materialize_fn["text_encoder_3"] = self._load_text_encoder( T5EncoderModel, weight_dtypes.text_encoder_3, weight_dtypes.fallback_train_dtype, base_model_name, "text_encoder_3", + stream_from_disk=stream_from_disk, ) else: - text_encoder_3 = None + model.text_encoder_3 = None - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + model.vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - transformer = self._load_diffusers_sub_module( + model.transformer, model.materialize_fn["transformer"] = self._load_transformer( SD3Transformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, + weight_dtypes, base_model_name, - "transformer", + "", quantization, + stream_from_disk=stream_from_disk, ) - model.model_type = model_type - model.tokenizer_1 = tokenizer_1 - model.orig_tokenizer_1 = copy.deepcopy(tokenizer_1) - model.tokenizer_2 = tokenizer_2 - model.orig_tokenizer_2 = copy.deepcopy(tokenizer_2) - model.tokenizer_3 = tokenizer_3 - model.orig_tokenizer_3 = copy.deepcopy(tokenizer_3) - model.noise_scheduler = noise_scheduler - model.text_encoder_1 = text_encoder_1 - model.text_encoder_2 = text_encoder_2 - model.text_encoder_3 = text_encoder_3 - model.vae = vae - model.transformer = transformer - def __load_safetensors( self, model: StableDiffusion3Model, @@ -251,6 +238,7 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] @@ -258,7 +246,7 @@ def load( self.__load_internal( model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, model_names.include_text_encoder, model_names.include_text_encoder_2, - model_names.include_text_encoder_3, quantization, + model_names.include_text_encoder_3, quantization, stream_from_disk, ) return except Exception: @@ -268,12 +256,17 @@ def load( self.__load_diffusers( model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, model_names.include_text_encoder, model_names.include_text_encoder_2, - model_names.include_text_encoder_3, quantization, + model_names.include_text_encoder_3, quantization, stream_from_disk, ) return except Exception: stacktraces.append(traceback.format_exc()) + if stream_from_disk: + # the single-file loader below builds a full pipeline via from_single_file, which can't stream; fall + # back to loading it fully into RAM. + print(f"Warning: 'stream from disk' is not supported for single-file {model_type}; loading fully into RAM.") + try: self.__load_safetensors( model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, diff --git a/modules/modelLoader/stableDiffusionXL/StableDiffusionXLModelLoader.py b/modules/modelLoader/stableDiffusionXL/StableDiffusionXLModelLoader.py index afbab6581..77ffeeb4b 100644 --- a/modules/modelLoader/stableDiffusionXL/StableDiffusionXLModelLoader.py +++ b/modules/modelLoader/stableDiffusionXL/StableDiffusionXLModelLoader.py @@ -49,9 +49,11 @@ def __load_internal( base_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): - self.__load_diffusers(model, model_type, weight_dtypes, base_model_name, vae_model_name, quantization) + self.__load_diffusers( + model, model_type, weight_dtypes, base_model_name, vae_model_name, quantization, stream_from_disk) else: raise Exception("not an internal model") @@ -63,78 +65,67 @@ def __load_diffusers( base_model_name: str, vae_model_name: str, quantization: QuantizationConfig, + stream_from_disk: bool, ): - tokenizer_1 = CLIPTokenizer.from_pretrained( + model.tokenizer_1 = CLIPTokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) + model.orig_tokenizer_1 = copy.deepcopy(model.tokenizer_1) - tokenizer_2 = CLIPTokenizer.from_pretrained( + model.tokenizer_2 = CLIPTokenizer.from_pretrained( base_model_name, subfolder="tokenizer_2", ) + model.orig_tokenizer_2 = copy.deepcopy(model.tokenizer_2) noise_scheduler = DDIMScheduler.from_pretrained( base_model_name, subfolder="scheduler", ) - noise_scheduler = create.create_noise_scheduler( + model.noise_scheduler = create.create_noise_scheduler( noise_scheduler=NoiseScheduler.DDIM, original_noise_scheduler=noise_scheduler, ) - text_encoder_1 = self._load_transformers_sub_module( + model.text_encoder_1, model.materialize_fn["text_encoder_1"] = self._load_text_encoder( CLIPTextModel, weight_dtypes.text_encoder, weight_dtypes.train_dtype, base_model_name, "text_encoder", + stream_from_disk=stream_from_disk, ) - text_encoder_2 = self._load_transformers_sub_module( + model.text_encoder_2, model.materialize_fn["text_encoder_2"] = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder_2, weight_dtypes.train_dtype, base_model_name, "text_encoder_2", + stream_from_disk=stream_from_disk, ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.fallback_train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.fallback_train_dtype, - base_model_name, - "vae", - ) + model.vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.fallback_train_dtype, + base_model_name, + vae_model_name, + ) - unet = self._load_diffusers_sub_module( + # the SDXL UNet has no single-file transformer helper and lives in the "unet" subfolder, so it loads via + # _load_diffusers_sub_module directly. + model.unet, model.materialize_fn["unet"] = self._load_diffusers_sub_module( UNet2DConditionModel, weight_dtypes.unet, weight_dtypes.train_dtype, base_model_name, "unet", quantization, + stream_from_disk=stream_from_disk, ) - model.model_type = model_type - model.tokenizer_1 = tokenizer_1 - model.orig_tokenizer_1 = copy.deepcopy(tokenizer_1) - model.tokenizer_2 = tokenizer_2 - model.orig_tokenizer_2 = copy.deepcopy(tokenizer_2) - model.noise_scheduler = noise_scheduler - model.text_encoder_1 = text_encoder_1 - model.text_encoder_2 = text_encoder_2 - model.vae = vae - model.unet = unet - def __load_ckpt( self, model: StableDiffusionXLModel, @@ -248,6 +239,7 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] @@ -255,17 +247,26 @@ def load( model.sd_config_filename = self._get_sd_config_name(model_type, model_names.base_model) try: - self.__load_internal(model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, quantization) + self.__load_internal( + model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, quantization, + stream_from_disk) return except Exception: stacktraces.append(traceback.format_exc()) try: - self.__load_diffusers(model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, quantization) + self.__load_diffusers( + model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, quantization, + stream_from_disk) return except Exception: stacktraces.append(traceback.format_exc()) + if stream_from_disk: + # the single-file loaders below build a full pipeline via from_single_file, which can't stream; fall + # back to loading it fully into RAM. + print(f"Warning: 'stream from disk' is not supported for single-file {model_type}; loading fully into RAM.") + try: self.__load_safetensors(model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, quantization) return diff --git a/modules/modelLoader/wuerstchen/WuerstchenModelLoader.py b/modules/modelLoader/wuerstchen/WuerstchenModelLoader.py index 188107a2c..c0d948880 100644 --- a/modules/modelLoader/wuerstchen/WuerstchenModelLoader.py +++ b/modules/modelLoader/wuerstchen/WuerstchenModelLoader.py @@ -62,20 +62,20 @@ def __load_diffusers( quantization: QuantizationConfig, ): if model_type.is_wuerstchen_v2(): - decoder_tokenizer = CLIPTokenizer.from_pretrained( + model.decoder_tokenizer = CLIPTokenizer.from_pretrained( decoder_model_name, subfolder="tokenizer", ) if model_type.is_stable_cascade(): - decoder_tokenizer = None + model.decoder_tokenizer = None - decoder_noise_scheduler = DDPMWuerstchenScheduler.from_pretrained( + model.decoder_noise_scheduler = DDPMWuerstchenScheduler.from_pretrained( decoder_model_name, subfolder="scheduler", ) if model_type.is_wuerstchen_v2(): - decoder_text_encoder = self._load_transformers_sub_module( + model.decoder_text_encoder, _ = self._load_text_encoder( CLIPTextModel, weight_dtypes.decoder_text_encoder, weight_dtypes.train_dtype, @@ -83,10 +83,10 @@ def __load_diffusers( "text_encoder", ) if model_type.is_stable_cascade(): - decoder_text_encoder = None + model.decoder_text_encoder = None if model_type.is_wuerstchen_v2(): - decoder_decoder = self._load_diffusers_sub_module( + model.decoder_decoder = self._load_diffusers_sub_module( WuerstchenDiffNeXt, weight_dtypes.decoder, weight_dtypes.train_dtype, @@ -94,7 +94,7 @@ def __load_diffusers( "decoder", ) elif model_type.is_stable_cascade(): - decoder_decoder = self._load_diffusers_sub_module( + model.decoder_decoder = self._load_diffusers_sub_module( StableCascadeUNet, weight_dtypes.decoder, weight_dtypes.train_dtype, @@ -102,7 +102,7 @@ def __load_diffusers( "decoder", ) - decoder_vqgan = self._load_diffusers_sub_module( + model.decoder_vqgan = self._load_diffusers_sub_module( PaellaVQModel, weight_dtypes.decoder_vqgan, weight_dtypes.train_dtype, @@ -111,7 +111,7 @@ def __load_diffusers( ) if model_type.is_wuerstchen_v2(): - effnet_encoder = self._load_diffusers_sub_module( + model.effnet_encoder = self._load_diffusers_sub_module( WuerstchenEfficientNetEncoder, weight_dtypes.effnet_encoder, weight_dtypes.fallback_train_dtype, @@ -121,12 +121,12 @@ def __load_diffusers( # TODO: this is a temporary workaround until the effnet weights are available in diffusers format effnet_encoder = WuerstchenEfficientNetEncoder(affine_batch_norm=False) effnet_encoder.load_state_dict(load_file(effnet_encoder_model_name)) - effnet_encoder = self._convert_diffusers_sub_module_to_dtype( + model.effnet_encoder = self._convert_diffusers_sub_module_to_dtype( effnet_encoder, weight_dtypes.effnet_encoder, weight_dtypes.fallback_train_dtype ) if model_type.is_wuerstchen_v2(): - prior_prior = self._load_diffusers_sub_module( + model.prior_prior = self._load_diffusers_sub_module( WuerstchenPrior, weight_dtypes.prior, weight_dtypes.train_dtype, @@ -145,11 +145,11 @@ def __load_diffusers( prior_config = json.load(config_file) prior_prior = StableCascadeUNet(**prior_config) prior_prior.load_state_dict(convert_stable_cascade_ckpt_to_diffusers(load_file(prior_prior_model_name))) - prior_prior = self._convert_diffusers_sub_module_to_dtype( + model.prior_prior = self._convert_diffusers_sub_module_to_dtype( prior_prior, weight_dtypes.prior, weight_dtypes.fallback_train_dtype, quantization, ) else: - prior_prior = self._load_diffusers_sub_module( + model.prior_prior = self._load_diffusers_sub_module( StableCascadeUNet, weight_dtypes.prior, weight_dtypes.fallback_train_dtype, @@ -158,13 +158,14 @@ def __load_diffusers( quantization, ) - prior_tokenizer = CLIPTokenizer.from_pretrained( + model.prior_tokenizer = CLIPTokenizer.from_pretrained( prior_model_name, subfolder="tokenizer", ) + model.orig_prior_tokenizer = copy.deepcopy(model.prior_tokenizer) if model_type.is_wuerstchen_v2(): - prior_text_encoder = self._load_transformers_sub_module( + model.prior_text_encoder, _ = self._load_text_encoder( CLIPTextModel, weight_dtypes.text_encoder, weight_dtypes.train_dtype, @@ -172,7 +173,7 @@ def __load_diffusers( "text_encoder", ) elif model_type.is_stable_cascade(): - prior_text_encoder = self._load_transformers_sub_module( + model.prior_text_encoder, _ = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder, weight_dtypes.train_dtype, @@ -180,24 +181,11 @@ def __load_diffusers( "text_encoder", ) - prior_noise_scheduler = DDPMWuerstchenScheduler.from_pretrained( + model.prior_noise_scheduler = DDPMWuerstchenScheduler.from_pretrained( prior_model_name, subfolder="scheduler", ) - model.model_type = model_type - model.decoder_tokenizer = decoder_tokenizer - model.decoder_noise_scheduler = decoder_noise_scheduler - model.decoder_text_encoder = decoder_text_encoder - model.decoder_decoder = decoder_decoder - model.decoder_vqgan = decoder_vqgan - model.effnet_encoder = effnet_encoder - model.prior_tokenizer = prior_tokenizer - model.orig_prior_tokenizer = copy.deepcopy(prior_tokenizer) - model.prior_text_encoder = prior_text_encoder - model.prior_noise_scheduler = prior_noise_scheduler - model.prior_prior = prior_prior - def load( self, model: WuerstchenModel, @@ -205,9 +193,15 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] + if stream_from_disk: + # not supported: Stable Cascade loads its prior (single-file override) and effnet encoder by + # constructing the module and calling load_state_dict directly, which can't stream from a meta skeleton. + print(f"Warning: 'stream from disk' is not supported for {model_type}; loading the model fully into RAM.") + prior_model_name = model_names.base_model prior_prior_model_name = model_names.prior_model effnet_encoder_model_name = model_names.effnet_encoder_model diff --git a/modules/modelSampler/AnimaSampler.py b/modules/modelSampler/AnimaSampler.py index d5de18ac1..b3a8a9764 100644 --- a/modules/modelSampler/AnimaSampler.py +++ b/modules/modelSampler/AnimaSampler.py @@ -10,15 +10,15 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch from diffusers import VaeImageProcessor import numpy as np -from tqdm import tqdm @factory.register(BaseModelSampler, ModelType.ANIMA) @@ -37,100 +37,139 @@ def __init__( self.image_processor = VaeImageProcessor(vae_scale_factor=8) @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("text_encoder") + batch_size = 2 if sample_config.cfg_scale > 1.0 else 1 + combined_prompt_embedding = self.model.encode_text( + text=[sample_config.prompt, sample_config.negative_prompt] if sample_config.cfg_scale > 1.0 else sample_config.prompt, + batch_size=batch_size, + train_device=self.train_device, + ) + + return { + "batch_size": batch_size, + "combined_prompt_embedding": combined_prompt_embedding, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + batch_size: int, + combined_prompt_embedding: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.model.transformer + vae_scale_factor = 8 + num_latent_channels = 16 + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + # prepare latent image + latent_image = torch.randn( + size=(1, num_latent_channels, 1, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + + sigmas = np.linspace(1.0, 1.0 / diffusion_steps, diffusion_steps) + noise_scheduler.set_timesteps(sigmas=sigmas, device=self.train_device) + timesteps = noise_scheduler.timesteps + + padding_mask = latent_image.new_zeros( + 1, 1, height, width, dtype=transformer.dtype, + ) + + # denoising loop + extra_step_kwargs = {} + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image] * batch_size) + expanded_timestep = timestep.expand(batch_size) / noise_scheduler.config.num_train_timesteps + noise_pred = transformer( + hidden_states=latent_model_input.to(dtype=transformer.dtype), + timestep=expanded_timestep, + encoder_hidden_states=combined_prompt_embedding.to(dtype=transformer.dtype), + padding_mask=padding_mask, + return_dict=False, + )[0] + + if cfg_scale > 1.0: + noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + + latent_image = noise_scheduler.step( + noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs + )[0] + + on_update_progress(i + 1, len(timesteps)) + + return { + "latent_image": latent_image, + } + + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - - transformer = self.model.transformer - vae = self.model.vae - vae_scale_factor = 8 - num_latent_channels = 16 - - # prepare prompt - self.model.materialize_only("text_encoder") - - batch_size = 2 if cfg_scale > 1.0 else 1 - combined_prompt_embedding = self.model.encode_text( - text=[prompt, negative_prompt] if cfg_scale > 1.0 else prompt, - batch_size=batch_size, - train_device=self.train_device, - ) + self.model.materialize_only("vae") + vae = self.model.vae - # prepare latent image - latent_image = torch.randn( - size=(1, num_latent_channels, 1, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32, - ) + latents = self.model.unscale_latents(latent_image) + image = vae.decode(latents, return_dict=False)[0][:, :, 0] + + do_denormalize = [True] * image.shape[0] + image = self.image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - sigmas = np.linspace(1.0, 1.0 / diffusion_steps, diffusion_steps) - noise_scheduler.set_timesteps(sigmas=sigmas, device=self.train_device) - timesteps = noise_scheduler.timesteps + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) - padding_mask = latent_image.new_zeros( - 1, 1, height, width, dtype=transformer.dtype, + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) + + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, ) - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator - - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * batch_size) - expanded_timestep = timestep.expand(batch_size) / noise_scheduler.config.num_train_timesteps - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=transformer.dtype), - timestep=expanded_timestep, - encoder_hidden_states=combined_prompt_embedding.to(dtype=transformer.dtype), - padding_mask=padding_mask, - return_dict=False, - )[0] - - if cfg_scale > 1.0: - noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - # decode - self.model.materialize_only("vae") - - latents = self.model.unscale_latents(latent_image) - image = vae.decode(latents, return_dict=False)[0][:, :, 0] - - do_denormalize = [True] * image.shape[0] - image = self.image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -141,22 +180,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/BaseModelSampler.py b/modules/modelSampler/BaseModelSampler.py index 5a5bdd15c..049fb87c6 100644 --- a/modules/modelSampler/BaseModelSampler.py +++ b/modules/modelSampler/BaseModelSampler.py @@ -73,10 +73,72 @@ def sample( ): pass + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + # Item-major fallback for samplers not split into pipeline stages: run each sample + # end to end, cycling this model's parts on-device once per sample instead of once + # per batch. Progress restarts at each sample, matching how these samplers reported + # when the trainer still sampled one at a time. Samplers that define stages override + # this with run_staged_pipeline. + sampler_outputs = [] + for sample_config, destination in zip(sample_configs, destinations, strict=True): + self.sample( + sample_config, destination, + image_format, video_format, audio_format, + on_sample=sampler_outputs.append, + on_update_progress=on_update_progress, + ) + return sampler_outputs + @staticmethod def quantize_resolution(resolution: int, quantization: int) -> int: return round(resolution / quantization) * quantization + @staticmethod + def batch_progress_callback( + sample_configs: list[SampleConfig], + on_update_progress: Callable[[int, int], None], + ) -> Callable[[int, int], None]: + # denoise reports its own per-sample step count; the staged pipeline denoises + # every sample before any decode, so translate those into one continuous bar + # spanning the whole batch rather than restarting at each sample + total_steps = sum(sample_config.diffusion_steps for sample_config in sample_configs) + completed_steps = 0 + + def batch_progress(step: int, sample_steps: int): + nonlocal completed_steps + on_update_progress(completed_steps + step, total_steps) + if step == sample_steps: + completed_steps += sample_steps + + return batch_progress + + @staticmethod + def build_video_sampler_output(video: torch.Tensor) -> ModelSamplerOutput: + # `video` is the video processor's output, [B, F, C, H, W] with values in [0, 1]. Only the first batch + # item is kept, and a single frame becomes an image rather than a one-frame video. + video = video.cpu().float() + + if video.shape[1] == 1: + image = video[0, 0].permute(1, 2, 0).numpy() # [H, W, C] + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=Image.fromarray((image * 255).round().astype("uint8")), + ) + else: + frames = video[0].permute(0, 2, 3, 1) # [F, H, W, C] + return ModelSamplerOutput( + file_type=FileType.VIDEO, + data=(frames.clamp(0, 1) * 255).round().to(dtype=torch.int8), + ) + @staticmethod def save_sampler_output( sampler_output: ModelSamplerOutput, diff --git a/modules/modelSampler/ChromaSampler.py b/modules/modelSampler/ChromaSampler.py index 23ef4aee5..5eb8d5728 100644 --- a/modules/modelSampler/ChromaSampler.py +++ b/modules/modelSampler/ChromaSampler.py @@ -10,13 +10,12 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.CHROMA_1) class ChromaSampler(BaseModelSampler): @@ -34,122 +33,161 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - text_encoder_layer_skip: int = 0, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, - ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - image_processor = self.pipeline.image_processor - transformer = self.pipeline.transformer - vae = self.pipeline.vae - vae_scale_factor = 8 - num_latent_channels = 16 - - # prepare prompt - self.model.materialize_only("text_encoder") - - combined_prompt_embedding, text_attention_mask = self.model.encode_text( - text=[prompt, negative_prompt], - batch_size = 2, - train_device=self.train_device, - text_encoder_layer_skip=text_encoder_layer_skip, - ) + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("text_encoder") + combined_prompt_embedding, text_attention_mask = self.model.encode_text( + text=[sample_config.prompt, sample_config.negative_prompt], + batch_size=2, + train_device=self.train_device, + text_encoder_layer_skip=sample_config.text_encoder_1_layer_skip, + ) - # prepare latent image - latent_image = torch.randn( - size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32, - ) + return { + "combined_prompt_embedding": combined_prompt_embedding, + "text_attention_mask": text_attention_mask, + } - image_ids = self.model.prepare_latent_image_ids( - height // vae_scale_factor, - width // vae_scale_factor, - self.train_device, - self.model.train_dtype.torch_dtype() - ) + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + combined_prompt_embedding: torch.Tensor, + text_attention_mask: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + vae_scale_factor = 8 + num_latent_channels = 16 + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + # prepare latent image + latent_image = torch.randn( + size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) - latent_image = self.model.pack_latents(latent_image) - - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) - timesteps = noise_scheduler.timesteps - - # denoising loop - extra_step_kwargs = {} - #TODO always True for FlowMatchEulerDiscreteScheduler - remove and pass directly? - #If so, also remove for other models - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator #TODO purpose? - - text_ids = torch.zeros(combined_prompt_embedding.shape[1], 3, device=self.train_device) - - image_seq_len = latent_image.shape[1] - image_attention_mask = torch.full((2, image_seq_len), True, dtype=torch.bool, device=text_attention_mask.device) - attention_mask = torch.cat([text_attention_mask, image_attention_mask], dim=1) - - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * 2) - expanded_timestep = timestep.expand(2) - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), - timestep=expanded_timestep / 1000, - encoder_hidden_states=combined_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), - txt_ids=text_ids.to(dtype=self.model.train_dtype.torch_dtype()), - img_ids=image_ids.to(dtype=self.model.train_dtype.torch_dtype()), - attention_mask=attention_mask, - joint_attention_kwargs=None, - return_dict=True - ).sample - - noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - latent_image = self.model.unpack_latents( - latent_image, - height // vae_scale_factor, - width // vae_scale_factor, - ) + image_ids = self.model.prepare_latent_image_ids( + height // vae_scale_factor, + width // vae_scale_factor, + self.train_device, + self.model.train_dtype.torch_dtype() + ) + + latent_image = self.model.pack_latents(latent_image) + + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) + timesteps = noise_scheduler.timesteps + + # denoising loop + extra_step_kwargs = {} + #TODO always True for FlowMatchEulerDiscreteScheduler - remove and pass directly? + #If so, also remove for other models + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator #TODO purpose? + + text_ids = torch.zeros(combined_prompt_embedding.shape[1], 3, device=self.train_device) + + image_seq_len = latent_image.shape[1] + image_attention_mask = torch.full((2, image_seq_len), True, dtype=torch.bool, device=text_attention_mask.device) + attention_mask = torch.cat([text_attention_mask, image_attention_mask], dim=1) + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image] * 2) + expanded_timestep = timestep.expand(2) + noise_pred = transformer( + hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), + timestep=expanded_timestep / 1000, + encoder_hidden_states=combined_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), + txt_ids=text_ids.to(dtype=self.model.train_dtype.torch_dtype()), + img_ids=image_ids.to(dtype=self.model.train_dtype.torch_dtype()), + attention_mask=attention_mask, + joint_attention_kwargs=None, + return_dict=True + ).sample + + noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + + # compute the previous noisy sample x_t -> x_t-1 + latent_image = noise_scheduler.step( + noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs + )[0] + + on_update_progress(i + 1, len(timesteps)) + + latent_image = self.model.unpack_latents( + latent_image, + height // vae_scale_factor, + width // vae_scale_factor, + ) + + return { + "latent_image": latent_image, + } + + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, + ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae + + latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor + image = vae.decode(latents, return_dict=False)[0] + + do_denormalize = [True] * image.shape[0] + image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - # decode - self.model.materialize_only("vae") + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) - latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor - image = vae.decode(latents, return_dict=False)[0] + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) - do_denormalize = [True] * image.shape[0] - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, + ) - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -160,23 +198,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - text_encoder_layer_skip=sample_config.text_encoder_1_layer_skip, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/ErnieSampler.py b/modules/modelSampler/ErnieSampler.py index a9cb57e0a..f2848f6f8 100644 --- a/modules/modelSampler/ErnieSampler.py +++ b/modules/modelSampler/ErnieSampler.py @@ -9,14 +9,14 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch import numpy as np from PIL import Image as PILImage -from tqdm import tqdm @factory.register(BaseModelSampler, ModelType.ERNIE) @@ -35,96 +35,137 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, - ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - vae = self.pipeline.vae - - vae_scale_factor = 8 - num_latent_channels = 32 + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("text_encoder") + batch_size = 2 if sample_config.cfg_scale > 1.0 else 1 + text_bth, text_lens = self.model.encode_text( + train_device=self.train_device, + text=[sample_config.prompt, sample_config.negative_prompt] if batch_size == 2 else sample_config.prompt, + ) - # encode text - self.model.materialize_only("text_encoder") + return { + "batch_size": batch_size, + "text_bth": text_bth, + "text_lens": text_lens, + } - batch_size = 2 if cfg_scale > 1.0 else 1 - text_bth, text_lens = self.model.encode_text( - train_device=self.train_device, - text=[prompt, negative_prompt] if batch_size == 2 else prompt, - ) - dtype = self.model.train_dtype.torch_dtype() - - # prepare latents - latent_image = torch.randn( - size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32, - ) - latent_image = self.model.patchify_latents(latent_image) + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + batch_size: int, + text_bth: torch.Tensor, + text_lens: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + vae_scale_factor = 8 + num_latent_channels = 32 + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + dtype = self.model.train_dtype.torch_dtype() + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + # prepare latents + latent_image = torch.randn( + size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + latent_image = self.model.patchify_latents(latent_image) - sigmas = np.linspace(1.0, 1 / diffusion_steps, diffusion_steps) - noise_scheduler.set_timesteps(sigmas=sigmas, device=self.train_device) - timesteps = noise_scheduler.timesteps + sigmas = np.linspace(1.0, 1 / diffusion_steps, diffusion_steps) + noise_scheduler.set_timesteps(sigmas=sigmas, device=self.train_device) + timesteps = noise_scheduler.timesteps - self.model.materialize_only("transformer") - transformer = self.pipeline.transformer + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image] * batch_size) + expanded_timestep = timestep.expand(latent_model_input.shape[0]) - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * batch_size) - expanded_timestep = timestep.expand(latent_model_input.shape[0]) + noise_pred = transformer( + hidden_states=latent_model_input.to(dtype=dtype), + timestep=expanded_timestep, + text_bth=text_bth, + text_lens=text_lens, + return_dict=False, + )[0] - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=dtype), - timestep=expanded_timestep, - text_bth=text_bth, - text_lens=text_lens, - return_dict=False, - )[0] + if batch_size == 2: + noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - if batch_size == 2: - noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + latent_image = noise_scheduler.step(noise_pred, timestep, latent_image, + return_dict=False)[0] - latent_image = noise_scheduler.step(noise_pred, timestep, latent_image, - return_dict=False)[0] + on_update_progress(i + 1, len(timesteps)) - on_update_progress(i + 1, len(timesteps)) + return { + "latent_image": latent_image, + } - self.model.materialize_only("vae") + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, + ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + vae = self.pipeline.vae + + # unscale and unpatchify + latents = self.model.unscale_latents(latent_image) + latents = self.model.unpatchify_latents(latents) + + image = vae.decode(latents, return_dict=False)[0] + # no VaeImageProcessor — pipeline does this manually + image = (image.clamp(-1, 1) + 1) / 2 + image = image.cpu().permute(0, 2, 3, 1).float().numpy() + image = [PILImage.fromarray((img * 255).astype(np.uint8)) for img in image] + + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) - # unscale and unpatchify - latents = self.model.unscale_latents(latent_image) - latents = self.model.unpatchify_latents(latents) + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) - image = vae.decode(latents, return_dict=False)[0] - # no VaeImageProcessor — pipeline does this manually - image = (image.clamp(-1, 1) + 1) / 2 - image = image.cpu().permute(0, 2, 3, 1).float().numpy() - image = [PILImage.fromarray((img * 255).astype(np.uint8)) for img in image] + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, + ) - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -135,22 +176,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/Flux2Sampler.py b/modules/modelSampler/Flux2Sampler.py index 7ecbd5c83..50c97e468 100644 --- a/modules/modelSampler/Flux2Sampler.py +++ b/modules/modelSampler/Flux2Sampler.py @@ -1,5 +1,6 @@ import copy import inspect +import math from collections.abc import Callable from modules.model.Flux2Model import Flux2Model @@ -10,15 +11,19 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch from diffusers.pipelines.flux2.pipeline_flux2 import compute_empirical_mu import numpy as np -from tqdm import tqdm + +VAE_SCALE_FACTOR = 8 +NUM_LATENT_CHANNELS = 32 +PATCH_SIZE = 2 @factory.register(BaseModelSampler, ModelType.FLUX_2) @@ -37,119 +42,159 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( + self, + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("text_encoder") + transformer = self.pipeline.transformer + + batch_size = 2 if sample_config.cfg_scale > 1.0 and not transformer.config.guidance_embeds else 1 + prompt_embedding = self.model.encode_text( + text=[sample_config.prompt, sample_config.negative_prompt] if batch_size == 2 else sample_config.prompt, + train_device=self.train_device, + text_encoder_sequence_length=sample_config.text_encoder_1_sequence_length, + ) + + return { + "batch_size": batch_size, + "prompt_embedding": prompt_embedding, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + batch_size: int, + prompt_embedding: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + text_ids = self.model.prepare_text_ids(prompt_embedding) + + # prepare latent image + latent_image = torch.randn( + size=(1, NUM_LATENT_CHANNELS, height // VAE_SCALE_FACTOR, width // VAE_SCALE_FACTOR), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + + latent_image = self.model.patchify_latents(latent_image) + image_ids = self.model.prepare_latent_image_ids(latent_image) + + latent_image = self.model.pack_latents(latent_image) + image_seq_len = latent_image.shape[1] + # the override is a shift factor, the same quantity the other flow-matching samplers pass as log(shift) + mu = math.log(sample_config.override_shift) if sample_config.override_shift \ + else compute_empirical_mu(image_seq_len, diffusion_steps) + + # prepare timesteps + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + #TODO for other models, too? This is different than with sigmas=None + sigmas = np.linspace(1.0, 1 / diffusion_steps, diffusion_steps) + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=mu, sigmas=sigmas) + timesteps = noise_scheduler.timesteps + + extra_step_kwargs = {} #TODO remove + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator + + guidance = (torch.tensor([sample_config.cfg_scale], device=self.train_device, dtype=self.model.train_dtype.torch_dtype()) + if transformer.config.guidance_embeds else None) + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image] * batch_size) + expanded_timestep = timestep.expand(latent_model_input.shape[0]) + + noise_pred = transformer( + hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), + timestep=expanded_timestep / 1000, + guidance=guidance, + encoder_hidden_states=prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), + txt_ids=text_ids, + img_ids=image_ids, + joint_attention_kwargs=None, + return_dict=True + ).sample + + if batch_size == 2: + noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) + noise_pred = noise_pred_negative + sample_config.cfg_scale * (noise_pred_positive - noise_pred_negative) + + latent_image = noise_scheduler.step(noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs)[0] + + on_update_progress(i + 1, len(timesteps)) + + return { + "latent_image": latent_image, + "height": height, + "width": width, + } + + @torch.no_grad() + def __decode( self, - prompt: str, - negative_prompt: str, height: int, width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - text_encoder_sequence_length: int | None = None, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, + latent_image: torch.Tensor, ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - image_processor = self.pipeline.image_processor - transformer = self.pipeline.transformer - vae = self.pipeline.vae - - vae_scale_factor = 8 - num_latent_channels = 32 - patch_size = 2 - - # prepare prompt - self.model.materialize_only("text_encoder") - - batch_size = 2 if cfg_scale > 1.0 and not transformer.config.guidance_embeds else 1 - prompt_embedding = self.model.encode_text( - text=[prompt, negative_prompt] if batch_size == 2 else prompt, - train_device=self.train_device, - text_encoder_sequence_length=text_encoder_sequence_length, - ) + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae + + latent_image = self.model.unpack_latents( + latent_image, + height // VAE_SCALE_FACTOR // PATCH_SIZE, + width // VAE_SCALE_FACTOR // PATCH_SIZE, + ) + latents = self.model.unscale_latents(latent_image) + latents = self.model.unpatchify_latents(latents) - # prepare latent image - latent_image = torch.randn( - size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32, - ) + image = vae.decode(latents, return_dict=False)[0] + image = image_processor.postprocess(image, output_type='pil') - latent_image = self.model.patchify_latents(latent_image) - image_ids = self.model.prepare_latent_image_ids(latent_image) - - latent_image = self.model.pack_latents(latent_image) - image_seq_len = latent_image.shape[1] - mu = compute_empirical_mu(image_seq_len, diffusion_steps) - - # prepare timesteps - #TODO for other models, too? This is different than with sigmas=None - sigmas = np.linspace(1.0, 1 / diffusion_steps, diffusion_steps) - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=mu, sigmas=sigmas) - timesteps = noise_scheduler.timesteps - - # denoising loop - extra_step_kwargs = {} #TODO remove - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator - - text_ids = self.model.prepare_text_ids(prompt_embedding) - - self.model.materialize_only("transformer") - guidance = (torch.tensor([cfg_scale], device=self.train_device, dtype=self.model.train_dtype.torch_dtype()) - if transformer.config.guidance_embeds else None) - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * batch_size) - expanded_timestep = timestep.expand(latent_model_input.shape[0]) - - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), - timestep=expanded_timestep / 1000, - guidance=guidance, - encoder_hidden_states=prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), - txt_ids=text_ids, - img_ids=image_ids, - joint_attention_kwargs=None, - return_dict=True - ).sample - - if batch_size == 2: - noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - latent_image = noise_scheduler.step(noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs)[0] - - on_update_progress(i + 1, len(timesteps)) - - self.model.materialize_only("vae") - - latent_image = self.model.unpack_latents( - latent_image, - height // vae_scale_factor // patch_size, - width // vae_scale_factor // patch_size, - ) - latents = self.model.unscale_latents(latent_image) - latents = self.model.unpatchify_latents(latents) + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) - image = vae.decode(latents, return_dict=False)[0] + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) - image = image_processor.postprocess(image, output_type='pil') + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, + ) - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -160,23 +205,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - text_encoder_sequence_length=sample_config.text_encoder_1_sequence_length, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/FluxSampler.py b/modules/modelSampler/FluxSampler.py index fc2b2e0c8..2325d04d3 100644 --- a/modules/modelSampler/FluxSampler.py +++ b/modules/modelSampler/FluxSampler.py @@ -11,16 +11,15 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat from modules.util.image_util import load_image +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch from torch import nn from torchvision.transforms import transforms -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.FLUX_DEV_1) @factory.register(BaseModelSampler, ModelType.FLUX_FILL_DEV_1) @@ -38,348 +37,283 @@ def __init__( self.model_type = model_type self.pipeline = model.create_pipeline() + def __create_erode_kernel(self, device, dtype=torch.float32): + kernel_radius = 2 + + kernel_size = kernel_radius * 2 + 1 + kernel_weights = torch.ones(1, 1, kernel_size, kernel_size, dtype=dtype) / (kernel_size * kernel_size) + kernel = nn.Conv2d( + in_channels=1, out_channels=1, kernel_size=kernel_size, bias=False, padding_mode='replicate', + padding=kernel_radius + ).to(dtype) + kernel.weight.data = kernel_weights + kernel.requires_grad_(False) + kernel.to(device) + return kernel + + # only present for conditioning (inpainting) model types: VAE-encode the conditioning image + mask @torch.no_grad() - def __sample_base( + def __cond_encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - text_encoder_1_layer_skip: int = 0, - text_encoder_2_layer_skip: int = 0, - text_encoder_2_sequence_length: int | None = None, - transformer_attention_mask: bool = False, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, - ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - image_processor = self.pipeline.image_processor - transformer = self.pipeline.transformer - vae = self.pipeline.vae - vae_scale_factor = 8 - num_latent_channels = 16 - - # prepare prompt - self.model.materialize_only_text_encoders() - - prompt_embedding, pooled_prompt_embedding = self.model.encode_text( - text=prompt, - train_device=self.train_device, - text_encoder_1_layer_skip=text_encoder_1_layer_skip, - text_encoder_2_layer_skip=text_encoder_2_layer_skip, - text_encoder_2_sequence_length=text_encoder_2_sequence_length, - apply_attention_mask=transformer_attention_mask, + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("vae") + vae = self.pipeline.vae + vae_scale_factor = 8 + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + + if sample_config.sample_inpainting: + t = transforms.Compose([ + transforms.ToTensor(), + transforms.Resize( + (height, width), interpolation=transforms.InterpolationMode.BILINEAR, antialias=True + ), + ]) + + image = load_image(sample_config.base_image_path, convert_mode="RGB") + image = t(image).to( + dtype=self.model.train_dtype.torch_dtype(), + device=self.train_device, ) - # prepare latent image - latent_image = torch.randn( - size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, + mask = load_image(sample_config.mask_image_path, convert_mode='L') + mask = t(mask).to( + dtype=self.model.train_dtype.torch_dtype(), device=self.train_device, - dtype=torch.float32, ) - image_ids = self.model.prepare_latent_image_ids( + erode_kernel = self.__create_erode_kernel(self.train_device, dtype=self.model.train_dtype.torch_dtype()) + eroded_mask = erode_kernel(mask) + eroded_mask = (eroded_mask > 0.5).to(dtype=self.model.train_dtype.torch_dtype()) + + image = (image * 2.0) - 1.0 + conditioning_image = (image * (1 - eroded_mask)) + conditioning_image = conditioning_image.unsqueeze(0) + + latent_conditioning_image = vae.encode(conditioning_image).latent_dist.mode() + latent_conditioning_image = (latent_conditioning_image - vae.config.shift_factor) \ + * vae.config.scaling_factor + + latent_conditioning_image = self.model.pack_latents(latent_conditioning_image) + + # batch_size, height, 8, width, 8 + mask = mask.view( + mask.shape[0], height // vae_scale_factor, + vae_scale_factor, width // vae_scale_factor, - self.train_device, - self.model.train_dtype.torch_dtype() + vae_scale_factor, ) - - shift = self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) - latent_image = self.model.pack_latents(latent_image) - - # prepare timesteps - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=math.log(shift)) - timesteps = noise_scheduler.timesteps - - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator - - text_ids = torch.zeros(prompt_embedding.shape[1], 3, device=self.train_device) - - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image]) - expanded_timestep = timestep.expand(latent_model_input.shape[0]) - - # handle guidance - if transformer.config.guidance_embeds: - guidance = torch.tensor([cfg_scale], device=self.train_device) - guidance = guidance.expand(latent_model_input.shape[0]) - else: - guidance = None - - # predict the noise residual - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), - timestep=expanded_timestep / 1000, - guidance=guidance.to(dtype=self.model.train_dtype.torch_dtype()), - pooled_projections=pooled_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), - encoder_hidden_states=prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), - txt_ids=text_ids, - img_ids=image_ids, - joint_attention_kwargs=None, - return_dict=True - ).sample - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - latent_image = self.model.unpack_latents( - latent_image, + # batch_size, 8, 8, height, width + mask = mask.permute(0, 2, 4, 1, 3) + # batch_size, 8*8, height, width + mask = mask.reshape( + mask.shape[0], + vae_scale_factor * vae_scale_factor, height // vae_scale_factor, width // vae_scale_factor, ) - # decode - self.model.materialize_only("vae") - - latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor - image = vae.decode(latents, return_dict=False)[0] + latent_mask = self.model.pack_latents(mask) + else: + conditioning_image = torch.zeros( + (1, 3, height, width), + dtype=self.model.train_dtype.torch_dtype(), + device=self.train_device, + ) + latent_conditioning_image = vae.encode(conditioning_image).latent_dist.mode() + latent_conditioning_image = (latent_conditioning_image - vae.config.shift_factor) \ + * vae.config.scaling_factor - do_denormalize = [True] * image.shape[0] #TODO remove and test, from Flux and other models. True is the default - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) + latent_conditioning_image = self.model.pack_latents(latent_conditioning_image) - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + latent_mask = torch.ones( + size=(1, (height // vae_scale_factor // 2) * (width // vae_scale_factor // 2), 256), + dtype=self.model.train_dtype.torch_dtype(), + device=self.train_device ) - def __create_erode_kernel(self, device, dtype=torch.float32): - kernel_radius = 2 - - kernel_size = kernel_radius * 2 + 1 - kernel_weights = torch.ones(1, 1, kernel_size, kernel_size, dtype=dtype) / (kernel_size * kernel_size) - kernel = nn.Conv2d( - in_channels=1, out_channels=1, kernel_size=kernel_size, bias=False, padding_mode='replicate', - padding=kernel_radius - ).to(dtype) - kernel.weight.data = kernel_weights - kernel.requires_grad_(False) - kernel.to(device) - return kernel + return { + "latent_conditioning_image": latent_conditioning_image, + "latent_mask": latent_mask, + } @torch.no_grad() - def __sample_inpainting( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - sample_inpainting: bool = False, - base_image_path: str = "", - mask_image_path: str = "", - text_encoder_1_layer_skip: int = 0, - text_encoder_2_layer_skip: int = 0, - text_encoder_2_sequence_length: int | None = None, - transformer_attention_mask: bool = False, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, - ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - image_processor = self.pipeline.image_processor - transformer = self.pipeline.transformer - vae = self.pipeline.vae - vae_scale_factor = 8 - num_latent_channels = 16 - - # prepare conditioning image - self.model.materialize_only("vae") - - if sample_inpainting: - t = transforms.Compose([ - transforms.ToTensor(), - transforms.Resize( - (height, width), interpolation=transforms.InterpolationMode.BILINEAR, antialias=True - ), - ]) - - image = load_image(base_image_path, convert_mode="RGB") - image = t(image).to( - dtype=self.model.train_dtype.torch_dtype(), - device=self.train_device, - ) - - mask = load_image(mask_image_path, convert_mode='L') - mask = t(mask).to( - dtype=self.model.train_dtype.torch_dtype(), - device=self.train_device, - ) + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only_text_encoders() + prompt_embedding, pooled_prompt_embedding = self.model.encode_text( + text=sample_config.prompt, + train_device=self.train_device, + text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, + text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, + text_encoder_2_sequence_length=sample_config.text_encoder_2_sequence_length, + apply_attention_mask=sample_config.transformer_attention_mask, + ) - erode_kernel = self.__create_erode_kernel(self.train_device, dtype=self.model.train_dtype.torch_dtype()) - eroded_mask = erode_kernel(mask) - eroded_mask = (eroded_mask > 0.5).to(dtype=self.model.train_dtype.torch_dtype()) + return { + "prompt_embedding": prompt_embedding, + "pooled_prompt_embedding": pooled_prompt_embedding, + } - image = (image * 2.0) - 1.0 - conditioning_image = (image * (1 - eroded_mask)) - conditioning_image = conditioning_image.unsqueeze(0) + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + prompt_embedding: torch.Tensor, + pooled_prompt_embedding: torch.Tensor, + on_update_progress: Callable[[int, int], None], + latent_conditioning_image: torch.Tensor | None = None, + latent_mask: torch.Tensor | None = None, + ) -> dict: + self.model.materialize_only("transformer") + # conditioning tensors are only present for inpainting model types (their cond-encode stage ran) + is_inpainting = latent_conditioning_image is not None + transformer = self.pipeline.transformer + vae_scale_factor = 8 + num_latent_channels = 16 + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) - latent_conditioning_image = vae.encode(conditioning_image).latent_dist.mode() - latent_conditioning_image = (latent_conditioning_image - vae.config.shift_factor) \ - * vae.config.scaling_factor + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - latent_conditioning_image = self.model.pack_latents(latent_conditioning_image) + # prepare latent image + latent_image = torch.randn( + size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) - # batch_size, height, 8, width, 8 - mask = mask.view( - mask.shape[0], - height // vae_scale_factor, - vae_scale_factor, - width // vae_scale_factor, - vae_scale_factor, - ) - # batch_size, 8, 8, height, width - mask = mask.permute(0, 2, 4, 1, 3) - # batch_size, 8*8, height, width - mask = mask.reshape( - mask.shape[0], - vae_scale_factor * vae_scale_factor, - height // vae_scale_factor, - width // vae_scale_factor, - ) + image_ids = self.model.prepare_latent_image_ids( + height // vae_scale_factor, + width // vae_scale_factor, + self.train_device, + self.model.train_dtype.torch_dtype() + ) - latent_mask = self.model.pack_latents(mask) - else: - conditioning_image = torch.zeros( - (1, 3, height, width), - dtype=self.model.train_dtype.torch_dtype(), - device=self.train_device, - ) - latent_conditioning_image = vae.encode(conditioning_image).latent_dist.mode() - latent_conditioning_image = (latent_conditioning_image - vae.config.shift_factor) \ - * vae.config.scaling_factor + shift = sample_config.override_shift \ + or self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) + latent_image = self.model.pack_latents(latent_image) - latent_conditioning_image = self.model.pack_latents(latent_conditioning_image) + # prepare timesteps + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=math.log(shift)) + timesteps = noise_scheduler.timesteps - latent_mask = torch.ones( - size=(1, (height // vae_scale_factor // 2) * (width // vae_scale_factor // 2), 256), - dtype=self.model.train_dtype.torch_dtype(), - device=self.train_device - ) + # denoising loop + extra_step_kwargs = {} + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator - # prepare prompt - self.model.materialize_only_text_encoders() + text_ids = torch.zeros(prompt_embedding.shape[1], 3, device=self.train_device) - prompt_embedding, pooled_prompt_embedding = self.model.encode_text( - text=prompt, - train_device=self.train_device, - text_encoder_1_layer_skip=text_encoder_1_layer_skip, - text_encoder_2_layer_skip=text_encoder_2_layer_skip, - text_encoder_2_sequence_length=text_encoder_2_sequence_length, - apply_attention_mask=transformer_attention_mask, - ) + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image]) + if is_inpainting: + latent_model_input = torch.concat( + [latent_model_input, latent_conditioning_image, latent_mask], -1 + ) + expanded_timestep = timestep.expand(latent_model_input.shape[0]) - # prepare latent image - latent_image = torch.randn( - size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32, - ) + # handle guidance + if transformer.config.guidance_embeds: + guidance = torch.tensor([cfg_scale], device=self.train_device) + guidance = guidance.expand(latent_model_input.shape[0]) + else: + guidance = None + + # predict the noise residual + noise_pred = transformer( + hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), + timestep=expanded_timestep / 1000, + guidance=guidance.to(dtype=self.model.train_dtype.torch_dtype()), + pooled_projections=pooled_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), + encoder_hidden_states=prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), + txt_ids=text_ids.to(dtype=self.model.train_dtype.torch_dtype()) if is_inpainting else text_ids, + img_ids=image_ids.to(dtype=self.model.train_dtype.torch_dtype()) if is_inpainting else image_ids, + joint_attention_kwargs=None, + return_dict=True + ).sample + + # compute the previous noisy sample x_t -> x_t-1 + latent_image = noise_scheduler.step( + noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs + )[0] + + on_update_progress(i + 1, len(timesteps)) + + latent_image = self.model.unpack_latents( + latent_image, + height // vae_scale_factor, + width // vae_scale_factor, + ) - image_ids = self.model.prepare_latent_image_ids( - height // vae_scale_factor, - width // vae_scale_factor, - self.train_device, - self.model.train_dtype.torch_dtype() - ) + return { + "latent_image": latent_image, + } - shift = self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) - latent_image = self.model.pack_latents(latent_image) - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=math.log(shift)) - timesteps = noise_scheduler.timesteps + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, + ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator + latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor + image = vae.decode(latents, return_dict=False)[0] - text_ids = torch.zeros(prompt_embedding.shape[1], 3, device=self.train_device) + do_denormalize = [True] * image.shape[0] + image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image]) - latent_model_input = torch.concat( - [latent_model_input, latent_conditioning_image, latent_mask], -1 - ) - expanded_timestep = timestep.expand(latent_model_input.shape[0]) - - # handle guidance - if transformer.config.guidance_embeds: - guidance = torch.tensor([cfg_scale], device=self.train_device) - guidance = guidance.expand(latent_model_input.shape[0]) - else: - guidance = None - - # predict the noise residual - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), - timestep=expanded_timestep / 1000, - guidance=guidance.to(dtype=self.model.train_dtype.torch_dtype()), - pooled_projections=pooled_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), - encoder_hidden_states=prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), - txt_ids=text_ids.to(dtype=self.model.train_dtype.torch_dtype()), - img_ids=image_ids.to(dtype=self.model.train_dtype.torch_dtype()), - joint_attention_kwargs=None, - return_dict=True - ).sample - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - latent_image = self.model.unpack_latents( - latent_image, - height // vae_scale_factor, - width // vae_scale_factor, - ) + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) - # decode - self.model.materialize_only("vae") + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + stages = [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)] + # conditioning (inpainting) model types VAE-encode a conditioning image first + if self.model_type.has_conditioning_image_input(): + stages.insert(0, ("encoding conditioning image", self.__cond_encode)) - latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor - image = vae.decode(latents, return_dict=False)[0] + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) - do_denormalize = [True] * image.shape[0] - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + stages, + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, + ) - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -390,47 +324,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - if self.model_type.has_conditioning_image_input(): - sampler_output = self.__sample_inpainting( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - sample_inpainting=sample_config.sample_inpainting, - base_image_path=sample_config.base_image_path, - mask_image_path=sample_config.mask_image_path, - text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, - text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, - text_encoder_2_sequence_length=sample_config.text_encoder_2_sequence_length, - transformer_attention_mask=sample_config.transformer_attention_mask, - on_update_progress=on_update_progress, - ) - else: - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, - text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, - text_encoder_2_sequence_length=sample_config.text_encoder_2_sequence_length, - transformer_attention_mask=sample_config.transformer_attention_mask, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/HiDreamSampler.py b/modules/modelSampler/HiDreamSampler.py index c5b30723b..bfa09dbac 100644 --- a/modules/modelSampler/HiDreamSampler.py +++ b/modules/modelSampler/HiDreamSampler.py @@ -10,13 +10,12 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.HI_DREAM_FULL) class HiDreamSampler(BaseModelSampler): @@ -34,124 +33,164 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - text_encoder_3_layer_skip: int = 0, - transformer_attention_mask: bool = False, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only_text_encoders() + text_encoder_3_prompt_embedding, text_encoder_4_prompt_embedding, pooled_prompt_embedding = \ + self.model.combine_text_encoder_output( + *self.model.encode_text( + text=sample_config.prompt, + train_device=self.train_device, + text_encoder_3_layer_skip=sample_config.text_encoder_3_layer_skip, + apply_attention_mask=sample_config.transformer_attention_mask, + )) + + negative_text_encoder_3_prompt_embedding, negative_text_encoder_4_prompt_embedding, negative_pooled_prompt_embedding = \ + self.model.combine_text_encoder_output( + *self.model.encode_text( + text=sample_config.negative_prompt, + train_device=self.train_device, + text_encoder_3_layer_skip=sample_config.text_encoder_3_layer_skip, + apply_attention_mask=sample_config.transformer_attention_mask, + )) + + combined_text_encoder_3_prompt_embedding = torch.cat( + [negative_text_encoder_3_prompt_embedding, text_encoder_3_prompt_embedding], dim=0) + combined_text_encoder_4_prompt_embedding = torch.cat( + [negative_text_encoder_4_prompt_embedding, text_encoder_4_prompt_embedding], dim=1) + combined_pooled_prompt_embedding = torch.cat( + [negative_pooled_prompt_embedding, pooled_prompt_embedding], dim=0) + + return { + "combined_text_encoder_3_prompt_embedding": combined_text_encoder_3_prompt_embedding, + "combined_text_encoder_4_prompt_embedding": combined_text_encoder_4_prompt_embedding, + "combined_pooled_prompt_embedding": combined_pooled_prompt_embedding, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + combined_text_encoder_3_prompt_embedding: torch.Tensor, + combined_text_encoder_4_prompt_embedding: torch.Tensor, + combined_pooled_prompt_embedding: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + vae_scale_factor = 8 + num_latent_channels = 16 + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + # prepare latent image + latent_image = torch.randn( + size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) + timesteps = noise_scheduler.timesteps + + # denoising loop + extra_step_kwargs = {} + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image] * 2) + expanded_timestep = timestep.expand(latent_model_input.shape[0]) + + with self.model.transformer_autocast_context: + # predict the noise residual + noise_pred = transformer( + hidden_states=latent_model_input.to(dtype=self.model.transformer_train_dtype.torch_dtype()), + timesteps=expanded_timestep, + encoder_hidden_states_t5=combined_text_encoder_3_prompt_embedding \ + .to(dtype=self.model.transformer_train_dtype.torch_dtype()), + encoder_hidden_states_llama3=combined_text_encoder_4_prompt_embedding \ + .to(dtype=self.model.transformer_train_dtype.torch_dtype()), + pooled_embeds=combined_pooled_prompt_embedding \ + .to(dtype=self.model.transformer_train_dtype.torch_dtype()), + return_dict=True + ).sample + noise_pred = -noise_pred + + # cfg + noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + + # compute the previous noisy sample x_t -> x_t-1 + latent_image = noise_scheduler.step( + noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs + )[0] + + on_update_progress(i + 1, len(timesteps)) + + return { + "latent_image": latent_image, + } + + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae + + latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor + image = vae.decode(latents, return_dict=False)[0] + + do_denormalize = [True] * image.shape[0] + image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) + + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) + + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) + with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - image_processor = self.pipeline.image_processor - transformer = self.pipeline.transformer - vae = self.pipeline.vae - vae_scale_factor = 8 - num_latent_channels = 16 - - # prepare prompt - self.model.materialize_only_text_encoders() - - text_encoder_3_prompt_embedding, text_encoder_4_prompt_embedding, pooled_prompt_embedding = \ - self.model.combine_text_encoder_output( - *self.model.encode_text( - text=prompt, - train_device=self.train_device, - text_encoder_3_layer_skip=text_encoder_3_layer_skip, - apply_attention_mask=transformer_attention_mask, - )) - - negative_text_encoder_3_prompt_embedding, negative_text_encoder_4_prompt_embedding, negative_pooled_prompt_embedding = \ - self.model.combine_text_encoder_output( - *self.model.encode_text( - text=negative_prompt, - train_device=self.train_device, - text_encoder_3_layer_skip=text_encoder_3_layer_skip, - apply_attention_mask=transformer_attention_mask, - )) - - combined_text_encoder_3_prompt_embedding = torch.cat( - [negative_text_encoder_3_prompt_embedding, text_encoder_3_prompt_embedding], dim=0) - combined_text_encoder_4_prompt_embedding = torch.cat( - [negative_text_encoder_4_prompt_embedding, text_encoder_4_prompt_embedding], dim=1) - combined_pooled_prompt_embedding = torch.cat( - [negative_pooled_prompt_embedding, pooled_prompt_embedding], dim=0) - - # prepare latent image - latent_image = torch.randn( - size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32, + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, ) - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) - timesteps = noise_scheduler.timesteps - - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator - - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * 2) - expanded_timestep = timestep.expand(latent_model_input.shape[0]) - - with self.model.transformer_autocast_context: - # predict the noise residual - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=self.model.transformer_train_dtype.torch_dtype()), - timesteps=expanded_timestep, - encoder_hidden_states_t5=combined_text_encoder_3_prompt_embedding \ - .to(dtype=self.model.transformer_train_dtype.torch_dtype()), - encoder_hidden_states_llama3=combined_text_encoder_4_prompt_embedding \ - .to(dtype=self.model.transformer_train_dtype.torch_dtype()), - pooled_embeds=combined_pooled_prompt_embedding \ - .to(dtype=self.model.transformer_train_dtype.torch_dtype()), - return_dict=True - ).sample - noise_pred = -noise_pred - - # cfg - noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - # decode - self.model.materialize_only("vae") - - latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor - image = vae.decode(latents, return_dict=False)[0] - - do_denormalize = [True] * image.shape[0] - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -162,24 +201,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - text_encoder_3_layer_skip=sample_config.text_encoder_3_layer_skip, - transformer_attention_mask=sample_config.transformer_attention_mask, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/HunyuanVideoSampler.py b/modules/modelSampler/HunyuanVideoSampler.py index 10b22bfc9..ed84be493 100644 --- a/modules/modelSampler/HunyuanVideoSampler.py +++ b/modules/modelSampler/HunyuanVideoSampler.py @@ -7,17 +7,14 @@ from modules.util import factory from modules.util.config.SampleConfig import SampleConfig from modules.util.enum.AudioFormat import AudioFormat -from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch -from PIL import Image -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.HUNYUAN_VIDEO) class HunyuanVideoSampler(BaseModelSampler): @@ -35,134 +32,156 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - num_frames: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - text_encoder_1_layer_skip: int = 0, - text_encoder_2_layer_skip: int = 0, - transformer_attention_mask: bool = False, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, - ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only_text_encoders() + prompt_embedding, pooled_prompt_embedding, prompt_attention_mask = self.model.encode_text( + text=sample_config.prompt, + train_device=self.train_device, + text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, + text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, + ) + + return { + "prompt_embedding": prompt_embedding, + "pooled_prompt_embedding": pooled_prompt_embedding, + "prompt_attention_mask": prompt_attention_mask, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + prompt_embedding: torch.Tensor, + pooled_prompt_embedding: torch.Tensor, + prompt_attention_mask: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + vae_temporal_scale_factor = 4 + vae_spacial_scale_factor = 8 + num_latent_channels = 16 + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + num_frames = self.quantize_resolution(sample_config.frames - 1, 4) + 1 + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + # prepare latent image + num_latent_frames = (num_frames - 1) // vae_temporal_scale_factor + 1 + latent_image = torch.randn( + size=( + 1, # batch size + num_latent_channels, + num_latent_frames, + height // vae_spacial_scale_factor, + width // vae_spacial_scale_factor + ), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + + # prepare timesteps + noise_scheduler.set_timesteps( + num_inference_steps=diffusion_steps, + device=self.train_device, + ) + timesteps = noise_scheduler.timesteps + + # denoising loop + extra_step_kwargs = {} + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image]) + expanded_timestep = timestep.expand(latent_model_input.shape[0]) + + # handle guidance + if transformer.config.guidance_embeds: + guidance = torch.tensor([cfg_scale * 1000.0], device=self.train_device) + guidance = guidance.expand(latent_model_input.shape[0]) else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - video_processor = self.pipeline.video_processor - transformer = self.pipeline.transformer - vae = self.pipeline.vae - vae_temporal_scale_factor = 4 - vae_spacial_scale_factor = 8 - num_latent_channels = 16 - - # prepare prompt - self.model.materialize_only_text_encoders() - - prompt_embedding, pooled_prompt_embedding, prompt_attention_mask = self.model.encode_text( - text=prompt, - train_device=self.train_device, - text_encoder_1_layer_skip=text_encoder_1_layer_skip, - text_encoder_2_layer_skip=text_encoder_2_layer_skip, - ) + guidance = None + + with self.model.transformer_autocast_context: + # predict the noise residual + noise_pred = transformer( + hidden_states=latent_model_input.to(dtype=self.model.transformer_train_dtype.torch_dtype()), + timestep=expanded_timestep, + guidance=guidance.to(dtype=self.model.transformer_train_dtype.torch_dtype()), + pooled_projections=pooled_prompt_embedding.to(dtype=self.model.transformer_train_dtype.torch_dtype()), + encoder_hidden_states=prompt_embedding.to(dtype=self.model.transformer_train_dtype.torch_dtype()), + encoder_attention_mask=prompt_attention_mask.to(dtype=self.model.transformer_train_dtype.torch_dtype()), + return_dict=True + ).sample + + # compute the previous noisy sample x_t -> x_t-1 + latent_image = noise_scheduler.step( + noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs + )[0] + + on_update_progress(i + 1, len(timesteps)) + + return { + "latent_image": latent_image, + } - # prepare latent image - num_latent_frames = (num_frames - 1) // vae_temporal_scale_factor + 1 - latent_image = torch.randn( - size=( - 1, # batch size - num_latent_channels, - num_latent_frames, - height // vae_spacial_scale_factor, - width // vae_spacial_scale_factor - ), - generator=generator, - device=self.train_device, - dtype=torch.float32, + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, + ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + video_processor = self.pipeline.video_processor + vae = self.pipeline.vae + + latents = latent_image / vae.config.scaling_factor + image = vae.decode(latents, return_dict=False)[0] + + image = video_processor.postprocess(image, output_type='pt') + + # postprocess keeps channels ahead of frames, so swap them to [B, F, C, H, W] + return self.build_video_sampler_output(image.permute(0, 2, 1, 3, 4)) + + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) + + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, ) - # prepare timesteps - noise_scheduler.set_timesteps( - num_inference_steps=diffusion_steps, - device=self.train_device, + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, + fps=self.model.NATIVE_FPS, ) - timesteps = noise_scheduler.timesteps - - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator - - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image]) - expanded_timestep = timestep.expand(latent_model_input.shape[0]) - - # handle guidance - if transformer.config.guidance_embeds: - guidance = torch.tensor([cfg_scale * 1000.0], device=self.train_device) - guidance = guidance.expand(latent_model_input.shape[0]) - else: - guidance = None - - with self.model.transformer_autocast_context: - # predict the noise residual - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=self.model.transformer_train_dtype.torch_dtype()), - timestep=expanded_timestep, - guidance=guidance.to(dtype=self.model.transformer_train_dtype.torch_dtype()), - pooled_projections=pooled_prompt_embedding.to(dtype=self.model.transformer_train_dtype.torch_dtype()), - encoder_hidden_states=prompt_embedding.to(dtype=self.model.transformer_train_dtype.torch_dtype()), - encoder_attention_mask=prompt_attention_mask.to(dtype=self.model.transformer_train_dtype.torch_dtype()), - return_dict=True - ).sample - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - # decode - self.model.materialize_only("vae") - - latents = latent_image / vae.config.scaling_factor - image = vae.decode(latents, return_dict=False)[0] - - image = video_processor.postprocess(image, output_type='pt') - - is_image = image.shape[2] == 1 - if is_image: - image = image.view((image.shape[0], image.shape[1], image.shape[3], image.shape[4])) - image = image.cpu().permute(0, 2, 3, 1).float().numpy() - image = (image * 255).round().astype("uint8") - image = Image.fromarray(image[0]) - - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image, - ) - else: - image = image.cpu().permute(0, 2, 3, 4, 1).float() - image = (image.clamp(0, 1) * 255).round().to(dtype=torch.int8) - image = image[0] - return ModelSamplerOutput( - file_type=FileType.VIDEO, - data=image, - ) + return sampler_outputs def sample( self, @@ -174,29 +193,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - num_frames=self.quantize_resolution(sample_config.frames - 1, 4) + 1, - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, - text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, - transformer_attention_mask=sample_config.transformer_attention_mask, - on_update_progress=on_update_progress, - ) - - fps = self.model.NATIVE_FPS - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - fps=fps, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/IdeogramSampler.py b/modules/modelSampler/IdeogramSampler.py index cb253db0e..52541c975 100644 --- a/modules/modelSampler/IdeogramSampler.py +++ b/modules/modelSampler/IdeogramSampler.py @@ -9,8 +9,9 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch @@ -18,7 +19,6 @@ import numpy as np from PIL import Image as PILImage -from tqdm import tqdm @factory.register(BaseModelSampler, ModelType.IDEOGRAM_4) @@ -37,172 +37,236 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, - ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - vae = self.pipeline.vae - transformer = self.pipeline.transformer - dtype = self.model.train_dtype.torch_dtype() - - # Ideogram uses asymmetric (dual-network) CFG: the negative branch is normally the unconditional_transformer - # run on the image tokens with zeroed text features, NOT a negative-prompt encode. negative_prompt is - # unused. If the unconditional transformer was not loaded, fall back to encoding an empty ("") prompt and - # running it through the conditional transformer instead, like standard CFG. - use_cfg = cfg_scale > 1.0 - use_unconditional_transformer = use_cfg and self.model.unconditional_transformer is not None - use_empty_prompt_negative = use_cfg and self.model.unconditional_transformer is None - - vae_scale_factor = 8 - patch_size = 2 - latent_dim = transformer.config.in_channels - grid_h = height // (vae_scale_factor * patch_size) - grid_w = width // (vae_scale_factor * patch_size) - num_image_tokens = grid_h * grid_w - - # build the packed [text][image] conditioning for a single text encode. Padding positions are masked - # out by segment_ids/indicator, so packing to the actual text length matches the 2048-pad pipeline. - def pack_conditioning(text_features: torch.Tensor, text_lengths: torch.Tensor) -> tuple: - max_text_tokens = text_features.shape[1] - position_ids, segment_ids, indicator = self.model.prepare_packed_ids( - text_lengths, grid_h, grid_w, max_text_tokens, self.train_device, - ) - llm_features = self.model.pack_llm_features(text_features, num_image_tokens).to(dtype) - text_z_padding = torch.zeros( - text_features.shape[0], max_text_tokens, latent_dim, dtype=dtype, device=self.train_device, - ) - return max_text_tokens, position_ids, segment_ids, indicator, llm_features, text_z_padding - - # encode text (conditional branch, and the empty-prompt negative branch if needed) - self.model.materialize_only("text_encoder") - text_features, text_lengths = self.model.encode_text( + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("text_encoder") + cfg_scale = sample_config.cfg_scale + + # Ideogram uses asymmetric (dual-network) CFG: the negative branch is normally the unconditional_transformer + # run on the image tokens with zeroed text features, NOT a negative-prompt encode. negative_prompt is + # unused. If the unconditional transformer was not loaded, fall back to encoding an empty ("") prompt and + # running it through the conditional transformer instead, like standard CFG. + use_cfg = cfg_scale > 1.0 + use_unconditional_transformer = use_cfg and self.model.unconditional_transformer is not None + use_empty_prompt_negative = use_cfg and self.model.unconditional_transformer is None + + # encode text (conditional branch, and the empty-prompt negative branch if needed) + text_features, text_lengths = self.model.encode_text( + train_device=self.train_device, + text=sample_config.prompt, + ) + + neg_text_features = None + neg_text_lengths = None + if use_empty_prompt_negative: + neg_text_features, neg_text_lengths = self.model.encode_text( train_device=self.train_device, - text=prompt, - ) - max_text_tokens, position_ids, segment_ids, indicator, llm_features, text_z_padding = pack_conditioning( - text_features, text_lengths, + text="", ) - if use_empty_prompt_negative: - neg_text_features, neg_text_lengths = self.model.encode_text( - train_device=self.train_device, - text="", - ) - ( - max_neg_text_tokens, neg_position_ids, neg_segment_ids, neg_indicator, neg_llm_features, - neg_text_z_padding, - ) = pack_conditioning(neg_text_features, neg_text_lengths) - del neg_text_features + return { + "use_cfg": use_cfg, + "use_unconditional_transformer": use_unconditional_transformer, + "use_empty_prompt_negative": use_empty_prompt_negative, + "text_features": text_features, + "text_lengths": text_lengths, + "neg_text_features": neg_text_features, + "neg_text_lengths": neg_text_lengths, + } - if use_unconditional_transformer: - # unconditional (image-only) branch: zeroed text features over the image-region slices of the layout - neg_position_ids = position_ids[:, max_text_tokens:] - neg_segment_ids = segment_ids[:, max_text_tokens:] - neg_indicator = indicator[:, max_text_tokens:] - neg_llm_features = torch.zeros( - text_features.shape[0], num_image_tokens, text_features.shape[-1], - dtype=dtype, device=self.train_device, - ) - - # packed (B, num_image_tokens, latent_dim) noise - latent_image = torch.randn( - size=(text_features.shape[0], num_image_tokens, latent_dim), - generator=generator, device=self.train_device, dtype=torch.float32, + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + use_cfg: bool, + use_unconditional_transformer: bool, + use_empty_prompt_negative: bool, + text_features: torch.Tensor, + text_lengths: torch.Tensor, + neg_text_features: torch.Tensor | None, + neg_text_lengths: torch.Tensor | None, + on_update_progress: Callable[[int, int], None], + ) -> dict: + # the unconditional transformer is only needed for asymmetric CFG's negative branch + transformer_parts = ("transformer", "unconditional_transformer") if use_unconditional_transformer else ("transformer",) + self.model.materialize_only(*transformer_parts) + transformer = self.pipeline.transformer + dtype = self.model.train_dtype.torch_dtype() + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + vae_scale_factor = 8 + patch_size = 2 + latent_dim = transformer.config.in_channels + grid_h = height // (vae_scale_factor * patch_size) + grid_w = width // (vae_scale_factor * patch_size) + num_image_tokens = grid_h * grid_w + + # build the packed [text][image] conditioning for a single text encode. Padding positions are masked + # out by segment_ids/indicator, so packing to the actual text length matches the 2048-pad pipeline. + def pack_conditioning(text_features: torch.Tensor, text_lengths: torch.Tensor) -> tuple: + max_text_tokens = text_features.shape[1] + position_ids, segment_ids, indicator = self.model.prepare_packed_ids( + text_lengths, grid_h, grid_w, max_text_tokens, self.train_device, ) + llm_features = self.model.pack_llm_features(text_features, num_image_tokens).to(dtype) + text_z_padding = torch.zeros( + text_features.shape[0], max_text_tokens, latent_dim, dtype=dtype, device=self.train_device, + ) + return max_text_tokens, position_ids, segment_ids, indicator, llm_features, text_z_padding + + max_text_tokens, position_ids, segment_ids, indicator, llm_features, text_z_padding = pack_conditioning( + text_features, text_lengths, + ) - # free before the denoising loop; closes the gap on the OOM observed in llm_cond_norm's fp32 variance - # upcast of neg_llm_features - del text_features + if use_empty_prompt_negative: + ( + max_neg_text_tokens, neg_position_ids, neg_segment_ids, neg_indicator, neg_llm_features, + neg_text_z_padding, + ) = pack_conditioning(neg_text_features, neg_text_lengths) + del neg_text_features + + if use_unconditional_transformer: + # unconditional (image-only) branch: zeroed text features over the image-region slices of the layout + neg_position_ids = position_ids[:, max_text_tokens:] + neg_segment_ids = segment_ids[:, max_text_tokens:] + neg_indicator = indicator[:, max_text_tokens:] + neg_llm_features = torch.zeros( + text_features.shape[0], num_image_tokens, text_features.shape[-1], + dtype=dtype, device=self.train_device, + ) - # resolution-aware logit-normal Euler schedule (pipeline overrides the scheduler's default sigmas) - schedule_mu = _resolution_aware_mu(height=height, width=width, base_mu=0.0) - sigmas = _logit_normal_sigmas(diffusion_steps, schedule_mu, std=1.5, device=self.train_device) - noise_scheduler.set_timesteps(sigmas=sigmas.tolist(), device=self.train_device) - timesteps = noise_scheduler.timesteps - num_train_timesteps = noise_scheduler.config.num_train_timesteps + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) - transformer_parts = ("transformer", "unconditional_transformer") if use_unconditional_transformer else ("transformer",) - self.model.materialize_only(*transformer_parts) + # packed (B, num_image_tokens, latent_dim) noise + latent_image = torch.randn( + size=(text_features.shape[0], num_image_tokens, latent_dim), + generator=generator, device=self.train_device, dtype=torch.float32, + ) - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - # scheduler stores num_train_timesteps-scaled timesteps; convert back to model time (0=noise, 1=data) - t_model = (1.0 - timestep.float() / num_train_timesteps).expand(latent_image.shape[0]) + # free before the denoising loop; closes the gap on the OOM observed in llm_cond_norm's fp32 variance + # upcast of neg_llm_features + del text_features + + # resolution-aware logit-normal Euler schedule (pipeline overrides the scheduler's default sigmas) + schedule_mu = _resolution_aware_mu(height=height, width=width, base_mu=0.0) + sigmas = _logit_normal_sigmas(diffusion_steps, schedule_mu, std=1.5, device=self.train_device) + noise_scheduler.set_timesteps(sigmas=sigmas.tolist(), device=self.train_device) + timesteps = noise_scheduler.timesteps + num_train_timesteps = noise_scheduler.config.num_train_timesteps + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + # scheduler stores num_train_timesteps-scaled timesteps; convert back to model time (0=noise, 1=data) + t_model = (1.0 - timestep.float() / num_train_timesteps).expand(latent_image.shape[0]) + + pos_z = torch.cat([text_z_padding, latent_image.to(dtype)], dim=1) + pos_out = transformer( + hidden_states=pos_z, + timestep=t_model, + encoder_hidden_states=llm_features, + position_ids=position_ids, + segment_ids=segment_ids, + indicator=indicator, + return_dict=False, + )[0] + pos_v = pos_out[:, max_text_tokens:].float() - pos_z = torch.cat([text_z_padding, latent_image.to(dtype)], dim=1) - pos_out = transformer( - hidden_states=pos_z, + if use_unconditional_transformer: + neg_v = self.model.unconditional_transformer( + hidden_states=latent_image.to(dtype), timestep=t_model, - encoder_hidden_states=llm_features, - position_ids=position_ids, - segment_ids=segment_ids, - indicator=indicator, + encoder_hidden_states=neg_llm_features, + position_ids=neg_position_ids, + segment_ids=neg_segment_ids, + indicator=neg_indicator, + return_dict=False, + )[0].float() + elif use_empty_prompt_negative: + neg_z = torch.cat([neg_text_z_padding, latent_image.to(dtype)], dim=1) + neg_out = transformer( + hidden_states=neg_z, + timestep=t_model, + encoder_hidden_states=neg_llm_features, + position_ids=neg_position_ids, + segment_ids=neg_segment_ids, + indicator=neg_indicator, return_dict=False, )[0] - pos_v = pos_out[:, max_text_tokens:].float() - - if use_unconditional_transformer: - neg_v = self.model.unconditional_transformer( - hidden_states=latent_image.to(dtype), - timestep=t_model, - encoder_hidden_states=neg_llm_features, - position_ids=neg_position_ids, - segment_ids=neg_segment_ids, - indicator=neg_indicator, - return_dict=False, - )[0].float() - elif use_empty_prompt_negative: - neg_z = torch.cat([neg_text_z_padding, latent_image.to(dtype)], dim=1) - neg_out = transformer( - hidden_states=neg_z, - timestep=t_model, - encoder_hidden_states=neg_llm_features, - position_ids=neg_position_ids, - segment_ids=neg_segment_ids, - indicator=neg_indicator, - return_dict=False, - )[0] - neg_v = neg_out[:, max_neg_text_tokens:].float() - - v = neg_v + cfg_scale * (pos_v - neg_v) if use_cfg else pos_v - - latent_image = noise_scheduler.step(-v, timestep, latent_image, return_dict=False)[0] - - on_update_progress(i + 1, len(timesteps)) - - self.model.materialize_only("vae") - - # bn-denormalize the packed latents and unpatchify back to (B, C, H, W) before VAE decode - latents = self.model.unscale_latents(latent_image) - latents = self.model.unpatchify_latents(latents, grid_h, grid_w) - - image = vae.decode(latents.to(vae.dtype), return_dict=False)[0] - # no VaeImageProcessor — match the pipeline's manual postprocess - image = (image.clamp(-1, 1) + 1) / 2 - image = image.cpu().permute(0, 2, 3, 1).float().numpy() - image = [PILImage.fromarray((img * 255).astype(np.uint8)) for img in image] - - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + neg_v = neg_out[:, max_neg_text_tokens:].float() + + v = neg_v + cfg_scale * (pos_v - neg_v) if use_cfg else pos_v + + latent_image = noise_scheduler.step(-v, timestep, latent_image, return_dict=False)[0] + + on_update_progress(i + 1, len(timesteps)) + + return { + "latent_image": latent_image, + "grid_h": grid_h, + "grid_w": grid_w, + } + + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, + grid_h: int, + grid_w: int, + ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + vae = self.pipeline.vae + + # bn-denormalize the packed latents and unpatchify back to (B, C, H, W) before VAE decode + latents = self.model.unscale_latents(latent_image) + latents = self.model.unpatchify_latents(latents, grid_h, grid_w) + + image = vae.decode(latents.to(vae.dtype), return_dict=False)[0] + # no VaeImageProcessor — match the pipeline's manual postprocess + image = (image.clamp(-1, 1) + 1) / 2 + image = image.cpu().permute(0, 2, 3, 1).float().numpy() + image = [PILImage.fromarray((img * 255).astype(np.uint8)) for img in image] + + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) + + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) + + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, ) + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, + ) + + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -213,22 +277,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/Krea2Sampler.py b/modules/modelSampler/Krea2Sampler.py index b83205a93..890a68df4 100644 --- a/modules/modelSampler/Krea2Sampler.py +++ b/modules/modelSampler/Krea2Sampler.py @@ -11,15 +11,14 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch from diffusers import Krea2Pipeline -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.KREA_2) class Krea2Sampler(BaseModelSampler): @@ -37,109 +36,152 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("text_encoder") + batch_size = 2 if sample_config.cfg_scale > 1.0 else 1 + combined_prompt_embedding, text_attention_mask = self.model.encode_text( + text=[sample_config.prompt, sample_config.negative_prompt] if sample_config.cfg_scale > 1.0 else sample_config.prompt, + batch_size=batch_size, + train_device=self.train_device, + ) + + return { + "batch_size": batch_size, + "combined_prompt_embedding": combined_prompt_embedding, + "text_attention_mask": text_attention_mask, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + batch_size: int, + combined_prompt_embedding: torch.Tensor, + text_attention_mask: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + vae_scale_factor = 8 + num_latent_channels = 16 + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + # prepare latent image + latent_image = torch.randn( + size=(1, num_latent_channels, 1, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + + shift = sample_config.override_shift \ + or self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) + latent_image = self.model.pack_latents(latent_image) + + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=math.log(shift)) + timesteps = noise_scheduler.timesteps + + # build position ids: text tokens at origin, image tokens at latent-grid coords + text_seq_len = combined_prompt_embedding.shape[1] + grid_height = height // vae_scale_factor // 2 # patch_size = 2 + grid_width = width // vae_scale_factor // 2 + position_ids = Krea2Pipeline.prepare_position_ids(text_seq_len, grid_height, grid_width, self.train_device) + + # denoising loop + extra_step_kwargs = {} + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image] * batch_size) + expanded_timestep = timestep.expand(batch_size) + noise_pred = transformer( + hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), + encoder_hidden_states=combined_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), + timestep=expanded_timestep / 1000, + position_ids=position_ids, + encoder_attention_mask=text_attention_mask, + return_dict=False, + )[0] + + if cfg_scale > 1.0: + noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + + latent_image = noise_scheduler.step(noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs)[0] + + on_update_progress(i + 1, len(timesteps)) + + latent_image = self.model.unpack_latents( + latent_image, + height // vae_scale_factor, + width // vae_scale_factor, + ) + + return { + "latent_image": latent_image, + } + + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - image_processor = self.pipeline.image_processor - - transformer = self.pipeline.transformer - vae = self.pipeline.vae - vae_scale_factor = 8 - num_latent_channels = 16 - - # prepare prompt - self.model.materialize_only("text_encoder") - - batch_size = 2 if cfg_scale > 1.0 else 1 - combined_prompt_embedding, text_attention_mask = self.model.encode_text( - text=[prompt, negative_prompt] if cfg_scale > 1.0 else prompt, - batch_size=batch_size, - train_device=self.train_device, - ) + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae - # prepare latent image - latent_image = torch.randn( - size=(1, num_latent_channels, 1, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32, - ) + latents = self.model.unscale_latents(latent_image) + image = vae.decode(latents, return_dict=False)[0].squeeze(-3) - shift = self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) - latent_image = self.model.pack_latents(latent_image) - - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=math.log(shift)) - timesteps = noise_scheduler.timesteps - - # build position ids: text tokens at origin, image tokens at latent-grid coords - text_seq_len = combined_prompt_embedding.shape[1] - grid_height = height // vae_scale_factor // 2 # patch_size = 2 - grid_width = width // vae_scale_factor // 2 - position_ids = Krea2Pipeline.prepare_position_ids(text_seq_len, grid_height, grid_width, self.train_device) - - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator - - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * batch_size) - expanded_timestep = timestep.expand(batch_size) - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), - encoder_hidden_states=combined_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), - timestep=expanded_timestep / 1000, - position_ids=position_ids, - encoder_attention_mask=text_attention_mask, - return_dict=False, - )[0] - - if cfg_scale > 1.0: - noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - latent_image = noise_scheduler.step(noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs)[0] - - on_update_progress(i + 1, len(timesteps)) - - latent_image = self.model.unpack_latents( - latent_image, - height // vae_scale_factor, - width // vae_scale_factor, - ) + do_denormalize = [True] * image.shape[0] + image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.materialize_only("vae") + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) - latents = self.model.unscale_latents(latent_image) - image = vae.decode(latents, return_dict=False)[0].squeeze(-3) + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) - do_denormalize = [True] * image.shape[0] - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, + ) - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -150,22 +192,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/LTXSampler.py b/modules/modelSampler/LTXSampler.py new file mode 100644 index 000000000..3e4878dd4 --- /dev/null +++ b/modules/modelSampler/LTXSampler.py @@ -0,0 +1,432 @@ +import copy +from collections.abc import Callable +from contextlib import nullcontext + +from modules.model.LTXModel import LTXModel +from modules.modelSampler.BaseModelSampler import BaseModelSampler, ModelSamplerOutput +from modules.util import factory +from modules.util.config.SampleConfig import SampleConfig +from modules.util.enum.AudioFormat import AudioFormat +from modules.util.enum.ImageFormat import ImageFormat +from modules.util.enum.ModelType import ModelType +from modules.util.enum.SamplingMethod import SamplingMethod +from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm + +import torch + +import numpy as np + +# The distilled expert's own fixed sigma schedule (Lightricks' DISTILLED_SIGMA_VALUES). Literal values, never +# run through calculate_timestep_shift: the reference passes them straight to the stage, bypassing the +# token-count shift. +DISTILLED_SIGMAS = (1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0) + +# The tail of that schedule (Lightricks' STAGE_2_DISTILLED_SIGMA_VALUES), what the expert runs when it takes +# over a trajectory already denoised down to its first value - which doubles as the hand-off point. +LOW_NOISE_SIGMAS = DISTILLED_SIGMAS[5:] + +# The sigma the full-model schedule ends on in both generations. 2.3's scheduler carries it as shift_terminal, +# 2.5's ships null but only because that scheduler is configured for the distilled DiT, which walks its own +# fixed sigma list instead. +SHIFT_TERMINAL = 0.1 + + +@factory.register(BaseModelSampler, ModelType.LTX_2) +class LTXSampler(BaseModelSampler): + def __init__( + self, + train_device: torch.device, + temp_device: torch.device, + model: LTXModel, + model_type: ModelType, + ): + super().__init__(train_device, temp_device) + + self.model = model + self.model_type = model_type + self.pipeline = model.create_pipeline() + + @torch.no_grad() + def __text_encode( + self, + sample_config: SampleConfig, + ) -> dict: + do_cfg = sample_config.cfg_scale > 1.0 + + self.model.materialize_only("text_encoder") + text_encoder_outputs, tokens_mask = self.model.encode_text_encoder( + text=[sample_config.negative_prompt, sample_config.prompt] if do_cfg else sample_config.prompt, + ) + + # park the TE outputs in CPU RAM: this stage runs over the whole batch before the next one starts, and + # the stacked outputs are large enough to fill VRAM at a higher batch size + text_encoder_outputs = [output.to("cpu") for output in text_encoder_outputs] + return { + "text_encoder_outputs": text_encoder_outputs, + "tokens_mask": tokens_mask.to("cpu"), + } + + @torch.no_grad() + def __connect( + self, + text_encoder_outputs: list[torch.Tensor], + tokens_mask: torch.Tensor, + ) -> dict: + self.model.materialize_only("connectors") + connector_prompt_embeds, connector_audio_prompt_embeds = self.model.encode_connectors( + text_encoder_outputs, tokens_mask, self.train_device, + ) + + return { + "connector_prompt_embeds": connector_prompt_embeds, + "connector_audio_prompt_embeds": connector_audio_prompt_embeds, + } + + @torch.no_grad() + def __run_denoise_loop( + self, + transformer, + autocast_context, + dtype: torch.dtype, + timesteps: torch.Tensor, + noise_scheduler, + audio_noise_scheduler, + latent_video: torch.Tensor, + latent_audio: torch.Tensor, + geometry: dict, + conditioning: dict, + do_cfg: bool, + cfg_scale: float, + on_update_progress: Callable[[int, int], None], + ) -> tuple[torch.Tensor, torch.Tensor]: + # run once per expert - both share an architecture, so only the module, its autocast context and its + # compute dtype differ + # with CFG the whole input side is a [negative, positive] batch: the embeds already arrive stacked that + # way, the latents and the coords are duplicated to match + video_coords = torch.cat([geometry["video_coords"]] * 2) if do_cfg else geometry["video_coords"] + audio_coords = torch.cat([geometry["audio_coords"]] * 2) if do_cfg else geometry["audio_coords"] + + for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): + latent_video_input = torch.cat([latent_video] * 2) if do_cfg else latent_video + latent_audio_input = torch.cat([latent_audio] * 2) if do_cfg else latent_audio + expanded_timestep = timestep.expand(latent_video_input.shape[0]) + + with autocast_context: + noise_pred_video, noise_pred_audio = transformer( + hidden_states=latent_video_input.to(dtype=dtype), + audio_hidden_states=latent_audio_input.to(dtype=dtype), + encoder_hidden_states=conditioning["connector_prompt_embeds"].to(dtype=dtype), + audio_encoder_hidden_states=conditioning["connector_audio_prompt_embeds"].to(dtype=dtype), + timestep=expanded_timestep, + sigma=expanded_timestep, + encoder_attention_mask=None, + audio_encoder_attention_mask=None, + num_frames=geometry["num_latent_frames"], + height=geometry["latent_height"], + width=geometry["latent_width"], + fps=geometry["frame_rate"], + audio_num_frames=geometry["audio_num_frames"], + video_coords=video_coords, + audio_coords=audio_coords, + return_dict=False, + ) + + noise_pred_video = noise_pred_video.float() + noise_pred_audio = noise_pred_audio.float() + + if do_cfg: + noise_pred_video_uncond, noise_pred_video_cond = noise_pred_video.chunk(2) + noise_pred_video = noise_pred_video_uncond \ + + cfg_scale * (noise_pred_video_cond - noise_pred_video_uncond) + + noise_pred_audio_uncond, noise_pred_audio_cond = noise_pred_audio.chunk(2) + noise_pred_audio = noise_pred_audio_uncond \ + + cfg_scale * (noise_pred_audio_cond - noise_pred_audio_uncond) + + latent_video = noise_scheduler.step(noise_pred_video, timestep, latent_video, return_dict=False)[0] + # audio branch is never decoded (video-only scope) - only stepped so the transformer keeps seeing a + # properly noised audio trajectory, matching what the model saw during audio-visual training + latent_audio = audio_noise_scheduler.step( + noise_pred_audio, timestep, latent_audio, return_dict=False, + )[0] + + on_update_progress(i + 1, len(timesteps)) + + return latent_video, latent_audio + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + connector_prompt_embeds: torch.Tensor, + connector_audio_prompt_embeds: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + do_cfg = cfg_scale > 1.0 + + num_frames = self.quantize_resolution(sample_config.frames - 1, 8) + 1 + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + # LTX-2 denoises the audio branch alongside the video branch even in video-only use, so it needs its + # own scheduler instance (mirrors LTX2Pipeline.__call__'s `audio_scheduler`) + audio_noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + transformer = self.pipeline.transformer + audio_vae = self.pipeline.audio_vae + + frame_rate = self.model.NATIVE_FPS + vae_spatial_scale_factor = self.pipeline.vae_spatial_compression_ratio + + num_latent_frames = (num_frames - 1) // self.pipeline.vae_temporal_compression_ratio + 1 + latent_height = self.quantize_resolution(sample_config.height, 32) // vae_spatial_scale_factor + latent_width = self.quantize_resolution(sample_config.width, 32) // vae_spatial_scale_factor + latent_video = torch.randn( + size=(1, transformer.config.in_channels, num_latent_frames, latent_height, latent_width), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + latent_video = self.model.pack_latents( + latent_video, + self.pipeline.transformer_spatial_patch_size, self.pipeline.transformer_temporal_patch_size, + ) + + # prepare audio latents: fed real noise only because the transformer takes audio inputs positionally, + # denoised in lockstep to mirror the reference pipeline but never decoded + audio_num_frames = round( + num_frames / frame_rate + * self.pipeline.audio_sampling_rate + / self.pipeline.audio_hop_length + / float(self.pipeline.audio_vae_temporal_compression_ratio) + ) + latent_audio = torch.randn( + size=(1, audio_vae.config.latent_channels, audio_num_frames, + audio_vae.config.mel_bins // self.pipeline.audio_vae_mel_compression_ratio), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + latent_audio = self.model.pack_audio_latents(latent_audio) + + if sample_config.override_shift: + shift = sample_config.override_shift + else: + shift = self.model.calculate_timestep_shift(num_latent_frames, latent_height, latent_width) + + # the reference's schedule: a 1/steps grid, bent towards the noisy end by the token-count shift, then + # stretched so its last sigma lands on SHIFT_TERMINAL + sigmas = np.linspace(1.0, 1.0 / diffusion_steps, diffusion_steps, dtype=np.float32) + sigmas = shift / (shift + (1.0 / sigmas - 1.0)) + sigmas = 1.0 - (1.0 - sigmas) * (1.0 - SHIFT_TERMINAL) / (1.0 - sigmas[-1]) + + # one schedule spans both stages, split at expert_start: this stage walks the steps before it, the + # expert stage continues on the same schedulers from there + if self.model.low_noise_transformer is not None: + if sample_config.sampling_method == SamplingMethod.HANDOFF_LOW_NOISE: + # the expert takes over at its own first sigma, so this stage keeps the steps above it + above_handoff = sigmas[sigmas > LOW_NOISE_SIGMAS[0]] + expert_start = len(above_handoff) + sigmas = np.append(above_handoff, LOW_NOISE_SIGMAS[:-1]) + elif sample_config.sampling_method == SamplingMethod.DISTILLED: + # the expert samples from noise on its own schedule, so it walks the whole trajectory + expert_start = 0 + sigmas = np.asarray(DISTILLED_SIGMAS[:-1], dtype=np.float32) + elif sample_config.sampling_method == SamplingMethod.STANDARD: + expert_start = len(sigmas) + else: + raise NotImplementedError(f"unsupported sampling method {sample_config.sampling_method}") + else: + # nothing to hand over to, so a method that wants the expert runs the trained transformer alone + expert_start = len(sigmas) + + # the schedule is installed verbatim, so the scheduler only does the Euler integration + for scheduler in (noise_scheduler, audio_noise_scheduler): + scheduler.register_to_config(use_dynamic_shifting=False, shift=1.0, shift_terminal=None) + scheduler.set_timesteps(sigmas=sigmas.astype(np.float32), device=self.train_device) + + timesteps = noise_scheduler.timesteps[:expert_start] + + # the steps walked here are not the configured count: a hand-off drops everything below its sigma + # to the expert, and a larger shift pushes more of the schedule above it, so more steps stay here. + # Printed in timestep units (sigma * 1000) so it lines up with the schedule. + tqdm.write(f"[sigmas] denoise: shift {shift:.3f}, {diffusion_steps} configured -> " + f"{len(timesteps)} steps, timesteps {[round(float(t), 1) for t in timesteps]}") + + video_coords = transformer.rope.prepare_video_coords( + latent_video.shape[0], num_latent_frames, latent_height, latent_width, latent_video.device, + fps=frame_rate, + ) + audio_coords = transformer.audio_rope.prepare_audio_coords( + latent_audio.shape[0], audio_num_frames, latent_audio.device, + ) + + geometry = { + "num_latent_frames": num_latent_frames, + "latent_height": latent_height, + "latent_width": latent_width, + "audio_num_frames": audio_num_frames, + "video_coords": video_coords, + "audio_coords": audio_coords, + "frame_rate": frame_rate, + } + conditioning = { + "connector_prompt_embeds": connector_prompt_embeds, + "connector_audio_prompt_embeds": connector_audio_prompt_embeds, + } + if len(timesteps) > 0: + # skipped in DISTILLED, where this stage takes no step - the 19B trained transformer then never + # reaches the train device at all + self.model.materialize_only("transformer") + latent_video, latent_audio = self.__run_denoise_loop( + transformer, self.model.transformer_autocast_context, + self.model.transformer_train_dtype.torch_dtype(), timesteps, + noise_scheduler, audio_noise_scheduler, latent_video, latent_audio, + geometry, conditioning, do_cfg, cfg_scale, on_update_progress, + ) + + if do_cfg: + # __text_encode stacks the embeds as [negative, positive]; the low noise expert stage runs + # unguided, so it is handed the conditional half alone + conditioning = {name: tensor.chunk(2)[1] for name, tensor in conditioning.items()} + + return { + "latent_video": latent_video, + "latent_audio": latent_audio, + "num_latent_frames": num_latent_frames, + "latent_height": latent_height, + "latent_width": latent_width, + "geometry": geometry, + "conditioning": conditioning, + "noise_scheduler": noise_scheduler, + "audio_noise_scheduler": audio_noise_scheduler, + "expert_timesteps": noise_scheduler.timesteps[expert_start:], + } + + @torch.no_grad() + def __denoise_low_noise( + self, + latent_video: torch.Tensor, + latent_audio: torch.Tensor, + noise_scheduler, + audio_noise_scheduler, + expert_timesteps: torch.Tensor, + geometry: dict, + conditioning: dict, + on_update_progress: Callable[[int, int], None], + ) -> dict: + # the split left nothing over, so the trained transformer already walked the whole schedule + if len(expert_timesteps) == 0: + return {"latent_video": latent_video} + + # printed next to the other stage's line to make the whole trajectory visible at once + tqdm.write(f"[sigmas] low noise expert: {len(expert_timesteps)} steps, " + f"timesteps {[round(float(t), 1) for t in expert_timesteps]}") + + self.model.materialize_only("low_noise_transformer") + dtype = self.model.low_noise_transformer_train_dtype.torch_dtype() + + if self.model.transformer_lora is None: + lora_context = nullcontext() + else: + # the same LoRA weights, rebound onto the expert for this stage. A LoRA travels with its stem, so + # materializing the expert evicted it along with the transformer - bring it back on its own, or + # the expert's forward mixes cuda and cpu tensors. The next materialize of either part takes it + # along again. + self.model.transformer_lora.to(device=self.train_device) + lora_context = self.model.transformer_lora.retargeted(self.model.low_noise_transformer) + + with lora_context: + latent_video, _ = self.__run_denoise_loop( + self.model.low_noise_transformer, self.model.low_noise_transformer_autocast_context, dtype, + expert_timesteps, noise_scheduler, audio_noise_scheduler, + latent_video, latent_audio, geometry, conditioning, False, 1.0, on_update_progress, + ) + + return {"latent_video": latent_video} + + @torch.no_grad() + def __decode( + self, + latent_video: torch.Tensor, + num_latent_frames: int, + latent_height: int, + latent_width: int, + ) -> ModelSamplerOutput: + # evict the transformer and materialize only the vae + self.model.materialize_only("vae") + vae = self.pipeline.vae + patch_size = self.pipeline.transformer_spatial_patch_size + patch_size_t = self.pipeline.transformer_temporal_patch_size + + latent_video = self.model.unpack_latents( + latent_video, num_latent_frames, latent_height, latent_width, patch_size, patch_size_t, + ) + latent_video = self.model.unscale_latents(latent_video) + latent_video = latent_video.to(dtype=vae.dtype) + + # a timestep-conditioned decoder would need a per-item timestep for its temb; both LTX-2 VAEs decode + # unconditionally, so none is built + assert not vae.config.timestep_conditioning + + video = vae.decode(latent_video, None, return_dict=False)[0] + # postprocess_video permutes frames ahead of channels, landing on [B, F, C, H, W] + video = self.pipeline.video_processor.postprocess_video(video, output_type='pt') + + return self.build_video_sampler_output(video) + + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) + + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + [("encoding text", self.__text_encode), ("connecting", self.__connect), + ("denoising", self.__denoise), ("denoising low noise", self.__denoise_low_noise), + ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, + ) + + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, + fps=self.model.NATIVE_FPS, + ) + + return sampler_outputs + + def sample( + self, + sample_config: SampleConfig, + destination: str, + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ): + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], + image_format, video_format, audio_format, + on_update_progress=on_update_progress, + )[0] + + on_sample(sampler_output) diff --git a/modules/modelSampler/PixArtAlphaSampler.py b/modules/modelSampler/PixArtAlphaSampler.py index f8c38f135..f4a3760bd 100644 --- a/modules/modelSampler/PixArtAlphaSampler.py +++ b/modules/modelSampler/PixArtAlphaSampler.py @@ -9,13 +9,12 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.PIXART_ALPHA) @factory.register(BaseModelSampler, ModelType.PIXART_SIGMA) @@ -34,124 +33,162 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - text_encoder_layer_skip: int = 0, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("text_encoder") + prompt_embedding, tokens_attention_mask = self.model.encode_text( + text=sample_config.prompt, + train_device=self.train_device, + text_encoder_layer_skip=sample_config.text_encoder_1_layer_skip, + ) + + negative_prompt_embedding, negative_tokens_attention_mask = self.model.encode_text( + text=sample_config.negative_prompt, + train_device=self.train_device, + text_encoder_layer_skip=sample_config.text_encoder_1_layer_skip, + ) + + combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) + combined_prompt_attention_mask = torch.cat([negative_tokens_attention_mask, tokens_attention_mask]) + + return { + "combined_prompt_embedding": combined_prompt_embedding, + "combined_prompt_attention_mask": combined_prompt_attention_mask, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + combined_prompt_embedding: torch.Tensor, + combined_prompt_attention_mask: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + vae_scale_factor = self.pipeline.vae_scale_factor + height = self.quantize_resolution(sample_config.height, 16) + width = self.quantize_resolution(sample_config.width, 16) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = create.create_noise_scheduler(sample_config.noise_scheduler, self.pipeline.scheduler, diffusion_steps) + + # prepare timesteps + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) + timesteps = noise_scheduler.timesteps + + # prepare latent image + num_channels_latents = transformer.config.in_channels + latent_image = torch.randn( + size=(1, num_channels_latents, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32 + ) * noise_scheduler.init_noise_sigma + + # denoising loop + extra_step_kwargs = {} + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator + + batch_size = latent_image.shape[0] * 2 + height = latent_image.shape[2] * 8 + width = latent_image.shape[3] * 8 + resolution = torch.tensor([height, width]).repeat(batch_size, 1) + aspect_ratio = torch.tensor([float(height / width)]).repeat(batch_size, 1) + resolution = resolution \ + .to(dtype=self.model.train_dtype.torch_dtype(), device=self.train_device) + aspect_ratio = aspect_ratio \ + .to(dtype=self.model.train_dtype.torch_dtype(), device=self.train_device) + added_cond_kwargs = {"resolution": resolution, "aspect_ratio": aspect_ratio} + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image] * 2) + latent_model_input = noise_scheduler.scale_model_input(latent_model_input, timestep) + + # predict the noise residual + noise_pred = transformer( + latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), + encoder_hidden_states=combined_prompt_embedding.to( + dtype=self.model.train_dtype.torch_dtype()), + encoder_attention_mask=combined_prompt_attention_mask.to( + dtype=self.model.train_dtype.torch_dtype()), + timestep=timestep.expand(latent_model_input.shape[0]), + added_cond_kwargs=added_cond_kwargs, + ).sample + + # extract mean + noise_pred = noise_pred.chunk(2, dim=1)[0] + + # cfg + noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + + # compute the previous noisy sample x_t -> x_t-1 + latent_image = noise_scheduler.step( + noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs + )[0] + + on_update_progress(i + 1, len(timesteps)) + + return { + "latent_image": latent_image, + } + + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae + + latent_image = latent_image.to(dtype=vae.dtype) + image = vae.decode(latent_image / vae.config.scaling_factor, return_dict=False)[0] + + do_denormalize = [True] * image.shape[0] + image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) + + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) + + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) + with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = create.create_noise_scheduler(noise_scheduler, self.pipeline.scheduler, diffusion_steps) - image_processor = self.pipeline.image_processor - transformer = self.pipeline.transformer - vae = self.pipeline.vae - vae_scale_factor = self.pipeline.vae_scale_factor - - # prepare prompt - self.model.materialize_only("text_encoder") - - prompt_embedding, tokens_attention_mask = self.model.encode_text( - text=prompt, - train_device=self.train_device, - text_encoder_layer_skip=text_encoder_layer_skip, + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, ) - negative_prompt_embedding, negative_tokens_attention_mask = self.model.encode_text( - text=negative_prompt, - train_device=self.train_device, - text_encoder_layer_skip=text_encoder_layer_skip, + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) - combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) - combined_prompt_attention_mask = torch.cat([negative_tokens_attention_mask, tokens_attention_mask]) - - # prepare timesteps - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) - timesteps = noise_scheduler.timesteps - - # prepare latent image - num_channels_latents = transformer.config.in_channels - latent_image = torch.randn( - size=(1, num_channels_latents, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32 - ) * noise_scheduler.init_noise_sigma - - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator - - batch_size = latent_image.shape[0] * 2 - height = latent_image.shape[2] * 8 - width = latent_image.shape[3] * 8 - resolution = torch.tensor([height, width]).repeat(batch_size, 1) - aspect_ratio = torch.tensor([float(height / width)]).repeat(batch_size, 1) - resolution = resolution \ - .to(dtype=self.model.train_dtype.torch_dtype(), device=self.train_device) - aspect_ratio = aspect_ratio \ - .to(dtype=self.model.train_dtype.torch_dtype(), device=self.train_device) - added_cond_kwargs = {"resolution": resolution, "aspect_ratio": aspect_ratio} - - # denoising loop - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * 2) - latent_model_input = noise_scheduler.scale_model_input(latent_model_input, timestep) - - # predict the noise residual - noise_pred = transformer( - latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), - encoder_hidden_states=combined_prompt_embedding.to( - dtype=self.model.train_dtype.torch_dtype()), - encoder_attention_mask=combined_prompt_attention_mask.to( - dtype=self.model.train_dtype.torch_dtype()), - timestep=timestep.expand(latent_model_input.shape[0]), - added_cond_kwargs=added_cond_kwargs, - ).sample - - # extract mean - noise_pred = noise_pred.chunk(2, dim=1)[0] - - # cfg - noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - # decode - self.model.materialize_only("vae") - - latent_image = latent_image.to(dtype=vae.dtype) - image = vae.decode(latent_image / vae.config.scaling_factor, return_dict=False)[0] - - do_denormalize = [True] * image.shape[0] - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], - ) + return sampler_outputs def sample( self, @@ -163,23 +200,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 16), - width=self.quantize_resolution(sample_config.width, 16), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - text_encoder_layer_skip=sample_config.text_encoder_1_layer_skip, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/QwenSampler.py b/modules/modelSampler/QwenSampler.py index c18eece7e..e19c66cb9 100644 --- a/modules/modelSampler/QwenSampler.py +++ b/modules/modelSampler/QwenSampler.py @@ -11,13 +11,12 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.QWEN) class QwenSampler(BaseModelSampler): @@ -35,121 +34,165 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("text_encoder") + cfg_scale = sample_config.cfg_scale + + #unlike other models, Qwen benefits from CFG but is still quite good at CFG 1. Optimize for that: + batch_size = 2 if cfg_scale > 1.0 else 1 + combined_prompt_embedding, text_attention_mask = self.model.encode_text( + text=[sample_config.prompt, sample_config.negative_prompt] if cfg_scale > 1.0 else sample_config.prompt, + batch_size=batch_size, + train_device=self.train_device, + ) + + return { + "batch_size": batch_size, + "combined_prompt_embedding": combined_prompt_embedding, + "text_attention_mask": text_attention_mask, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + batch_size: int, + combined_prompt_embedding: torch.Tensor, + text_attention_mask: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + vae_scale_factor = 8 + num_latent_channels = 16 + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + # prepare latent image + latent_image = torch.randn( + size=(1, num_latent_channels, 1, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + + shift = sample_config.override_shift \ + or self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) + latent_image = self.model.pack_latents(latent_image) + + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=math.log(shift)) + timesteps = noise_scheduler.timesteps + + # denoising loop + extra_step_kwargs = {} + #TODO always True for FlowMatchEulerDiscreteScheduler - remove and pass directly? + #If so, also remove for other models + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator #TODO purpose? + + #FIXME list of lists is not according to type hint, but according to diffusers code + #https://github.com/huggingface/diffusers/issues/12295 + img_shapes = [[( + 1, #frame for future video model - not batch size + height // vae_scale_factor // 2, + width // vae_scale_factor // 2) + ]] * batch_size + + if torch.all(text_attention_mask): + text_attention_mask = None + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image] * batch_size) + expanded_timestep = timestep.expand(batch_size) + noise_pred = transformer( + hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), + timestep=expanded_timestep / 1000, + encoder_hidden_states=combined_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), + encoder_hidden_states_mask=text_attention_mask, + img_shapes=img_shapes, + return_dict=True, + ).sample + + if cfg_scale > 1.0: + noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + + # compute the previous noisy sample x_t -> x_t-1 + latent_image = noise_scheduler.step( + noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs + )[0] + + on_update_progress(i + 1, len(timesteps)) + + latent_image = self.model.unpack_latents( + latent_image, + height // vae_scale_factor, + width // vae_scale_factor, + ) + + return { + "latent_image": latent_image, + } + + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - image_processor = self.pipeline.image_processor - - transformer = self.pipeline.transformer - vae = self.pipeline.vae - vae_scale_factor = 8 - num_latent_channels = 16 - - # prepare prompt - self.model.materialize_only("text_encoder") - - #unlike other models, Qwen benefits from CFG but is still quite good at CFG 1. Optimize for that: - batch_size = 2 if cfg_scale > 1.0 else 1 - combined_prompt_embedding, text_attention_mask = self.model.encode_text( - text=[prompt, negative_prompt] if cfg_scale > 1.0 else prompt, - batch_size=batch_size, - train_device=self.train_device, - ) + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae - # prepare latent image - latent_image = torch.randn( - size=(1, num_latent_channels, 1, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32, - ) + latents = self.model.unscale_latents(latent_image) + image = vae.decode(latents, return_dict=False)[0].squeeze(-3) - shift = self.model.calculate_timestep_shift(latent_image.shape[-2], latent_image.shape[-1]) - latent_image = self.model.pack_latents(latent_image) - - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device, mu=math.log(shift)) - timesteps = noise_scheduler.timesteps - - # denoising loop - extra_step_kwargs = {} - #TODO always True for FlowMatchEulerDiscreteScheduler - remove and pass directly? - #If so, also remove for other models - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator #TODO purpose? - - #FIXME list of lists is not according to type hint, but according to diffusers code - #https://github.com/huggingface/diffusers/issues/12295 - img_shapes = [[( - 1, #frame for future video model - not batch size - height // vae_scale_factor // 2, - width // vae_scale_factor // 2) - ]] * batch_size - - if torch.all(text_attention_mask): - text_attention_mask = None - - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * batch_size) - expanded_timestep = timestep.expand(batch_size) - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), - timestep=expanded_timestep / 1000, - encoder_hidden_states=combined_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), - encoder_hidden_states_mask=text_attention_mask, - img_shapes=img_shapes, - return_dict=True, - ).sample - - if cfg_scale > 1.0: - noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - latent_image = self.model.unpack_latents( - latent_image, - height // vae_scale_factor, - width // vae_scale_factor, - ) + do_denormalize = [True] * image.shape[0] + image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - # decode - self.model.materialize_only("vae") + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) - latents = self.model.unscale_latents(latent_image) - image = vae.decode(latents, return_dict=False)[0].squeeze(-3) + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) - do_denormalize = [True] * image.shape[0] - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, + ) - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -160,22 +203,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/SanaSampler.py b/modules/modelSampler/SanaSampler.py index f4089a222..2d60f012c 100644 --- a/modules/modelSampler/SanaSampler.py +++ b/modules/modelSampler/SanaSampler.py @@ -10,13 +10,12 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.SANA) class SanaSampler(BaseModelSampler): @@ -34,109 +33,147 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - text_encoder_layer_skip: int = 0, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("text_encoder") + prompt_embedding, tokens_attention_mask = self.model.encode_text( + text=sample_config.prompt, + train_device=self.train_device, + text_encoder_layer_skip=sample_config.text_encoder_1_layer_skip, + ) + + negative_prompt_embedding, negative_tokens_attention_mask = self.model.encode_text( + text=sample_config.negative_prompt, + train_device=self.train_device, + text_encoder_layer_skip=sample_config.text_encoder_1_layer_skip, + ) + + combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) + combined_prompt_attention_mask = torch.cat([negative_tokens_attention_mask, tokens_attention_mask]) + + return { + "combined_prompt_embedding": combined_prompt_embedding, + "combined_prompt_attention_mask": combined_prompt_attention_mask, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + combined_prompt_embedding: torch.Tensor, + combined_prompt_attention_mask: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + vae_scale_factor = self.pipeline.vae_scale_factor + height = self.quantize_resolution(sample_config.height, 32) + width = self.quantize_resolution(sample_config.width, 32) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + # prepare timesteps + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) + timesteps = noise_scheduler.timesteps + + # prepare latent image + num_channels_latents = transformer.config.in_channels + latent_image = torch.randn( + size=(1, num_channels_latents, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32 + ) * noise_scheduler.init_noise_sigma + + # denoising loop + extra_step_kwargs = {} + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image] * 2) + + # predict the noise residual + noise_pred = transformer( + latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), + encoder_hidden_states=combined_prompt_embedding \ + .to(dtype=self.model.train_dtype.torch_dtype()), + encoder_attention_mask=combined_prompt_attention_mask \ + .to(dtype=self.model.train_dtype.torch_dtype()), + timestep=timestep.expand(latent_model_input.shape[0]), + ).sample + + # cfg + noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + + # compute the previous noisy sample x_t -> x_t-1 + latent_image = noise_scheduler.step( + noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs + )[0] + + on_update_progress(i + 1, len(timesteps)) + + return { + "latent_image": latent_image, + } + + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae + + latent_image = latent_image.to(dtype=vae.dtype) + with self.model.vae_autocast_context: + image = vae.decode(latent_image / vae.config.scaling_factor, return_dict=False)[0] + + do_denormalize = [True] * image.shape[0] + image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) + + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) + + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) + with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - image_processor = self.pipeline.image_processor - transformer = self.pipeline.transformer - vae = self.pipeline.vae - vae_scale_factor = self.pipeline.vae_scale_factor - - # prepare prompt - self.model.materialize_only("text_encoder") - - prompt_embedding, tokens_attention_mask = self.model.encode_text( - text=prompt, - train_device=self.train_device, - text_encoder_layer_skip=text_encoder_layer_skip, + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, ) - negative_prompt_embedding, negative_tokens_attention_mask = self.model.encode_text( - text=negative_prompt, - train_device=self.train_device, - text_encoder_layer_skip=text_encoder_layer_skip, + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) - combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) - combined_prompt_attention_mask = torch.cat([negative_tokens_attention_mask, tokens_attention_mask]) - - # prepare timesteps - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) - timesteps = noise_scheduler.timesteps - - # prepare latent image - num_channels_latents = transformer.config.in_channels - latent_image = torch.randn( - size=(1, num_channels_latents, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32 - ) * noise_scheduler.init_noise_sigma - - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator - - # denoising loop - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * 2) - - # predict the noise residual - noise_pred = transformer( - latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), - encoder_hidden_states=combined_prompt_embedding \ - .to(dtype=self.model.train_dtype.torch_dtype()), - encoder_attention_mask=combined_prompt_attention_mask \ - .to(dtype=self.model.train_dtype.torch_dtype()), - timestep=timestep.expand(latent_model_input.shape[0]), - ).sample - - # cfg - noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - # decode - self.model.materialize_only("vae") - - latent_image = latent_image.to(dtype=vae.dtype) - with self.model.vae_autocast_context: - image = vae.decode(latent_image / vae.config.scaling_factor, return_dict=False)[0] - - do_denormalize = [True] * image.shape[0] - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], - ) + return sampler_outputs def sample( self, @@ -148,23 +185,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 32), - width=self.quantize_resolution(sample_config.width, 32), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - text_encoder_layer_skip=sample_config.text_encoder_1_layer_skip, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/StableDiffusion3Sampler.py b/modules/modelSampler/StableDiffusion3Sampler.py index f21f34627..c6d7858f4 100644 --- a/modules/modelSampler/StableDiffusion3Sampler.py +++ b/modules/modelSampler/StableDiffusion3Sampler.py @@ -10,13 +10,12 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_3) @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_35) @@ -35,121 +34,157 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - text_encoder_1_layer_skip: int = 0, - text_encoder_2_layer_skip: int = 0, - text_encoder_3_layer_skip: int = 0, - transformer_attention_mask: bool = False, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only_text_encoders() + prompt_embedding, pooled_prompt_embedding = self.model.combine_text_encoder_output( + *self.model.encode_text( + text=sample_config.prompt, + train_device=self.train_device, + text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, + text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, + text_encoder_3_layer_skip=sample_config.text_encoder_3_layer_skip, + apply_attention_mask=sample_config.transformer_attention_mask, + )) + + negative_prompt_embedding, negative_pooled_prompt_embedding = self.model.combine_text_encoder_output( + *self.model.encode_text( + text=sample_config.negative_prompt, + train_device=self.train_device, + text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, + text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, + text_encoder_3_layer_skip=sample_config.text_encoder_3_layer_skip, + apply_attention_mask=sample_config.transformer_attention_mask, + )) + + combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding], dim=0) + combined_pooled_prompt_embedding = torch.cat( + [negative_pooled_prompt_embedding, pooled_prompt_embedding], dim=0) + + return { + "combined_prompt_embedding": combined_prompt_embedding, + "combined_pooled_prompt_embedding": combined_pooled_prompt_embedding, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + combined_prompt_embedding: torch.Tensor, + combined_pooled_prompt_embedding: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + vae_scale_factor = self.pipeline.vae_scale_factor + height = self.quantize_resolution(sample_config.height, 16) + width = self.quantize_resolution(sample_config.width, 16) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + # prepare timesteps + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) + timesteps = noise_scheduler.timesteps + + # prepare latent image + num_channels_latents = transformer.config.in_channels + latent_image = torch.randn( + size=(1, num_channels_latents, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + + # denoising loop + extra_step_kwargs = {} + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = torch.cat([latent_image] * 2) + expanded_timestep = timestep.expand(latent_model_input.shape[0]) + # Don't seem to scale the latents in SD3. + + # predict the noise residual + noise_pred = transformer( + hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), + timestep=expanded_timestep, + encoder_hidden_states=combined_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), + pooled_projections=combined_pooled_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), + return_dict=True + ).sample + + # cfg + noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + + # compute the previous noisy sample x_t -> x_t-1 + latent_image = noise_scheduler.step( + noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs + )[0] + + on_update_progress(i + 1, len(timesteps)) + + return { + "latent_image": latent_image, + } + + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae + + latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor + image = vae.decode(latents, return_dict=False)[0] + + do_denormalize = [True] * image.shape[0] + image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) + + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) + + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) + with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - image_processor = self.pipeline.image_processor - transformer = self.pipeline.transformer - vae = self.pipeline.vae - vae_scale_factor = self.pipeline.vae_scale_factor - - # prepare prompt - self.model.materialize_only_text_encoders() - - prompt_embedding, pooled_prompt_embedding = self.model.combine_text_encoder_output( - *self.model.encode_text( - text=prompt, - train_device=self.train_device, - text_encoder_1_layer_skip=text_encoder_1_layer_skip, - text_encoder_2_layer_skip=text_encoder_2_layer_skip, - text_encoder_3_layer_skip=text_encoder_3_layer_skip, - apply_attention_mask=transformer_attention_mask, - )) - - negative_prompt_embedding, negative_pooled_prompt_embedding = self.model.combine_text_encoder_output( - *self.model.encode_text( - text=negative_prompt, - train_device=self.train_device, - text_encoder_1_layer_skip=text_encoder_1_layer_skip, - text_encoder_2_layer_skip=text_encoder_2_layer_skip, - text_encoder_3_layer_skip=text_encoder_3_layer_skip, - apply_attention_mask=transformer_attention_mask, - )) - - combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding], dim=0) - combined_pooled_prompt_embedding = torch.cat( - [negative_pooled_prompt_embedding, pooled_prompt_embedding], dim=0) - - # prepare timesteps - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) - timesteps = noise_scheduler.timesteps - - # prepare latent image - num_channels_latents = transformer.config.in_channels - latent_image = torch.randn( - size=(1, num_channels_latents, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32, + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, ) - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator - - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * 2) - expanded_timestep = timestep.expand(latent_model_input.shape[0]) - # Don't seem to scale the latents in SD3. - - # predict the noise residual - noise_pred = transformer( - hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), - timestep=expanded_timestep, - encoder_hidden_states=combined_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), - pooled_projections=combined_pooled_prompt_embedding.to(dtype=self.model.train_dtype.torch_dtype()), - return_dict=True - ).sample - - # cfg - noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - # decode - self.model.materialize_only("vae") - - latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor - image = vae.decode(latents, return_dict=False)[0] - - do_denormalize = [True] * image.shape[0] - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -160,26 +195,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 16), - width=self.quantize_resolution(sample_config.width, 16), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, - text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, - text_encoder_3_layer_skip=sample_config.text_encoder_3_layer_skip, - transformer_attention_mask=sample_config.transformer_attention_mask, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/StableDiffusionSampler.py b/modules/modelSampler/StableDiffusionSampler.py index 791290edc..6f49ef737 100644 --- a/modules/modelSampler/StableDiffusionSampler.py +++ b/modules/modelSampler/StableDiffusionSampler.py @@ -12,13 +12,12 @@ from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat from modules.util.image_util import load_image +from modules.util.tqdm_util import tqdm import torch from torch import nn from torchvision.transforms import transforms -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_15) @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_15_INPAINTING) diff --git a/modules/modelSampler/StableDiffusionXLSampler.py b/modules/modelSampler/StableDiffusionXLSampler.py index 93d9f23d3..f3650cd1f 100644 --- a/modules/modelSampler/StableDiffusionXLSampler.py +++ b/modules/modelSampler/StableDiffusionXLSampler.py @@ -9,16 +9,15 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat from modules.util.image_util import load_image +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch from torch import nn from torchvision.transforms import transforms -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_XL_10_BASE) @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_XL_10_BASE_INPAINTING) @@ -36,158 +35,6 @@ def __init__( self.model_type = model_type self.pipeline = model.create_pipeline() - @torch.no_grad() - def __sample_base( - self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - cfg_rescale: float = 0.7, - text_encoder_1_layer_skip: int = 0, - text_encoder_2_layer_skip: int = 0, - force_last_timestep: bool = False, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, - ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = create.create_noise_scheduler(noise_scheduler, self.model.noise_scheduler, diffusion_steps) - image_processor = self.pipeline.image_processor - unet = self.pipeline.unet - vae = self.pipeline.vae - vae_scale_factor = self.pipeline.vae_scale_factor - - # prepare prompt - self.model.materialize_only_text_encoders() - - prompt_embedding, pooled_text_encoder_2_output = self.model.combine_text_encoder_output(*self.model.encode_text( - text=prompt, - train_device=self.train_device, - text_encoder_1_layer_skip=text_encoder_1_layer_skip, - text_encoder_2_layer_skip=text_encoder_2_layer_skip, - )) - - negative_prompt_embedding, negative_pooled_text_encoder_2_output = self.model.combine_text_encoder_output(*self.model.encode_text( - text=negative_prompt, - train_device=self.train_device, - text_encoder_1_layer_skip=text_encoder_1_layer_skip, - text_encoder_2_layer_skip=text_encoder_2_layer_skip, - )) - - combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) \ - .to(dtype=self.model.train_dtype.torch_dtype()) - - # prepare timesteps - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) - timesteps = noise_scheduler.timesteps - - if force_last_timestep: - last_timestep = torch.ones(1, device=self.train_device, dtype=torch.int64) \ - * (noise_scheduler.config.num_train_timesteps - 1) - - # add the final timestep to force predicting with zero snr if it's not already here - if timesteps[0] != last_timestep: - noise_scheduler.set_timesteps(diffusion_steps + 1, device=self.train_device) - timesteps = torch.cat([last_timestep, timesteps]) - - original_height = height - original_width = width - crops_coords_top = 0 - crops_coords_left = 0 - target_height = height - target_width = width - - add_time_ids = torch.tensor([ - original_height, - original_width, - crops_coords_top, - crops_coords_left, - target_height, - target_width - ]).unsqueeze(dim=0) - - add_time_ids = add_time_ids.to( - device=self.train_device, - ) - - # prepare latent image - num_channels_latents = unet.config.in_channels - latent_image = torch.randn( - size=(1, num_channels_latents, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=self.model.train_dtype.torch_dtype(), - ) * noise_scheduler.init_noise_sigma - - added_cond_kwargs = { - "text_embeds": torch.concat([negative_pooled_text_encoder_2_output, pooled_text_encoder_2_output], dim=0), - "time_ids": torch.concat([add_time_ids] * 2, dim=0), - } - - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator - - # denoising loop - self.model.materialize_only("unet") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = torch.cat([latent_image] * 2) - latent_model_input = noise_scheduler.scale_model_input(latent_model_input, timestep) - - # predict the noise residual - noise_pred = unet( - sample=latent_model_input, - timestep=timestep, - encoder_hidden_states=combined_prompt_embedding, - added_cond_kwargs=added_cond_kwargs, - )[0] - - # cfg - noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - if cfg_rescale > 0.0: - # From: Common Diffusion Noise Schedules and Sample Steps are Flawed (https://arxiv.org/abs/2305.08891) - std_positive = noise_pred_positive.std(dim=list(range(1, noise_pred_positive.ndim)), keepdim=True) - std_pred = noise_pred.std(dim=list(range(1, noise_pred.ndim)), keepdim=True) - noise_pred_rescaled = noise_pred * (std_positive / std_pred) - noise_pred = ( - cfg_rescale * noise_pred_rescaled + (1 - cfg_rescale) * noise_pred - ) - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - # decode - self.model.materialize_only("vae") - - latent_image = latent_image.to(dtype=self.model.vae_train_dtype.torch_dtype()) - with self.model.vae_autocast_context: - image = vae.decode(latent_image / vae.config.scaling_factor, return_dict=False)[0] - - do_denormalize = [True] * image.shape[0] - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], - ) - def __create_erode_kernel(self, device): kernel_radius = 2 @@ -202,231 +49,292 @@ def __create_erode_kernel(self, device): kernel.to(device) return kernel + # only present for conditioning (inpainting) model types: VAE-encode the conditioning image + mask @torch.no_grad() - def __sample_inpainting( + def __cond_encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - cfg_rescale: float = 0.7, - sample_inpainting: bool = False, - base_image_path: str = "", - mask_image_path: str = "", - text_encoder_1_layer_skip: int = 0, - text_encoder_2_layer_skip: int = 0, - force_last_timestep: bool = False, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, - ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = create.create_noise_scheduler(noise_scheduler, self.model.noise_scheduler, diffusion_steps) - image_processor = self.pipeline.image_processor - unet = self.pipeline.unet - vae = self.pipeline.vae - vae_scale_factor = self.pipeline.vae_scale_factor - - # prepare conditioning image - self.model.materialize_only("vae") - - with self.model.vae_autocast_context: - if sample_inpainting: - t = transforms.Compose([ - transforms.ToTensor(), - transforms.Resize( - (height, width), interpolation=transforms.InterpolationMode.BILINEAR, antialias=True - ), - ]) - - image = load_image(base_image_path, convert_mode="RGB") - image = t(image).to( - dtype=self.model.vae_train_dtype.torch_dtype(), - device=self.train_device, - ) - - mask = load_image(mask_image_path, convert_mode='L') - mask = t(mask).to( - dtype=self.model.train_dtype.torch_dtype(), - device=self.train_device, - ) - - erode_kernel = self.__create_erode_kernel(self.train_device) - eroded_mask = erode_kernel(mask) - eroded_mask = (eroded_mask > 0.5).float() - - image = (image * 2.0) - 1.0 - conditioning_image = (image * (1 - eroded_mask)) - conditioning_image = conditioning_image.unsqueeze(0) - - latent_conditioning_image = vae.encode( - conditioning_image).latent_dist.mode() * vae.config.scaling_factor - - rescale_mask = transforms.Resize( - (round(mask.shape[1] // 8), round(mask.shape[2] // 8)), - interpolation=transforms.InterpolationMode.BILINEAR, - antialias=True - ) - latent_mask = rescale_mask(mask) - latent_mask = (latent_mask > 0).float() - latent_mask = latent_mask.unsqueeze(0) - else: - conditioning_image = torch.zeros( - (1, 3, height, width), - dtype=self.model.vae_train_dtype.torch_dtype(), - device=self.train_device, - ) - conditioning_image = conditioning_image - latent_conditioning_image = vae.encode(conditioning_image).latent_dist.mode() * vae.config.scaling_factor - latent_mask = torch.ones( - size=(1, 1, latent_conditioning_image.shape[2], latent_conditioning_image.shape[3]), - dtype=self.model.train_dtype.torch_dtype(), - device=self.train_device - ) - - # prepare prompt - self.model.materialize_only_text_encoders() - - prompt_embedding, pooled_text_encoder_2_output = self.model.combine_text_encoder_output( - *self.model.encode_text( - text=prompt, - train_device=self.train_device, - text_encoder_1_layer_skip=text_encoder_1_layer_skip, - text_encoder_2_layer_skip=text_encoder_2_layer_skip, - )) - - negative_prompt_embedding, negative_pooled_text_encoder_2_output = self.model.combine_text_encoder_output( - *self.model.encode_text( - text=negative_prompt, - train_device=self.train_device, - text_encoder_1_layer_skip=text_encoder_1_layer_skip, - text_encoder_2_layer_skip=text_encoder_2_layer_skip, - )) - - combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) \ - .to(dtype=self.model.train_dtype.torch_dtype()) - - # prepare timesteps - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) - timesteps = noise_scheduler.timesteps - - if force_last_timestep: - last_timestep = torch.ones(1, device=self.train_device, dtype=torch.int64) \ - * (noise_scheduler.config.num_train_timesteps - 1) - - # add the final timestep to force predicting with zero snr if it's not already here - if timesteps[0] != last_timestep: - noise_scheduler.set_timesteps(diffusion_steps + 1, device=self.train_device) - timesteps = torch.cat([last_timestep, timesteps]) - - original_height = height - original_width = width - crops_coords_top = 0 - crops_coords_left = 0 - target_height = height - target_width = width - - add_time_ids = torch.tensor([ - original_height, - original_width, - crops_coords_top, - crops_coords_left, - target_height, - target_width - ]).unsqueeze(dim=0) - - add_time_ids = add_time_ids.to( - device=self.train_device, - ) + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("vae") + vae = self.pipeline.vae + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + + with self.model.vae_autocast_context: + if sample_config.sample_inpainting: + t = transforms.Compose([ + transforms.ToTensor(), + transforms.Resize( + (height, width), interpolation=transforms.InterpolationMode.BILINEAR, antialias=True + ), + ]) + + image = load_image(sample_config.base_image_path, convert_mode="RGB") + image = t(image).to( + dtype=self.model.vae_train_dtype.torch_dtype(), + device=self.train_device, + ) - # prepare latent image - num_channels_latents = latent_conditioning_image.shape[1] - latent_image = torch.randn( - size=(1, num_channels_latents, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=self.model.train_dtype.torch_dtype(), - ) + mask = load_image(sample_config.mask_image_path, convert_mode='L') + mask = t(mask).to( + dtype=self.model.train_dtype.torch_dtype(), + device=self.train_device, + ) + + erode_kernel = self.__create_erode_kernel(self.train_device) + eroded_mask = erode_kernel(mask) + eroded_mask = (eroded_mask > 0.5).float() - if sample_inpainting: - # SDXL inpainting is terrible at reconstructing from pure noise. - # This removes the last timestep to let the model know about the general image composition and brightness - timesteps = timesteps[1:] - latent_image = noise_scheduler.add_noise(latent_conditioning_image, latent_image, timesteps[:1]) + image = (image * 2.0) - 1.0 + conditioning_image = (image * (1 - eroded_mask)) + conditioning_image = conditioning_image.unsqueeze(0) + + latent_conditioning_image = vae.encode( + conditioning_image).latent_dist.mode() * vae.config.scaling_factor + + rescale_mask = transforms.Resize( + (round(mask.shape[1] // 8), round(mask.shape[2] // 8)), + interpolation=transforms.InterpolationMode.BILINEAR, + antialias=True + ) + latent_mask = rescale_mask(mask) + latent_mask = (latent_mask > 0).float() + latent_mask = latent_mask.unsqueeze(0) else: - latent_image = latent_image * noise_scheduler.init_noise_sigma + conditioning_image = torch.zeros( + (1, 3, height, width), + dtype=self.model.vae_train_dtype.torch_dtype(), + device=self.train_device, + ) + latent_conditioning_image = vae.encode(conditioning_image).latent_dist.mode() * vae.config.scaling_factor + latent_mask = torch.ones( + size=(1, 1, latent_conditioning_image.shape[2], latent_conditioning_image.shape[3]), + dtype=self.model.train_dtype.torch_dtype(), + device=self.train_device + ) + + return { + "latent_conditioning_image": latent_conditioning_image, + "latent_mask": latent_mask, + } + + @torch.no_grad() + def __encode( + self, + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only_text_encoders() + prompt_embedding, pooled_text_encoder_2_output = self.model.combine_text_encoder_output( + *self.model.encode_text( + text=sample_config.prompt, + train_device=self.train_device, + text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, + text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, + )) + + negative_prompt_embedding, negative_pooled_text_encoder_2_output = self.model.combine_text_encoder_output( + *self.model.encode_text( + text=sample_config.negative_prompt, + train_device=self.train_device, + text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, + text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, + )) + + combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) \ + .to(dtype=self.model.train_dtype.torch_dtype()) + + return { + "combined_prompt_embedding": combined_prompt_embedding, + "pooled_text_encoder_2_output": pooled_text_encoder_2_output, + "negative_pooled_text_encoder_2_output": negative_pooled_text_encoder_2_output, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + combined_prompt_embedding: torch.Tensor, + pooled_text_encoder_2_output: torch.Tensor, + negative_pooled_text_encoder_2_output: torch.Tensor, + on_update_progress: Callable[[int, int], None], + latent_conditioning_image: torch.Tensor | None = None, + latent_mask: torch.Tensor | None = None, + ) -> dict: + self.model.materialize_only("unet") + # conditioning tensors are only present for inpainting model types (their cond-encode stage ran) + is_inpainting = latent_conditioning_image is not None + unet = self.pipeline.unet + vae_scale_factor = self.pipeline.vae_scale_factor + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + cfg_scale = sample_config.cfg_scale + cfg_rescale = 0.7 if sample_config.force_last_timestep else 0.0 + force_last_timestep = sample_config.force_last_timestep + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = create.create_noise_scheduler(sample_config.noise_scheduler, self.model.noise_scheduler, diffusion_steps) + + # prepare timesteps + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) + timesteps = noise_scheduler.timesteps + + if force_last_timestep: + last_timestep = torch.ones(1, device=self.train_device, dtype=torch.int64) \ + * (noise_scheduler.config.num_train_timesteps - 1) + + # add the final timestep to force predicting with zero snr if it's not already here + if timesteps[0] != last_timestep: + noise_scheduler.set_timesteps(diffusion_steps + 1, device=self.train_device) + timesteps = torch.cat([last_timestep, timesteps]) + + original_height = height + original_width = width + crops_coords_top = 0 + crops_coords_left = 0 + target_height = height + target_width = width + + add_time_ids = torch.tensor([ + original_height, + original_width, + crops_coords_top, + crops_coords_left, + target_height, + target_width + ]).unsqueeze(dim=0) + + add_time_ids = add_time_ids.to( + device=self.train_device, + ) + + # prepare latent image + num_channels_latents = latent_conditioning_image.shape[1] if is_inpainting else unet.config.in_channels + latent_image = torch.randn( + size=(1, num_channels_latents, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=self.model.train_dtype.torch_dtype(), + ) + + if is_inpainting and sample_config.sample_inpainting: + # SDXL inpainting is terrible at reconstructing from pure noise. + # This removes the last timestep to let the model know about the general image composition and brightness + timesteps = timesteps[1:] + latent_image = noise_scheduler.add_noise(latent_conditioning_image, latent_image, timesteps[:1]) + else: + latent_image = latent_image * noise_scheduler.init_noise_sigma - added_cond_kwargs = { - "text_embeds": torch.concat([negative_pooled_text_encoder_2_output, pooled_text_encoder_2_output], dim=0), - "time_ids": torch.concat([add_time_ids] * 2, dim=0), - } + added_cond_kwargs = { + "text_embeds": torch.concat([negative_pooled_text_encoder_2_output, pooled_text_encoder_2_output], dim=0), + "time_ids": torch.concat([add_time_ids] * 2, dim=0), + } - # denoising loop - extra_step_kwargs = {} - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator + # denoising loop + extra_step_kwargs = {} + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator - # denoising loop - self.model.materialize_only("unet") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + if is_inpainting: latent_model_input = noise_scheduler.scale_model_input(latent_image, timestep) latent_model_input = torch.concat( [latent_model_input, latent_mask, latent_conditioning_image], 1 ) latent_model_input = torch.cat([latent_model_input] * 2) + else: + latent_model_input = torch.cat([latent_image] * 2) + latent_model_input = noise_scheduler.scale_model_input(latent_model_input, timestep) + + # predict the noise residual + noise_pred = unet( + sample=latent_model_input, + timestep=timestep, + encoder_hidden_states=combined_prompt_embedding, + added_cond_kwargs=added_cond_kwargs, + )[0] + + # cfg + noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + + if cfg_rescale > 0.0: + # From: Common Diffusion Noise Schedules and Sample Steps are Flawed (https://arxiv.org/abs/2305.08891) + std_positive = noise_pred_positive.std(dim=list(range(1, noise_pred_positive.ndim)), keepdim=True) + std_pred = noise_pred.std(dim=list(range(1, noise_pred.ndim)), keepdim=True) + noise_pred_rescaled = noise_pred * (std_positive / std_pred) + noise_pred = ( + cfg_rescale * noise_pred_rescaled + (1 - cfg_rescale) * noise_pred + ) + + # compute the previous noisy sample x_t -> x_t-1 + latent_image = noise_scheduler.step( + noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs + )[0] + + on_update_progress(i + 1, len(timesteps)) + + return { + "latent_image": latent_image, + } + + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, + ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae + + latent_image = latent_image.to(dtype=self.model.vae_train_dtype.torch_dtype()) + with self.model.vae_autocast_context: + image = vae.decode(latent_image / vae.config.scaling_factor, return_dict=False)[0] - # predict the noise residual - noise_pred = unet( - sample=latent_model_input, - timestep=timestep, - encoder_hidden_states=combined_prompt_embedding, - added_cond_kwargs=added_cond_kwargs, - )[0] - - # cfg - noise_pred_negative, noise_pred_positive = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - - if cfg_rescale > 0.0: - # From: Common Diffusion Noise Schedules and Sample Steps are Flawed (https://arxiv.org/abs/2305.08891) - std_positive = noise_pred_positive.std(dim=list(range(1, noise_pred_positive.ndim)), keepdim=True) - std_pred = noise_pred.std(dim=list(range(1, noise_pred.ndim)), keepdim=True) - noise_pred_rescaled = noise_pred * (std_positive / std_pred) - noise_pred = ( - cfg_rescale * noise_pred_rescaled + (1 - cfg_rescale) * noise_pred - ) - - # compute the previous noisy sample x_t -> x_t-1 - latent_image = noise_scheduler.step( - noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs - )[0] - - on_update_progress(i + 1, len(timesteps)) - - # decode - self.model.materialize_only("vae") - - latent_image = latent_image.to(dtype=self.model.vae_train_dtype.torch_dtype()) - with self.model.vae_autocast_context: - image = vae.decode(latent_image / vae.config.scaling_factor, return_dict=False)[0] - - do_denormalize = [True] * image.shape[0] - image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + do_denormalize = [True] * image.shape[0] + image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) + + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) + + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + stages = [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)] + # conditioning (inpainting) model types VAE-encode a conditioning image first + if self.model_type.has_conditioning_image_input(): + stages.insert(0, ("encoding conditioning image", self.__cond_encode)) + + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) + + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + stages, + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, ) + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, + ) + + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -437,47 +345,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - if self.model_type.has_conditioning_image_input(): - sampler_output = self.__sample_inpainting( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - cfg_rescale=0.7 if sample_config.force_last_timestep else 0.0, - sample_inpainting=sample_config.sample_inpainting, - base_image_path=sample_config.base_image_path, - mask_image_path=sample_config.mask_image_path, - text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, - text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, - force_last_timestep=sample_config.force_last_timestep, - on_update_progress=on_update_progress, - ) - else: - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - cfg_rescale=0.7 if sample_config.force_last_timestep else 0.0, - text_encoder_1_layer_skip=sample_config.text_encoder_1_layer_skip, - text_encoder_2_layer_skip=sample_config.text_encoder_2_layer_skip, - force_last_timestep=sample_config.force_last_timestep, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSampler/WuerstchenSampler.py b/modules/modelSampler/WuerstchenSampler.py index a1e9d8ba0..d5c833031 100644 --- a/modules/modelSampler/WuerstchenSampler.py +++ b/modules/modelSampler/WuerstchenSampler.py @@ -11,11 +11,11 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch from PIL import Image -from tqdm import tqdm @factory.register(BaseModelSampler, ModelType.WUERSTCHEN_2) diff --git a/modules/modelSampler/ZImageSampler.py b/modules/modelSampler/ZImageSampler.py index 0e001df2a..ccd69d800 100644 --- a/modules/modelSampler/ZImageSampler.py +++ b/modules/modelSampler/ZImageSampler.py @@ -10,13 +10,12 @@ from modules.util.enum.FileType import FileType from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.ModelType import ModelType -from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.staged_pipeline import run_staged_pipeline +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.Z_IMAGE) class ZImageSampler(BaseModelSampler): @@ -34,98 +33,138 @@ def __init__( self.pipeline = model.create_pipeline() @torch.no_grad() - def __sample_base( + def __encode( self, - prompt: str, - negative_prompt: str, - height: int, - width: int, - seed: int, - random_seed: bool, - diffusion_steps: int, - cfg_scale: float, - noise_scheduler: NoiseScheduler, - on_update_progress: Callable[[int, int], None] = lambda _, __: None, - ) -> ModelSamplerOutput: - with self.model.autocast_context: - generator = torch.Generator(device=self.train_device) - if random_seed: - generator.seed() - else: - generator.manual_seed(seed) - - noise_scheduler = copy.deepcopy(self.model.noise_scheduler) - image_processor = self.pipeline.image_processor - transformer = self.pipeline.transformer - vae = self.pipeline.vae - - vae_scale_factor = 8 - num_latent_channels = transformer.in_channels - #patch_size = 2 - - # prepare prompt - self.model.materialize_only("text_encoder") - - batch_size = 2 if cfg_scale > 1.0 else 1 - prompt_embedding = self.model.encode_text( - text=[prompt, negative_prompt] if cfg_scale > 1.0 else prompt, - batch_size=batch_size, - train_device=self.train_device, - ) + sample_config: SampleConfig, + ) -> dict: + self.model.materialize_only("text_encoder") + batch_size = 2 if sample_config.cfg_scale > 1.0 else 1 + prompt_embedding = self.model.encode_text( + text=[sample_config.prompt, sample_config.negative_prompt] if sample_config.cfg_scale > 1.0 else sample_config.prompt, + batch_size=batch_size, + train_device=self.train_device, + ) - # prepare latent image - latent_image = torch.randn( - size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), - generator=generator, - device=self.train_device, - dtype=torch.float32, - ) + return { + "batch_size": batch_size, + "prompt_embedding": prompt_embedding, + } + + @torch.no_grad() + def __denoise( + self, + sample_config: SampleConfig, + batch_size: int, + prompt_embedding: torch.Tensor, + on_update_progress: Callable[[int, int], None], + ) -> dict: + self.model.materialize_only("transformer") + transformer = self.pipeline.transformer + vae_scale_factor = 8 + num_latent_channels = transformer.in_channels + #patch_size = 2 + height = self.quantize_resolution(sample_config.height, 64) + width = self.quantize_resolution(sample_config.width, 64) + cfg_scale = sample_config.cfg_scale + diffusion_steps = sample_config.diffusion_steps + + generator = torch.Generator(device=self.train_device) + if sample_config.random_seed: + generator.seed() + else: + generator.manual_seed(sample_config.seed) + + noise_scheduler = copy.deepcopy(self.model.noise_scheduler) + + # prepare latent image + latent_image = torch.randn( + size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), + generator=generator, + device=self.train_device, + dtype=torch.float32, + ) + + # prepare timesteps + noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) + timesteps = noise_scheduler.timesteps + + # denoising loop + extra_step_kwargs = {} #TODO remove + if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): + extra_step_kwargs["generator"] = generator + + for i, timestep in enumerate(tqdm(timesteps, desc="steps", leave=False)): + latent_model_input = latent_image.unsqueeze(2).to(dtype=self.model.train_dtype.torch_dtype()) + latent_model_input = torch.cat([latent_model_input] * batch_size) + latent_model_input_list = list(latent_model_input.unbind(dim=0)) + timestep_model_input = timestep.unsqueeze(0) + assert timestep_model_input.ndim == 1 + output_list = transformer( + latent_model_input_list, + (1000 - timestep_model_input) / 1000, + prompt_embedding, + return_dict=True + ).sample + + noise_pred = - torch.stack(output_list, dim=0).squeeze(dim=2) - # prepare timesteps - noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) - timesteps = noise_scheduler.timesteps + if cfg_scale > 1.0: + noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) + noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) - # denoising loop - extra_step_kwargs = {} #TODO remove - if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): - extra_step_kwargs["generator"] = generator + latent_image = noise_scheduler.step(noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs)[0] - self.model.materialize_only("transformer") - for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): - latent_model_input = latent_image.unsqueeze(2).to(dtype=self.model.train_dtype.torch_dtype()) - latent_model_input = torch.cat([latent_model_input] * batch_size) - latent_model_input_list = list(latent_model_input.unbind(dim=0)) - timestep_model_input = timestep.unsqueeze(0) - assert timestep_model_input.ndim == 1 - output_list = transformer( - latent_model_input_list, - (1000 - timestep_model_input) / 1000, - prompt_embedding, - return_dict=True - ).sample + on_update_progress(i + 1, len(timesteps)) - noise_pred = - torch.stack(output_list, dim=0).squeeze(dim=2) + return { + "latent_image": latent_image, + } - if cfg_scale > 1.0: - noise_pred_positive, noise_pred_negative = noise_pred.chunk(2) - noise_pred = noise_pred_negative + cfg_scale * (noise_pred_positive - noise_pred_negative) + @torch.no_grad() + def __decode( + self, + latent_image: torch.Tensor, + ) -> ModelSamplerOutput: + self.model.materialize_only("vae") + image_processor = self.pipeline.image_processor + vae = self.pipeline.vae - latent_image = noise_scheduler.step(noise_pred, timestep, latent_image, return_dict=False, **extra_step_kwargs)[0] + latents = self.model.unscale_latents(latent_image) + image = vae.decode(latents, return_dict=False)[0] - on_update_progress(i + 1, len(timesteps)) + image = image_processor.postprocess(image, output_type='pil') - self.model.materialize_only("vae") + return ModelSamplerOutput( + file_type=FileType.IMAGE, + data=image[0], + ) - latents = self.model.unscale_latents(latent_image) - image = vae.decode(latents, return_dict=False)[0] + def sample_all( + self, + sample_configs: list[SampleConfig], + destinations: list[str], + image_format: ImageFormat | None = None, + video_format: VideoFormat | None = None, + audio_format: AudioFormat | None = None, + on_update_progress: Callable[[int, int], None] = lambda _, __: None, + ) -> list[ModelSamplerOutput]: + batch_progress = self.batch_progress_callback(sample_configs, on_update_progress) - image = image_processor.postprocess(image, output_type='pil') + with self.model.autocast_context: + sampler_outputs = run_staged_pipeline( + [("encoding", self.__encode), ("denoising", self.__denoise), ("decoding", self.__decode)], + {"sample_config": sample_configs}, + {"on_update_progress": batch_progress}, + ) - return ModelSamplerOutput( - file_type=FileType.IMAGE, - data=image[0], + for sampler_output, destination in zip(sampler_outputs, destinations, strict=True): + self.save_sampler_output( + sampler_output, destination, + image_format, video_format, audio_format, ) + return sampler_outputs + def sample( self, sample_config: SampleConfig, @@ -136,22 +175,11 @@ def sample( on_sample: Callable[[ModelSamplerOutput], None] = lambda _: None, on_update_progress: Callable[[int, int], None] = lambda _, __: None, ): - sampler_output = self.__sample_base( - prompt=sample_config.prompt, - negative_prompt=sample_config.negative_prompt, - height=self.quantize_resolution(sample_config.height, 64), - width=self.quantize_resolution(sample_config.width, 64), - seed=sample_config.seed, - random_seed=sample_config.random_seed, - diffusion_steps=sample_config.diffusion_steps, - cfg_scale=sample_config.cfg_scale, - noise_scheduler=sample_config.noise_scheduler, - on_update_progress=on_update_progress, - ) - - self.save_sampler_output( - sampler_output, destination, + # single-sample entry point: a staged batch of one + sampler_output = self.sample_all( + [sample_config], [destination], image_format, video_format, audio_format, - ) + on_update_progress=on_update_progress, + )[0] on_sample(sampler_output) diff --git a/modules/modelSaver/LTXFineTuneModelSaver.py b/modules/modelSaver/LTXFineTuneModelSaver.py new file mode 100644 index 000000000..2add9f56e --- /dev/null +++ b/modules/modelSaver/LTXFineTuneModelSaver.py @@ -0,0 +1,11 @@ +from modules.model.LTXModel import LTXModel +from modules.modelSaver.GenericFineTuneModelSaver import make_fine_tune_model_saver +from modules.modelSaver.ltx2.LTXModelSaver import LTXModelSaver +from modules.util.enum.ModelType import ModelType + +LTXFineTuneModelSaver = make_fine_tune_model_saver( + ModelType.LTX_2, + model_class=LTXModel, + model_saver_class=LTXModelSaver, + embedding_saver_class=None, +) diff --git a/modules/modelSaver/LTXLoRAModelSaver.py b/modules/modelSaver/LTXLoRAModelSaver.py new file mode 100644 index 000000000..f66c9e9b0 --- /dev/null +++ b/modules/modelSaver/LTXLoRAModelSaver.py @@ -0,0 +1,11 @@ +from modules.model.LTXModel import LTXModel +from modules.modelSaver.GenericLoRAModelSaver import make_lora_model_saver +from modules.modelSaver.ltx2.LTXLoRASaver import LTXLoRASaver +from modules.util.enum.ModelType import ModelType + +LTXLoRAModelSaver = make_lora_model_saver( + ModelType.LTX_2, + model_class=LTXModel, + lora_saver_class=LTXLoRASaver, + embedding_saver_class=None, +) diff --git a/modules/modelSaver/ltx2/LTXLoRASaver.py b/modules/modelSaver/ltx2/LTXLoRASaver.py new file mode 100644 index 000000000..d12af3baf --- /dev/null +++ b/modules/modelSaver/ltx2/LTXLoRASaver.py @@ -0,0 +1,22 @@ +from modules.model.LTXModel import LTXModel +from modules.modelSaver.mixin.LoRASaverMixin import LoRASaverMixin + +from torch import Tensor + + +class LTXLoRASaver( + LoRASaverMixin, +): + def __init__(self): + super().__init__() + + def _get_state_dict( + self, + model: LTXModel, + ) -> dict[str, Tensor]: + state_dict = {} + if model.transformer_lora is not None: + state_dict |= model.transformer_lora.state_dict() + if model.lora_state_dict is not None: + state_dict |= model.lora_state_dict + return state_dict diff --git a/modules/modelSaver/ltx2/LTXModelSaver.py b/modules/modelSaver/ltx2/LTXModelSaver.py new file mode 100644 index 000000000..b776b9007 --- /dev/null +++ b/modules/modelSaver/ltx2/LTXModelSaver.py @@ -0,0 +1,72 @@ +import os.path +from pathlib import Path + +from modules.model.LTXModel import LTXModel +from modules.modelSaver.mixin.DtypeModelSaverMixin import DtypeModelSaverMixin +from modules.util.enum.ModelFormat import ModelFormat + +import torch + +from safetensors.torch import save_file + + +class LTXModelSaver( + DtypeModelSaverMixin, +): + def __init__(self): + super().__init__() + + def __save_diffusers( + self, + model: LTXModel, + destination: str, + dtype: torch.dtype | None, + ): + pipeline = model.create_pipeline() + pipeline.to("cpu") + save_pipeline = self._copy_pipeline_to_dtype(pipeline, dtype, pipeline.tokenizer) + + os.makedirs(Path(destination).absolute(), exist_ok=True) + save_pipeline.save_pretrained(destination) + + if dtype is not None: + del save_pipeline + + def __save_safetensors( + self, + model: LTXModel, + destination: str, + dtype: torch.dtype | None, + ): + state_dict = model.transformer.state_dict() + + save_state_dict = self._convert_state_dict_dtype(state_dict, dtype) + self._convert_state_dict_to_contiguous(save_state_dict) + + os.makedirs(Path(destination).parent.absolute(), exist_ok=True) + + save_file(save_state_dict, destination, self._create_safetensors_header(model, save_state_dict)) + + def __save_internal( + self, + model: LTXModel, + destination: str, + ): + self.__save_diffusers(model, destination, None) + + def save( + self, + model: LTXModel, + output_model_format: ModelFormat, + output_model_destination: str, + dtype: torch.dtype | None, + ): + match output_model_format: + case ModelFormat.DIFFUSERS: + self.__save_diffusers(model, output_model_destination, dtype) + case ModelFormat.ORIGINAL_TRANSFORMER: + self.__save_safetensors(model, output_model_destination, dtype) + case ModelFormat.INTERNAL: + self.__save_internal(model, output_model_destination) + case _: + raise NotImplementedError(f"Unsupported output format: {output_model_format}") diff --git a/modules/modelSetup/BaseLTXSetup.py b/modules/modelSetup/BaseLTXSetup.py new file mode 100644 index 000000000..9468c3bb1 --- /dev/null +++ b/modules/modelSetup/BaseLTXSetup.py @@ -0,0 +1,211 @@ +from abc import ABCMeta + +import modules.util.multi_gpu_util as multi +from modules.model.LTXModel import LTXModel +from modules.modelSetup.BaseModelSetup import BaseModelSetup +from modules.modelSetup.mixin.ModelSetupDebugMixin import ModelSetupDebugMixin +from modules.modelSetup.mixin.ModelSetupDiffusionLossMixin import ModelSetupDiffusionLossMixin +from modules.modelSetup.mixin.ModelSetupEmbeddingMixin import ModelSetupEmbeddingMixin +from modules.modelSetup.mixin.ModelSetupFlowMatchingMixin import ModelSetupFlowMatchingMixin +from modules.modelSetup.mixin.ModelSetupNoiseMixin import ModelSetupNoiseMixin +from modules.util.checkpointing_util import ( + enable_checkpointing_for_gemma3_encoder_layers, + enable_checkpointing_for_gemma4_encoder_layers, + enable_checkpointing_for_ltx_connectors, + enable_checkpointing_for_ltx_transformer, +) +from modules.util.config.TrainConfig import TrainConfig +from modules.util.TrainProgress import TrainProgress + +import torch +from torch import Tensor + +from transformers import Gemma3ForConditionalGeneration, Gemma4UnifiedForConditionalGeneration + + +class BaseLTXSetup( + BaseModelSetup, + ModelSetupDiffusionLossMixin, + ModelSetupDebugMixin, + ModelSetupNoiseMixin, + ModelSetupFlowMatchingMixin, + ModelSetupEmbeddingMixin, + metaclass=ABCMeta +): + # The leading dot keeps LoRA on the video submodules - the audio counterparts are "audio_attn1" etc. + # + # The "lightricks-*" entries are quantization filters, not LoRA target sets (the two dropdowns share this + # dict). Each reproduces one Lightricks prequantized release, read off that checkpoint's safetensors + # header. Spelled as exclusions because regex has no numeric ranges, so they encode the 48-block count; + # "to_gate_logits" needs the ".*" because a Linear directly in a ModuleList is filtered under its parent's + # name. The dict form with regex=True is how a preset carries the "Use Regex" toggle. + LAYER_PRESETS = { + "video-attn-mlp": [".attn1", ".attn2", ".ff"], + "video-attn-only": [".attn1", ".attn2"], + "blocks": ["transformer_blocks"], + "lightricks-2.3-fp8": {"patterns": [r"^transformer_blocks\.(?![01]\.|4[67]\.)"], "regex": True}, + "lightricks-2.5-int8-convrot": { + "patterns": [r"^transformer_blocks\.\d+\.(?!.*to_gate_logits)"], "regex": True}, + "lightricks-2.5-nvfp4": { + "patterns": [r"^transformer_blocks\.(?!4[2-7]\.)\d+\.(?!.*to_gate_logits)"], "regex": True}, + "full": [], + } + + def setup_optimizations( + self, + model: LTXModel, + config: TrainConfig, + ): + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_ltx_transformer, disable_fp16_autocast=True) + self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) + if isinstance(model.text_encoder, Gemma3ForConditionalGeneration): + enable_checkpointing_for_text_encoder = enable_checkpointing_for_gemma3_encoder_layers + elif isinstance(model.text_encoder, Gemma4UnifiedForConditionalGeneration): + enable_checkpointing_for_text_encoder = enable_checkpointing_for_gemma4_encoder_layers + else: + raise NotImplementedError(f"no checkpointing wrapper for text encoder {type(model.text_encoder)}") + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_text_encoder, disable_fp16_autocast=True) + self._setup_model_part(model, config, "connectors", config.connectors, enable_checkpointing_for_ltx_connectors, disable_fp16_autocast=True) + if model.low_noise_transformer is not None: + self._setup_model_part(model, config, "low_noise_transformer", config.low_noise_transformer, enable_checkpointing_for_ltx_transformer, disable_fp16_autocast=True) + self._set_attention_backend(model.low_noise_transformer, config.attention_mechanism, mask=True) + self._setup_model_part(model, config, "vae", config.vae) + + model.vae.enable_tiling() + + def predict( + self, + model: LTXModel, + batch: dict, + config: TrainConfig, + train_progress: TrainProgress, + *, + deterministic: bool = False, + ) -> dict: + with model.autocast_context: + batch_seed = 0 if deterministic else train_progress.global_step * multi.world_size() + multi.rank() + generator = torch.Generator(device=config.train_device) + generator.manual_seed(batch_seed) + + # only the video conditioning is cached: the audio branch is isolated (isolate_modalities below) + # and discarded, so zeros of the right shape are fed for it instead + connector_prompt_embeds = model.encode_text( + train_device=self.train_device, + connector_video_embeds=batch['connector_video_embeds'], + text_encoder_dropout_probability=config.text_encoder.dropout_probability if not deterministic else None, + ) + + latent_image = batch['latent_image'].float() + # a video sample carries a frame dim, an image sample does not, and a dataset can hold both + if latent_image.ndim == 4: + latent_image = latent_image.unsqueeze(2) + batch_size, _, num_latent_frames, latent_height, latent_width = latent_image.shape + + scaled_latent_image = model.scale_latents(latent_image) + latent_noise = self._create_noise(scaled_latent_image, config, generator) + + shift = model.calculate_timestep_shift(num_latent_frames, latent_height, latent_width) + timestep = self._get_timestep_discrete( + model.noise_scheduler.config['num_train_timesteps'], + deterministic, + generator, + batch_size, + config, + shift=shift if config.dynamic_timestep_shifting else config.timestep_shift, + ) + + scaled_noisy_latent_image, sigma = self._add_noise_discrete( + scaled_latent_image, + latent_noise, + timestep, + model.noise_scheduler.timesteps, + ) + + patch_size = model.transformer.config.patch_size + patch_size_t = model.transformer.config.patch_size_t + packed_noisy_latent_image = model.pack_latents(scaled_noisy_latent_image, patch_size, patch_size_t) + + # video-only scope: isolate_modalities=True skips the cross-modality attention blocks, so the audio + # inputs the transformer takes positionally only need the right shape - hence zeros. Without it the + # video LoRA would take gradients from nonsense audio content every step. + frame_rate = model.NATIVE_FPS + pixel_num_frames = (num_latent_frames - 1) * model.vae.temporal_compression_ratio + 1 + duration_s = pixel_num_frames / frame_rate + audio_latents_per_second = ( + model.audio_vae.config.sample_rate + / model.audio_vae.config.mel_hop_length + / float(model.audio_vae.temporal_compression_ratio) + ) + audio_num_frames = round(duration_s * audio_latents_per_second) + latent_mel_bins = model.audio_vae.config.mel_bins // model.audio_vae.mel_compression_ratio + audio_shape = (batch_size, audio_num_frames, model.audio_vae.config.latent_channels * latent_mel_bins) + with model.transformer_autocast_context: + dtype = model.transformer_train_dtype.torch_dtype() + predicted_flow, _ = model.transformer( + hidden_states=packed_noisy_latent_image.to(dtype=dtype), + audio_hidden_states=torch.zeros(audio_shape, device=self.train_device, dtype=dtype), + encoder_hidden_states=connector_prompt_embeds.to(dtype=dtype), + audio_encoder_hidden_states=torch.zeros( + (batch_size, connector_prompt_embeds.shape[1], model.connectors.config.audio_hidden_dim), + device=self.train_device, dtype=dtype, + ), + timestep=timestep, + sigma=timestep, + encoder_attention_mask=None, + audio_encoder_attention_mask=None, + num_frames=num_latent_frames, + height=latent_height, + width=latent_width, + fps=frame_rate, + audio_num_frames=audio_num_frames, + isolate_modalities=True, + return_dict=False, + ) + + # unpack, to make the shape match the mask shape of masked training. LTX's unpack also unfolds the + # patches, so it lands on [B, C, F, H, W] in one step + predicted_flow = model.unpack_latents( + predicted_flow, num_latent_frames, latent_height, latent_width, patch_size, patch_size_t, + ) + + flow = latent_noise - scaled_latent_image + model_output_data = { + 'loss_type': 'target', + 'timestep': timestep, + 'predicted': predicted_flow, + 'target': flow, + } + + if config.debug_mode: + with torch.no_grad(): + predicted_scaled_latent_image = scaled_noisy_latent_image - predicted_flow * sigma + self._save_tokens('7-prompt', batch['tokens'], model.tokenizer, config, train_progress) + self._save_latent('1-noise', latent_noise, config, train_progress) + self._save_latent('2-noisy_image', scaled_noisy_latent_image, config, train_progress) + self._save_latent('3-predicted_flow', predicted_flow, config, train_progress) + self._save_latent('4-flow', flow, config, train_progress) + self._save_latent('5-predicted_image', predicted_scaled_latent_image, config, train_progress) + self._save_latent('6-image', scaled_latent_image, config, train_progress) + + return model_output_data + + def calculate_loss( + self, + model: LTXModel, + batch: dict, + data: dict, + config: TrainConfig, + ) -> Tensor: + return self._flow_matching_losses( + batch=batch, + data=data, + config=config, + train_device=self.train_device, + sigmas=model.noise_scheduler.sigmas, + ).mean() + + def prepare_text_caching(self, model: LTXModel, config: TrainConfig): + # the connector output is what gets cached, so the connectors run in the caching pass too + model.materialize_only("text_encoder", "connectors") + model.eval() diff --git a/modules/modelSetup/BaseModelSetup.py b/modules/modelSetup/BaseModelSetup.py index 19be0d81b..669b2badd 100644 --- a/modules/modelSetup/BaseModelSetup.py +++ b/modules/modelSetup/BaseModelSetup.py @@ -235,6 +235,13 @@ def _setup_model_part_requires_grad( not self.__stop_model_part_training_elapsed(unique_name, config, train_progress) model.requires_grad_(train_model_part) + # a streamed part (loaded as a meta skeleton) with cache-in-ram off is dropped to meta and re-streamed from + # the checkpoint on every reload, so training it would discard the update. Refuse the combination early. + if train_model_part and not config.cache_in_ram and any(p.is_meta for p in model.parameters()): + raise ValueError( + f"'{unique_name}' is trained with 'stream from disk' on and 'cache in ram' off -- the trained " + f"weights would be re-streamed from the checkpoint and lost. Enable 'cache in ram' for this part.") + #even if frozen parameters are not passed to the optimizer, required_grad has to be False. #otherwise, gradients accumulate in param.grad and waste vram if unique_name in self.frozen_parameters: @@ -255,10 +262,12 @@ def _setup_model_part( if module is None: return + materialize_fn = model.materialize_fn.get(attr) + if checkpointing_fn is not None: conductor = checkpointing_fn(module, config, config_part) if conductor is not None: - setattr(model, f"{attr}_offload_conductor", conductor) + model.offload_conductor[attr] = conductor if disable_fp16_autocast: autocast_context, train_dtype = disable_fp16_autocast_context( @@ -268,7 +277,10 @@ def _setup_model_part( else: train_dtype = model.train_dtype - quantize_layers(module, self.train_device, train_dtype, config, compress=config_part.weight_dtype.is_compressed()) + # a streamed module (materialize_fn set) stays on meta until materialized and is quantized per-materialize, + # so there is nothing to quantize here; a non-streamed module is quantized now. + if materialize_fn is None: + quantize_layers(module, self.train_device, train_dtype, config) @staticmethod def _set_attention_backend(component, attn: AttentionMechanism, mask: bool): diff --git a/modules/modelSetup/LTXFineTuneSetup.py b/modules/modelSetup/LTXFineTuneSetup.py new file mode 100644 index 000000000..b44b048dd --- /dev/null +++ b/modules/modelSetup/LTXFineTuneSetup.py @@ -0,0 +1,87 @@ +from modules.model.LTXModel import LTXModel +from modules.modelSetup.BaseLTXSetup import BaseLTXSetup +from modules.modelSetup.BaseModelSetup import BaseModelSetup +from modules.util import factory +from modules.util.config.TrainConfig import TrainConfig +from modules.util.enum.ModelType import ModelType +from modules.util.enum.TrainingMethod import TrainingMethod +from modules.util.ModuleFilter import ModuleFilter +from modules.util.NamedParameterGroup import NamedParameterGroupCollection +from modules.util.optimizer_util import init_model_parameters +from modules.util.TrainProgress import TrainProgress + + +@factory.register(BaseModelSetup, ModelType.LTX_2, TrainingMethod.FINE_TUNE) +class LTXFineTuneSetup( + BaseLTXSetup, +): + def create_parameters( + self, + model: LTXModel, + config: TrainConfig, + ) -> NamedParameterGroupCollection: + parameter_group_collection = NamedParameterGroupCollection() + self._create_model_part_parameters( + parameter_group_collection, "transformer", model.transformer, config.transformer, + freeze=ModuleFilter.create(config), debug=config.debug_mode, + ) + return parameter_group_collection + + def __setup_requires_grad( + self, + model: LTXModel, + config: TrainConfig, + ): + self._setup_model_part_requires_grad("transformer", model.transformer, config.transformer, model.train_progress) + model.vae.requires_grad_(False) + model.audio_vae.requires_grad_(False) + model.vocoder.requires_grad_(False) + model.text_encoder.requires_grad_(False) + model.connectors.requires_grad_(False) + + def setup_model( + self, + model: LTXModel, + config: TrainConfig, + ): + params = self.create_parameters(model, config) + self.__setup_requires_grad(model, config) + init_model_parameters(model, params, self.train_device) + + def setup_train_device( + self, + model: LTXModel, + config: TrainConfig, + ): + vae_on_train_device = not config.latent_caching + text_encoder_on_train_device = not config.latent_caching + + parts = ["transformer"] + if text_encoder_on_train_device: + # the connectors run alongside the TE in the dataloader, so they are needed exactly when it is + parts.append("text_encoder") + parts.append("connectors") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) + # keep the VAE latent stats on the train device: predict() normalizes with them every step, + # and .to(cuda) from an offloaded VAE would block-sync the stream each step. + model.vae.latents_mean = model.vae.latents_mean.to(self.train_device) + model.vae.latents_std = model.vae.latents_std.to(self.train_device) + + model.text_encoder.eval() + model.connectors.eval() + model.vae.eval() + + if config.transformer.train: + model.transformer.train() + else: + model.transformer.eval() + + def after_optimizer_step( + self, + model: LTXModel, + config: TrainConfig, + train_progress: TrainProgress, + ): + self.__setup_requires_grad(model, config) diff --git a/modules/modelSetup/LTXLoRASetup.py b/modules/modelSetup/LTXLoRASetup.py new file mode 100644 index 000000000..f9103e59e --- /dev/null +++ b/modules/modelSetup/LTXLoRASetup.py @@ -0,0 +1,102 @@ +from modules.model.LTXModel import LTXModel +from modules.modelSetup.BaseLTXSetup import BaseLTXSetup +from modules.modelSetup.BaseModelSetup import BaseModelSetup +from modules.module.LoRAModule import LoRAModuleWrapper +from modules.util import factory +from modules.util.config.TrainConfig import TrainConfig +from modules.util.enum.ModelType import ModelType +from modules.util.enum.TrainingMethod import TrainingMethod +from modules.util.NamedParameterGroup import NamedParameterGroupCollection +from modules.util.optimizer_util import init_model_parameters +from modules.util.TrainProgress import TrainProgress + + +@factory.register(BaseModelSetup, ModelType.LTX_2, TrainingMethod.LORA) +class LTXLoRASetup( + BaseLTXSetup, +): + def create_parameters( + self, + model: LTXModel, + config: TrainConfig, + ) -> NamedParameterGroupCollection: + parameter_group_collection = NamedParameterGroupCollection() + self._create_model_part_parameters(parameter_group_collection, "transformer", model.transformer_lora, config.transformer) + return parameter_group_collection + + def __setup_requires_grad( + self, + model: LTXModel, + config: TrainConfig, + ): + model.text_encoder.requires_grad_(False) + model.connectors.requires_grad_(False) + model.transformer.requires_grad_(False) + if model.low_noise_transformer is not None: + model.low_noise_transformer.requires_grad_(False) + model.vae.requires_grad_(False) + model.audio_vae.requires_grad_(False) + model.vocoder.requires_grad_(False) + self._setup_model_part_requires_grad("transformer", model.transformer_lora, config.transformer, model.train_progress) + + def setup_model( + self, + model: LTXModel, + config: TrainConfig, + ): + model.transformer_lora = LoRAModuleWrapper( + model.transformer, "transformer", config, config.layer_filter.split(",") + ) + + if model.lora_state_dict: + model.transformer_lora.load_state_dict(model.lora_state_dict) + model.lora_state_dict = None + + model.transformer_lora.set_dropout(config.dropout_probability) + model.transformer_lora.to(dtype=config.lora_weight_dtype.torch_dtype()) + model.transformer_lora.hook_to_module() + + params = self.create_parameters(model, config) + self.__setup_requires_grad(model, config) + init_model_parameters(model, params, self.train_device) + + def setup_train_device( + self, + model: LTXModel, + config: TrainConfig, + ): + vae_on_train_device = not config.latent_caching + text_encoder_on_train_device = not config.latent_caching + + parts = ["transformer"] + if text_encoder_on_train_device: + # the connectors run alongside the TE in the dataloader, so they are needed exactly when it is + parts.append("text_encoder") + parts.append("connectors") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) + # keep the VAE latent stats on the train device: predict() normalizes with them every step, + # and .to(cuda) from an offloaded VAE would block-sync the stream each step. + model.vae.latents_mean = model.vae.latents_mean.to(self.train_device) + model.vae.latents_std = model.vae.latents_std.to(self.train_device) + + model.text_encoder.eval() + model.connectors.eval() + model.vae.eval() + if model.low_noise_transformer is not None: + # absent from `parts` above: only the sampler runs it, which materializes it on demand + model.low_noise_transformer.eval() + + if config.transformer.train: + model.transformer.train() + else: + model.transformer.eval() + + def after_optimizer_step( + self, + model: LTXModel, + config: TrainConfig, + train_progress: TrainProgress, + ): + self.__setup_requires_grad(model, config) diff --git a/modules/module/AdditionalEmbeddingWrapper.py b/modules/module/AdditionalEmbeddingWrapper.py index 573bcbb96..48cf0fd73 100644 --- a/modules/module/AdditionalEmbeddingWrapper.py +++ b/modules/module/AdditionalEmbeddingWrapper.py @@ -30,7 +30,12 @@ def __init__( self.is_applied = False self.orig_forward = self.orig_module.forward - self.orig_median_norm = torch.norm(self.orig_module.weight, dim=1).median().item() + # orig_median_norm is only read by normalize_embeddings(), which only touches learned embeddings. A text + # encoder left on meta (streamed but not materialized, because none of its embeddings are trained) never + # reaches that path, so skip the norm read that would otherwise fail on a meta tensor (#69). + self.orig_median_norm = None + if not self.orig_module.weight.is_meta: + self.orig_median_norm = torch.norm(self.orig_module.weight, dim=1).median().item() def forward(self, x, *args, **kwargs): # ensure that the original weights only contain as many embeddings as the unmodified tokenizer can create diff --git a/modules/module/BaseImageCaptionModel.py b/modules/module/BaseImageCaptionModel.py index 2dfcf4a87..39cf79703 100644 --- a/modules/module/BaseImageCaptionModel.py +++ b/modules/module/BaseImageCaptionModel.py @@ -6,9 +6,9 @@ from modules.util import path_util from modules.util.image_util import load_image +from modules.util.tqdm_util import tqdm from PIL import Image -from tqdm import tqdm class CaptionSample: diff --git a/modules/module/BaseImageMaskModel.py b/modules/module/BaseImageMaskModel.py index 017ec7dd3..f5cfbfabb 100644 --- a/modules/module/BaseImageMaskModel.py +++ b/modules/module/BaseImageMaskModel.py @@ -5,13 +5,13 @@ from modules.util import path_util from modules.util.image_util import load_image +from modules.util.tqdm_util import tqdm import torch from torch import Tensor from torchvision.transforms import transforms from PIL import Image -from tqdm import tqdm class MaskSample: diff --git a/modules/module/FusedModule.py b/modules/module/FusedModule.py index 1948c7270..f517c4102 100644 --- a/modules/module/FusedModule.py +++ b/modules/module/FusedModule.py @@ -95,6 +95,11 @@ def __init__(self, prefix: str, leaves: list[nn.Module], klass, additional_args: # recompose the base weight itself (delta_forward returns None) keep going through the slower, # generic self.module.forward(x) path. def _leaf_output(self, leaf_index: int, start: int, end: int, x, *args, **kwargs): + # a leaf that can fold the adapter into its own base matmul returns the finished output, so + # neither the full fused delta nor the separate add happens for it + fused = self.module.fused_leaf_forward(self.leaves[leaf_index], x, start, end) + if fused is not None: + return fused delta = self.module.delta_forward(x, *args, **kwargs) if delta is None: return self.module.forward(x)[..., start:end] diff --git a/modules/module/GenerateLossesModel.py b/modules/module/GenerateLossesModel.py index d4c821b74..ea90beddc 100644 --- a/modules/module/GenerateLossesModel.py +++ b/modules/module/GenerateLossesModel.py @@ -8,12 +8,11 @@ from modules.util import create from modules.util.config.TrainConfig import QuantizationConfig, TrainConfig from modules.util.torch_util import torch_gc +from modules.util.tqdm_util import tqdm from modules.util.TrainProgress import TrainProgress import torch -from tqdm import tqdm - class GenerateLossesModel: """Based on train args, writes a JSON instead of a model with filenames mapped to losses, diff --git a/modules/module/LoRAModule.py b/modules/module/LoRAModule.py index 4f20c56e4..a3dd6a3da 100644 --- a/modules/module/LoRAModule.py +++ b/modules/module/LoRAModule.py @@ -3,11 +3,12 @@ from abc import abstractmethod from collections import defaultdict from collections.abc import Mapping +from contextlib import contextmanager from typing import Any from modules.module.FusedModule import FusedModuleGroup, check_fusion_match, discover_fused_groups from modules.module.oft_utils import OFTRotationModule -from modules.module.quantized.LinearSVD import BaseLinearSVD +from modules.module.quantized.mixin.LoRAFusableLinearMixin import LoRAFusableLinearMixin from modules.util.config.TrainConfig import TrainConfig from modules.util.enum.ModelType import PeftType from modules.util.lokr_utils import factorization, make_kron, rebuild_tucker @@ -115,6 +116,13 @@ def delta_forward(self, x, *args, **kwargs) -> Tensor | None: # going through the fused forward. return None + def fused_leaf_forward(self, leaf: nn.Module, x, start: int, end: int) -> Tensor | None: + # Returns one leaf's complete output (base + this adapter's contribution) when the leaf can + # fold the adapter into its own base matmul, so no delta is computed or added separately. + # start/end are the leaf's row range in the fused output. Returns None (the default) when no + # such path exists, and the caller falls back to delta_forward. + return None + @property def orig_module(self) -> nn.Module: assert self._orig_module is not None @@ -570,11 +578,18 @@ def check_initialized(self): def forward(self, x, *args, **kwargs): self.check_initialized() - if isinstance(self.orig_module, BaseLinearSVD): - return self.orig_module.forward_with_lora(x, self.lora_down, self.lora_up, self.dropout, self.alpha) + if isinstance(self.orig_module, LoRAFusableLinearMixin): + return self.orig_module.forward_with_lora(x, self.lora_down.weight, self.lora_up.weight, self.dropout, self.alpha) return self.orig_forward(x) + self.delta_forward(x, *args, **kwargs) + def fused_leaf_forward(self, leaf: nn.Module, x, start: int, end: int) -> Tensor | None: + self.check_initialized() + if not isinstance(leaf, LoRAFusableLinearMixin): + return None + #the fused adapter's up spans all leaves' concatenated outputs; narrow it to this leaf's rows + return leaf.forward_with_lora(x, self.lora_down.weight, self.lora_up.weight[start:end], self.dropout, self.alpha) + def delta_forward(self, x, *args, **kwargs) -> Tensor | None: self.check_initialized() ld = self.lora_up(self.dropout(self.lora_down(x))) @@ -781,6 +796,11 @@ def delta_forward(self, x, *args, **kwargs) -> Tensor | None: # DoRA scales the recomposed weight, so there is no delta term; back to None from LoRAModule's. return None + def fused_leaf_forward(self, leaf: nn.Module, x, start: int, end: int) -> Tensor | None: + # same reason as delta_forward: the leaf folds an additive term into its own base matmul, + # and DoRA recomposes the weight instead + return None + def forward(self, x, *args, **kwargs): self.check_initialized() A = self.lora_down.weight @@ -1092,6 +1112,42 @@ def modules(self) -> list[nn.Module]: return modules + @contextmanager + def retargeted(self, orig_module: nn.Module): + # Temporarily applies this LoRA to a different base module of the same architecture. + # TODO: temporary workaround until a LoRA can be attached to more than one base module; rebinds this + # one back and forth instead. The PeftBase objects and their weights are shared, not copied. + previous = self.orig_module + self.remove_hook_from_module() + self.__retarget(orig_module) + self.hook_to_module() + try: + yield + finally: + self.remove_hook_from_module() + self.__retarget(previous) + self.hook_to_module() + + def __retarget(self, orig_module: nn.Module): + # Binds by name against the tree as it stands right now, normalized the same way __create_modules did + # so a checkpointing wrapper on either base maps to the same child. Eviction moves weights to meta and + # never replaces a module, so an evicted component binds to the layers a later materialize fills. + children = {name.replace(".checkpoint.", "."): child for name, child in orig_module.named_modules()} + for name, lora_module in self.lora_modules.items(): + if isinstance(lora_module, FusedModuleGroup): + # a fused group also owns a _FusedLinear built over the old leaves, so rebinding it means + # rebuilding that too; nothing needs it yet, so refuse rather than half-rebind + raise NotImplementedError("retargeting a fused LoRA module is not supported") + target = children.get(name) + if target is None: + raise ValueError(f"retarget: {orig_module.__class__.__name__} has no module named {name}") + if get_weight_shape(target) != get_weight_shape(lora_module.orig_module): + raise ValueError( + f"retarget: shape mismatch for {name}: " + f"{get_weight_shape(lora_module.orig_module)} vs {get_weight_shape(target)}") + lora_module._orig_module = [target] + self.orig_module = orig_module + def hook_to_module(self): """ Hooks the LoRA into the module without changing its weights diff --git a/modules/module/quantized/LinearFp8.py b/modules/module/quantized/LinearFp8.py index 116929353..59abd8798 100644 --- a/modules/module/quantized/LinearFp8.py +++ b/modules/module/quantized/LinearFp8.py @@ -17,17 +17,26 @@ def __init__(self, *args, **kwargs): self.is_quantized = False self.fp8_dtype = torch.float8_e4m3fn - self._scale = torch.tensor(1.0, dtype=torch.float) - self.register_buffer("scale", self._scale) + self.register_buffer("scale", torch.tensor(1.0, dtype=torch.float)) self.compute_dtype = None def original_weight_shape(self) -> tuple[int, ...]: return self.weight.shape + def mark_needs_requantization(self): + self.is_quantized = False + + def predict_offload_bytes(self) -> int: + # weight quantizes to float8_e4m3fn (1 byte/elem, same shape); bias is left unchanged. Matches + # get_offload_tensors (weight + optional bias); the scalar scale buffer is not offload-counted. + weight_bytes = self.weight.numel() + bias_bytes = self.bias.numel() * self.bias.element_size() if self.bias is not None else 0 + return weight_bytes + bias_bytes + def unquantized_weight(self, dtype: torch.dtype, device: torch.device) -> torch.Tensor: # 'scale' is not offloaded, so it can sit on the train device while 'weight' is parked on the temp device - if self._scale is not None: - return self.weight.detach().to(dtype) * self._scale.to(dtype=dtype, device=self.weight.device) + if self.scale is not None: + return self.weight.detach().to(dtype) * self.scale.to(dtype=dtype, device=self.weight.device) else: return self.weight.detach().to(dtype=dtype) @@ -43,19 +52,22 @@ def quantize(self, device: torch.device | None = None): weight = weight.to(device=device) abs_max = weight.abs().max() - self._scale.copy_(torch.clamp(abs_max, min=1e-12) / torch.finfo(self.fp8_dtype).max) - weight = weight.div_(self._scale).to(dtype=self.fp8_dtype) + scale = torch.clamp(abs_max, min=1e-12) / torch.finfo(self.fp8_dtype).max + weight = weight.div_(scale).to(dtype=self.fp8_dtype) if device is not None: weight = weight.to(device=orig_device) + + # keep the scale on the weight's device (see LinearW8A8.quantize) + self.scale = scale.detach().to(orig_device) self.weight.data = weight def forward(self, x: torch.Tensor) -> torch.Tensor: weight = self.weight.detach() weight = weight.to(dtype=self.compute_dtype if self.compute_dtype is not None else x.dtype) - if self._scale is not None: - weight = weight.mul_(self._scale) + if self.scale is not None: + weight = weight.mul_(self.scale) x = nn.functional.linear(x, weight, self.bias) return x diff --git a/modules/module/quantized/LinearGGUFA8.py b/modules/module/quantized/LinearGGUFA8.py index 8b015c564..d11606d53 100644 --- a/modules/module/quantized/LinearGGUFA8.py +++ b/modules/module/quantized/LinearGGUFA8.py @@ -1,5 +1,8 @@ +from modules.module.quantized.mixin.LoRAFusableLinearMixin import LoRAFusableLinearMixin from modules.util.mm_8bit import mm_8bit as mm_8bit +from modules.util.mm_8bit import rowcol_scaled_lora_mm_8bit, rowcol_scaled_mm_8bit from modules.util.quantization_util import ( + quantize_axiswise, quantize_fp8_axiswise, quantize_int8_axiswise, ) @@ -14,70 +17,133 @@ UNQUANTIZED_TYPES = [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16, gguf.GGMLQuantizationType.BF16] +#unlike LinearW8A8, whose weight is quantized once and tensorwise, the GGUF weight is dequantized and +#requantized axiswise per pass (it can be, since the dequant happens anyway). That leaves two scale +#vectors varying along different axes - per-token on M and per-output-channel on N - which cannot be +#collapsed into one, so this layer uses the rowcol_ variants of the epilogue-scaled kernels. + +#the weight reaches these dequantized, so unlike LinearW8A8 they are told which 8-bit dtype to +#requantize to. Only the torch forward differs beyond the quantizer: int8 and fp8 have separate +#torch mms with separate scaling @torch.no_grad() -def int8_forward_axiswise(x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: - x_8, x_scale = quantize_int8_axiswise(x, dim=-1) - w_8, w_scale = quantize_int8_axiswise(weight, dim=-1) - res = torch._int_mm(x_8, w_8.T) - res_scaled = res.float().mul_(w_scale.T).mul_(x_scale).to(compute_dtype) +def forward_axiswise_postscaled_torch(dtype: torch.dtype, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: + if dtype == torch.int8: + x_8, x_scale = quantize_int8_axiswise(x, dim=-1) + w_8, w_scale = quantize_int8_axiswise(weight, dim=-1) + res = torch._int_mm(x_8, w_8.T) + res_scaled = res.float().mul_(w_scale.T).mul_(x_scale).to(compute_dtype) + else: + x_8, x_scale = quantize_fp8_axiswise(x, dim=-1) + w_8, w_scale = quantize_fp8_axiswise(weight, dim=-1) + one = torch.ones(1, device=x.device) + res = torch._scaled_mm(x_8, w_8.T, scale_a=one, scale_b=one, out_dtype=torch.float) + res_scaled = res.mul_(w_scale.T).mul_(x_scale).to(compute_dtype) #much faster than scaled by _scaled_mm if bias is not None: res_scaled.add_(bias) return res_scaled @torch.no_grad() -def fp8_forward_axiswise(x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: - x_8, x_scale = quantize_fp8_axiswise(x, dim=-1) - w_8, w_scale = quantize_fp8_axiswise(weight, dim=-1) - one = torch.ones(1, device=x.device) - res = torch._scaled_mm(x_8, w_8.T, scale_a=one, scale_b=one, out_dtype=torch.float) - res_scaled = res.mul_(w_scale.T).mul_(x_scale).to(compute_dtype) #much faster than scaled by _scaled_mm +def backward_axiswise_postscaled_triton(dtype: torch.dtype, output: Tensor, weight: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=dtype) + w_8, w_scale = quantize_axiswise(weight, dim=0, dtype=dtype) + mm_res = mm_8bit(output_8.contiguous(), w_8) + return mm_res.float().mul_(w_scale).mul_(output_scale).to(output.dtype) + + +@torch.no_grad() +def forward_axiswise_epiloguescaled_triton(dtype: torch.dtype, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: + x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=dtype) + w_8, w_scale = quantize_axiswise(weight, dim=-1, dtype=dtype) + #rowcol_scaled_mm_8bit folds the per-token scale (axis 0) and the per-channel weight scale + #(axis 1) into the epilogue and returns compute_dtype directly + res_scaled = rowcol_scaled_mm_8bit(x_8, w_8.T, x_scale, w_scale, compute_dtype) if bias is not None: res_scaled.add_(bias) return res_scaled @torch.no_grad() -def int8_backward_axiswise(output: Tensor, weight: Tensor) -> Tensor: - output_8, output_scale = quantize_int8_axiswise(output, dim=-1) - w_8, w_scale = quantize_int8_axiswise(weight, dim=0) - mm_res = mm_8bit(output_8.contiguous(), w_8) - return mm_res.float().mul_(w_scale).mul_(output_scale).to(output.dtype) +def backward_axiswise_epiloguescaled_triton(dtype: torch.dtype, output: Tensor, weight: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=dtype) + w_8, w_scale = quantize_axiswise(weight, dim=0, dtype=dtype) + return rowcol_scaled_mm_8bit(output_8.contiguous(), w_8, output_scale, w_scale, output.dtype) + @torch.no_grad() -def fp8_backward_axiswise(output: Tensor, weight: Tensor) -> Tensor: - output_8, output_scale = quantize_fp8_axiswise(output, dim=-1) - w_8, w_scale = quantize_fp8_axiswise(weight, dim=0) - mm_res = mm_8bit(output_8.contiguous(), w_8) - return mm_res.float().mul_(w_scale).mul_(output_scale).to(output.dtype) +def forward_axiswise_lora_epiloguescaled_triton(dtype: torch.dtype, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype, x_down: Tensor, lora_up: Tensor) -> Tensor: + x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=dtype) + w_8, w_scale = quantize_axiswise(weight, dim=-1, dtype=dtype) + res_scaled = rowcol_scaled_lora_mm_8bit(x_8, w_8.T, x_scale, w_scale, x_down, lora_up, compute_dtype) + if bias is not None: + res_scaled.add_(bias) + return res_scaled + +@torch.no_grad() +def backward_axiswise_lora_epiloguescaled_triton(dtype: torch.dtype, output: Tensor, weight: Tensor, grad_x_down_pre: Tensor, lora_down: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=dtype) + w_8, w_scale = quantize_axiswise(weight, dim=0, dtype=dtype) + return rowcol_scaled_lora_mm_8bit(output_8.contiguous(), w_8, output_scale, w_scale, grad_x_down_pre, lora_down.to(grad_x_down_pre.dtype), output.dtype) + + +forward_axiswise = forward_axiswise_epiloguescaled_triton +backward_axiswise = backward_axiswise_epiloguescaled_triton +forward_axiswise_lora = forward_axiswise_lora_epiloguescaled_triton +backward_axiswise_lora = backward_axiswise_lora_epiloguescaled_triton -class LinearGGUFIntA8RequantFunction(torch.autograd.Function): + +class LinearGGUFA8RequantFunction(torch.autograd.Function): + #`weight` is the dequantized GGUF weight, so saving it lets backward requantize instead of + #dequantizing the GGUF blocks a second time @staticmethod - def forward(ctx, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: + def forward(ctx, x: Tensor, weight: Tensor, dtype: torch.dtype, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: ctx.save_for_backward(weight) - #axiswise performs better than tensorwise in tests, even though - #it requires another requant during backward - but requant is cheap - return int8_forward_axiswise(x, weight, bias, compute_dtype) + ctx.dtype = dtype + return forward_axiswise(dtype, x, weight, bias, compute_dtype) @staticmethod def backward(ctx, output: Tensor): - if ctx.needs_input_grad != (True, False, False, False): + if ctx.needs_input_grad != (True, False, False, False, False): raise NotImplementedError("GGUF cannot be used for full finetuning") + weight, = ctx.saved_tensors - return int8_backward_axiswise(output, weight), None, None, None + return backward_axiswise(ctx.dtype, output, weight), None, None, None, None + -class LinearGGUFFpA8RequantFunction(torch.autograd.Function): +class LinearGGUFA8RequantLoRAFunction(torch.autograd.Function): + #LinearGGUFA8RequantFunction plus a low-rank update: it owns the down-projection and the + #dropout, so the LoRA dgrad folds into the backward epilogue and the (M, out) product is + #never materialized @staticmethod - def forward(ctx, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: - ctx.save_for_backward(weight) - return fp8_forward_axiswise(x, weight, bias, compute_dtype) + def forward(ctx, x: Tensor, weight: Tensor, dtype: torch.dtype, bias: Tensor | None, compute_dtype: torch.dtype, lora_down: Tensor, lora_up: Tensor, dropout_mask: Tensor | None) -> Tensor: + x_down = torch.nn.functional.linear(x, lora_down) + if dropout_mask is not None: + x_down = x_down * dropout_mask + ctx.dtype = dtype + ctx.save_for_backward(weight, x.to(x_down.dtype), lora_down, lora_up, dropout_mask, x_down) + return forward_axiswise_lora(dtype, x, weight, bias, compute_dtype, x_down, lora_up) @staticmethod - def backward(ctx, output: Tensor): - if ctx.needs_input_grad != (True, False, False, False): + def backward(ctx, grad_output: Tensor): + if ctx.needs_input_grad[1:5] != (False, False, False, False): raise NotImplementedError("GGUF cannot be used for full finetuning") - weight, = ctx.saved_tensors - return fp8_backward_axiswise(output, weight), None, None, None -class LinearGGUFA8(GGUFLinear): + weight, x, lora_down, lora_up, dropout_mask, x_down = ctx.saved_tensors + needs_x, needs_down, needs_up = ctx.needs_input_grad[0], ctx.needs_input_grad[5], ctx.needs_input_grad[6] + #grad_x_down_pre (post-dropout-backward) feeds both the grad_x epilogue fold and grad_lora_down + grad_x_down_pre = None + if needs_x or needs_down: + grad_x_down_pre = grad_output @ lora_up.T + if dropout_mask is not None: + grad_x_down_pre = grad_x_down_pre * dropout_mask + grad_x = backward_axiswise_lora(ctx.dtype, grad_output, weight, grad_x_down_pre, lora_down) if needs_x else None + grad_lora_down = (grad_x_down_pre.T @ x).to(lora_down.dtype) if needs_down else None + grad_lora_up = x_down.T @ grad_output if needs_up else None + return grad_x, None, None, None, None, grad_lora_down, grad_lora_up, None + + +class LinearGGUFA8( + GGUFLinear, + LoRAFusableLinearMixin, +): def __init__(self, dtype: torch.dtype, *args, **kwargs): super().__init__(*args, **kwargs) @@ -87,14 +153,38 @@ def __init__(self, dtype: torch.dtype, *args, **kwargs): def forward(self, x_orig: torch.Tensor) -> torch.Tensor: assert not self.weight.requires_grad x = x_orig.reshape(-1, x_orig.shape[-1]) - w = dequantize_gguf_tensor(self.weight.detach()) + w = dequantize_gguf_tensor(self.weight.detach()) + #the 8-bit path only pays for itself above a few tokens, and an unquantized GGUF type has + #nothing to requantize - dequantize_gguf_tensor already hands back the dense weight if x.shape[0] > 16 and self.weight.quant_type not in UNQUANTIZED_TYPES: - if self._dtype == torch.int8: - y = LinearGGUFIntA8RequantFunction.apply(x, w, self.bias, self.compute_dtype) - else: - y = LinearGGUFFpA8RequantFunction.apply(x, w, self.bias, self.compute_dtype) + #axiswise performs better than tensorwise in tests, even though + #it requires another requant during backward - but requant is cheap + y = LinearGGUFA8RequantFunction.apply(x, w, self._dtype, self.bias, self.compute_dtype) else: y = torch.nn.functional.linear(x, w, self.bias) return y.reshape(x_orig.shape[:-1] + (y.shape[-1], )) + + def forward_with_lora(self, x_orig: torch.Tensor, lora_down: torch.Tensor, lora_up: torch.Tensor, dropout: torch.nn.Dropout, alpha: torch.Tensor) -> torch.Tensor: + lora_rank = lora_down.shape[0] + + x = x_orig.reshape(-1, x_orig.shape[-1]) + if x.shape[0] <= 16 or self.weight.quant_type in UNQUANTIZED_TYPES: + ld = torch.nn.functional.linear(dropout(torch.nn.functional.linear(x_orig, lora_down)), lora_up) + return LinearGGUFA8.forward(self, x_orig) + ld * (alpha / lora_rank) + + #the cast matches lora_up (kept in its own dtype, e.g. f32) to the down-projection's autocast + #dtype, which the epilogue dot needs; dropout(ones) yields the scaled 0/(1/(1-p)) mask, drawn + #here rather than inside LinearGGUFA8RequantLoRAFunction so the RNG stays inductor-native + lora_up_scaled = (lora_up * (alpha / lora_rank)).T.to(self.compute_dtype) + dropout_mask = dropout(torch.ones(x.shape[0], lora_rank, device=x.device, dtype=self.compute_dtype)) if (dropout.training and dropout.p > 0) else None + y = self._fused_lora_forward(x, lora_down, lora_up_scaled, dropout_mask) + + return y.reshape(x_orig.shape[:-1] + (y.shape[-1], )) + + #see LinearW8A8._fused_lora_forward for what the fused node buys and what the operands are + def _fused_lora_forward(self, x: Tensor, lora_down: Tensor, lora_up: Tensor, dropout_mask: Tensor | None) -> Tensor: + assert not self.weight.requires_grad + w = dequantize_gguf_tensor(self.weight.detach()) + return LinearGGUFA8RequantLoRAFunction.apply(x, w, self._dtype, self.bias, self.compute_dtype, lora_down, lora_up, dropout_mask) diff --git a/modules/module/quantized/LinearNf4.py b/modules/module/quantized/LinearNf4.py index 2a4bfbf17..718b65856 100644 --- a/modules/module/quantized/LinearNf4.py +++ b/modules/module/quantized/LinearNf4.py @@ -38,7 +38,21 @@ def __init__(self, *args, **kwargs): self.quant_state = None def original_weight_shape(self) -> tuple[int, ...]: - return self.weight.shape + # self.weight is repacked to a flat [N, 1] uint8 layout once quantized; self.shape keeps the original. + return self.shape + + def mark_needs_requantization(self): + self.is_quantized = False + + def predict_offload_bytes(self) -> int: + # nf4 packs the weight to 4-bit (2 values per uint8), and with double quant (compress_statistics) stores + # quant_state.absmax as one uint8 per block_size elements. Matches get_offload_tensors (packed weight + + # quant_state.absmax + optional bias); the small code/offset/nested-absmax buffers are not offload-counted. + numel = self.shape.numel() + weight_bytes = (numel + 1) // 2 + absmax_bytes = (numel + self.block_size - 1) // self.block_size + bias_bytes = self.bias.numel() * self.bias.element_size() if self.bias is not None else 0 + return weight_bytes + absmax_bytes + bias_bytes def unquantized_weight(self, dtype: torch.dtype, device: torch.device) -> torch.Tensor: if self.is_quantized: diff --git a/modules/module/quantized/LinearSVD.py b/modules/module/quantized/LinearSVD.py index 16e2f2650..f652f5f6a 100644 --- a/modules/module/quantized/LinearSVD.py +++ b/modules/module/quantized/LinearSVD.py @@ -1,6 +1,7 @@ -from abc import abstractmethod +import os from contextlib import suppress +from modules.module.quantized.mixin.LoRAFusableLinearMixin import LoRAFusableLinearMixin from modules.module.quantized.mixin.QuantizedLinearMixin import QuantizedLinearMixin from modules.module.quantized.mixin.QuantizedModuleMixin import QuantizedModuleMixin @@ -10,14 +11,11 @@ class BaseLinearSVD( QuantizedModuleMixin, QuantizedLinearMixin, + LoRAFusableLinearMixin, ): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - @abstractmethod - def forward_with_lora(self, x: torch.Tensor, lora_down: torch.nn.Linear, lora_up: torch.nn.Linear, dropout: torch.nn.Dropout, alpha: float) -> torch.Tensor: - pass - def _get_tensor_hash(t: torch.Tensor) -> str: t = t.flatten().to(torch.float32) vals = torch.stack([ @@ -51,6 +49,19 @@ def unquantized_weight(self, dtype: torch.dtype, device: torch.device) -> torch. else: return super().unquantized_weight(dtype, device) + def mark_needs_requantization(self): + # reset both the SVD split flag and the parent's base-weight flag so the next quantize() re-runs fully. + self.__svd_is_quantized = False + super().mark_needs_requantization() + + def predict_offload_bytes(self) -> int: + # the residual quantized weight (base quant type) plus the low-rank factors svd_up (out x rank) and + # svd_down (rank x in), both in svd_dtype. Sized from the meta skeleton -- the factors don't exist yet. + out_features, in_features = self.original_weight_shape() + svd_bytes = (out_features * self.rank + self.rank * in_features) \ + * torch.empty((), dtype=self.svd_dtype).element_size() + return super().predict_offload_bytes() + svd_bytes + @torch.no_grad() def quantize(self, device: torch.device | None = None): if self.__svd_is_quantized: @@ -73,11 +84,17 @@ def quantize(self, device: torch.device | None = None): U, S, Vh = torch.linalg.svd(W, full_matrices=False) if self.cache_dir is not None: + # write to a per-process temp then atomically rename in: under multi-GPU every rank quantizes + # concurrently and writes the same hash-named file, so a plain torch.save races and a reader can + # pick up a half-written file. os.replace is atomic on the same filesystem, so a concurrent reader + # sees either no file or a complete one, and multiple writers just overwrite with identical content. + tmp_filename = filename + f".tmp.{os.getpid()}" torch.save(( U[:, :self.max_cache_rank].clone(), S[:self.max_cache_rank].clone(), Vh[:self.max_cache_rank, :].clone(), - ), filename) + ), tmp_filename) + os.replace(tmp_filename, filename) U_r = U[:, :self.rank] S_r = S[:self.rank] @@ -105,20 +122,44 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: x_up = torch.nn.functional.linear(x_down, self.svd_up) return x_up + super().forward(x) - def forward_with_lora(self, x: torch.Tensor, lora_down: torch.nn.Linear, lora_up: torch.nn.Linear, dropout: torch.nn.Dropout, alpha: float) -> torch.Tensor: + def forward_with_lora(self, x: torch.Tensor, lora_down: torch.Tensor, lora_up: torch.Tensor, dropout: torch.nn.Dropout, alpha: torch.Tensor) -> torch.Tensor: assert self.__svd_is_quantized assert not self.svd_down.requires_grad and not self.svd_up.requires_grad - assert lora_down.bias is None and lora_up.bias is None - lora_rank = lora_down.weight.shape[0] - down_merged = torch.cat([lora_down.weight, self.svd_down], dim=0) + lora_rank = lora_down.shape[0] + flat = x.reshape(-1, x.shape[-1]) + from modules.module.quantized.LinearW8A8 import LinearW8A8 + #the fused kernel needs a LinearW8A8 base and enough tokens to beat the dense merge below + if flat.shape[0] > 16 and isinstance(self, LinearW8A8): + y = self.__fused_forward_with_lora(flat, lora_down, lora_up, dropout, alpha) + return y.reshape(x.shape[:-1] + (y.shape[-1], )) + + #autocast casts the merged matrix to the compute dtype for the matmul anyway, so cast the + #halves before the cat instead of building an f32 buffer to narrow again + down_merged = torch.cat([lora_down.to(self.compute_dtype), self.svd_down.to(self.compute_dtype)], dim=0) x_down = torch.nn.functional.linear(x, down_merged) if dropout.p > 0.0 and self.training: x_down[..., :lora_rank] = dropout(x_down[..., :lora_rank]) - lora_up_scaled = lora_up.weight * (alpha / lora_rank) - up_merged = torch.cat([lora_up_scaled, self.svd_up], dim=1) + lora_up_scaled = (lora_up * (alpha / lora_rank)).to(self.compute_dtype) + up_merged = torch.cat([lora_up_scaled, self.svd_up.to(self.compute_dtype)], dim=1) x_up = torch.nn.functional.linear(x_down, up_merged) return x_up + super().forward(x) + #fuse both low-rank branches into the residual mm's epilogue: concat the LoRA and SVD + #down/up into one rank r_lora + r_svd operand, alpha folded into the LoRA up-block only + #and a dropout mask that's 1 on the svd columns, so only LoRA drops. x is 2D (flattened) + def __fused_forward_with_lora(self, x: torch.Tensor, lora_down: torch.Tensor, lora_up: torch.Tensor, dropout: torch.nn.Dropout, alpha: torch.Tensor) -> torch.Tensor: + lora_rank = lora_down.shape[0] + + down_merged = torch.cat([lora_down.to(self.compute_dtype), self.svd_down.to(self.compute_dtype)], dim=0) + up_fused = torch.cat([(lora_up * (alpha / lora_rank)).to(self.compute_dtype), self.svd_up.to(self.compute_dtype)], dim=1).T + dropout_mask = None + if dropout.training and dropout.p > 0: + svd_rank = self.svd_down.shape[0] + lora_mask = dropout(torch.ones(x.shape[0], lora_rank, device=x.device, dtype=self.compute_dtype)) + dropout_mask = torch.cat([lora_mask, torch.ones(x.shape[0], svd_rank, device=x.device, dtype=self.compute_dtype)], dim=1) + + return self._fused_lora_forward(x, down_merged, up_fused, dropout_mask) + return LinearSVD diff --git a/modules/module/quantized/LinearW8A8.py b/modules/module/quantized/LinearW8A8.py index 8bbb90e5a..7a18a27b1 100644 --- a/modules/module/quantized/LinearW8A8.py +++ b/modules/module/quantized/LinearW8A8.py @@ -1,14 +1,16 @@ - from modules.module.quantized.mixin.CompressedWeightMixin import CompressedWeightMixin +from modules.module.quantized.mixin.LoRAFusableLinearMixin import LoRAFusableLinearMixin from modules.module.quantized.mixin.QuantizedLinearMixin import QuantizedLinearMixin from modules.module.quantized.mixin.QuantizedModuleMixin import QuantizedModuleMixin from modules.util.mm_8bit import mm_8bit as mm_8bit +from modules.util.mm_8bit import scaled_lora_mm_8bit, scaled_mm_8bit from modules.util.quantization_util import ( dequantize, + quantize_axiswise, quantize_fp8_axiswise, - quantize_fp8_tensorwise, + quantize_fp8_tensorwise_chunked, quantize_int8_axiswise, - quantize_int8_tensorwise, + quantize_int8_tensorwise_chunked, ) import torch @@ -16,40 +18,64 @@ @torch.no_grad() -def int8_forward_tokenwise(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: - x_8, x_scale = quantize_int8_axiswise(x, dim=-1) - res = torch._int_mm(x_8, weight.T) - res_scaled = res.float().mul_(weight_scale * x_scale).to(compute_dtype) +def forward_tokenwise_postscaled_torch(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: + if weight.dtype == torch.int8: + x_8, x_scale = quantize_int8_axiswise(x, dim=-1) + res = torch._int_mm(x_8, weight.T) + res_scaled = res.float().mul_(weight_scale * x_scale).to(compute_dtype) + else: + x_8, x_scale = quantize_fp8_axiswise(x, dim=-1) + one = torch.tensor(1.0, device=x.device) + res = torch._scaled_mm(x_8, weight.T, scale_a=one, scale_b=weight_scale.float(), out_dtype=torch.float) + res_scaled = res.mul_(x_scale).to(compute_dtype) #much faster than scaled by _scaled_mm if bias is not None: res_scaled.add_(bias) return res_scaled @torch.no_grad() -def fp8_forward_tokenwise(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: - x_8, x_scale = quantize_fp8_axiswise(x, dim=-1) - one = torch.tensor(1.0, device=x.device) - res = torch._scaled_mm(x_8, weight.T, scale_a=one, scale_b=weight_scale.float(), out_dtype=torch.float) - res_scaled = res.mul_(x_scale).to(compute_dtype) #much faster than scaled by _scaled_mm +def forward_tokenwise_epiloguescaled_triton(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: + x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=weight.dtype) + res_scaled = scaled_mm_8bit(x_8, weight.T, weight_scale * x_scale, compute_dtype) if bias is not None: res_scaled.add_(bias) return res_scaled @torch.no_grad() -def int8_backward_axiswise(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: - output_8, output_scale = quantize_int8_axiswise(output, dim=-1) +def backward_tokenwise_postscaled_triton(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=weight.dtype) #almost always, grad outputs are already contiguous and this is a no-op. But there are some grad outputs from SDXL that are non-contiguous: mm_res = mm_8bit(output_8.contiguous(), weight) return mm_res.float().mul_(weight_scale * output_scale).to(output.dtype) @torch.no_grad() -def fp8_backward_axiswise(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: - output_8, output_scale = quantize_fp8_axiswise(output, dim=-1) - mm_res = mm_8bit(output_8.contiguous(), weight) - return mm_res.float().mul_(weight_scale * output_scale).to(output.dtype) +def backward_tokenwise_epiloguescaled_triton(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=weight.dtype) + return scaled_mm_8bit(output_8.contiguous(), weight, weight_scale * output_scale, output.dtype) + + +forward_tokenwise = forward_tokenwise_epiloguescaled_triton +backward_tokenwise = backward_tokenwise_epiloguescaled_triton + + +@torch.no_grad() +def forward_tokenwise_lora_epiloguescaled_triton(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype, x_down: Tensor, lora_up: Tensor) -> Tensor: + x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=weight.dtype) + res_scaled = scaled_lora_mm_8bit(x_8, weight.T, weight_scale * x_scale, x_down, lora_up, compute_dtype) + if bias is not None: + res_scaled.add_(bias) + return res_scaled + +@torch.no_grad() +def backward_tokenwise_lora_epiloguescaled_triton(output: Tensor, weight: Tensor, weight_scale: Tensor, grad_x_down_pre: Tensor, lora_down: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=weight.dtype) + return scaled_lora_mm_8bit(output_8.contiguous(), weight, weight_scale * output_scale, grad_x_down_pre, lora_down.to(grad_x_down_pre.dtype), output.dtype) + + +forward_tokenwise_lora = forward_tokenwise_lora_epiloguescaled_triton +backward_tokenwise_lora = backward_tokenwise_lora_epiloguescaled_triton class LinearW8A8Function(torch.autograd.Function): - #int8 and fp8 differ only in which math runs, and the quantized weight already carries the dtype @staticmethod def forward(ctx, x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: # `weight` is the decompressed weight, so saving it keeps a full-size copy alive until backward. @@ -59,10 +85,7 @@ def forward(ctx, x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | # TODO once offloading uses non-reentrant checkpointing, consider not saving the decompressed # weight and decoding in backward() instead. ctx.save_for_backward(weight, weight_scale) - if weight.dtype == torch.int8: - return int8_forward_tokenwise(x, weight, weight_scale, bias, compute_dtype) - else: - return fp8_forward_tokenwise(x, weight, weight_scale, bias, compute_dtype) + return forward_tokenwise(x, weight, weight_scale, bias, compute_dtype) @staticmethod def backward(ctx, output: Tensor): @@ -70,10 +93,40 @@ def backward(ctx, output: Tensor): raise NotImplementedError("Int/Float A8W8 cannot be used for full finetuning") weight, weight_scale = ctx.saved_tensors - if weight.dtype == torch.int8: - return int8_backward_axiswise(output, weight, weight_scale), None, None, None, None - else: - return fp8_backward_axiswise(output, weight, weight_scale), None, None, None, None + return backward_tokenwise(output, weight, weight_scale), None, None, None, None + + +class LinearW8A8LoRAFunction(torch.autograd.Function): + #LinearW8A8Function plus a low-rank update: it owns the down-projection and the dropout, so the LoRA + #dgrad folds into the backward epilogue and the (M, out) product is never materialized. The weight is + #the already decompressed one, as in LinearW8A8Function + @staticmethod + def forward(ctx, x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype, lora_down: Tensor, lora_up: Tensor, dropout_mask: Tensor | None) -> Tensor: + x_down = torch.nn.functional.linear(x, lora_down) + if dropout_mask is not None: + x_down = x_down * dropout_mask + #backward runs outside autocast, so x is saved in x_down's dtype: grad_lora_down is a + #compute-dtype product there, as it is in the unfused path where autocast casts x + ctx.save_for_backward(weight, weight_scale, x.to(x_down.dtype), lora_down, lora_up, dropout_mask, x_down) + return forward_tokenwise_lora(x, weight, weight_scale, bias, compute_dtype, x_down, lora_up) + + @staticmethod + def backward(ctx, grad_output: Tensor): + if ctx.needs_input_grad[1:5] != (False, False, False, False): + raise NotImplementedError("Int/Float A8W8 cannot be used for full finetuning") + + weight, weight_scale, x, lora_down, lora_up, dropout_mask, x_down = ctx.saved_tensors + needs_x, needs_down, needs_up = ctx.needs_input_grad[0], ctx.needs_input_grad[5], ctx.needs_input_grad[6] + #grad_x_down_pre (post-dropout-backward) feeds both the grad_x epilogue fold and grad_lora_down + grad_x_down_pre = None + if needs_x or needs_down: + grad_x_down_pre = grad_output @ lora_up.T + if dropout_mask is not None: + grad_x_down_pre = grad_x_down_pre * dropout_mask + grad_x = backward_tokenwise_lora(grad_output, weight, weight_scale, grad_x_down_pre, lora_down) if needs_x else None + grad_lora_down = (grad_x_down_pre.T @ x).to(lora_down.dtype) if needs_down else None + grad_lora_up = x_down.T @ grad_output if needs_up else None + return grad_x, None, None, None, None, grad_lora_down, grad_lora_up, None class LinearW8A8( @@ -81,43 +134,68 @@ class LinearW8A8( QuantizedModuleMixin, QuantizedLinearMixin, CompressedWeightMixin, + LoRAFusableLinearMixin, ): - def __init__(self, dtype: torch.dtype, *args, **kwargs): + is_quantized: bool + + def __init__(self, dtype: torch.dtype, compress: bool = False, *args, **kwargs): super().__init__(*args, **kwargs) assert dtype in [torch.int8, torch.float8_e4m3fn] self._dtype = dtype - self.__is_quantized = False + self.is_quantized = False self.compute_dtype = None self.register_buffer("scale", torch.tensor(1.0, dtype=torch.float32)) - self._init_compressed_state() + self._init_compressed_state(compress) def original_weight_shape(self) -> tuple[int, ...]: if self._compressed: return self._weight_shape return self.weight.shape + def mark_needs_requantization(self): + self.is_quantized = False + self.mark_needs_recompression() + + def predict_offload_bytes(self) -> int: + # weight quantizes tensorwise to int8/float8_e4m3fn (both 1 byte/elem, same shape); bias is left + # unchanged. Matches get_offload_tensors (weight + optional bias); the scalar scale buffer is not + # offload-counted. _dtype is asserted int8/float8_e4m3fn in __init__, so 1 byte/elem is exact. + # a compressed weight offloads as its blob, so the measured length replaces the element count once it exists. + # With compression on the element count is not an approximation but a wrong answer -- it would size every + # offload arena ~40% too large -- so a caller that sizes before __measure_compressed_sizes has run is a bug + # in the sizing order rather than something to paper over. + if self.compress and self._compressed_bytes is None: + raise RuntimeError( + "offload sizing for a compressed weight whose blob length has not been measured yet; " + "LayerOffloadConductor.__measure_compressed_sizes has to run before the arenas are sized") + weight_bytes = self._compressed_bytes if self._compressed_bytes is not None else self.weight.numel() + bias_bytes = self.bias.numel() * self.bias.element_size() if self.bias is not None else 0 + return weight_bytes + bias_bytes + def unquantized_weight(self, dtype: torch.dtype, device: torch.device) -> torch.Tensor: + if not self.is_quantized: + return self.weight.detach().to(dtype) weight = self._decompress(self.weight.detach()) if self._compressed else self.weight.detach() # 'scale' is not offloaded, so it can sit on the train device while 'weight' is parked on the temp device return dequantize(weight, self.scale.to(device=weight.device)).to(dtype) @torch.no_grad() def quantize(self, device: torch.device | None = None): - if self.__is_quantized: + if self.is_quantized: return - self.__is_quantized = True + self.is_quantized = True weight = self.weight.detach() orig_device = weight.device if device is not None: weight = weight.to(device=device) if self._dtype == torch.int8: - weight, scale = quantize_int8_tensorwise(weight) + weight, scale = quantize_int8_tensorwise_chunked(weight) else: - weight, scale = quantize_fp8_tensorwise(weight) + weight, scale = quantize_fp8_tensorwise_chunked(weight) if device is not None: weight = weight.to(device=orig_device) @@ -125,14 +203,15 @@ def quantize(self, device: torch.device | None = None): self.requires_grad_(False) self.weight.data = weight - self.scale.copy_(scale) + # keep the scale on the weight's device so the batched int8/fp8 path finds it co-located there + self.scale = scale.detach().to(orig_device) if self.compress: self._compress_weight(device=device) def forward(self, x_orig: torch.Tensor) -> torch.Tensor: assert not self.weight.requires_grad - assert self.__is_quantized + assert self.is_quantized x = x_orig.reshape(-1, x_orig.shape[-1]) weight = self._decompress(self.weight.detach()) if self._compressed else self.weight @@ -145,6 +224,30 @@ def forward(self, x_orig: torch.Tensor) -> torch.Tensor: return y.reshape(x_orig.shape[:-1] + (y.shape[-1], )) + def forward_with_lora(self, x_orig: torch.Tensor, lora_down: torch.Tensor, lora_up: torch.Tensor, dropout: torch.nn.Dropout, alpha: torch.Tensor) -> torch.Tensor: + assert self.is_quantized + lora_rank = lora_down.shape[0] + + x = x_orig.reshape(-1, x_orig.shape[-1]) + if x.shape[0] <= 16: + ld = torch.nn.functional.linear(dropout(torch.nn.functional.linear(x_orig, lora_down)), lora_up) + return LinearW8A8.forward(self, x_orig) + ld * (alpha / lora_rank) + + #the cast matches lora_up (kept in its own dtype, e.g. f32) to the down-projection's autocast + #dtype, which the epilogue dot needs; dropout(ones) yields the scaled 0/(1/(1-p)) mask, drawn + #here rather than inside LinearW8A8LoRAFunction so the RNG stays inductor-native + lora_up_scaled = (lora_up * (alpha / lora_rank)).T.to(self.compute_dtype) + dropout_mask = dropout(torch.ones(x.shape[0], lora_rank, device=x.device, dtype=self.compute_dtype)) if (dropout.training and dropout.p > 0) else None + y = self._fused_lora_forward(x, lora_down, lora_up_scaled, dropout_mask) + + return y.reshape(x_orig.shape[:-1] + (y.shape[-1], )) + + #x is 2D, lora_down is (r, in_features), lora_up is (r, out_features) with alpha already folded in + def _fused_lora_forward(self, x: Tensor, lora_down: Tensor, lora_up: Tensor, dropout_mask: Tensor | None) -> Tensor: + assert not self.weight.requires_grad + weight = self._decompress(self.weight.detach()) if self._compressed else self.weight + return LinearW8A8LoRAFunction.apply(x, weight, self.scale, self.bias, self.compute_dtype, lora_down, lora_up, dropout_mask) + def run_benchmark(fn, desc, steps=10000, warmup=500, compile=False): if compile: fn = torch.compile(fn, fullgraph=True) @@ -158,7 +261,7 @@ def run_benchmark(fn, desc, steps=10000, warmup=500, compile=False): @torch.no_grad() -def benchmark_int8(m, k, n, device = 'cuda'): +def benchmark_int8(m, k, n, device = 'cuda', steps = 10000): x = torch.randn(m,k, device=device, dtype=torch.bfloat16) x_8 = torch.ones (m,k, device=device, dtype=torch.int8) y = torch.randn(m,n, device=device, dtype=torch.bfloat16) @@ -167,19 +270,21 @@ def benchmark_int8(m, k, n, device = 'cuda'): w_scale = torch.ones(1, device=device) - run_benchmark(lambda: torch._int_mm(x_8, w_8.T), "torch mm int") - run_benchmark(lambda: mm_8bit(x_8, w_8.T), "triton mm int") + run_benchmark(lambda: torch._int_mm(x_8, w_8.T), "torch mm int", steps=steps) + run_benchmark(lambda: mm_8bit(x_8, w_8.T), "triton mm int", steps=steps) def torch_backward(a, b): torch._int_mm(a, b.T.contiguous().T) - run_benchmark(lambda: torch_backward(y_8, w_8), "torch mm backward int8") - run_benchmark(lambda: mm_8bit(y_8, w_8), "triton mm backward int8") + run_benchmark(lambda: torch_backward(y_8, w_8), "torch mm backward int8", steps=steps) + run_benchmark(lambda: mm_8bit(y_8, w_8), "triton mm backward int8", steps=steps) - run_benchmark(lambda: int8_forward_tokenwise(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "torch forward int", compile=True) - run_benchmark(lambda: int8_backward_axiswise(y, w_8, w_scale), "triton backward int", compile=True) + run_benchmark(lambda: forward_tokenwise_postscaled_torch(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "torch forward int", steps=steps, compile=True) + run_benchmark(lambda: backward_tokenwise_postscaled_triton(y, w_8, w_scale), "triton backward int", steps=steps, compile=True) + run_benchmark(lambda: forward_tokenwise_epiloguescaled_triton(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "triton scaled forward int", steps=steps, compile=True) + run_benchmark(lambda: backward_tokenwise_epiloguescaled_triton(y, w_8, w_scale), "triton scaled backward int", steps=steps, compile=True) @torch.no_grad() -def benchmark_fp8(m, k, n, device = 'cuda'): +def benchmark_fp8(m, k, n, device = 'cuda', steps = 10000): x = torch.randn(m,k, device=device, dtype=torch.bfloat16) x_8 = torch.ones (m,k, device=device, dtype=torch.float8_e4m3fn) y = torch.randn(m,n, device=device, dtype=torch.bfloat16) @@ -188,16 +293,93 @@ def benchmark_fp8(m, k, n, device = 'cuda'): w_scale = torch.ones(1, device=device, dtype=torch.bfloat16) one_scale = torch.ones(1, device=device) - run_benchmark(lambda: torch._scaled_mm(x_8, w_8.T, out_dtype=torch.bfloat16, scale_a=one_scale.float(), scale_b=w_scale.float()), "torch mm fp8") - run_benchmark(lambda: mm_8bit(x_8, w_8.T), "triton mm fp8") + run_benchmark(lambda: torch._scaled_mm(x_8, w_8.T, out_dtype=torch.bfloat16, scale_a=one_scale.float(), scale_b=w_scale.float()), "torch mm fp8", steps=steps) + run_benchmark(lambda: mm_8bit(x_8, w_8.T), "triton mm fp8", steps=steps) def torch_backward(a, b): torch._scaled_mm(a, b.T.contiguous().T, out_dtype=torch.bfloat16, scale_a=one_scale.float(), scale_b=w_scale.float()) - run_benchmark(lambda: torch_backward(y_8, w_8), "torch mm backward fp8") - run_benchmark(lambda: mm_8bit(y_8, w_8), "triton mm backward fp8") - run_benchmark(lambda: fp8_forward_tokenwise(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "torch forward fp8", compile=True) - run_benchmark(lambda: fp8_backward_axiswise(y, w_8, w_scale), "triton backward fp8", compile=True) + run_benchmark(lambda: torch_backward(y_8, w_8), "torch mm backward fp8", steps=steps) + run_benchmark(lambda: mm_8bit(y_8, w_8), "triton mm backward fp8", steps=steps) + + run_benchmark(lambda: forward_tokenwise_postscaled_torch(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "torch forward fp8", steps=steps, compile=True) + run_benchmark(lambda: forward_tokenwise_epiloguescaled_triton(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "triton scaled forward fp8", steps=steps, compile=True) + run_benchmark(lambda: backward_tokenwise_epiloguescaled_triton(y, w_8, w_scale), "triton scaled backward fp8", steps=steps, compile=True) + + +@torch.no_grad() +def benchmark_lora(m, k, n, r, dtype, device='cuda', steps=1000): + is_int8 = (dtype == torch.int8) + x = torch.randn(m, k, device=device, dtype=torch.bfloat16) + if is_int8: + w_8 = torch.randint(-127, 127, (n, k), device=device, dtype=torch.int8) + else: + w_8 = torch.randn(n, k, device=device).to(torch.float8_e4m3fn) + w_scale = torch.full((1,), 0.01, device=device) + down_w = torch.randn(r, k, device=device, dtype=torch.bfloat16) * 0.02 + up_w = torch.randn(n, r, device=device, dtype=torch.bfloat16) * 0.02 + + baseline = forward_tokenwise + fused = forward_tokenwise_lora_epiloguescaled_triton + + def run_unfused(): + y = baseline(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16) + return y + (x @ down_w.T) @ up_w.T + + def run_fused(): + x_down = x @ down_w.T + return fused(x, w_8, w_scale, None, torch.bfloat16, x_down, up_w.T) + + name = "int8" if is_int8 else "fp8" + diff = (run_unfused().float() - run_fused().float()).abs().max().item() + ref = run_unfused().float().abs().mean().item() + print(f"lora {name} m={m} k={k} n={n} r={r}: max abs diff fused vs unfused = {diff:.4f} (mean magnitude {ref:.3f})") + run_benchmark(lambda: baseline(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), f"no-lora {name} baseline", steps=steps, compile=True) + run_benchmark(run_unfused, f"unfused lora {name}", steps=steps, compile=True) + run_benchmark(run_fused, f"fused lora {name}", steps=steps, compile=True) + + +@torch.no_grad() +def benchmark_lora_backward(m, k, n, r, dtype, device='cuda', steps=1000): + #isolates the grad_x path that the backward fusion changes: unfused = main bwd mm + + #standalone (M,K) lora dgrad + merge add; fused = the same lora dgrad folded into the + #bwd mm epilogue. grad_x_down (= grad_y @ up.T) is shared and timed in both; the + #grad_lora_up / grad_lora_down gemms are identical in both variants and excluded. + is_int8 = (dtype == torch.int8) + grad_y = torch.randn(m, n, device=device, dtype=torch.bfloat16) + if is_int8: + w_8 = torch.randint(-127, 127, (n, k), device=device, dtype=torch.int8) + else: + w_8 = torch.randn(n, k, device=device).to(torch.float8_e4m3fn) + w_scale = torch.full((1,), 0.01, device=device) + down_w = torch.randn(r, k, device=device, dtype=torch.bfloat16) * 0.02 + up_w = torch.randn(n, r, device=device, dtype=torch.bfloat16) * 0.02 + + main_bwd = backward_tokenwise + fused_bwd = backward_tokenwise_lora + + def run_unfused(): + grad_x = main_bwd(grad_y, w_8, w_scale) + grad_x_down = grad_y @ up_w + return grad_x + grad_x_down @ down_w + + def run_fused(): + grad_x_down = grad_y @ up_w + return fused_bwd(grad_y, w_8, w_scale, grad_x_down, down_w) + + name = "int8" if is_int8 else "fp8" + diff = (run_unfused().float() - run_fused().float()).abs().max().item() + ref = run_unfused().float().abs().mean().item() + print(f"lora-bwd {name} m={m} k={k} n={n} r={r}: max abs diff fused vs unfused = {diff:.4f} (mean magnitude {ref:.3f})") + run_benchmark(run_unfused, f"unfused lora-bwd {name}", steps=steps, compile=True) + run_benchmark(run_fused, f"fused lora-bwd {name}", steps=steps, compile=True) if __name__ == "__main__": + #ragged shape: M and the backward's contraction over 3088 are both non-%128, so the + #masked-loop fallback and, with r=12, the padded rank tile are covered benchmark_int8(2 * 1024 + 50, 3072, 3072 + 16) benchmark_fp8(2 * 1024 + 50, 3072, 3072 + 16) + benchmark_lora(2 * 1024 + 50, 3072, 3072 + 16, r=12, dtype=torch.int8) + benchmark_lora_backward(2 * 1024 + 50, 3072, 3072 + 16, r=12, dtype=torch.int8) + #a real FLUX.2-klein-9B shape at 512px batch 2: M = 2*(1024 image + 512 text) tokens + #through a double block's 4096 -> 4096 attention projection + benchmark_lora(2 * (1024 + 512), 4096, 4096, r=16, dtype=torch.int8, steps=1000) diff --git a/modules/module/quantized/mixin/CompressedWeightMixin.py b/modules/module/quantized/mixin/CompressedWeightMixin.py index f114cd07d..3d5730b28 100644 --- a/modules/module/quantized/mixin/CompressedWeightMixin.py +++ b/modules/module/quantized/mixin/CompressedWeightMixin.py @@ -7,12 +7,13 @@ class CompressedWeightMixin(metaclass=ABCMeta): - def _init_compressed_state(self): - self.compress = False + def _init_compressed_state(self, compress: bool): + self.compress = compress self._compressed = False self._weight_shape = None self._uncompressed_bytes = 0 self._compressed_dtype = None + self._compressed_bytes = None def _decompress(self, blob: torch.Tensor) -> torch.Tensor: # decoding only runs on the GPU. DoRA calls this during initialization, when the weight can @@ -25,10 +26,21 @@ def _decompress(self, blob: torch.Tensor) -> torch.Tensor: def uncompressed_bytes(self) -> int: # bytes the weight occupies decompressed; weight.nbytes is the stored size and drops to the # blob length once compressed - if not self._compressed: + if self._compressed_bytes is None: return self.weight.nbytes return self._uncompressed_bytes + def compressed_bytes(self) -> int | None: + # None until the layer has been compressed once. Measured rather than read off self.weight, which no + # longer holds the blob after an eviction. + return self._compressed_bytes + + def mark_needs_recompression(self): + # a re-quantize rebuilds the weight from the checkpoint and drops the blob; without this, _compress_weight's + # early-out would leave an uncompressed weight flagged compressed and forward() would decode non-blob bytes. + # _compressed_bytes stays: the length belongs to the checkpoint weight, not to this materialize. + self._compressed = False + @torch.no_grad() def _compress_weight(self, device: torch.device | None = None): if self._compressed: @@ -43,6 +55,7 @@ def _compress_weight(self, device: torch.device | None = None): self._weight_shape = tuple(gpu_weight.shape) self._compressed_dtype = gpu_weight.dtype blob, self._uncompressed_bytes = nvcomp_util.compress(gpu_weight.contiguous()) + self._compressed_bytes = blob.numel() self._compressed = True if device is not None: diff --git a/modules/module/quantized/mixin/LoRAFusableLinearMixin.py b/modules/module/quantized/mixin/LoRAFusableLinearMixin.py new file mode 100644 index 000000000..784b9de0d --- /dev/null +++ b/modules/module/quantized/mixin/LoRAFusableLinearMixin.py @@ -0,0 +1,9 @@ +from abc import ABCMeta, abstractmethod + +import torch + + +class LoRAFusableLinearMixin(metaclass=ABCMeta): + @abstractmethod + def forward_with_lora(self, x: torch.Tensor, lora_down: torch.Tensor, lora_up: torch.Tensor, dropout: torch.nn.Dropout, alpha: torch.Tensor) -> torch.Tensor: + pass diff --git a/modules/module/quantized/mixin/QuantizedLinearMixin.py b/modules/module/quantized/mixin/QuantizedLinearMixin.py index cac28d442..aacc1dee1 100644 --- a/modules/module/quantized/mixin/QuantizedLinearMixin.py +++ b/modules/module/quantized/mixin/QuantizedLinearMixin.py @@ -15,3 +15,15 @@ def original_weight_shape(self) -> tuple[int, ...]: @abstractmethod def unquantized_weight(self, dtype: torch.dtype, device: torch.device) -> torch.Tensor: pass + + @abstractmethod + def mark_needs_requantization(self): + # reset the concrete class's is-quantized flag so the next materialize re-quantizes. Called by streaming + # eviction, which discards the packed weights back to meta. + pass + + def predict_offload_bytes(self) -> int: + # post-quantization offload footprint, predicted from the unpacked skeleton shape while the module is still a + # meta skeleton (the real packed tensors don't exist yet). + raise NotImplementedError( + f"{type(self).__name__} does not implement predict_offload_bytes (disk-offload conductor sizing)") diff --git a/modules/trainer/BaseTrainer.py b/modules/trainer/BaseTrainer.py index 20fc5eb7a..87a96d7d0 100644 --- a/modules/trainer/BaseTrainer.py +++ b/modules/trainer/BaseTrainer.py @@ -97,7 +97,11 @@ def _start_tensorboard(self): if self.config.tensorboard_expose: tensorboard_args.append("--bind_all") - self.tensorboard_subprocess = subprocess.Popen(tensorboard_args) + # discard the child's banner and notices; the UI already shows the tensorboard URL. + # Popen still raises if the executable is missing. + self.tensorboard_subprocess = subprocess.Popen( + tensorboard_args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) def _stop_tensorboard(self): self.tensorboard_subprocess.kill() diff --git a/modules/trainer/GenericTrainer.py b/modules/trainer/GenericTrainer.py index a6547afbe..653f3dcac 100644 --- a/modules/trainer/GenericTrainer.py +++ b/modules/trainer/GenericTrainer.py @@ -11,7 +11,7 @@ from modules.dataLoader.BaseDataLoader import BaseDataLoader from modules.model.BaseModel import BaseModel from modules.modelLoader.BaseModelLoader import BaseModelLoader -from modules.modelSampler.BaseModelSampler import BaseModelSampler, ModelSamplerOutput +from modules.modelSampler.BaseModelSampler import BaseModelSampler from modules.modelSaver.BaseModelSaver import BaseModelSaver from modules.modelSetup.BaseModelSetup import BaseModelSetup from modules.trainer.BaseTrainer import BaseTrainer @@ -32,6 +32,7 @@ from modules.util.profiling_util import PeakMemoryRecorder, TorchMemoryRecorder, TorchProfiler from modules.util.time_util import get_string_timestamp from modules.util.torch_util import torch_gc +from modules.util.tqdm_util import tqdm from modules.util.TrainProgress import TrainProgress import torch @@ -41,8 +42,6 @@ from torch.utils.tensorboard import SummaryWriter from torchvision.transforms.functional import pil_to_tensor -from tqdm import tqdm - # OT_DEBUG_PROFILES=1 dumps a CUDA memory snapshot for the first two steps, where the allocator is still # growing, and a profiler trace at steps 10 and 40, past compilation and warmup. _DEBUG_PROFILES = os.environ.get("OT_DEBUG_PROFILES") == "1" @@ -137,6 +136,8 @@ def start(self): model_names=model_names, weight_dtypes=self.config.weight_dtypes(), quantization=self.config.quantization, + stream_from_disk=self.config.stream_from_disk, + cache_in_ram=self.config.cache_in_ram(), ) self.model.train_config = self.config @@ -204,11 +205,14 @@ def __prune_backups(self, backups_to_keep: int): def __enqueue_sample_during_training(self, fun: Callable): self.sample_queue.append(fun) - def __execute_sample_during_training(self): + def __execute_sample_during_training(self) -> bool: + # returns whether any samples were run + sampled = bool(self.sample_queue) with PeakMemoryRecorder("sampling", enabled=False): for fun in self.sample_queue: fun() self.sample_queue = [] + return sampled def __sample_loop( self, @@ -219,12 +223,15 @@ def __sample_loop( folder_postfix: str = "", is_custom_sample: bool = False, ): - for i, sample_config in multi.distributed( - [(i, sample_config) for i, sample_config in enumerate(sample_config_list) if sample_config.enabled], - distribute=not self.config.samples_to_tensorboard and not ema_applied - ): - try: - safe_prompt = path_util.safe_filename(sample_config.prompt) + on_update_progress = self.callbacks.on_update_sample_custom_progress if is_custom_sample else self.callbacks.on_update_sample_default_progress + + try: + jobs = [] # (sample_config, destination, index, prompt_label) for this rank's share + for i, sample_config in multi.distributed( + [(i, sample_config) for i, sample_config in enumerate(sample_config_list) if sample_config.enabled], + distribute=not self.config.samples_to_tensorboard and not ema_applied + ): + prompt_label = path_util.safe_filename(sample_config.prompt) if is_custom_sample: sample_dir = os.path.join( @@ -236,47 +243,50 @@ def __sample_loop( sample_dir = os.path.join( self.config.workspace_dir, "samples", - f"{str(i)} - {safe_prompt}{folder_postfix}", + f"{str(i)} - {prompt_label}{folder_postfix}", ) + # custom samples all share the samples/custom dir (regular ones each get their own + # {i} - prompt dir), so add the index here - without it a batch that finishes within the + # same second shares one filename (timestamp + filename_string are identical) and overwrites. + sample_suffix = f"-{str(i)}" if is_custom_sample else "" sample_path = os.path.join( sample_dir, - f"{self.config.save_filename_prefix}{get_string_timestamp()}-training-sample-{train_progress.filename_string()}" + f"{self.config.save_filename_prefix}{get_string_timestamp()}-training-sample-{train_progress.filename_string()}{sample_suffix}" ) - def on_sample_default(sampler_output: ModelSamplerOutput): - if self.config.samples_to_tensorboard and sampler_output.file_type == FileType.IMAGE: - self.tensorboard.add_image( - f"sample{str(i)} - {safe_prompt}", pil_to_tensor(sampler_output.data), # noqa: B023 - train_progress.global_step - ) - self.callbacks.on_sample_default(sampler_output) - - def on_sample_custom(sampler_output: ModelSamplerOutput): - self.callbacks.on_sample_custom(sampler_output) + sample_config = copy.copy(sample_config) + sample_config.from_train_config(self.config) - on_sample = on_sample_custom if is_custom_sample else on_sample_default - on_update_progress = self.callbacks.on_update_sample_custom_progress if is_custom_sample else self.callbacks.on_update_sample_default_progress + jobs.append((sample_config, sample_path, i, prompt_label)) + if jobs: self.model.eval() - sample_config = copy.copy(sample_config) - sample_config.from_train_config(self.config) - - self.model_sampler.sample( - sample_config=sample_config, - destination=sample_path, - image_format=self.config.sample_image_format, - video_format=self.config.sample_video_format, - audio_format=self.config.sample_audio_format, - on_sample=on_sample, + sampler_outputs = self.model_sampler.sample_all( + [sample_config for sample_config, _, _, _ in jobs], + [destination for _, destination, _, _ in jobs], + self.config.sample_image_format, + self.config.sample_video_format, + self.config.sample_audio_format, on_update_progress=on_update_progress, ) - except Exception: - traceback.print_exc() - tqdm.write("Error during sampling, proceeding without sampling") - torch_gc() + for (_, _, i, prompt_label), sampler_output in zip(jobs, sampler_outputs, strict=True): + if is_custom_sample: + self.callbacks.on_sample_custom(sampler_output) + else: + if self.config.samples_to_tensorboard and sampler_output.file_type == FileType.IMAGE: + self.tensorboard.add_image( + f"sample{str(i)} - {prompt_label}", pil_to_tensor(sampler_output.data), + train_progress.global_step + ) + self.callbacks.on_sample_default(sampler_output) + except Exception: + traceback.print_exc() + tqdm.write("Error during sampling, proceeding without sampling") + + torch_gc() def __sample_during_training( self, @@ -664,7 +674,9 @@ def train(self): torch.clear_autocast_cache() self.model.optimizer.train() - torch_gc() + # no torch_gc here: setup_train_device above ends in materialize_only(), which collects whenever a + # part actually moved. On an epoch that changed nothing (the usual case once the caches are warm) + # there is nothing to reclaim, and this collection cost ~350ms of the epoch gap on a large model. if lr_scheduler is None: lr_scheduler = create.create_lr_scheduler( @@ -717,7 +729,16 @@ def sample_commands_fun(): torch_gc() if not has_gradient: - self.__execute_sample_during_training() + if self.__execute_sample_during_training(): + # a custom sample (UI "sample now") can arrive while that sample was running; it only + # reaches us on the command channel, so keep sampling until none are left instead of + # making it wait for the next training step + while True: + multi.sync_commands(self.commands) + sample_commands = self.commands.get_and_reset_sample_custom_commands() + if not sample_commands: + break + self.__sample_during_training(train_progress, train_device, sample_commands) backup = self.commands.get_and_reset_backup_command() save = self.commands.get_and_reset_save_command() if multi.is_master() and (backup or save): diff --git a/modules/ui/BaseModelTabView.py b/modules/ui/BaseModelTabView.py index 3b48a46f0..0609098a5 100644 --- a/modules/ui/BaseModelTabView.py +++ b/modules/ui/BaseModelTabView.py @@ -36,11 +36,14 @@ def build_content(self, frame, controller, ui_state): allow_override_transformer=controller.supports_override_transformer(), has_unconditional_transformer="unconditional_transformer" in parts, has_text_encoder=not model_type.has_multiple_text_encoders(), + allow_override_text_encoder=controller.supports_override_text_encoder(), has_text_encoder_1=model_type.has_multiple_text_encoders(), has_text_encoder_2="text_encoder_2" in parts, has_text_encoder_3="text_encoder_3" in parts, has_text_encoder_4="text_encoder_4" in parts, allow_override_text_encoder_4="text_encoder_4" in parts, + has_connectors="connectors" in parts, + has_low_noise_transformer="low_noise_transformer" in parts, has_vae="vae" in parts, include_compressed=model_type.supports_compression(), ) @@ -111,6 +114,11 @@ def __create_base_dtype_components(self, frame, row: int, ui_state) -> int: mode="dir", placeholder=HF_HUB_CACHE, ) + # stream from disk + self.components.label(frame, row, 3, "Stream From Disk", + tooltip="Uses the streaming model loader to stream frozen weights from disk to VRAM on demand, greatly reducing RAM usage. Only turn off if you hit compatibility issues.") + self.components.switch(frame, row, 4, ui_state, "stream_from_disk") + row += 1 # base model @@ -143,23 +151,17 @@ def __create_base_components( allow_override_transformer: bool = False, has_unconditional_transformer: bool = False, allow_override_text_encoder_4: bool = False, + allow_override_text_encoder: bool = False, has_text_encoder: bool = False, has_text_encoder_1: bool = False, has_text_encoder_2: bool = False, has_text_encoder_3: bool = False, has_text_encoder_4: bool = False, + has_connectors: bool = False, + has_low_noise_transformer: bool = False, has_vae: bool = False, include_compressed: bool = False, ) -> int: - if has_unet: - # unet weight dtype - self.components.label(frame, row, 3, "UNet Data Type", - tooltip="The unet weight data type") - self.components.options_kv(frame, row, 4, self.__create_dtype_options(include_a8=True, include_compressed=include_compressed), - ui_state, "unet.weight_dtype") - - row += 1 - if has_prior: if allow_override_prior: # prior model @@ -178,35 +180,22 @@ def __create_base_components( row += 1 - if has_transformer: - if allow_override_transformer: - # transformer model - self.components.label(frame, row, 0, "Override Transformer / GGUF", - tooltip="Can be used to override the transformer in the base model. Safetensors and GGUF files are supported, local and on Huggingface. If a GGUF file is used, the DataType must also be set to GGUF") - self.components.path_entry( - frame, row, 1, ui_state, "transformer.model_name", - mode="file", path_modifier=path_util.json_path_modifier - ) - - # transformer weight dtype - self.components.label(frame, row, 3, "Transformer Data Type", - tooltip="The transformer weight data type") - self.components.options_kv(frame, row, 4, self.__create_dtype_options(include_gguf=True, include_a8=True, include_compressed=include_compressed), - ui_state, "transformer.weight_dtype") - - row += 1 - - if has_unconditional_transformer: - # unconditional transformer weight dtype - self.components.label(frame, row, 3, "Unconditional Transformer Data Type", - tooltip="The weight data type of the unconditional transformer, used for the negative branch of CFG during sampling") - self.components.options_kv(frame, row, 4, self.__create_dtype_options(include_a8=True, include_compressed=include_compressed), - ui_state, "unconditional_transformer.weight_dtype") + if has_transformer and allow_override_transformer: + # transformer model + self.components.label(frame, row, 0, "Override Transformer / GGUF", + tooltip="Can be used to override the transformer in the base model. Safetensors and GGUF files are supported, local and on Huggingface. If a GGUF file is used, the DataType must also be set to GGUF") + self.components.path_entry( + frame, row, 1, ui_state, "transformer.model_name", + mode="file", path_modifier=path_util.json_path_modifier + ) row += 1 presets = controller.get_presets() + # Quantization Layer Filter (col 0/1) is a tall widget (preset row, custom entry row, regex row). + # UNet/Transformer Data Type and the quantization Fallback Data Type share this row too (col 3/4), + # lined up so the data type row matches the preset row, and the fallback row matches the entry row. self.components.label(frame, row, 0, "Quantization") self.components.layer_filter_entry(frame, row, 1, ui_state, preset_var_name="quantization.layer_filter_preset", presets=presets, @@ -219,7 +208,38 @@ def __create_base_components( frame_color="transparent", ) - # SVDQuant - create vertical grids to match the size of layer_filter_entry + if has_unet or has_transformer: + dtype_label_frame = self.components.inline_frame(frame, row, 3) + dtype_entry_frame = self.components.inline_frame(frame, row, 4) + + if has_unet: + self.components.label(dtype_label_frame, 0, 0, "UNet Data Type", + tooltip="The unet weight data type") + self.components.options_kv(dtype_entry_frame, 0, 0, self.__create_dtype_options(include_a8=True, include_compressed=include_compressed), + ui_state, "unet.weight_dtype") + else: + self.components.label(dtype_label_frame, 0, 0, "Transformer Data Type", + tooltip="The transformer weight data type") + self.components.options_kv(dtype_entry_frame, 0, 0, self.__create_dtype_options(include_gguf=True, include_a8=True, include_compressed=include_compressed), + ui_state, "transformer.weight_dtype") + + self.components.label(dtype_label_frame, 1, 0, "Fallback Data Type", + tooltip="The weight data type used for layers excluded by the quantization layer filter. Can itself be a quantized type.") + self.components.options_kv(dtype_entry_frame, 1, 0, self.__create_dtype_options(include_a8=True, include_compressed=include_compressed), + ui_state, "quantization.fallback_dtype") + + row += 1 + + if has_unconditional_transformer: + # unconditional transformer weight dtype + self.components.label(frame, row, 3, "Unconditional Transformer Data Type", + tooltip="The weight data type of the unconditional transformer, used for the negative branch of CFG during sampling") + self.components.options_kv(frame, row, 4, self.__create_dtype_options(include_a8=True, include_compressed=include_compressed), + ui_state, "unconditional_transformer.weight_dtype") + + row += 1 + + # SVDQuant svd_label_frame, svd_entry_frame = self._make_svd_frames(frame, row) self.components.label(svd_label_frame, 0, 0, "SVDQuant", tooltip="What datatype to use for SVDQuant weights decomposition.") @@ -284,6 +304,29 @@ def __create_base_components( row += 1 + if has_connectors: + self.components.label(frame, row, 3, "Connectors Data Type", + tooltip="The weight data type of the LTX connectors, the frozen network that turns text encoder output into the transformer's conditioning") + self.components.options_kv(frame, row, 4, self.__create_dtype_options(include_a8=True), + ui_state, "connectors.weight_dtype") + + row += 1 + + if has_low_noise_transformer: + self.components.label(frame, row, 0, "Low Noise Expert", + tooltip="Directory or Hugging Face repository of the distilled LTX transformer in diffusers format, or a single safetensors or GGUF file of the distilled transformer. Used only for the low-noise steps when sampling. Leave empty to sample with the trained transformer alone.") + self.components.path_entry( + frame, row, 1, ui_state, "low_noise_transformer.model_name", + mode="file", path_modifier=path_util.json_path_modifier + ) + + self.components.label(frame, row, 3, "Low Noise Expert Data Type", + tooltip="The weight data type of the distilled low-noise expert") + self.components.options_kv(frame, row, 4, self.__create_dtype_options(include_gguf=True, include_a8=True, include_compressed=include_compressed), + ui_state, "low_noise_transformer.weight_dtype") + + row += 1 + if has_vae: # base model self.components.label(frame, row, 0, "VAE Override", diff --git a/modules/ui/BaseSampleFrameView.py b/modules/ui/BaseSampleFrameView.py index eacd8c0ac..13d727a27 100644 --- a/modules/ui/BaseSampleFrameView.py +++ b/modules/ui/BaseSampleFrameView.py @@ -1,4 +1,5 @@ from modules.util.enum.NoiseScheduler import NoiseScheduler +from modules.util.enum.SamplingMethod import SamplingMethod class BaseSampleFrameView: @@ -9,6 +10,7 @@ def build_content(self, top_frame, bottom_frame, ui_state, controller, include_p is_flow_matching = controller.is_flow_matching() is_inpainting_model = controller.is_inpainting_model() is_video_model = controller.is_video_model() + is_audio_model = controller.is_audio_model() if include_prompt: # prompt self.components.label(top_frame, 0, 0, "prompt:") @@ -34,6 +36,7 @@ def build_content(self, top_frame, bottom_frame, ui_state, controller, include_p tooltip="Number of frames to generate. Only used when generating videos.") self.components.entry(bottom_frame, 1, 1, ui_state, "frames") + if is_audio_model: # length self.components.label(bottom_frame, 1, 2, "length:", tooltip="Length in seconds of audio output.") @@ -71,6 +74,29 @@ def build_content(self, top_frame, bottom_frame, ui_state, controller, include_p self.components.label(bottom_frame, 4, 0, "steps:") self.components.entry(bottom_frame, 4, 1, ui_state, "diffusion_steps") + if controller.model_type.has_dynamic_timestep_shift(): + self.components.label(bottom_frame, 4, 2, "override shift:", + tooltip="Timestep shift for this sample, as the multiplicative factor " + "(not mu). Empty uses the shift the model derives from the " + "token count.") + self.components.entry(bottom_frame, 4, 3, ui_state, "override_shift") + + if controller.supports_multiple_sampling_methods(): + self.components.label(bottom_frame, 6, 0, "sampling method:", + tooltip="How this sample is generated. 'base model' walks the whole " + "schedule with the transformer being trained. 'handoff to " + "low-noise expert' hands over to the distilled low-noise expert " + "below the expert's first sigma, as the reference pipeline does. " + "'distilled' runs the expert's own short schedule from noise, so " + "the trained transformer never runs. The trained LoRA applies to " + "the expert as well. Both expert methods fall back to 'base " + "model' when no expert model is loaded.") + self.components.options_kv(bottom_frame, 6, 1, [ + ("base model", SamplingMethod.STANDARD), + ("handoff to low-noise expert", SamplingMethod.HANDOFF_LOW_NOISE), + ("distilled", SamplingMethod.DISTILLED), + ], ui_state, "sampling_method") + # inpainting if is_inpainting_model: self.components.label(bottom_frame, 5, 0, "inpainting:", diff --git a/modules/ui/BaseSampleWindowView.py b/modules/ui/BaseSampleWindowView.py index 1a9c6c0c6..ed047d7b1 100644 --- a/modules/ui/BaseSampleWindowView.py +++ b/modules/ui/BaseSampleWindowView.py @@ -1,8 +1,28 @@ +class BaseSampleWindowView: + def __init__(self, components): + self.components = components + # every image this window has shown, in arrival order, so the user can + # scroll back through a batch (or the whole session) with the nav arrows + self._gallery = [] + self._gallery_index = -1 + def gallery_add(self, image): + # a freshly produced image always becomes the shown one + self._gallery.append(image) + self._gallery_index = len(self._gallery) - 1 + def gallery_step(self, delta): + self._gallery_index = max(0, min(len(self._gallery) - 1, self._gallery_index + delta)) + @property + def gallery_current(self): + return self._gallery[self._gallery_index] if self._gallery else None -class BaseSampleWindowView: - def __init__(self, components): - pass + @property + def gallery_index(self): + return self._gallery_index + + @property + def gallery_count(self): + return len(self._gallery) diff --git a/modules/ui/BaseTimestepDistributionWindowView.py b/modules/ui/BaseTimestepDistributionWindowView.py index 1f29e7079..281e4bf5d 100644 --- a/modules/ui/BaseTimestepDistributionWindowView.py +++ b/modules/ui/BaseTimestepDistributionWindowView.py @@ -42,5 +42,5 @@ def build_content(self, frame, controller, ui_state): # dynamic timestep shifting self.components.label(frame, 6, 0, "Dynamic Timestep Shifting", - tooltip="Dynamically shift the timestep distribution based on resolution. If enabled, the shifting parameters are taken from the model's scheduler configuration and Timestep Shift is ignored. Dynamic Timestep Shifting is not shown in the preview. For Ideogram, the shifting instead follows the model's own resolution-aware sampling schedule. Note: For Z-Image, the dynamic shifting parameters are likely wrong and unknown. Use with care or set your own, fixed shift.", wide_tooltip=True) + tooltip="Dynamically shift the timestep distribution based on resolution. If enabled, the shifting parameters are taken from the model's scheduler configuration and Timestep Shift is ignored. Dynamic Timestep Shifting is not shown in the preview. For Ideogram, the shifting instead follows the model's own resolution-aware sampling schedule. Note: For Z-Image, the dynamic shifting parameters are likely wrong and unknown. Use with care or set your own, fixed shift. Note: For LTX, the shift grows exponentially with the token count and is not capped, so a video puts nearly every timestep at high noise - a 10s 720p clip gives a shift above 30000. Use with care or set your own, fixed shift.", wide_tooltip=True) self.components.switch(frame, 6, 1, ui_state, "dynamic_timestep_shifting") diff --git a/modules/ui/BaseTrainingTabView.py b/modules/ui/BaseTrainingTabView.py index 29a54bb1e..2151b62a5 100644 --- a/modules/ui/BaseTrainingTabView.py +++ b/modules/ui/BaseTrainingTabView.py @@ -66,6 +66,8 @@ def build(self, column_0, column_1, column_2, controller, ui_state): self.__setup_ernie_ui(column_0, column_1, column_2, controller, ui_state) elif model_type.is_ideogram(): self.__setup_ideogram_ui(column_0, column_1, column_2, controller, ui_state) + elif model_type.is_ltx_2(): + self.__setup_ltx_2_ui(column_0, column_1, column_2, controller, ui_state) def __setup_stable_diffusion_ui(self, column_0, column_1, column_2, controller, ui_state): self.__create_base_frame(column_0, 0, controller, ui_state) @@ -247,6 +249,20 @@ def __setup_ideogram_ui(self, column_0, column_1, column_2, controller, ui_state self.__create_loss_frame(column_2, 2, controller, ui_state) self.__create_layer_frame(column_2, 3, controller, ui_state) + def __setup_ltx_2_ui(self, column_0, column_1, column_2, controller, ui_state): + self.__create_base_frame(column_0, 0, controller, ui_state) + self.__create_text_encoder_frame(column_0, 1, ui_state, supports_clip_skip=False, supports_training=False) + self.__create_connectors_frame(column_0, 2, ui_state) + self.__create_low_noise_transformer_frame(column_0, 3, ui_state) + + self.__create_base2_frame(column_1, 0, controller, ui_state, video_training_enabled=True) + self.__create_transformer_frame(column_1, 1, ui_state, supports_guidance_scale=False, supports_force_attention_mask=False) + self.__create_noise_frame(column_1, 2, ui_state, supports_dynamic_timestep_shifting=True) + + self.__create_masked_frame(column_2, 1, ui_state) + self.__create_loss_frame(column_2, 2, controller, ui_state) + self.__create_layer_frame(column_2, 3, controller, ui_state) + def __setup_sana_ui(self, column_0, column_1, column_2, controller, ui_state): self.__create_base_frame(column_0, 0, controller, ui_state) self.__create_text_encoder_frame(column_0, 1, ui_state) @@ -455,12 +471,22 @@ def __create_offloading_widgets(self, frame, row, ui_state, part, supports_check self.components.entry(frame, row, 1, ui_state, f"{part}.offload_fraction") row += 1 + self.components.label(frame, row, 0, "Simplex Offloading", + tooltip="Holds this component's weights in a single RAM buffer, so an offloaded layer never has to be copied back to RAM. Faster, but costs RAM for the whole component instead of only its offloaded layers. Not available for a fully fine-tuned component.") + self.components.switch(frame, row, 1, ui_state, f"{part}.simplex_offloading") + row += 1 + if supports_activation_offloading: self.components.label(frame, row, 0, "Offload Activations", tooltip="Offloads this component's activations to CPU during training to reduce VRAM usage") self.components.switch(frame, row, 1, ui_state, f"{part}.activation_offloading") row += 1 + self.components.label(frame, row, 0, "Cache In RAM", + tooltip="Keeps this model part's streamed weights in RAM between uses instead of re-reading them from disk on every use, trading RAM for loading speed. Only has an effect when \"Stream From Disk\" (model page) is enabled.") + self.components.switch(frame, row, 1, ui_state, f"{part}.cache_in_ram") + row += 1 + return row def __create_text_encoder_frame(self, master, row, ui_state, supports_clip_skip=True, supports_training=True, @@ -711,6 +737,32 @@ def __create_unconditional_transformer_frame(self, master, row, ui_state): row = self.__create_offloading_widgets(frame, row, ui_state, "unconditional_transformer", supports_checkpointing=False) + def __create_connectors_frame(self, master, row, ui_state): + frame = self.components.section_frame(master, row) + row = 0 + + self.components.label(frame, row, 0, "Connectors") + row += 1 + + row = self.__create_offloading_widgets( + frame, row, ui_state, "connectors", supports_checkpointing=False) + + def __create_low_noise_transformer_frame(self, master, row, ui_state): + frame = self.components.section_frame(master, row) + row = 0 + + # include low noise expert + self.components.label(frame, row, 0, "Include Low Noise Expert", + tooltip="Loads the distilled transformer named on the model tab and hands the " + "low-noise sampling steps over to it. If disabled, or if no model is given, " + "the trained transformer samples the whole schedule on its own") + self.components.switch(frame, row, 1, ui_state, "low_noise_transformer.include") + row += 1 + + row = self.__create_offloading_widgets( + frame, row, ui_state, "low_noise_transformer", supports_checkpointing=False, + supports_activation_offloading=False) + def __create_noise_frame(self, master, row, ui_state, supports_generalized_offset_noise: bool = False, supports_dynamic_timestep_shifting: bool = False): @@ -769,7 +821,7 @@ def __create_noise_frame(self, master, row, ui_state, if supports_dynamic_timestep_shifting: # dynamic timestep shifting self.components.label(frame, 9, 0, "Dynamic Timestep Shifting", - tooltip="Dynamically shift the timestep distribution based on resolution. If enabled, the shifting parameters are taken from the model's scheduler configuration and Timestep Shift is ignored. For Ideogram, the shifting instead follows the model's own resolution-aware sampling schedule. Note: For Z-Image, the dynamic shifting parameters are likely wrong and unknown. Use with care or set your own, fixed shift.", wide_tooltip=True) + tooltip="Dynamically shift the timestep distribution based on resolution. If enabled, the shifting parameters are taken from the model's scheduler configuration and Timestep Shift is ignored. For Ideogram, the shifting instead follows the model's own resolution-aware sampling schedule. Note: For Z-Image, the dynamic shifting parameters are likely wrong and unknown. Use with care or set your own, fixed shift. Note: For LTX, the shift grows exponentially with the token count and is not capped, so a video puts nearly every timestep at high noise - a 10s 720p clip gives a shift above 30000. Use with care or set your own, fixed shift.", wide_tooltip=True) self.components.switch(frame, 9, 1, ui_state, "dynamic_timestep_shifting") def __create_masked_frame(self, master, row, ui_state): diff --git a/modules/ui/CtkSampleWindowView.py b/modules/ui/CtkSampleWindowView.py index 3b67f03d5..707c966f2 100644 --- a/modules/ui/CtkSampleWindowView.py +++ b/modules/ui/CtkSampleWindowView.py @@ -59,7 +59,18 @@ def __init__( ) image_label = ctk.CTkLabel(master=self, text="", image=self.image, height=512, width=512) - image_label.grid(row=1, column=1, rowspan=3, sticky="nsew") + image_label.grid(row=1, column=1, rowspan=2, sticky="nsew") + + # gallery navigation, on the same row as the sample button + self.nav_frame = ctk.CTkFrame(self, fg_color="transparent") + self.prev_button = ctk.CTkButton(self.nav_frame, text="◀", width=40, command=lambda: self.__step_gallery(-1)) + self.counter_label = ctk.CTkLabel(self.nav_frame, text="") + self.next_button = ctk.CTkButton(self.nav_frame, text="▶", width=40, command=lambda: self.__step_gallery(1)) + self.prev_button.grid(row=0, column=0, padx=5) + self.counter_label.grid(row=0, column=1, padx=5) + self.next_button.grid(row=0, column=2, padx=5) + self.nav_frame.grid(row=3, column=1) + self.__render_gallery() self.progress = self.components.progress(self, 2, 0) self.components.button(self, 3, 0, "sample", @@ -71,12 +82,25 @@ def __init__( def __update_preview(self, sampler_output: ModelSamplerOutput): if sampler_output.file_type == FileType.IMAGE: - image = sampler_output.data + self.gallery_add(sampler_output.data) + self.__render_gallery() + + def __step_gallery(self, delta: int): + self.gallery_step(delta) + self.__render_gallery() + + def __render_gallery(self): + image = self.gallery_current + if image is not None: self.image.configure( light_image=image, size=(image.width, image.height), ) + self.counter_label.configure(text=f"{self.gallery_index + 1} / {self.gallery_count}") + self.prev_button.configure(state="normal" if self.gallery_index > 0 else "disabled") + self.next_button.configure(state="normal" if self.gallery_index < self.gallery_count - 1 else "disabled") + def __update_progress(self, progress: int, max_progress: int): self.progress.set(progress / max_progress) self.update() diff --git a/modules/ui/ModelTabController.py b/modules/ui/ModelTabController.py index e78432a39..14ce85156 100644 --- a/modules/ui/ModelTabController.py +++ b/modules/ui/ModelTabController.py @@ -26,8 +26,16 @@ def supports_override_transformer(self) -> bool: or model_type.is_qwen() or model_type.is_anima() or model_type.is_hunyuan_video() + or model_type.is_ltx_2() ) + def supports_override_text_encoder(self) -> bool: + model_type = self.train_config.model_type + # Only LTX-2's loader reads a text encoder override right now (useful because the + # diffusers-format checkpoint bundles a float32 copy of the stock Gemma-3-12B text encoder, + # roughly double the size of the official bf16 google/gemma-3-* release). + return model_type.is_ltx_2() + def get_output_formats(self) -> list[tuple[str, ModelFormat]]: labels = { ModelFormat.SAFETENSORS: "Safetensors", diff --git a/modules/ui/PySide6SampleWindowView.py b/modules/ui/PySide6SampleWindowView.py index 6f49a62c8..32073c04f 100644 --- a/modules/ui/PySide6SampleWindowView.py +++ b/modules/ui/PySide6SampleWindowView.py @@ -14,7 +14,7 @@ from PIL.ImageQt import ImageQt from PySide6.QtCore import Qt, QTimer from PySide6.QtGui import QPixmap -from PySide6.QtWidgets import QDialog, QGridLayout, QLabel, QProgressBar, QPushButton +from PySide6.QtWidgets import QDialog, QGridLayout, QHBoxLayout, QLabel, QProgressBar, QPushButton, QWidget class PySide6SampleWindowView(BaseSampleWindowView, QDialog): @@ -48,7 +48,20 @@ def __init__(self, parent, controller: SampleWindowController): self._image_label.setFixedSize(512, 512) self._image_label.setAlignment(Qt.AlignCenter) self._image_label.setStyleSheet("background: black;") - outer.addWidget(self._image_label, 1, 1, 3, 1) + outer.addWidget(self._image_label, 1, 1, 2, 1) + + # gallery navigation, on the same row as the sample button + self._nav_widget = QWidget(self) + nav_layout = QHBoxLayout(self._nav_widget) + self._prev_button = QPushButton("◀", self._nav_widget) + self._prev_button.clicked.connect(lambda: self._step_gallery(-1)) + self._counter_label = QLabel("", self._nav_widget) + self._next_button = QPushButton("▶", self._nav_widget) + self._next_button.clicked.connect(lambda: self._step_gallery(1)) + nav_layout.addWidget(self._prev_button) + nav_layout.addWidget(self._counter_label) + nav_layout.addWidget(self._next_button) + outer.addWidget(self._nav_widget, 3, 1, alignment=Qt.AlignCenter) self._progress = QProgressBar(self) self._progress.setRange(0, 1000) @@ -78,6 +91,8 @@ def _run(): sample_btn.clicked.connect(_on_sample) outer.addWidget(sample_btn, 3, 0) + self._render_gallery() + def schedule_on_main_thread(self, fn): QTimer.singleShot(0, self, fn) @@ -89,9 +104,24 @@ def _update_preview(self, sampler_output: ModelSamplerOutput): self.schedule_on_main_thread(lambda: self._do_update_preview(image)) def _do_update_preview(self, image): - pixmap = QPixmap.fromImage(ImageQt(image.convert("RGBA"))) - self._image_label.setFixedSize(pixmap.size()) - self._image_label.setPixmap(pixmap) + # gallery mutation runs on the main thread, so state stays consistent + self.gallery_add(image) + self._render_gallery() + + def _step_gallery(self, delta): + self.gallery_step(delta) + self._render_gallery() + + def _render_gallery(self): + image = self.gallery_current + if image is not None: + pixmap = QPixmap.fromImage(ImageQt(image.convert("RGBA"))) + self._image_label.setFixedSize(pixmap.size()) + self._image_label.setPixmap(pixmap) + + self._counter_label.setText(f"{self.gallery_index + 1} / {self.gallery_count}") + self._prev_button.setEnabled(self.gallery_index > 0) + self._next_button.setEnabled(self.gallery_index < self.gallery_count - 1) def _update_progress(self, progress: int, max_progress: int): # Called from training thread — dispatch to main thread diff --git a/modules/ui/SampleFrameController.py b/modules/ui/SampleFrameController.py index 4a6e3498b..2dbd0e482 100644 --- a/modules/ui/SampleFrameController.py +++ b/modules/ui/SampleFrameController.py @@ -16,5 +16,11 @@ def is_inpainting_model(self) -> bool: def is_video_model(self) -> bool: return self.model_type.is_video_model() + def is_audio_model(self) -> bool: + return self.model_type.is_audio_model() + def supports_negative_prompt(self) -> bool: return self.model_type.supports_negative_prompt() + + def supports_multiple_sampling_methods(self) -> bool: + return self.model_type.has_low_noise_expert() diff --git a/modules/ui/SampleWindowController.py b/modules/ui/SampleWindowController.py index b388bca5a..4dccbd2d1 100644 --- a/modules/ui/SampleWindowController.py +++ b/modules/ui/SampleWindowController.py @@ -8,6 +8,7 @@ from modules.util import create from modules.util.callbacks.TrainCallbacks import TrainCallbacks from modules.util.commands.TrainCommands import TrainCommands +from modules.util.compile_util import init_compile from modules.util.config.SampleConfig import SampleConfig from modules.util.config.TrainConfig import TrainConfig from modules.util.enum.EMAMode import EMAMode @@ -90,6 +91,8 @@ def load_model(self) -> BaseModel: model_names=model_names, weight_dtypes=self.initial_train_config.weight_dtypes(), quantization=self.initial_train_config.quantization, + stream_from_disk=self.initial_train_config.stream_from_disk, + cache_in_ram=self.initial_train_config.cache_in_ram(), ) model.train_config = self.initial_train_config @@ -119,6 +122,8 @@ def do_sample(self, on_sample, on_update_progress): self.model = self.load_model() self.model_sampler = self.create_sampler(self.model) + init_compile() + sample.from_train_config(self.current_train_config) sample_dir = os.path.join( diff --git a/modules/ui/TopBarController.py b/modules/ui/TopBarController.py index 052614a73..5f864cd24 100644 --- a/modules/ui/TopBarController.py +++ b/modules/ui/TopBarController.py @@ -44,6 +44,7 @@ def get_model_types(self) -> list[tuple[str, ModelType]]: ("Z-Image", ModelType.Z_IMAGE), ("Ernie Image", ModelType.ERNIE), ("Ideogram 4", ModelType.IDEOGRAM_4), + ("LTX 2", ModelType.LTX_2), ] def get_training_methods(self, model_type: ModelType) -> list[tuple[str, TrainingMethod]]: diff --git a/modules/ui/TrainUIController.py b/modules/ui/TrainUIController.py index 7aa550cee..895be42fa 100644 --- a/modules/ui/TrainUIController.py +++ b/modules/ui/TrainUIController.py @@ -108,8 +108,11 @@ def _start_always_on_tensorboard(self): if self.train_config.tensorboard_expose: tensorboard_args.append("--bind_all") + # discard the child's banner and notices; the UI already shows the tensorboard URL. try: - self.always_on_tensorboard_subprocess = subprocess.Popen(tensorboard_args) + self.always_on_tensorboard_subprocess = subprocess.Popen( + tensorboard_args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) except Exception: self.always_on_tensorboard_subprocess = None diff --git a/modules/util/LayerOffloadConductor.py b/modules/util/LayerOffloadConductor.py index f3e737c0e..934993c6f 100644 --- a/modules/util/LayerOffloadConductor.py +++ b/modules/util/LayerOffloadConductor.py @@ -1,18 +1,27 @@ import math import random +from collections import deque +from collections.abc import Callable from typing import Any +from modules.module.quantized.mixin.CompressedWeightMixin import CompressedWeightMixin from modules.util.config.TrainConfig import TrainConfig, TrainModelPartConfig -from modules.util.quantization_util import get_offload_tensor_bytes, offload_quantized +from modules.util.disk_stream import _is_evicted, evict_to_meta +from modules.util.enum.DataType import DataType +from modules.util.quantization_util import ( + get_offload_tensor_bytes, + get_offload_tensors, + is_quantized_module, + offload_quantized, + report_compression, +) from modules.util.torch_util import ( + create_mem_pool, create_stream_context, device_equals, - get_tensor_data, + mem_pool_context, pin_tensor_, - replace_tensors_, - tensors_match_device, tensors_record_stream, - tensors_to_device_, torch_gc, unpin_tensor_, ) @@ -20,26 +29,57 @@ import torch from torch import nn +from tqdm import tqdm + MESSAGES = [] +# Only relevant with activation offloading. Each layer-call the CPU enqueues ahead of the GPU floats one +# layer's worth of activations - the forward's copy stays alive until its D2H runs, the backward's reload +# destination is allocated when the CPU reaches the block - so the run-ahead has to be capped for peak VRAM +# to be predictable at all. Saturation needs only enough queued work to cover CPU-side jitter (tens of ms) +# against a layer-call of tens to hundreds of ms, so a small cap costs no throughput. 0 = unbounded. +# +# Also the number of trailing layers whose activations are not offloaded: the backward consumes N-1..0, so +# the offloads still draining when the forward ends - at most this many - are the ones needed first. Keeping +# those layers resident is free at this cap and saves a D2H/H2D round-trip of the same bytes. +MAX_LAYER_CALLS_IN_FLIGHT = 2 + + def log(msg: str = ''): pass # print(msg) # MESSAGES.append(msg) -def clone_tensor_allocator(tensor: torch.Tensor) -> torch.Tensor: +def flat_storage_view(tensor: torch.Tensor) -> torch.Tensor | None: + # Contiguous 1-D view over a tensor's whole storage region, or None if the strides leave gaps so that the + # elements between the first and last are not all part of this tensor. A permuted view of a freshly + # allocated tensor (a transpose, a head-split) is dense and yields a view; a slice of something larger + # does not. Copying through this view moves the bytes as-is instead of gathering them element by element. + span = 1 + sum((size - 1) * stride for size, stride in zip(tensor.shape, tensor.stride(), strict=True)) + if span != tensor.numel(): + return None + return torch.as_strided(tensor, (tensor.numel(),), (1,), tensor.storage_offset()) + + +def clone_tensor_allocator(tensor: torch.Tensor, non_blocking: bool = False) -> torch.Tensor: # clones a tensor into a new memory location to remove all memory dependencies between tensors return tensor.clone() -def ceil_16(number: int) -> int: - return number + (16 - (number % 16)) % 16 +# allocate_like places each cached tensor at an aligned offset, wasting up to this many bytes per tensor. +# also the reserved size at the start of each cache tensor (see allocate_like); must stay >= 2 so no view +# ever lands at storage_offset 0 or 1, the two values torch.compile bakes into separate specialized graphs +TENSOR_ALIGNMENT_BYTES = 16 + +def align_up(number: int) -> int: + return number + (TENSOR_ALIGNMENT_BYTES - (number % TENSOR_ALIGNMENT_BYTES)) % TENSOR_ALIGNMENT_BYTES -def floor_16(number: int) -> int: - return number - (number % 16) + +def align_down(number: int) -> int: + return number - (number % TENSOR_ALIGNMENT_BYTES) class StaticLayerTensorAllocator: @@ -72,17 +112,19 @@ def allocate_like(self, source_tensor: torch.Tensor) -> torch.Tensor: # never hand out views at storage_offset 0: torch.compile creates a 0/1-specialized # symbol for the storage_offset of any tensor with a dynamic dim (compressed weights), # so an offset-0 view needs its own graph while one "2 <= offset" guard covers all - # other placements. keeping every view at byte offset >= 16 avoids those recompiles. - cache_tensor_allocation_end = max(ceil_16(self.__allocation_end % cache_tensor_size), 16) + # other placements. keeping every view past the first alignment slot avoids those + # recompiles, and costs each tensor at most its alignment budget (the first tensor in a + # cache tensor previously wasted 0 of it) + cache_tensor_allocation_end = max(align_up(self.__allocation_end % cache_tensor_size), TENSOR_ALIGNMENT_BYTES) if cache_tensor_allocation_end + num_bytes > cache_tensor_size: # move to the start of the next cache tensor cache_tensor_index += 1 - cache_tensor_allocation_end = 16 + cache_tensor_allocation_end = TENSOR_ALIGNMENT_BYTES if cache_tensor_index * cache_tensor_size + cache_tensor_allocation_end + num_bytes > total_cache_bytes: # move to the first cache tensor cache_tensor_index = 0 - cache_tensor_allocation_end = 16 + cache_tensor_allocation_end = TENSOR_ALIGNMENT_BYTES self.__allocation_end = cache_tensor_index * cache_tensor_size + cache_tensor_allocation_end self.__layer_allocator.ensure_allocation(cache_tensor_index) @@ -95,9 +137,10 @@ def allocate_like(self, source_tensor: torch.Tensor) -> torch.Tensor: cache_tensor_index = self.__allocation_start // cache_tensor_size cache_tensor_allocation_start = self.__allocation_start % cache_tensor_size - # "< 16" instead of "< 0": the first 16 bytes of every cache tensor are reserved so no - # view lands at storage_offset 0 (see the forward-direction comment above) - if cache_tensor_allocation_start - num_bytes < 16: + # "< TENSOR_ALIGNMENT_BYTES" instead of "< 0": the first alignment slot of every cache + # tensor is reserved so no view lands at storage_offset 0 (see the forward-direction + # comment above) + if cache_tensor_allocation_start - num_bytes < TENSOR_ALIGNMENT_BYTES: # move to the end of the previous cache tensor cache_tensor_index -= 1 cache_tensor_allocation_start = cache_tensor_size @@ -106,7 +149,7 @@ def allocate_like(self, source_tensor: torch.Tensor) -> torch.Tensor: cache_tensor_index = len(self.__layer_allocator.cache_tensors) - 1 cache_tensor_allocation_start = cache_tensor_size - new_allocation_start = floor_16(cache_tensor_allocation_start - num_bytes) + new_allocation_start = align_down(cache_tensor_allocation_start - num_bytes) self.__layer_allocator.ensure_allocation(cache_tensor_index) cache_tensor = self.__layer_allocator.cache_tensors[cache_tensor_index] allocated_tensor = cache_tensor[new_allocation_start:new_allocation_start + num_bytes] @@ -116,6 +159,12 @@ def allocate_like(self, source_tensor: torch.Tensor) -> torch.Tensor: return allocated_tensor.view(dtype=source_tensor.dtype).view(size=source_tensor.shape) + def place(self, source_tensor: torch.Tensor, non_blocking: bool = False) -> torch.Tensor: + # place functor: allocate a fresh cache slot and copy the source into it. + new_tensor = self.allocate_like(source_tensor) + new_tensor.copy_(source_tensor.data, non_blocking=non_blocking) + return new_tensor + def deallocate(self, deallocate_forward): if deallocate_forward: log(f"{self.__layer_allocator.device}/deallocating layer {self.__layer_index}, allocation_start {self.__allocation_end:_}") @@ -159,31 +208,58 @@ def __init__( self.__tensor_allocators = [] - def allocate_cache(self, layers: list[nn.Module], target_bytes: int): + self.__mem_pool = None + + def allocate_cache(self, layers: list[nn.Module], target_bytes: int, streaming: bool, cache_in_ram: bool): if not self.__allocate_statically or any(x is not None for x in self.cache_tensors): return log(f"allocating cache on device {self.device}") + # keep the cache tensor in its own MemPool to avoid fragmenting the next cycle's allocation + if self.__mem_pool is None: + self.__mem_pool = create_mem_pool(self.device) + self.__max_tensor_bytes = 0 self.__layer_bytes = [] + total_tensors = 0 # count of individual offload tensors == number of allocate_like calls == alignment slots for layer in layers: layer_tensor_bytes = [get_offload_tensor_bytes(x) for x in layer.modules()] + total_tensors += sum(len(get_offload_tensors(x)) for x in layer.modules()) self.__max_tensor_bytes = max(self.__max_tensor_bytes, *layer_tensor_bytes) self.__layer_bytes.append(sum(layer_tensor_bytes)) cache_bytes = target_bytes - num_cache_tensors = min( - # no more than 10% overhead - math.ceil(int(cache_bytes * 0.10) / self.__max_tensor_bytes), - # at least twice self.__max_tensor_bytes for each tensor - math.ceil(cache_bytes / (self.__max_tensor_bytes * 2)), - # no more than 10 cache tensors - 10 - ) - # add self.__max_tensor_bytes to ensure even the largest tensors can be allocated in the remaining space - # add 4kb for the alignment overhead - self.cache_tensor_size = math.ceil(cache_bytes / num_cache_tensors) + self.__max_tensor_bytes + 4096 + if self.device.type == "cuda": + # single cache tensor on the GPU: a large cuda allocation is page-mapped (assembled from scattered + # physical pages), so one buffer allocates as readily as many and packs with no inter-chunk tail waste. + # The GPU cache is filled one layer at a time from the CPU, so the destination buffer and a full + # resident source never coexist on the device -- no peak-doubling to guard against here. + num_cache_tensors = 1 + elif streaming and not cache_in_ram: + # host/pinned cache, disk-streaming with cache_in_ram off: layers stream+quantize straight from the + # checkpoint and evict back to meta, so no resident copy ever coexists with the pinned cache -- none of the + # peak-doubling that justifies chunking below. A single large pinned buffer is fine: pin_tensor_ page-locks + # the existing scattered pages in place, and the CPU allocator has no pool to fragment. Same as the GPU cache. + num_cache_tensors = 1 + else: + # host/pinned cache, resident model (classic offload, or streaming with cache_in_ram on): the chunks are + # allocated lazily (per ensure_allocation) to cap peak host RAM while the resident model is copied into + # the pinned cache (and, on evict, cloned back out of it), which a single eager buffer would roughly double. + num_cache_tensors = min( + # no more than 10% overhead + math.ceil(int(cache_bytes * 0.10) / self.__max_tensor_bytes), + # at least twice self.__max_tensor_bytes for each tensor + math.ceil(cache_bytes / (self.__max_tensor_bytes * 2)), + # no more than 10 cache tensors + 10 + ) + # the alignment budget must cover EVERY tensor packed into a cache tensor: allocate_like wastes up to + # TENSOR_ALIGNMENT_BYTES per tensor and the ring wrap is unguarded, so a fixed total would silently + # overwrite live weights once a cache tensor holds enough tensors. Size it from the actual tensor count. + alignment_bytes = TENSOR_ALIGNMENT_BYTES * total_tensors + # add self.__max_tensor_bytes so even the largest tensor fits in the space left after a ring wrap + self.cache_tensor_size = math.ceil(cache_bytes / num_cache_tensors) + self.__max_tensor_bytes + alignment_bytes self.__tensor_allocators = [None] * len(layers) self.cache_tensors = [None] * num_cache_tensors @@ -194,15 +270,19 @@ def ensure_allocation(self, cache_tensor_index: int): if self.cache_tensors[cache_tensor_index] is None: torch_gc() - self.cache_tensors[cache_tensor_index] = \ - torch.zeros((self.cache_tensor_size,), dtype=torch.int8, device=self.device) + # create the cache tensor inside the MemPool so it lands in the pool's isolated segments. the buffers + # are allocated lazily here (allocate_cache only sizes them), so the pool context wraps this + # allocation rather than allocate_cache. + with mem_pool_context(self.__mem_pool): + self.cache_tensors[cache_tensor_index] = \ + torch.zeros((self.cache_tensor_size,), dtype=torch.int8, device=self.device) log(f"tensor {cache_tensor_index} not allocated, allocating {self.cache_tensor_size} bytes") if self.__is_pinned: pin_tensor_(self.cache_tensors[cache_tensor_index]) - def deallocate_cache(self): + def deactivate_cache(self): if not self.__allocate_statically: return @@ -212,6 +292,27 @@ def deallocate_cache(self): self.cache_tensors = [None] * len(self.cache_tensors) self.__tensor_allocators = [None] * len(self.__tensor_allocators) + # the loop above leaves `cache_tensor` bound to the last tensor; clear it so that stray reference can't + # keep the MemPool alive through the torch_gc below + cache_tensor = None + + # drop the MemPool once its tensors are freed so its now-empty segments return to the driver for the + # default pool; a fresh one is created on the next allocate_cache. + if self.__mem_pool is not None: + self.__mem_pool = None + torch_gc() + + def free(self): + self.deactivate_cache() + + @property + def mem_pool(self): + # the MemPool holding this allocator's cache tensor(s); also used to keep the conductor's resident non-layer + # remainder out of the default pool. allocate_cache creates it before the materialize layer loop; create it + # here too in case a caller reaches for it first. deactivate_cache drops it (static allocators only). + if self.__mem_pool is None: + self.__mem_pool = create_mem_pool(self.device) + return self.__mem_pool def get_allocator(self, layer_index: int, allocate_forward: bool) -> StaticLayerTensorAllocator | None: if self.__allocate_statically: @@ -227,6 +328,101 @@ def deallocate_layer(self, layer_index: int, deallocate_forward: bool): self.__tensor_allocators[layer_index] = None +class FullModelTensorAllocator: + # sibling of StaticLayerTensorAllocator: place() returns the tensor's permanent CPU slot with no copy, so an + # offload is a pure pointer swap. deallocate is a no-op -- the slot is permanent. + def __init__(self, layer_allocator: 'FullModelLayerAllocator'): + self.__layer_allocator = layer_allocator + + def place(self, source_tensor: torch.Tensor, non_blocking: bool = False) -> torch.Tensor: + return self.__layer_allocator.slot_for(source_tensor) + + def deallocate(self, deallocate_forward: bool): + pass + + +class FullModelLayerAllocator: + # Temp/CPU-side sibling of StaticLayerAllocator, selected in simplex mode. A single flat pinned buffer holds + # every layer's packed weights; offload is a pointer swap into that buffer, so no device->host copy ever runs + # on the hot path. The buffer is filled once, at materialize, by an explicit GPU->CPU copy. Its lifetime + # follows cache_in_ram: on it survives every evict (fill runs once for the model's life), off it is freed with + # the weights at the evict to meta and rebuilt from a fresh stream on the next materialize. + device: torch.device + + def __init__(self, device: torch.device): + assert device.type == "cpu", "FullModelLayerAllocator is CPU-only" + self.device = device + self.__buffer = None # single flat int8 buffer holding every layer's packed weights + self.__fill_offset = 0 # running fill cursor into __buffer (advances only during materialize) + self.__slots = {} # frozen param tensor -> its permanent view into __buffer + self.__is_buffer_pinned = False + + @property + def filled(self) -> bool: + # True once the buffer holds the model's weights -- distinguishes a cold materialize from a warm re-activate. + return len(self.__slots) > 0 + + def allocate_cache(self, layers: list[nn.Module], target_bytes: int, streaming: bool, cache_in_ram: bool): + # create the full-model buffer if this is a cold activate (never filled, or freed by a cache_in_ram-off + # evict); (re)pin on every activate. target_bytes is ignored -- every layer is resident for as long as the + # part is materialized, so the buffer is sized to the whole model's footprint plus alignment. + if self.__buffer is None: + total_tensors = 0 + total_bytes = 0 + for layer in layers: + for module in layer.modules(): + total_bytes += get_offload_tensor_bytes(module) + total_tensors += len(get_offload_tensors(module)) + buffer_bytes = total_bytes + TENSOR_ALIGNMENT_BYTES * total_tensors + torch_gc() + self.__buffer = torch.zeros((buffer_bytes,), dtype=torch.int8, device=self.device) + self.__fill_offset = 0 + + if not self.__is_buffer_pinned: + pin_tensor_(self.__buffer) + self.__is_buffer_pinned = True + + def fill(self, tensor: torch.Tensor) -> torch.Tensor: + # carve this tensor's permanent slot out of the flat buffer and copy its quantized weight into it. + num_bytes = tensor.numel() * tensor.element_size() + slot = self.__buffer[self.__fill_offset:self.__fill_offset + num_bytes] \ + .view(dtype=tensor.dtype).view(size=tensor.shape) + self.__fill_offset = align_up(self.__fill_offset + num_bytes) + slot.copy_(tensor.data) + self.__slots[tensor] = slot + return slot + + def slot_for(self, tensor: torch.Tensor) -> torch.Tensor: + return self.__slots[tensor] + + def repoint_all(self): + # point every frozen weight at its permanent CPU slot -- rescues the layers whose .data viewed the + # now-freed GPU ring. + for tensor, slot in self.__slots.items(): + tensor.data = slot + + def get_allocator(self, layer_index: int, allocate_forward: bool) -> FullModelTensorAllocator: + return FullModelTensorAllocator(self) + + def deallocate_layer(self, layer_index: int, deallocate_forward: bool): + pass # slots are permanent -- an offloaded layer's weight stays in the buffer for the next onload + + def deactivate_cache(self): + # unpins but keeps the buffer and its data resident, so the frozen weights survive to re-activate. + if self.__is_buffer_pinned: + unpin_tensor_(self.__buffer) + self.__is_buffer_pinned = False + + def free(self): + # releases the buffer and every slot. Reached from the evict to meta -- the cache_in_ram-off eviction and the + # error rollback. The slots must go with it: the evict re-registers the weights as new meta parameters, so + # every key in the identity-keyed slot map is stale afterwards. + self.deactivate_cache() + self.__buffer = None + self.__fill_offset = 0 + self.__slots = {} + + class StaticActivationAllocator: __device: torch.device __allocate_statically: bool @@ -252,9 +448,13 @@ def __init__( self.__allocated_bytes = 0 self.__max_allocated_bytes = 0 + @property + def allocated_bytes(self) -> int: + return self.__allocated_bytes + def reserve_cache(self, tensors: list[torch.Tensor]): num_bytes = sum(tensor.element_size() * tensor.numel() for tensor in tensors) \ - + len(tensors) * 16 # add enough padding for alignment + + len(tensors) * TENSOR_ALIGNMENT_BYTES # add enough padding for alignment if num_bytes == 0: return @@ -273,8 +473,16 @@ def reserve_cache(self, tensors: list[torch.Tensor]): self.__current_cache_tensor_offset = 0 if not cache_found: - torch_gc() - cache_tensor = torch.zeros((num_bytes,), dtype=torch.int8, device=self.__device) + try: + cache_tensor = torch.zeros((num_bytes,), dtype=torch.int8, device=self.__device) + except (torch.OutOfMemoryError, MemoryError): + # collect only when the allocation actually needs the room. This runs once per activation that + # overflows the current cache tensor -- on the first step of a run that is every one of them -- + # and torch_gc's synchronize + gc.collect + empty_cache measured ~350 ms each, 78 s over the 230 + # allocations one LTX step makes, while none of them ever needed the room. A CUDA cache tensor + # raises OutOfMemoryError, a host one MemoryError. + torch_gc() + cache_tensor = torch.zeros((num_bytes,), dtype=torch.int8, device=self.__device) log(f"{self.__device}/allocating activations cache {num_bytes:_}, total: {self.__allocated_bytes:_}, max: {self.__max_allocated_bytes:_}") if self.__is_pinned: @@ -290,7 +498,7 @@ def allocate_like(self, source_tensor: torch.Tensor) -> torch.Tensor: cache_tensor = self.__cache_tensors[self.__current_cache_tensor] allocated_tensor = \ cache_tensor[self.__current_cache_tensor_offset:self.__current_cache_tensor_offset + num_bytes] - self.__current_cache_tensor_offset += ceil_16(num_bytes) + self.__current_cache_tensor_offset += align_up(num_bytes) return allocated_tensor.view(dtype=source_tensor.dtype).view(size=source_tensor.shape) @@ -527,6 +735,23 @@ def get_layers_to_load( return [x for x in layers if x < layer_index] + [x for x in layers if x >= layer_index] +class _BoundaryActivation: + # Offloaded copy of a saved activation on the boundary path. Copy semantics (never mutate the saved + # tensor in place - it may be shared across blocks, e.g. a conditioning embedding). cpu holds the + # temp-device copy; gpu the reloaded train-device copy; event marks the reload transfer. + def __init__(self): + self.cpu = None + self.gpu = None + self.event = None + # the source tensor's stride, restored on reload. A saved activation is often a permuted view + # (an attention output, a head-split), and the compiled backward asserts on exact strides, so + # handing it back contiguous fails assert_size_stride rather than silently computing wrong. + self.stride = None + # whether the source's storage region was dense, so both legs could move it as flat storage rather + # than reordering elements through a copy kernel. + self.dense = False + + class LayerOffloadConductor: __module: nn.Module @@ -534,7 +759,6 @@ class LayerOffloadConductor: __layer_device_map: list[torch.device | None] __layer_offload_fraction: float - __layer_activations_included_offload_param_indices_map: list[list[int]] __train_device: torch.device __temp_device: torch.device @@ -548,31 +772,35 @@ class LayerOffloadConductor: __activations_transfer_stream: torch.Stream | None __train_device_layer_allocator: StaticLayerAllocator - __temp_device_layer_allocator: StaticLayerAllocator + __temp_device_layer_allocator: StaticLayerAllocator | FullModelLayerAllocator __temp_device_activations_allocator: StaticActivationAllocator __layer_train_event_map: list[SyncEvent] __layer_transfer_event_map: list[SyncEvent] - __activations_map: dict[int, Any] - __call_index_layer_index_map: dict[int, int] - __activations_transfer_event_map: dict[int, SyncEvent] __offload_strategy = LayerOffloadStrategy | None __is_forward_pass: bool - __keep_graph: bool + __backward_follows: bool - __is_active: bool + __materialized: bool + + __simplex_active: bool __deferred_layers: list[int] __config: TrainConfig + __disk_remainder_materialized: bool # whether the non-layer remainder (embedders/norms/proj) has been streamed since the last evict + __disk_layer_key_prefixes: list[str] # per-layer (indexed like __layers) checkpoint-absolute path, so a single layer subtree can be streamed on its own + __disk_module_name_by_id: dict[int, str] # module-name snapshot taken pre-wrapping, used to build the key prefixes above + def __init__( self, module: nn.Module, config: TrainConfig, part: TrainModelPartConfig, + simplex: bool = False, ): super().__init__() @@ -582,7 +810,6 @@ def __init__( self.__layer_device_map = [] self.__layer_offload_fraction = part.offload_fraction - self.__layer_activations_included_offload_param_indices_map = [] self.__train_device = torch.device(config.train_device) self.__temp_device = torch.device(config.temp_device) @@ -600,111 +827,242 @@ def __init__( self.__layer_transfer_stream = None self.__activations_transfer_stream = None + self.__simplex_active = simplex + if self.__simplex_active: + print(f"simplex full-model-buffer offload activated for {type(self.__module).__name__}") + self.__train_device_layer_allocator = StaticLayerAllocator(self.__train_device) - self.__temp_device_layer_allocator = StaticLayerAllocator(self.__temp_device) + self.__temp_device_layer_allocator = FullModelLayerAllocator(self.__temp_device) \ + if self.__simplex_active else StaticLayerAllocator(self.__temp_device) self.__temp_device_activations_allocator = StaticActivationAllocator(self.__temp_device) self.__layer_train_event_map = [] self.__layer_transfer_event_map = [] - self.__activations_map = {} - self.__call_index_layer_index_map = {} - self.__activations_transfer_event_map = {} + self.__boundary_activations = {} self.__offload_strategy = None self.__is_forward_pass = False - self.__keep_graph = False + self.__backward_follows = False + self.__inflight_call_events = deque() + self.__inflight_transfer_events = deque() + self.__warned_non_dense = False - self.__is_active = False + self.__materialized = False self.__deferred_layers = [] self.__config = config + self.__disk_remainder_materialized = False + self.__disk_layer_key_prefixes = [] + self.__disk_module_name_by_id = {id(m): name for name, m in module.named_modules()} + def offload_activated(self) -> bool: return self.__offload_activations or self.__offload_layers - def evict(self): + def offloads_activations(self) -> bool: + return self.__offload_activations + + def evict(self, to_meta: bool = False) -> bool: + # returns whether anything was actually evicted, so the caller can skip its gc when nothing moved. + # Nothing is resident while __materialized is False, so every transfer wait, device walk and collection + # below would be pure overhead - and materialize_only() evicts every part it doesn't want on each call, + # so in steady state most of these are repeat evictions of an already evicted part. The rollback in + # materialize() calls __evict_to_temp/__evict_to_meta directly and so is unaffected by this guard: it + # has to run precisely when __materialized is still False but weights are resident. + if not self.__materialized: + return False + torch_gc() self.__wait_all_layer_transfers() - self.__wait_all_activation_transfers() log("to temp device") + if to_meta: + self.__evict_to_meta() + else: + self.__evict_to_temp() + return True + + def __evict_to_temp(self): + # move every layer and the non-layer remainder back to the temp device and free the static caches (the + # non-disk eviction path). Also the rollback for a resident conductor whose materialize() raised partway. # deallocate the cache before to take advantage of the gc - self.__train_device_layer_allocator.deallocate_cache() - self.__temp_device_layer_allocator.deallocate_cache() + self.__train_device_layer_allocator.deactivate_cache() + self.__temp_device_layer_allocator.deactivate_cache() self.__temp_device_activations_allocator.deallocate_cache() self.__module_to_device_except_layers(self.__temp_device) - for layer_index, layer in enumerate(self.__layers): - self.__layers[layer_index].to(self.__temp_device) - for module in layer.modules(): - offload_quantized(module, self.__temp_device, allocator=clone_tensor_allocator) - self.__layer_device_map[layer_index] = None - - self.__is_active = False - - torch_gc() + if self.__simplex_active: + # every frozen weight already lives in the permanent CPU buffer; repoint the layers that were + # resident in the now-freed GPU ring back to their dormant CPU slots. + self.__temp_device_layer_allocator.repoint_all() + for layer_index in range(len(self.__layers)): + self.__layer_device_map[layer_index] = None + else: + for layer_index, layer in enumerate(self.__layers): + self.__layers[layer_index].to(self.__temp_device) + for module in layer.modules(): + offload_quantized(module, self.__temp_device, place=clone_tensor_allocator) + self.__layer_device_map[layer_index] = None + + self.__materialized = False + + def materialize( + self, train_dtype: DataType | None = None, name: str | None = None, + materialize_fn: Callable | None = None, cache_in_ram: bool = True) -> bool: + # returns whether anything was actually materialized. Already-materialized is the steady state at every + # epoch boundary, where setup_train_device re-states what it wants without anything having moved: the + # allocators would find their caches allocated and every layer already mapped, so the whole body below + # is a walk that changes nothing. The torch_gc is deliberately inside the guard rather than at the top: + # it is a pre-allocation collect (the layer ring and the full-model buffer are allocated below), so it + # is only worth its cost when an allocation actually follows. + if self.__materialized: + return False - def materialize(self): torch_gc() self.__wait_all_layer_transfers() - self.__wait_all_activation_transfers() - - log("to train device") - self.__offload_strategy = LayerOffloadStrategy(self.__layers, self.__layer_offload_fraction) + streaming = materialize_fn is not None - self.__train_device_layer_allocator.allocate_cache( - self.__layers, self.__offload_strategy.max_loaded_bytes) - self.__temp_device_layer_allocator.allocate_cache( - self.__layers, self.__offload_strategy.max_offloaded_bytes) - self.__module_to_device_except_layers(self.__train_device) + log("to train device") - # move all layers to the train device, then move offloadable tensors back to the temp device - for layer_index, layer in enumerate(self.__layers): - if self.__layer_device_map[layer_index] is None: + try: + self.__measure_compressed_sizes(materialize_fn, train_dtype, name) + + self.__offload_strategy = LayerOffloadStrategy(self.__layers, self.__layer_offload_fraction) + + if self.__simplex_active and not streaming: + raise NotImplementedError( + "the simplex full-model-buffer offload requires a disk-streamed component; a single-file override " + "with 'Stream From Disk' enabled is not supported yet") + + self.__train_device_layer_allocator.allocate_cache( + self.__layers, self.__offload_strategy.max_loaded_bytes, streaming=streaming, cache_in_ram=cache_in_ram) + self.__temp_device_layer_allocator.allocate_cache( + self.__layers, self.__offload_strategy.max_offloaded_bytes, streaming=streaming, cache_in_ram=cache_in_ram) + # place the resident non-layer remainder onto the train device. When streaming, route it into the conductor + # pool: on a warm cache_in_ram re-activate it comes from cpu/temp and lands there directly (no default-pool + # copy to relocate); on a cold stream it is still meta here and gets skipped, then streamed below. + self.__module_to_device_except_layers( + self.__train_device, + pool=self.__train_device_layer_allocator.mem_pool if streaming else None) + + cold_layers = sum(1 for i, layer in enumerate(self.__layers) + if self.__layer_device_map[i] is None and _is_evicted(layer)) if streaming else 0 + disk_bar = tqdm(total=cold_layers, unit="layer", desc=f"streaming {name}", leave=False) \ + if cold_layers > 0 else None + + already_filled = self.__temp_device_layer_allocator.filled if self.__simplex_active else False + + # per-layer materialize helpers, shared by the simplex and static paths below + def bring_to_train_device(layer, layer_index): + # get the layer's weights onto the train device: stream+quantize from the checkpoint if the layer + # was evicted to disk, otherwise a plain device move of the still-resident weights. + if streaming and _is_evicted(layer): + materialize_fn( + layer, self.__train_device, train_dtype, self.__disk_layer_key_prefixes[layer_index]) + if disk_bar is not None: + disk_bar.update(1) + else: + layer.to(self.__train_device) + + def copy_into_gpu_ring(layer, layer_index): + # copy the layer's weights into its GPU ring cache slot and mark it train-device resident + allocator = self.__train_device_layer_allocator.get_allocator(layer_index, allocate_forward=True) + for module in layer.modules(): + offload_quantized(module, self.__train_device, place=allocator.place) + self.__layer_device_map[layer_index] = self.__train_device + + for layer_index, layer in enumerate(self.__layers): + if self.__layer_device_map[layer_index] is not None: + continue log(f"layer {layer_index} to train device") - layer.to(self.__train_device) - - if layer_index in self.__offload_strategy.initial_loaded_layers: - allocator = self.__train_device_layer_allocator.get_allocator( - layer_index, allocate_forward=True) - for module in layer.modules(): - offload_quantized(module, self.__train_device, allocator=allocator.allocate_like) - self.__layer_device_map[layer_index] = self.__train_device + + if self.__simplex_active: + if not already_filled: + # cold materialize: bring the layer in, then copy each weight into its slot in the + # full-model CPU buffer -- a GPU->CPU copy that runs once per fill of the buffer (once for + # the model's life with cache_in_ram on, once per materialize with it off). + bring_to_train_device(layer, layer_index) + for module in layer.modules(): + for tensor in get_offload_tensors(module): + tensor.data = self.__temp_device_layer_allocator.fill(tensor) + else: + # warm re-activate (cache_in_ram): the weights still live in the CPU buffer; the evict only + # freed the GPU ring their .data viewed, so re-point each weight at its surviving slot -- + # no copy, no stream. + for module in layer.modules(): + for tensor in get_offload_tensors(module): + tensor.data = self.__temp_device_layer_allocator.slot_for(tensor) + + if layer_index in self.__offload_strategy.initial_loaded_layers: + # dual residency: also copy the CPU slot into the GPU ring; the CPU slot stays filled but + # dormant until this layer is offloaded again. + copy_into_gpu_ring(layer, layer_index) + else: + self.__layer_device_map[layer_index] = self.__temp_device else: - allocator = self.__temp_device_layer_allocator.get_allocator(layer_index, allocate_forward=True) - for module in layer.modules(): - offload_quantized(module, self.__temp_device, allocator=allocator.allocate_like) - self.__layer_device_map[layer_index] = self.__temp_device + bring_to_train_device(layer, layer_index) + if layer_index in self.__offload_strategy.initial_loaded_layers: + copy_into_gpu_ring(layer, layer_index) + else: + # copy into the pinned CPU cache slot for an offloaded layer + allocator = self.__temp_device_layer_allocator.get_allocator(layer_index, allocate_forward=True) + for module in layer.modules(): + offload_quantized(module, self.__temp_device, place=allocator.place) + self.__layer_device_map[layer_index] = self.__temp_device if self.__async_transfer: event = SyncEvent(self.__train_stream.record_event(), f"train on {self.__train_device}") self.__layer_train_event_map[layer_index] = event - self.__is_active = True - - torch_gc() + if disk_bar is not None: + disk_bar.close() + + if streaming and not self.__disk_remainder_materialized: + # the non-layer remainder (embedders/norms/proj) is still meta the first time; stream it to the train + # device now, where it stays resident. dest_pool routes the non-quantized weights straight into the + # conductor pool so no model weight sits in the default pool (which the optimizer state and quantize + # transients draw from). Quantized remainder weights pack in the default pool -- their dequant scratch + # stays out of the pool -- and are relocated into it just below, once small. + materialize_fn(self.__module, self.__train_device, train_dtype, "", + dest_pool=self.__train_device_layer_allocator.mem_pool) + self.__disk_remainder_materialized = True + self.__relocate_quantized_remainder_to_pool() + except Exception: + # a materialize that fails partway (typically OOM) leaves layers/cache tensors resident while + # __materialized is still False, so a later evict() would skip them and strand that VRAM. Force the unit + # back to its pre-materialize state, keyed on the actual weight state: a parameter still on meta means a + # cold disk-stream was in flight, so meta is the only valid target (re-stream next time, lossless since + # frozen); otherwise roll back to the temp device and keep the resident quantized copy. + if any(parameter.is_meta for parameter in self.__module.parameters()): + self.__evict_to_meta() + else: + self.__evict_to_temp() + # the rollback helpers no longer gc, and no caller gc's a failed materialize -- reclaim the stranded VRAM + # here before re-raising (evict() instead relies on BaseModel.evict()'s trailing gc). + torch_gc() + raise - def add_layer(self, layer: nn.Module, included_offload_param_indices: list[int] = None): - if included_offload_param_indices is None: - included_offload_param_indices = [] + self.__materialized = True + return True + def add_layer(self, layer: nn.Module): self.__layers.append(layer) self.__layer_device_map.append(None) self.__layer_train_event_map.append(SyncEvent()) self.__layer_transfer_event_map.append(SyncEvent()) + # checkpoint-absolute path of this layer, for the per-layer disk stream (empty for a layer built outside self.__module) + self.__disk_layer_key_prefixes.append(self.__disk_module_name_by_id.get(id(layer), "")) - self.__layer_activations_included_offload_param_indices_map.append(included_offload_param_indices) - - def start_forward(self, keep_graph: bool): + def start_forward(self, backward_follows: bool): log("starting forward") - if not self.__is_active: + if not self.__materialized: return if self.__async_transfer: @@ -713,96 +1071,112 @@ def start_forward(self, keep_graph: bool): self.__clear_activations() self.__is_forward_pass = True - self.__keep_graph = keep_graph - - def before_layer(self, layer_index: int, call_index: int, activations: Any) -> Any: - log() - log(f"before layer {layer_index}, {call_index}") - - if not self.__is_active: - return activations - - self.__call_index_layer_index_map[call_index] = layer_index - - if torch.is_grad_enabled() and self.__is_forward_pass: - # Offloading can only be used with the use_reentrant=True checkpointing variant. - # Gradients are only enabled during the back pass. - log("starting backward") - self.__is_forward_pass = False - - if self.__offload_activations and not self.__is_forward_pass: - self.__wait_activations_transfer(call_index) - - tensor_indices = self.__layer_activations_included_offload_param_indices_map[layer_index] - - if call_index in self.__activations_map: - # during the back pass, replace activations with saved acitvations - replace_tensors_(activations, self.__activations_map.pop(call_index), tensor_indices) - - # if current activations are not on train_device, move them now - if not tensors_match_device( - activations, self.__train_device, - tensor_indices): - log(f"activations for layer {layer_index} not loaded to train device, transferring now") - self.__schedule_activations_to_device( - activations, self.__train_device, call_index, wait_train_stream=False) - self.__wait_activations_transfer(call_index) - - # schedule previous activations to the train device - if call_index - 1 in self.__activations_map: - self.__schedule_activations_to_device( - self.__activations_map[call_index - 1], self.__train_device, call_index - 1, - wait_train_stream=False) - - # schedule loading of the next layer and offloading of the previous layer - if self.__offload_layers: - self.__wait_layer_transfer(layer_index) - - self.__schedule_deferred_layers_to_temp(except_layer=layer_index) - for i in self.__offload_strategy.get_layers_to_offload( - layer_index=layer_index, - is_forward=self.__is_forward_pass, - is_next_forward=not self.__keep_graph, - loaded_layers=self.__get_loaded_layers(), - ): - self.__schedule_layer_to(i, self.__temp_device, is_forward=self.__is_forward_pass) - - for i in self.__offload_strategy.get_layers_to_load( - layer_index=layer_index, - is_forward=self.__is_forward_pass, - is_next_forward=not self.__keep_graph, - loaded_layers=self.__get_loaded_layers(), - ): - self.__schedule_layer_to(i, self.__train_device, is_forward=self.__is_forward_pass) - - return activations - - def after_layer(self, layer_index: int, call_index: int, activations: Any): - log(f"after layer {layer_index}, {call_index}") - - if not self.__is_active: + self.__backward_follows = backward_follows + # events from the previous step refer to work the GPU has long finished; carrying them over would + # make the first calls of this pass wait on stale entries. + self.__inflight_call_events.clear() + self.__inflight_transfer_events.clear() + + def __schedule_layer_offload(self, layer_index: int, is_forward: bool, is_next_forward: bool): + # windowed layer load/offload schedule around layer_index, in the direction the caller is + # running: the boundaries know whether they are in the forward or the backward pass. + if not self.__offload_layers: return - # record stream - if self.__async_transfer: - tensors_record_stream(self.__train_stream, activations) + self.__wait_layer_transfer(layer_index) + + self.__schedule_deferred_layers_to_temp(except_layer=layer_index) + for i in self.__offload_strategy.get_layers_to_offload( + layer_index=layer_index, + is_forward=is_forward, + is_next_forward=is_next_forward, + loaded_layers=self.__get_loaded_layers(), + ): + self.__schedule_layer_to(i, self.__temp_device, is_forward=is_forward) + + for i in self.__offload_strategy.get_layers_to_load( + layer_index=layer_index, + is_forward=is_forward, + is_next_forward=is_next_forward, + loaded_layers=self.__get_loaded_layers(), + ): + self.__schedule_layer_to(i, self.__train_device, is_forward=is_forward) + + def before_layer(self, layer_index: int, is_forward: bool): + # called before a block runs: LoadBoundary in the forward pass, EvictBoundary before + # the backward recompute. Waits for this layer's transfer and slides the load window. + if not self.__materialized: + return - # save activations during the forward pass to make them accessible during the backward pass - if self.__offload_activations and self.__keep_graph and self.__is_forward_pass: - log(f"saving layer {call_index} activations for back pass") - self.__activations_map[call_index] = activations - self.__schedule_activations_to_device(activations, self.__temp_device, call_index, wait_train_stream=True) + # block until compute and activation transfers have caught up to within MAX_LAYER_CALLS_IN_FLIGHT, + # so the floating activations stay bounded. + if MAX_LAYER_CALLS_IN_FLIGHT > 0: + queues = (("compute", self.__inflight_call_events), ("transfer", self.__inflight_transfer_events)) + for name, queue in queues: + while len(queue) > MAX_LAYER_CALLS_IN_FLIGHT: + queue.popleft().synchronize(f"layer-calls-in-flight cap ({name})") + + self.__schedule_layer_offload(layer_index, is_forward, not self.__backward_follows) + + def after_layer(self, layer_index: int, activations: Any): + # called after a block runs: EvictBoundary in the forward pass, LoadBoundary after the + # backward. Records the block's output/grad on the train stream and marks compute done, + # so a later offload of this layer waits for it. + if not self.__materialized or not self.__async_transfer: + return + tensors_record_stream(self.__train_stream, activations) + event = SyncEvent(self.__train_stream.record_event(), f"train on {self.__train_device}") + self.__layer_train_event_map[layer_index] = event + # the same event, queued in call order, is what the run-ahead cap waits on + self.__inflight_call_events.append(event) + + # Second marker, on the activations transfer stream, which trails the train stream by its own queue + # (measured ~106ms median, ~12 layer-calls of float against a cap of 2). An offloaded activation's + # source block is pinned until its D2H actually executes, so this is the distance that holds memory. + if self.__offload_activations: + self.__inflight_transfer_events.append( + SyncEvent(self.__activations_transfer_stream.record_event(), "activations transfer")) + + def __measure_compressed_sizes(self, materialize_fn: Callable | None, train_dtype: DataType, name: str): + # a compressed weight's blob length is data-dependent, so the arenas cannot be sized before every layer has + # been compressed once, and sizing them from the uncompressed footprint would cancel the saving. Stream each + # layer, compress it, keep the measured length and drop it back to meta: one extra streaming pass, no peak + # memory. The lengths outlive evict_to_meta and ANS is deterministic on the same bytes, so this runs once. + if materialize_fn is None: + return + # a resident layer's weights are live, so get_offload_tensor_bytes already measures its real tensors + unsized = [i for i, layer in enumerate(self.__layers) + if _is_evicted(layer) and any(isinstance(m, CompressedWeightMixin) + and m.compress and m.compressed_bytes() is None + for m in layer.modules())] + if not unsized: + return - if self.__async_transfer: - event = SyncEvent(self.__train_stream.record_event(), f"train on {self.__train_device}") - self.__layer_train_event_map[layer_index] = event + for layer_index in tqdm(unsized, unit="layer", desc=f"measuring compressed size of {name}", leave=False): + layer = self.__layers[layer_index] + materialize_fn(layer, self.__train_device, train_dtype, self.__disk_layer_key_prefixes[layer_index]) + evict_to_meta(layer) + torch_gc() + # the streamed path never reaches quantize_layers' report, so emit the same line here + report_compression(self.__module) def __get_loaded_layers(self) -> list[int]: return [i for i in range(len(self.__layers)) if device_equals(self.__layer_device_map[i], self.__train_device)] + def __evict_to_meta(self): + evict_to_meta(self.__module) + for layer_index in range(len(self.__layers)): + self.__layer_device_map[layer_index] = None + self.__disk_remainder_materialized = False + self.__train_device_layer_allocator.free() + self.__temp_device_layer_allocator.free() + self.__temp_device_activations_allocator.deallocate_cache() + self.__materialized = False + def __module_to_device_except_layers( self, device: torch.device, + pool=None, ): sub_module_parameters = set(sum([list(x.parameters()) for x in self.__layers], [])) @@ -810,14 +1184,39 @@ def convert(t): if t in sub_module_parameters or t.is_meta: return t + if pool is not None: + # place the (already-final) non-layer remainder weight straight into the conductor's pool instead of + # the default pool, which the optimizer state and quantize transients allocate from -- a weight left + # there fragments it and strands the region when the remainder is evicted. A weight from cpu/temp (warm + # cache_in_ram re-activate) lands in the pool directly; one already on the train device is relocated + # with a clone. + with mem_pool_context(pool): + return t.clone() if device_equals(t.device, device) else t.to(device=device) + return t.to(device=device) self.__module._apply(convert) + def __relocate_quantized_remainder_to_pool(self): + # the cold remainder stream packs quantized non-layer weights (e.g. a tied lm_head) in the default pool so + # their dequant scratch never enters the conductor pool. Copy just the packed weights into the pool now, so + # no model weight is left in the default pool (where the optimizer state and quantize transients would + # fragment/strand it). Small: the packed weights are a fraction of their fp size. Non-quantized remainder + # weights were streamed straight into the pool (dest_pool) and are not touched here. Layer modules are + # excluded -- they own their static cache slots -- matching __module_to_device_except_layers' scope. + pool = self.__train_device_layer_allocator.mem_pool + + def pool_clone(tensor, non_blocking=False): + with mem_pool_context(pool): + return tensor.clone() + + layer_modules = {module for layer in self.__layers for module in layer.modules()} + for module in self.__module.modules(): + if module not in layer_modules and is_quantized_module(module): + offload_quantized(module, self.__train_device, place=pool_clone) + def __clear_activations(self): - self.__activations_map.clear() - self.__call_index_layer_index_map.clear() - self.__activations_transfer_event_map.clear() + self.__boundary_activations.clear() self.__temp_device_activations_allocator.deallocate() def __wait_all_layer_train(self): @@ -828,11 +1227,6 @@ def __wait_all_layer_transfers(self): for layer_index in range(len(self.__layers)): self.__wait_layer_transfer(layer_index) - def __wait_all_activation_transfers(self): - call_indices = list(self.__activations_transfer_event_map.keys()) - for call_index in call_indices: - self.__wait_activations_transfer(call_index) - def __wait_layer_train(self, layer_index: int): self.__layer_train_event_map[layer_index] \ .wait(self.__layer_transfer_stream, f"wait layer train {layer_index}") @@ -844,12 +1238,6 @@ def __wait_layer_transfer(self, layer_index: int): .wait(self.__train_stream, f"wait layer transfer {layer_index}") self.__layer_transfer_event_map[layer_index] = SyncEvent() - def __wait_activations_transfer(self, call_index: int): - event = self.__activations_transfer_event_map.pop(call_index, None) - - if event is not None: - event.wait(self.__train_stream, f"wait activations transfer {call_index}") - def __schedule_layer_to( self, layer_index: int, @@ -870,7 +1258,7 @@ def __schedule_layer_to( else self.__temp_device_layer_allocator allocator = layer_allocator.get_allocator(layer_index, is_forward) - allocator_fn = allocator.allocate_like if allocator is not None else None + place_fn = allocator.place if allocator is not None else None if not is_forward and device_equals(device, self.__temp_device): layer = self.__layers[layer_index] @@ -895,7 +1283,7 @@ def __schedule_layer_to( self.__wait_layer_train(layer_index) layer = self.__layers[layer_index] for module in layer.modules(): - offload_quantized(module, device, non_blocking=self.__async_transfer, allocator=allocator_fn) + offload_quantized(module, device, non_blocking=self.__async_transfer, place=place_fn) layer_deallocator.deallocate_layer(layer_index, deallocate_forward=is_forward) @@ -920,40 +1308,92 @@ def __schedule_deferred_layers_to_temp( continue self.__schedule_layer_to(layer_index, device=self.__temp_device, is_forward=False) - def __schedule_activations_to_device( - self, - activations: Any, - device: torch.device, - call_index: int, - wait_train_stream: bool, - ): - log(f"schedule {call_index} activations to {str(device)}") - layer_index = self.__call_index_layer_index_map[call_index] + def pack_activation(self, layer_index: int, tensor: torch.Tensor): + # Copies rather than moving in place, so a tensor still referenced by other blocks (e.g. a shared + # conditioning embedding) is never corrupted. The caller selects which tensors to offload. + if not self.__materialized or not self.__offload_activations \ + or not device_equals(tensor.device, self.__train_device) \ + or layer_index >= len(self.__layers) - MAX_LAYER_CALLS_IN_FLIGHT: + return tensor + + handle = _BoundaryActivation() + handle.stride = tensor.stride() + # Transfer the storage as-is when the tensor is dense, so both sides of the copy are contiguous and + # equal-length and the transfer stays a DMA. Non-dense tensors have no flat view and fall back to the + # logical copy, which reorders elements through a kernel but is always correct. + source = flat_storage_view(tensor) + handle.dense = source is not None + if source is None: + source = tensor + if not self.__warned_non_dense: + # not expected with the current save set (SDPA and mm outputs are fresh allocations), but a + # chunk or slice would land here, so say so rather than silently paying the kernel path. + # Warn rather than raise: a save set that widens to include a view should get slower, not + # abort a training run. + self.__warned_non_dense = True + print("non-dense activation offloaded via the logical copy path: " + f"layer {layer_index}, shape {tuple(tensor.shape)}, stride {tensor.stride()}") - activations_allocator = self.__temp_device_activations_allocator \ - if device_equals(device, self.__temp_device) \ - else None + with create_stream_context(self.__activations_transfer_stream): + if self.__async_transfer: + self.__activations_transfer_stream.wait_stream(self.__train_stream) + self.__temp_device_activations_allocator.reserve_cache([tensor]) + handle.cpu = self.__temp_device_activations_allocator.allocate_like(tensor) + destination = handle.cpu.view(-1) if handle.dense else handle.cpu + destination.copy_(source, non_blocking=self.__async_transfer) + if self.__async_transfer: + tensors_record_stream(self.__activations_transfer_stream, tensor) # source alive until copied - allocator_fn = activations_allocator.allocate_like if activations_allocator is not None else None + self.__boundary_activations.setdefault(layer_index, []).append(handle) + return handle - event = None - if wait_train_stream and self.__async_transfer: - event = SyncEvent(self.__train_stream.record_event(), f"train before activations transfer {call_index}") + def __reload_activation(self, handle: '_BoundaryActivation'): + if handle.gpu is not None: + return + # Allocate under the train stream: the allocator's free lists are per-stream, so a destination + # homed on the transfer stream can never be reused by compute and forms a second segment pool. + # Dense tensors allocate contiguous and get the recorded layout as_strided over them, so the DMA + # moves storage without reordering. empty_strided is the non-dense fallback. + with create_stream_context(self.__train_stream): + if handle.dense: + flat = torch.empty(handle.cpu.numel(), dtype=handle.cpu.dtype, device=self.__train_device) + handle.gpu = torch.as_strided(flat, handle.cpu.shape, handle.stride) + else: + handle.gpu = torch.empty_strided( + handle.cpu.shape, handle.stride, dtype=handle.cpu.dtype, device=self.__train_device) - with create_stream_context(self.__activations_transfer_stream): - tensor_indices = self.__layer_activations_included_offload_param_indices_map[layer_index] + # the allocator may have just reclaimed this block from still-queued train work, which the H2D on + # the transfer stream would otherwise overwrite. Recorded, not waited on here, so the transfer + # still overlaps the current layer's compute. + event = SyncEvent(self.__train_stream.record_event(), "train before activation reload") \ + if self.__async_transfer else None + with create_stream_context(self.__activations_transfer_stream): if event is not None: event.wait(self.__activations_transfer_stream) - - tensors = get_tensor_data(activations, tensor_indices) - if activations_allocator is not None: - activations_allocator.reserve_cache(tensors) - tensors_to_device_(activations, device, tensor_indices, non_blocking=self.__async_transfer, allocator=allocator_fn) - + if handle.dense: + flat_storage_view(handle.gpu).copy_(handle.cpu.view(-1), non_blocking=self.__async_transfer) + else: + handle.gpu.copy_(handle.cpu, non_blocking=self.__async_transfer) if self.__async_transfer: - tensors_record_stream(self.__activations_transfer_stream, tensors) - self.__activations_transfer_event_map[call_index] = \ - SyncEvent(self.__activations_transfer_stream.record_event(), f"transfer to {device}") + tensors_record_stream(self.__activations_transfer_stream, handle.gpu) + handle.event = SyncEvent(self.__activations_transfer_stream.record_event()) - del tensors + def prefetch_activations(self, layer_index: int): + # reload a block's offloaded activations one block ahead so unpack only waits on the transfer + if not self.__materialized or not self.__offload_activations: + return + for handle in self.__boundary_activations.get(layer_index, []): + self.__reload_activation(handle) + + def unpack_activation(self, handle: Any): + if not isinstance(handle, _BoundaryActivation): + return handle # was not offloaded + self.__reload_activation(handle) # no-op if already prefetched + if self.__async_transfer and handle.event is not None: + handle.event.wait(self.__train_stream) + tensors_record_stream(self.__train_stream, handle.gpu) + gpu = handle.gpu + handle.gpu = None + handle.event = None + return gpu diff --git a/modules/util/ModelNames.py b/modules/util/ModelNames.py index a4b7310ff..daaec2e19 100644 --- a/modules/util/ModelNames.py +++ b/modules/util/ModelNames.py @@ -14,8 +14,10 @@ def __init__( base_model: str = "", prior_model: str = "", transformer_model: str = "", + low_noise_transformer_model: str = "", effnet_encoder_model: str = "", decoder_model: str = "", + text_encoder_model: str = "", text_encoder_4: str = "", vae_model: str = "", lora: str = "", @@ -26,12 +28,15 @@ def __init__( include_text_encoder_3: bool = True, include_text_encoder_4: bool = True, include_unconditional_transformer: bool = True, + include_low_noise_transformer: bool = True, ): self.base_model = base_model self.prior_model = prior_model self.transformer_model = transformer_model + self.low_noise_transformer_model = low_noise_transformer_model self.effnet_encoder_model = effnet_encoder_model self.decoder_model = decoder_model + self.text_encoder_model = text_encoder_model self.text_encoder_4 = text_encoder_4 self.vae_model = vae_model self.lora = lora @@ -42,6 +47,7 @@ def __init__( self.include_text_encoder_3 = include_text_encoder_3 self.include_text_encoder_4 = include_text_encoder_4 self.include_unconditional_transformer = include_unconditional_transformer + self.include_low_noise_transformer = include_low_noise_transformer def all_embedding(self): if self.embedding is not None: diff --git a/modules/util/ModelWeightDtypes.py b/modules/util/ModelWeightDtypes.py index dbfdfb74f..b3f0d02be 100644 --- a/modules/util/ModelWeightDtypes.py +++ b/modules/util/ModelWeightDtypes.py @@ -16,6 +16,8 @@ def __init__( text_encoder_2: DataType, text_encoder_3: DataType, text_encoder_4: DataType, + connectors: DataType, + low_noise_transformer: DataType, vae: DataType, effnet_encoder: DataType, decoder: DataType, @@ -35,6 +37,8 @@ def __init__( self.text_encoder_2 = text_encoder_2 self.text_encoder_3 = text_encoder_3 self.text_encoder_4 = text_encoder_4 + self.connectors = connectors + self.low_noise_transformer = low_noise_transformer self.vae = vae self.effnet_encoder = effnet_encoder self.decoder = decoder @@ -53,6 +57,8 @@ def all_dtypes(self) -> list: self.text_encoder_2, self.text_encoder_3, self.text_encoder_4, + self.connectors, + self.low_noise_transformer, self.vae, self.effnet_encoder, self.decoder, diff --git a/modules/util/checkpointing_util.py b/modules/util/checkpointing_util.py index 2df40065e..4305efc05 100644 --- a/modules/util/checkpointing_util.py +++ b/modules/util/checkpointing_util.py @@ -6,7 +6,6 @@ from modules.util.compile_util import init_compile from modules.util.config.TrainConfig import TrainConfig, TrainModelPartConfig from modules.util.LayerOffloadConductor import LayerOffloadConductor -from modules.util.torch_util import add_dummy_grad_fn_, has_grad_fn import torch from torch import nn @@ -20,6 +19,8 @@ ) from transformers.models.clip.modeling_clip import CLIPEncoderLayer from transformers.models.gemma2.modeling_gemma2 import Gemma2DecoderLayer +from transformers.models.gemma3.modeling_gemma3 import Gemma3DecoderLayer +from transformers.models.gemma4_unified.modeling_gemma4_unified import Gemma4UnifiedTextDecoderLayer from transformers.models.llama.modeling_llama import LlamaDecoderLayer from transformers.models.mistral.modeling_mistral import MistralDecoderLayer from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import Qwen2_5_VLDecoderLayer @@ -45,6 +46,13 @@ def _kwargs_to_args(fun: Callable, args: tuple[Any, ...], kwargs: dict[str, Any] return tuple(parameters) +def _view_key(tensor: torch.Tensor) -> tuple: + # Identifies a tensor by the memory it reads: first-element address, extents, steps, element type. + # Two tensors agreeing on all four are the same view of the same data, since live storages cannot + # overlap. Used instead of identity because autograd hands back aliases, not the original objects. + return tensor.data_ptr(), tensor.shape, tensor.stride(), tensor.dtype + + def __get_args_indices(fun: Callable, arg_names: list[str]) -> list[int]: signature = dict(inspect.signature(fun).parameters) indices = [] @@ -56,22 +64,13 @@ def __get_args_indices(fun: Callable, arg_names: list[str]) -> list[int]: return indices -__current_call_index = 0 - - -def _generate_call_index() -> int: - global __current_call_index - __current_call_index += 1 - return __current_call_index - - class BaseCheckpointLayer(torch.nn.Module): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) class CheckpointLayer(BaseCheckpointLayer): - def __init__(self, orig_module: nn.Module, orig_forward, train_device: torch.device, checkpointing: bool = True): + def __init__(self, orig_module: nn.Module, orig_forward, checkpointing: bool = True): super().__init__() assert (orig_module is None or orig_forward is None) and not (orig_module is None and orig_forward is None) @@ -79,20 +78,13 @@ def __init__(self, orig_module: nn.Module, orig_forward, train_device: torch.dev self.orig_forward = orig_forward self.checkpointing = checkpointing - # dummy tensor that requires grad is needed for checkpointing to work when training a LoRA - self.dummy = torch.zeros((1,), device=train_device, requires_grad=True) - def __orig(self, *args, **kwargs): return self.orig_forward(*args, **kwargs) if self.checkpoint is None else self.checkpoint(*args, **kwargs) - def __checkpointing_forward(self, dummy: torch.Tensor, *args, **kwargs): - return self.__orig(*args, **kwargs) - def forward(self, *args, **kwargs): if self.checkpointing and torch.is_grad_enabled(): return torch.utils.checkpoint.checkpoint( - self.__checkpointing_forward, - self.dummy, + self.__orig, *args, **kwargs, use_reentrant=False @@ -100,18 +92,88 @@ def forward(self, *args, **kwargs): else: return self.__orig(*args, **kwargs) -class OffloadCheckpointLayer(BaseCheckpointLayer): - def __init__(self, orig_module: nn.Module, orig_forward, train_device: torch.device, conductor: LayerOffloadConductor, layer_index: int, checkpointing: bool): + +class LoadBoundary(torch.autograd.Function): + # Identity op on a block's INPUT tensors, driving the conductor only: forward calls before_layer, + # backward calls after_layer. Sitting on the input makes its backward fire LAST in the block's + # backward. `dummy` keeps the node alive when the inputs carry no grad, see _apply_boundary. + @staticmethod + def forward(ctx, conductor, layer_index, dummy, *tensors): + ctx.conductor = conductor + ctx.layer_index = layer_index + conductor.before_layer(layer_index, is_forward=True) + return tensors + @staticmethod + def backward(ctx, *grads): + # calling after_layer here, rather than after the forward, makes the train event cover the block's + # backward kernels, so this layer's offload transfer cannot overlap them. + ctx.conductor.after_layer(ctx.layer_index, list(grads)) + return (None, None, None, *grads) + + +class EvictBoundary(torch.autograd.Function): + # Identity op on a block's grad-carrying OUTPUT tensors: forward calls after_layer, backward calls + # before_layer. Sitting on the output makes its backward fire FIRST - before the checkpoint + # rematerializes the block, so the weights are in by then. + @staticmethod + def forward(ctx, conductor, layer_index, dummy, *tensors): + ctx.conductor = conductor + ctx.layer_index = layer_index + conductor.after_layer(layer_index, list(tensors)) + return tensors + @staticmethod + def backward(ctx, *grads): + # workaround for https://github.com/pytorch/pytorch/issues/186537. This is the block's first + # backward node, and it runs on the autograd worker thread - which is where AOTAutograd + # compiles the backward graph, and which does not inherit the main thread's dynamo config. + init_compile() + ctx.conductor.before_layer(ctx.layer_index, is_forward=False) + ctx.conductor.prefetch_activations(ctx.layer_index - 1) + return (None, None, None, *grads) + + +def _apply_boundary(boundary, conductor, layer_index, values: tuple, dummy: torch.Tensor | None = None) -> tuple: + # Wrap all grad-requiring tensors in `values` in a single boundary Function, so its backward fires + # exactly once around the checkpoint's recompute and no gradient can reach the checkpoint without + # crossing the boundary. Non-tensor and grad-free entries pass through untouched. + indices = [i for i, v in enumerate(values) if isinstance(v, torch.Tensor) and v.requires_grad] + if not indices and dummy is not None: + # Nothing entering the block carries grad, because nothing trainable sits in front of it. A + # Function whose inputs all lack grad creates no autograd node, so the boundary would not be + # applied at all and the block would run against weights still on the temp device. Force the first + # float tensor through against `dummy`, a grad-requiring leaf: the node then exists, and the + # block's input gradient - discarded at the boundary - is what fires the backward duty. + indices = [i for i, v in enumerate(values) if isinstance(v, torch.Tensor) and v.is_floating_point()][:1] + if not indices: + return values + wrapped = boundary.apply(conductor, layer_index, dummy, *(values[i] for i in indices)) + values = list(values) + for j, i in enumerate(indices): + values[i] = wrapped[j] + return tuple(values) + + +class BoundaryOffloadCheckpointLayer(BaseCheckpointLayer): + # Weight movement driven by the LoadBoundary / EvictBoundary autograd Functions, which run eagerly + # outside the compiled region. The block's forward and backward each stay a single traced graph, which + # is what makes this path compatible with torch.compile's cudagraph trees. Activation offloading, when + # enabled, rides on saved_tensors_hooks around the block. + def __init__(self, orig_module: nn.Module, orig_forward, conductor: LayerOffloadConductor, layer_index: int, checkpointing: bool, included_offload_param_indices: list[int], compile: bool): super().__init__() assert (orig_module is None or orig_forward is None) and not (orig_module is None and orig_forward is None) self.checkpoint = orig_module self.orig_forward = orig_forward - - self.dummy = torch.zeros((1,), device=train_device, requires_grad=True) self.conductor = conductor self.layer_index = layer_index self.checkpointing = checkpointing + self.included_offload_param_indices = included_offload_param_indices + self.dummy = None + # compile the block together with its checkpoint, not the bare block: dynamo then traces the + # checkpoint as a higher-order op and the min-cut partitioner prunes the recompute to what the + # backward needs. With the checkpoint outside the compiled region the backward re-runs the whole + # block. Traceable at all only because no conductor call happens in here. + self.run_block = torch.compile(self.__checkpointed_block, fullgraph=True) if compile else self.__checkpointed_block def __deepcopy__(self, memo): # conductor holds torch.cuda.Stream/Event objects that cannot be deep-copied or pickled. @@ -120,57 +182,78 @@ def __deepcopy__(self, memo): cls = self.__class__ result = cls.__new__(cls) memo[id(self)] = result + # run_block is shared for the same reason as conductor: it closes over this instance, and + # deepcopy is only used at save time where it is never invoked. for key, value in self.__dict__.items(): - result.__dict__[key] = value if key == "conductor" else copy.deepcopy(value, memo) + result.__dict__[key] = value if key in ("conductor", "run_block") else copy.deepcopy(value, memo) return result - def __checkpointing_forward(self, dummy: torch.Tensor, call_id: int, *args): - init_compile() # workaround for https://github.com/pytorch/pytorch/issues/186537 - if self.layer_index == 0 and not torch.is_grad_enabled(): - self.conductor.start_forward(True) + def __orig(self, *args): + return self.orig_forward(*args) if self.checkpoint is None else self.checkpoint(*args) + + def __checkpointed_block(self, *args): + if self.checkpointing: + # recompute the block in the backward, pruned by the min-cut partitioner when compiled + return torch.utils.checkpoint.checkpoint(self.__orig, *args, use_reentrant=False) + # no checkpointing: run once, autograd keeps every activation the backward needs + return self.__orig(*args) + + def __run_block(self, args): + def run(): + return self.run_block(*args) + + if not self.conductor.offloads_activations(): + return run() + # Offload only the declared activation args. The declaration carries cross-block knowledge the + # min-cut partitioner cannot have, seeing one block at a time: a tensor shared by every block + # (rotary embeddings, masks) is saved once per block, and offloading those copies frees nothing + # because the original stays live for the remaining blocks. + # + # Matched by view rather than by identity, because autograd saves an alias of the arg rather than + # the arg object, so id() misses. Grad-agnostic: a LoRA layer filter can leave an offloaded + # activation grad-free. + layer_index = self.layer_index + targets = {_view_key(args[i]) for i in self.included_offload_param_indices + if i < len(args) and isinstance(args[i], torch.Tensor)} + + def pack(t): + return self.conductor.pack_activation(layer_index, t) if _view_key(t) in targets else t + + with torch.autograd.graph.saved_tensors_hooks(pack, self.conductor.unpack_activation): + return run() - args = self.conductor.before_layer(self.layer_index, call_id, args) - output = self.orig_forward(*args) if self.checkpoint is None else self.checkpoint(*args) + def forward(self, *args, **kwargs): + args = _kwargs_to_args(self.orig_forward if self.checkpoint is None else self.checkpoint.forward, args, kwargs) - self.conductor.after_layer(self.layer_index, call_id, args) + if not torch.is_grad_enabled(): + # inference / frozen: no backward flows, so no boundaries are needed. Schedule, run + # and record inline. + if self.layer_index == 0: + self.conductor.start_forward(backward_follows=False) + self.conductor.before_layer(self.layer_index, is_forward=True) + # still through run_block: under no_grad the checkpoint inside it is a passthrough, and this + # keeps sampling on the compiled path (dynamo traces a separate inference variant) + output = self.run_block(*args) + self.conductor.after_layer(self.layer_index, list(args)) + return output - # make sure at least one of the output tensors has a grad_fn so the output of the checkpoint has a grad_fn. - # this can only happen if a checkpointed block has no trainable parameters, because of a layer filter - # was used. Adding a dummy grad function is a workaround required by use_reentrant==True checkpointing: - if torch.is_grad_enabled() and not has_grad_fn(output): - output = add_dummy_grad_fn_(output) + if self.layer_index == 0: + self.conductor.start_forward(backward_follows=True) - return output + if self.dummy is None: + device = next((v.device for v in args if isinstance(v, torch.Tensor)), None) + if device is not None: + self.dummy = torch.zeros((1,), device=device, requires_grad=True) - def forward(self, *args, **kwargs): - call_id = _generate_call_index() - args = _kwargs_to_args(self.orig_forward if self.checkpoint is None else self.checkpoint.forward, args, kwargs) - if torch.is_grad_enabled(): - # a backward will flow through this layer (grad enabled), so offloading needs use_reentrant=True - # checkpointing to move the offloaded tensors back during recompute. Fail loud rather than silently - # enabling checkpointing the part disabled. Under no_grad (e.g. sampling a frozen part) the branch - # below offloads without checkpointing. - if not self.checkpointing: - raise NotImplementedError("offloading requires gradient checkpointing") - return torch.utils.checkpoint.checkpoint( - self.__checkpointing_forward, - self.dummy, - call_id, - *args, - use_reentrant=True - ) - else: - if self.layer_index == 0: - self.conductor.start_forward(False) + args = _apply_boundary(LoadBoundary, self.conductor, self.layer_index, args, self.dummy) + output = self.__run_block(args) + output_tuple = output if isinstance(output, tuple) else (output,) + output_tuple = _apply_boundary(EvictBoundary, self.conductor, self.layer_index, output_tuple) + return output_tuple if isinstance(output, tuple) else output_tuple[0] - args = self.conductor.before_layer(self.layer_index, call_id, args) - output = self.orig_forward(*args) if self.checkpoint is None else self.checkpoint(*args) - self.conductor.after_layer(self.layer_index, call_id, args) - return output def create_checkpoint( orig_module: nn.Module, - train_device: torch.device, include_from_offload_param_names: list[str] = None, conductor: LayerOffloadConductor | None = None, checkpointing: bool = True, @@ -182,32 +265,43 @@ def create_checkpoint( included_offload_param_indices = __get_args_indices(orig_module.forward, include_from_offload_param_names) if conductor is not None: - conductor.add_layer(orig_module, included_offload_param_indices) + conductor.add_layer(orig_module) if conductor is not None and conductor.offload_activated(): - # offloading is structurally coupled to use_reentrant=True checkpointing during the back pass: the - # recompute is the only thing firing before_layer/after_layer in the backward direction, so both layer - # and activation offloading need checkpointing to move tensors back for backward. That coupling only - # matters when a backward actually flows, so OffloadCheckpointLayer.forward enforces it per call (fail - # loud under grad, offload freely under no_grad) instead of rejecting the frozen/inference case here. + # Compiled, the boundary layer compiles the block together with its checkpoint, so orig_module must + # not be compiled separately here - that would put the checkpoint back outside the compiled region + # and force a full recompute. if compile: - layer = OffloadCheckpointLayer(orig_module=orig_module, orig_forward=None, train_device=train_device, conductor=conductor, layer_index=layer_index, checkpointing=checkpointing) - #don't compile the checkpointing layer - offloading cannot be compiled: - orig_module.compile(fullgraph=True) - return layer + return BoundaryOffloadCheckpointLayer( + orig_module=orig_module, + orig_forward=None, + conductor=conductor, + layer_index=layer_index, + checkpointing=checkpointing, + included_offload_param_indices=included_offload_param_indices, + compile=True, + ) else: #only patch forward() if possible. Inserting layers is necessary for torch.compile, but causes issues with at least 1 text encoder model. we don't compile text encoders - layer = OffloadCheckpointLayer(orig_module=None, orig_forward=orig_module.forward, train_device=train_device, conductor=conductor, layer_index=layer_index, checkpointing=checkpointing) + layer = BoundaryOffloadCheckpointLayer( + orig_module=None, + orig_forward=orig_module.forward, + conductor=conductor, + layer_index=layer_index, + checkpointing=checkpointing, + included_offload_param_indices=included_offload_param_indices, + compile=False, + ) orig_module.forward = layer.forward return orig_module else: if compile: - layer = CheckpointLayer(orig_module=orig_module, orig_forward=None, train_device=train_device, checkpointing=checkpointing) + layer = CheckpointLayer(orig_module=orig_module, orig_forward=None, checkpointing=checkpointing) #do compile the checkpointing layer - slightly faster layer.compile(fullgraph=True) return layer else: - layer = CheckpointLayer(orig_module=None, orig_forward=orig_module.forward, train_device=train_device, checkpointing=checkpointing) + layer = CheckpointLayer(orig_module=None, orig_forward=orig_module.forward, checkpointing=checkpointing) orig_module.forward = layer.forward return orig_module @@ -216,7 +310,6 @@ def _create_checkpoints_for_module_list( include_from_offload_param_names: list[str], conductor: LayerOffloadConductor, checkpointing: bool, - train_device: torch.device, layer_index: int, compile: bool, ) -> int: @@ -225,7 +318,7 @@ def _create_checkpoints_for_module_list( if isinstance(module_list[i], BaseCheckpointLayer): continue module_list[i] = create_checkpoint( - layer, train_device, + layer, include_from_offload_param_names, conductor, checkpointing, layer_index, compile=compile, ) @@ -245,20 +338,56 @@ def enable_checkpointing( lists, # if there are multiple entries in this list, they must be in the exact order they are executed - otherwise offloading fails supports_offloading: bool = True, ) -> LayerOffloadConductor | None: + # A full fine-tune updates the base weights, but meta-eviction (stream_from_disk + cache_in_ram off) re-streams + # them from the checkpoint on each use, discarding those updates. Reject that combo. + if config.stream_from_disk and config.part_trained_in_place(part) and not part.cache_in_ram: + raise NotImplementedError( + "a fully fine-tuned component cannot stream from disk without keeping it cached in RAM: it re-streams " + "weights from the checkpoint on each use, discarding training updates. Enable 'Cache In RAM' for this " + "component") + + # the full-model-buffer offload (simplex) is opt-in per part. It needs a never-changing base (the RAM buffer is + # filled once per materialize and never written back, so in-place weight updates would be lost), a disk-streamed + # part (the buffer is filled from the streamed weights) and layer offloading (the buffer only pays off as the + # offload target of a layer ring). It does not need cache_in_ram: with it off, the buffer is freed on evict and + # rebuilt from a fresh stream on the next materialize. Reject an unusable combination instead of silently + # downgrading to the static path. + # only when the component actually offloads layers. Simplex is a layer-offloading mode, so with layer offloading + # off there is no conductor to give it to and the switch does nothing -- which makes it a stale setting rather + # than a bad one, and the combinations below are only worth rejecting for a component that would use it. + simplex = part.simplex_offloading and supports_offloading and part.offload_fraction > 0 + if simplex: + if config.part_trained_in_place(part): + raise NotImplementedError( + "a fully fine-tuned component cannot use 'Simplex Offloading': its weights live in a RAM buffer that " + "is filled from the checkpoint, so in-place training updates would be discarded. Disable 'Simplex " + "Offloading' for this component") + if not config.stream_from_disk: + raise NotImplementedError( + "'Simplex Offloading' requires 'Stream From Disk': the RAM buffer is filled from the streamed " + "weights. Enable 'Stream From Disk' on the model page, or disable 'Simplex Offloading' for this " + "component") + if not (supports_offloading and part.offload_fraction > 0): + raise NotImplementedError( + "'Simplex Offloading' is a layer-offloading mode and needs this component to offload layers: set " + "'Layer Offload Fraction' above 0, or disable 'Simplex Offloading' for this component") + if not part.checkpointing_or_offloading_enabled() and not compile: return None # a conductor exists iff this part actually offloads: the user enabled it (part.offloading_enabled()) and the # architecture can be driven by the conductor (supports_offloading). offload = supports_offloading and part.offloading_enabled() - conductor = LayerOffloadConductor(model, config, part) if offload else None + conductor = LayerOffloadConductor(model, config, part, simplex=simplex) if offload else None checkpointing = part.checkpointing_enabled() - # a trained part always has grad flowing through it, so offloading without checkpointing is guaranteed to hit - # OffloadCheckpointLayer.forward's fail-loud path. Reject it here so the misconfiguration surfaces at setup - # instead of the first training step. Frozen parts (part.train == False) are left to the per-call check: e.g. - # Ideogram's unconditional transformer runs only under no_grad during sampling, so it offloads without - # checkpointing there, while a frozen denoiser/TE still fails loud when a trained embedding routes grad through it. + # Offloading requires checkpointing. A block's backward reads its weights through SavedVariables taken + # during the forward, which hold a shallow copy - repointing param.data at a reloaded buffer never + # redirects them - while the conductor recycles weight buffers between layers, so by then that buffer + # can already hold a different layer's weights: plausible but wrong gradients, not an error. The + # recompute re-reads the weights after the layer has been loaded back in. Compiled blocks escape this + # today (AOTAutograd passes parameters as graph inputs, read at call time), but that is a calling + # convention and not a guarantee. Frozen parts run no backward through the block and are exempt. if offload and not checkpointing and part.train: raise NotImplementedError("offloading requires gradient checkpointing") @@ -273,7 +402,6 @@ def enable_checkpointing( param_names, conductor, checkpointing, - torch.device(config.train_device), layer_index, compile = compile, ) @@ -288,7 +416,6 @@ def enable_checkpointing( param_names, conductor, checkpointing, - torch.device(config.train_device), layer_index, compile = compile, ) @@ -355,6 +482,26 @@ def enable_checkpointing_for_mistral_encoder_layers( ]) +def enable_checkpointing_for_gemma3_encoder_layers( + model: nn.Module, + config: TrainConfig, + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, False, [ + (Gemma3DecoderLayer, []), # no activation offloading: this encoder is never trained + ]) + + +def enable_checkpointing_for_gemma4_encoder_layers( + model: nn.Module, + config: TrainConfig, + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, False, [ + (Gemma4UnifiedTextDecoderLayer, []), # no activation offloading: this encoder is never trained + ]) + + def enable_checkpointing_for_qwen25vl_encoder_layers( model: nn.Module, @@ -424,6 +571,26 @@ def enable_checkpointing_for_qwen_transformer( (model.transformer_blocks, ["hidden_states", "encoder_hidden_states"]), ]) +def enable_checkpointing_for_ltx_transformer( + model: nn.Module, + config: TrainConfig, + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + # LTX2VideoTransformerBlock.forward returns (hidden_states, audio_hidden_states) + return enable_checkpointing(model, config, part, config.compile, [ + (model.transformer_blocks, ["hidden_states", "audio_hidden_states"]), + ]) + +def enable_checkpointing_for_ltx_connectors( + model: nn.Module, + config: TrainConfig, + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, False, [ + (model.video_connector.transformer_blocks, ["hidden_states"]), + (model.audio_connector.transformer_blocks, ["hidden_states"]), + ]) + def enable_checkpointing_for_z_image_transformer( model: nn.Module, config: TrainConfig, diff --git a/modules/util/compile_util.py b/modules/util/compile_util.py index 0a09f1023..ffada3b2a 100644 --- a/modules/util/compile_util.py +++ b/modules/util/compile_util.py @@ -1,9 +1,10 @@ +from modules.util.tqdm_util import tqdm + import torch import torch._dynamo.callback import torch.utils._sympy.functions from sympy import S -from tqdm import tqdm #code from https://github.com/pytorch/pytorch/blob/ed82d5fcfd80110565f69130f286c7bfec6db2dc/torch/utils/_sympy/functions.py#L481 @@ -95,7 +96,7 @@ def init_compile(): def _on_compile_start(args: "torch._dynamo.callback.CallbackArgs") -> None: frame_id, _, frame_compile_id = args.compile_id.partition("/") direction = "backward" if args.callback_trigger == torch._dynamo.callback.CallbackTrigger.LAZY_BACKWARD else "forward" - tqdm.write(f"[torch.compile] compiling kernel {frame_id} {direction} (variant #{frame_compile_id or 0})...") + tqdm.show_status(f"compiling kernel {frame_id} {direction} (variant #{frame_compile_id or 0})...") torch._dynamo.callback.on_compile_start(_on_compile_start) diff --git a/modules/util/config/SampleConfig.py b/modules/util/config/SampleConfig.py index 9b2b2c0b1..76ef1c52d 100644 --- a/modules/util/config/SampleConfig.py +++ b/modules/util/config/SampleConfig.py @@ -2,6 +2,7 @@ from modules.util.config.BaseConfig import BaseConfig from modules.util.enum.NoiseScheduler import NoiseScheduler +from modules.util.enum.SamplingMethod import SamplingMethod def _get_model_defaults(model_type) -> dict: @@ -17,6 +18,8 @@ def _get_model_defaults(model_type) -> dict: "cfg_scale": 7.0, "noise_scheduler": NoiseScheduler.DDIM, "negative_prompt": "", + # None keeps the sampler's derived shift; a model whose reference pipeline pins a fixed one sets it here + "override_shift": None, } if model_type is None: @@ -143,6 +146,14 @@ def _get_model_defaults(model_type) -> dict: "diffusion_steps": 25, "cfg_scale": 4.0, }) + elif model_type.is_ltx_2(): + defaults.update({ + "width": 960, + "height": 540, + "diffusion_steps": 40, + "cfg_scale": 4.0, + "override_shift": 7.8, + }) elif model_type.is_ideogram(): # Ideogram 4 recommends 48 flow-matching steps on a logit-normal schedule with # guidance held at 7.0 for the main steps (dropping to 3.0 for the final polish steps). @@ -170,6 +181,8 @@ class SampleConfig(BaseConfig): diffusion_steps: int cfg_scale: float noise_scheduler: NoiseScheduler + sampling_method: SamplingMethod + override_shift: float | None text_encoder_1_layer_skip: int text_encoder_1_sequence_length: int | None @@ -214,6 +227,8 @@ def default_values(model_type=None): data.append(("diffusion_steps", defaults["diffusion_steps"], int, False)) data.append(("cfg_scale", defaults["cfg_scale"], float, False)) data.append(("noise_scheduler", defaults["noise_scheduler"], NoiseScheduler, False)) + data.append(("sampling_method", SamplingMethod.HANDOFF_LOW_NOISE, SamplingMethod, False)) + data.append(("override_shift", defaults["override_shift"], float, True)) data.append(("text_encoder_1_layer_skip", 0, int, False)) data.append(("text_encoder_1_sequence_length", None, int, True)) diff --git a/modules/util/config/TrainConfig.py b/modules/util/config/TrainConfig.py index ba11de9d7..5c950c891 100644 --- a/modules/util/config/TrainConfig.py +++ b/modules/util/config/TrainConfig.py @@ -269,7 +269,9 @@ class TrainModelPartConfig(BaseConfig): guidance_scale: float gradient_checkpointing: bool offload_fraction: float + simplex_offloading: bool activation_offloading: bool + cache_in_ram: bool def __init__(self, data: list[(str, Any, type, bool)]): super().__init__(data) @@ -309,7 +311,9 @@ def default_values(): data.append(("guidance_scale", 1.0, float, False)) data.append(("gradient_checkpointing", True, bool, False)) data.append(("offload_fraction", 0.0, float, False)) + data.append(("simplex_offloading", False, bool, False)) data.append(("activation_offloading", False, bool, False)) + data.append(("cache_in_ram", True, bool, False)) return TrainModelPartConfig(data) @@ -349,6 +353,7 @@ class QuantizationConfig(BaseConfig): layer_filter: str layer_filter_preset: str layer_filter_regex: bool + fallback_dtype: DataType svd_dtype: DataType svd_rank: int cache_dir: str @@ -361,6 +366,7 @@ def default_values(): data.append(("layer_filter", "", str, False)) data.append(("layer_filter_preset", "full", str, False)) data.append(("layer_filter_regex", False, bool, False)) + data.append(("fallback_dtype", DataType.BFLOAT_16, DataType, False)) data.append(("svd_dtype", DataType.NONE, DataType, False)) data.append(("svd_rank", 16, int, False)) data.append(("cache_dir", None, str, True)) @@ -402,6 +408,7 @@ class TrainConfig(BaseConfig): async_offloading: bool force_circular_padding: bool compile: bool + stream_from_disk: bool # data settings concept_file_name: str @@ -499,6 +506,9 @@ class TrainConfig(BaseConfig): text_encoder_4: TrainModelPartConfig text_encoder_4_layer_skip: int + connectors: TrainModelPartConfig + low_noise_transformer: TrainModelPartConfig + # vae vae: TrainModelPartConfig @@ -882,6 +892,8 @@ def weight_dtypes(self) -> ModelWeightDtypes: self.text_encoder_2.weight_dtype, self.text_encoder_3.weight_dtype, self.text_encoder_4.weight_dtype, + self.connectors.weight_dtype, + self.low_noise_transformer.weight_dtype, self.vae.weight_dtype, self.effnet_encoder.weight_dtype, self.decoder.weight_dtype, @@ -891,13 +903,27 @@ def weight_dtypes(self) -> ModelWeightDtypes: self.embedding_weight_dtype, ) + def cache_in_ram(self) -> dict[str, bool]: + return {part: getattr(self, part).cache_in_ram for part in self.model_type.model_parts()} + + def part_trained_in_place(self, part: TrainModelPartConfig) -> bool: + # True iff a FINE_TUNE run updates this part's base weights. 'train' defaults True even for parts the + # architecture can't train (e.g. a frozen text encoder), so also require the model type to list the part as + # trainable. Gates the offload/streaming modes that would silently discard in-place weight updates. + if self.training_method != TrainingMethod.FINE_TUNE or not part.train: + return False + name = next((p for p in self.model_type.model_parts() if getattr(self, p) is part), None) + return name in self.model_type.trainable_parts() + def model_names(self) -> ModelNames: return ModelNames( base_model=self.base_model_name, prior_model=self.prior.model_name, transformer_model=self.transformer.model_name, + low_noise_transformer_model=self.low_noise_transformer.model_name, effnet_encoder_model=self.effnet_encoder.model_name, decoder_model=self.decoder.model_name, + text_encoder_model=self.text_encoder.model_name, text_encoder_4=self.text_encoder_4.model_name, vae_model=self.vae.model_name, lora=self.lora_model_name, @@ -910,6 +936,7 @@ def model_names(self) -> ModelNames: include_text_encoder_3=self.text_encoder_3.include, include_text_encoder_4=self.text_encoder_4.include, include_unconditional_transformer=self.unconditional_transformer.include, + include_low_noise_transformer=self.low_noise_transformer.include, ) def train_any_embedding(self) -> bool: @@ -1046,6 +1073,7 @@ def default_values() -> 'TrainConfig': data.append(("async_offloading", True, bool, False)) data.append(("force_circular_padding", False, bool, False)) data.append(("compile", False, bool, False)) + data.append(("stream_from_disk", True, bool, False)) # data settings data.append(("concept_file_name", "training_concepts/concepts.json", str, False)) @@ -1179,6 +1207,22 @@ def default_values() -> 'TrainConfig': data.append(("text_encoder_4", text_encoder_4, TrainModelPartConfig, False)) data.append(("text_encoder_4_layer_skip", 0, int, False)) + # connectors + connectors = TrainModelPartConfig.default_values() + connectors.model_name = "" + connectors.train = False + connectors.gradient_checkpointing = False + connectors.activation_offloading = False + data.append(("connectors", connectors, TrainModelPartConfig, False)) + + # low noise transformer + low_noise_transformer = TrainModelPartConfig.default_values() + low_noise_transformer.model_name = "" + low_noise_transformer.train = False + low_noise_transformer.gradient_checkpointing = False + low_noise_transformer.activation_offloading = False + data.append(("low_noise_transformer", low_noise_transformer, TrainModelPartConfig, False)) + # vae vae = TrainModelPartConfig.default_values() vae.model_name = "" diff --git a/modules/util/disk_stream.py b/modules/util/disk_stream.py new file mode 100644 index 000000000..1195e6fad --- /dev/null +++ b/modules/util/disk_stream.py @@ -0,0 +1,119 @@ +from collections.abc import Callable + +from modules.module.quantized.mixin.QuantizedLinearMixin import QuantizedLinearMixin +from modules.util.enum.DataType import DataType +from modules.util.quantization_util import report_compression +from modules.util.torch_util import torch_gc + +import torch +from torch import nn + +# A streamed sub-module keeps its base weights frozen (LoRA training streams too -- only the adapter trains, so the +# streamed base weights never diverge from disk; a fully fine-tuned part cannot stream, its in-place updates would be +# discarded). It is loaded as a meta skeleton and its real weights are streamed straight from the checkpoint to the +# compute device and quantized the first time it is used -- so the full unquantized module never lands in system RAM. +# Both load paths share +# this materialize step and differ only in how they evict the weights off the compute device afterwards, selected by +# cache_in_ram: +# - cache_in_ram off: discard the weights to meta; re-materialize by re-streaming from the checkpoint. Frees both +# VRAM and RAM. Lossless because the module is frozen -- its weights never diverge from disk. +# - cache_in_ram on: keep the streamed+quantized weights resident on the temp device; re-materialize by moving them +# back to the compute device. Frees VRAM only, but avoids re-reading the checkpoint on every use. + + +def _is_evicted(module: nn.Module) -> bool: + # the skeleton is fully on meta between uses; a single real parameter means it is currently materialized + for parameter in module.parameters(): + return parameter.is_meta + return True + + +def _current_device(module: nn.Module) -> torch.device: + for parameter in module.parameters(): + return parameter.device + for buffer in module.buffers(): + return buffer.device + return torch.device("meta") + + +def evict_to_meta(module: nn.Module): + for sub_module in module.modules(): + for name, parameter in list(sub_module.named_parameters(recurse=False)): + if parameter.is_meta: + continue + if name == "weight" and isinstance(sub_module, QuantizedLinearMixin): + # a quantized weight is stored in a packed layout (nf4 packs to a flat [N, 1] tensor); reset it to a + # meta tensor of the original unpacked shape so the next materialize can stream the checkpoint weight + # back into it and re-quantize. Its dtype is irrelevant (the stream overwrites it), so keep the current. + sub_module.register_parameter(name, nn.Parameter( + torch.empty(sub_module.original_weight_shape(), dtype=parameter.dtype, device="meta"), + requires_grad=False)) + else: + sub_module.register_parameter( + name, nn.Parameter(parameter.detach().to("meta"), requires_grad=False)) + for name, buffer in list(sub_module._buffers.items()): + # non-persistent buffers (e.g. rotary inv_freq) are config-derived constants, not disk weights; + # keep them resident rather than evict and re-derive them. + if name in sub_module._non_persistent_buffers_set: + continue + if buffer is not None and not buffer.is_meta: + sub_module._buffers[name] = buffer.to("meta") + # let the next materialize() re-quantize the freshly streamed weights + if isinstance(sub_module, QuantizedLinearMixin): + sub_module.mark_needs_requantization() + + +def stream_module_to( + module: nn.Module, + device: torch.device, + materialize_fn: Callable[[nn.Module, torch.device, DataType], None], + train_dtype: DataType, + cache_in_ram: bool, + name: str, + temp_device: torch.device, +) -> bool: + # module.to()-style entry point for a materialize-on-demand component; see the module-level comment for the + # materialize/evict semantics. Idempotent; train_dtype is used only when materializing. Returns whether any + # weight actually moved, so the caller can skip the collection that follows an eviction that did nothing. + if device.type not in ("meta", temp_device.type): + # target is the compute device -> materialize the module onto it + current = _current_device(module) + try: + if current.type == "meta": + # cold: stream+quantize the weights from the checkpoint onto the compute device + materialize_fn(module, device, train_dtype, part_name=name) + # quantize_layers reports the saving for a resident component; a streamed one never goes through it + report_compression(module) + return True + elif current.type == temp_device.type: + # warm (cache_in_ram): the quantized weights are staged resident on the temp device, move them back to + # the compute device. Dispatch on device *type* (not equality) so a module already on the compute + # device isn't dragged through module.to(), which would raise on the non-persistent buffers left on meta. + module.to(device=device) + return True + except Exception: + # a materialize that fails partway (typically OOM) leaves already-streamed weights resident on the compute + # device -- live model state torch_gc can't reclaim, which can cascade into a second OOM. Roll back along the + # inverse of the failed move: a meta origin re-streams next time (drop the partial fill back to meta), a cpu + # origin keeps its RAM copy (move back to the temp device). + if current.type == "meta": + evict_to_meta(module) + # reclaim the partial fill now: this rollback runs under BaseModel.materialize, which (unlike + # evict) has no trailing torch_gc, so the stranded VRAM would otherwise survive into the re-raise. + torch_gc() + else: + module.to(device=current) + raise + return False + elif not cache_in_ram: + if not _is_evicted(module): + evict_to_meta(module) + return True + else: + # cache_in_ram: stage the resident quantized weights on the temp device. Only when currently on the compute + # device -- a module still on meta (never materialized) has nothing resident to stage and .to() can't move meta, + # so it stays a no-op here and streams from the checkpoint on its first materialize. + if _current_device(module).type not in (device.type, "meta"): + module.to(device=device) + return True + return False diff --git a/modules/util/enum/ModelType.py b/modules/util/enum/ModelType.py index 8820892e8..e0d48a242 100644 --- a/modules/util/enum/ModelType.py +++ b/modules/util/enum/ModelType.py @@ -49,6 +49,8 @@ class ModelType(Enum): IDEOGRAM_4 = 'IDEOGRAM_4' + LTX_2 = 'LTX_2' + def __str__(self): return self.value @@ -129,6 +131,16 @@ def is_ernie(self): def is_ideogram(self): return self == ModelType.IDEOGRAM_4 + def is_ltx_2(self): + return self == ModelType.LTX_2 + + def has_dynamic_timestep_shift(self) -> bool: + return self.is_flux_1() \ + or self.is_flux_2() \ + or self.is_qwen() \ + or self.is_krea2() \ + or self.is_ltx_2() + def supports_negative_prompt(self) -> bool: # asymmetric dual-network CFG models drive the negative branch from a frozen unconditional network (or an # empty prompt), not a user-supplied negative prompt @@ -152,6 +164,9 @@ def has_depth_input(self): def has_multiple_text_encoders(self): return "text_encoder_2" in self.model_parts() + def has_low_noise_expert(self) -> bool: + return "low_noise_transformer" in self.model_parts() + def is_sd_v1(self): return self == ModelType.STABLE_DIFFUSION_15 \ or self == ModelType.STABLE_DIFFUSION_15_INPAINTING @@ -182,10 +197,16 @@ def is_flow_matching(self) -> bool: or self.is_hi_dream() \ or self.is_z_image() \ or self.is_ernie() \ - or self.is_ideogram() + or self.is_ideogram() \ + or self.is_ltx_2() def is_video_model(self) -> bool: - return self.is_hunyuan_video() #incase we add more video models in the future + return self.is_hunyuan_video() or self.is_ltx_2() + + def is_audio_model(self) -> bool: + # LTX-2 is audio-visual, but audio training/output is not yet implemented (video-only scope) - + # its audio branch is frozen and never decoded, so it does not belong here yet. + return False def supports_compression(self) -> bool: return not (self.is_stable_diffusion() or self.is_stable_diffusion_xl() or self.is_wuerstchen()) @@ -207,7 +228,7 @@ def supported_training_methods(self) -> tuple[TrainingMethod, ...]: or self.is_chroma(): return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA, TrainingMethod.EMBEDDING) if self.is_qwen() or self.is_z_image() or self.is_flux_2() or self.is_ernie() \ - or self.is_anima() or self.is_krea2() or self.is_ideogram(): + or self.is_anima() or self.is_krea2() or self.is_ideogram() or self.is_ltx_2(): return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA) raise ValueError(f"No supported training methods defined for model type {self}") @@ -219,6 +240,9 @@ def text_encoder_parts(self) -> tuple[str, ...]: # the text encoder components, named "text_encoder"/"text_encoder_2"/... by convention (see below). return tuple(part for part in _MODEL_PARTS[self] if part.startswith("text_encoder")) + def trainable_parts(self) -> tuple[str, ...]: + return _TRAINABLE_PARTS[self] + def supported_lora_formats(self) -> list[ModelFormat]: formats = [ ModelFormat.DIFFUSERS_LORA, @@ -314,6 +338,43 @@ def supported_output_formats(self, training_method: TrainingMethod) -> list[Mode ModelType.Z_IMAGE: ("transformer", "text_encoder", "vae"), ModelType.ERNIE: ("transformer", "text_encoder", "vae"), ModelType.IDEOGRAM_4: ("transformer", "text_encoder", "unconditional_transformer", "vae"), + ModelType.LTX_2: ("transformer", "text_encoder", "connectors", "low_noise_transformer", "vae"), +} + +# subset of _MODEL_PARTS the architecture allows a run to train, for both LoRA and fine-tuning -- the parts each setup +# routes through _setup_model_part_requires_grad. Parts omitted here (VAE everywhere; the text encoder on the newer +# transformer models; Ideogram's unconditional_transformer; Wuerstchen's decoder stack) are architecture-frozen. +_TRAINABLE_PARTS: dict[ModelType, tuple[str, ...]] = { + ModelType.STABLE_DIFFUSION_15: ("unet", "text_encoder"), + ModelType.STABLE_DIFFUSION_15_INPAINTING: ("unet", "text_encoder"), + ModelType.STABLE_DIFFUSION_20: ("unet", "text_encoder"), + ModelType.STABLE_DIFFUSION_20_BASE: ("unet", "text_encoder"), + ModelType.STABLE_DIFFUSION_20_INPAINTING: ("unet", "text_encoder"), + ModelType.STABLE_DIFFUSION_20_DEPTH: ("unet", "text_encoder"), + ModelType.STABLE_DIFFUSION_21: ("unet", "text_encoder"), + ModelType.STABLE_DIFFUSION_21_BASE: ("unet", "text_encoder"), + ModelType.STABLE_DIFFUSION_3: ("transformer", "text_encoder", "text_encoder_2", "text_encoder_3"), + ModelType.STABLE_DIFFUSION_35: ("transformer", "text_encoder", "text_encoder_2", "text_encoder_3"), + ModelType.STABLE_DIFFUSION_XL_10_BASE: ("unet", "text_encoder", "text_encoder_2"), + ModelType.STABLE_DIFFUSION_XL_10_BASE_INPAINTING: ("unet", "text_encoder", "text_encoder_2"), + ModelType.WUERSTCHEN_2: ("prior", "text_encoder"), + ModelType.STABLE_CASCADE_1: ("prior", "text_encoder"), + ModelType.PIXART_ALPHA: ("transformer", "text_encoder"), + ModelType.PIXART_SIGMA: ("transformer", "text_encoder"), + ModelType.FLUX_DEV_1: ("transformer", "text_encoder", "text_encoder_2"), + ModelType.FLUX_FILL_DEV_1: ("transformer", "text_encoder", "text_encoder_2"), + ModelType.FLUX_2: ("transformer",), + ModelType.ANIMA: ("transformer",), + ModelType.SANA: ("transformer", "text_encoder"), + ModelType.HUNYUAN_VIDEO: ("transformer", "text_encoder", "text_encoder_2"), + ModelType.HI_DREAM_FULL: ("transformer", "text_encoder", "text_encoder_2", "text_encoder_3", "text_encoder_4"), + ModelType.CHROMA_1: ("transformer", "text_encoder"), + ModelType.QWEN: ("transformer", "text_encoder"), + ModelType.KREA_2: ("transformer",), + ModelType.Z_IMAGE: ("transformer",), + ModelType.ERNIE: ("transformer",), + ModelType.IDEOGRAM_4: ("transformer",), + ModelType.LTX_2: ("transformer",), } diff --git a/modules/util/enum/SamplingMethod.py b/modules/util/enum/SamplingMethod.py new file mode 100644 index 000000000..4a8ff58f0 --- /dev/null +++ b/modules/util/enum/SamplingMethod.py @@ -0,0 +1,10 @@ +from enum import Enum + + +class SamplingMethod(Enum): + STANDARD = 'STANDARD' # the model's own way of generating a sample, with nothing swapped in + HANDOFF_LOW_NOISE = 'HANDOFF_LOW_NOISE' # the trained transformer denoises down to the expert's first sigma, the expert runs the tail + DISTILLED = 'DISTILLED' # the expert runs its own full schedule from noise; the trained transformer never runs + + def __str__(self): + return self.value diff --git a/modules/util/mm_8bit.py b/modules/util/mm_8bit.py index a4e74f4bf..e76a2155e 100644 --- a/modules/util/mm_8bit.py +++ b/modules/util/mm_8bit.py @@ -1,8 +1,15 @@ +import torch + try: - from modules.util.triton_mm_8bit import mm_8bit + from modules.util.triton_mm_8bit import ( + mm_8bit, + rowcol_scaled_lora_mm_8bit, + rowcol_scaled_mm_8bit, + scaled_lora_mm_8bit, + scaled_mm_8bit, + ) except ImportError as e: print(str(e) + ", continuing without triton") - import torch def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: assert a.shape[1] == b.shape[0], "Incompatible dimensions" assert a.is_contiguous(), "Matrix A must be contiguous" @@ -12,4 +19,17 @@ def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: return torch._int_mm(a, b) else: one = torch.ones(1, device=a.device) - return torch._scaled_mm(a, b.T.contiguous().T, scale_a=one, scale_b=one) + #out_dtype defaults to a's dtype, which would round the accumulator back to fp8 + return torch._scaled_mm(a, b.T.contiguous().T, scale_a=one, scale_b=one, out_dtype=torch.float32) + + def scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).to(out_dtype) + + def scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).add_(xd @ up).to(out_dtype) + + def rowcol_scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).mul_(scale_n.reshape(1, -1)).to(out_dtype) + + def rowcol_scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).mul_(scale_n.reshape(1, -1)).add_(xd @ up).to(out_dtype) diff --git a/modules/util/multi_gpu_util.py b/modules/util/multi_gpu_util.py index 962434b65..545a90ef8 100644 --- a/modules/util/multi_gpu_util.py +++ b/modules/util/multi_gpu_util.py @@ -3,11 +3,10 @@ from modules.util.bf16_stochastic_rounding import copy_stochastic_ from modules.util.commands.TrainCommands import TrainCommands from modules.util.enum.GradientReducePrecision import GradientReducePrecision +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - def is_enabled() -> bool: return torch.distributed.is_available() and torch.distributed.is_initialized() diff --git a/modules/util/optimizer/muon_util.py b/modules/util/optimizer/muon_util.py index 092070d45..0a49b6e9b 100644 --- a/modules/util/optimizer/muon_util.py +++ b/modules/util/optimizer/muon_util.py @@ -57,6 +57,10 @@ def build_muon_adam_key_fn( 'text_fusion.layerwise_blocks', 'text_fusion.refiner_blocks', ] + case ModelType.LTX_2: + default_patterns = [ + 'transformer_blocks', + ] case _: # Unmatched cases raise NotImplementedError(f"Default hidden layer patterns are not defined for model type: {model.model_type}") filters = [ModuleFilter(p, use_regex=False) for p in default_patterns] diff --git a/modules/util/quantization_util.py b/modules/util/quantization_util.py index d7229687d..40164f6c3 100644 --- a/modules/util/quantization_util.py +++ b/modules/util/quantization_util.py @@ -9,6 +9,7 @@ from modules.util.config.TrainConfig import QuantizationConfig, TrainConfig from modules.util.enum.DataType import DataType from modules.util.ModuleFilter import ModuleFilter +from modules.util.tqdm_util import tqdm import torch from torch import Tensor, nn @@ -16,7 +17,6 @@ from diffusers.quantizers.gguf.utils import GGUFLinear, dequantize_gguf_tensor import accelerate -from tqdm import tqdm try: from modules.module.quantized.LinearNf4 import LinearNf4 @@ -27,19 +27,33 @@ LinearNf4 = None def quantize_int8(x: Tensor, scale: float | Tensor) -> Tensor: - q = x.float().mul(1.0 / scale).round_().clamp_(-128.0, 127.0).to(torch.int8) - return q + xf = x.to(torch.float32, copy=True) + return xf.mul_(1.0 / scale).round_().clamp_(-128.0, 127.0).to(torch.int8) -def quantize_int8_tensorwise_get_scale(x: Tensor) -> float: - abs_max = x.abs().max() - scale = (abs_max.float() / 127.0).clamp(min=1e-30) - return scale +def quantize_int8_tensorwise_get_scale(x: Tensor) -> Tensor: + # max|x| == max(max x, -min x): one pass over x, no full-tensor abs() copy + min_val, max_val = torch.aminmax(x) + abs_max = torch.maximum(max_val, min_val.neg()) + return (abs_max.float() / 127.0).clamp(min=1e-30) -def quantize_int8_tensorwise(x: Tensor) -> tuple[Tensor, float]: +def quantize_int8_tensorwise(x: Tensor) -> tuple[Tensor, Tensor]: scale = quantize_int8_tensorwise_get_scale(x) q = quantize_int8(x, scale) return q, scale +# Quantizing a whole weight at once allocates a full-size fp32 intermediary, which spikes VRAM on small +# GPUs. The chunked variants work in row-blocks, bounding that transient to one block. Load-time only, +# so the Python loop costs nothing. +_QUANTIZE_CHUNK_ELEMENTS = 16 * 1024 * 1024 + +def quantize_int8_tensorwise_chunked(x: Tensor) -> tuple[Tensor, Tensor]: + scale = quantize_int8_tensorwise_get_scale(x) + q = torch.empty_like(x, dtype=torch.int8) + rows = max(1, _QUANTIZE_CHUNK_ELEMENTS // x[0].numel()) + for i in range(0, x.shape[0], rows): + q[i:i + rows] = quantize_int8(x[i:i + rows], scale) + return q, scale + def quantize_int8_axiswise_get_scale(x: Tensor, dim: int) -> Tensor: abs_max = x.abs().amax(dim=dim, keepdim=True) scale = (abs_max.float() / 127.0).clamp(min=1e-30) @@ -51,29 +65,46 @@ def quantize_int8_axiswise(x: Tensor, dim: int) -> tuple[Tensor, Tensor]: return q, scale def quantize_fp8(x: Tensor, scale: float | Tensor) -> Tensor: - q = x.float().mul(1.0 / scale).clamp_(-448.0, 448.0).to(torch.float8_e4m3fn) - return q + xf = x.to(torch.float32, copy=True) + return xf.mul_(1.0 / scale).clamp_(-448.0, 448.0).to(torch.float8_e4m3fn) -def quantize_fp8_tensorwise_get_scale(x: Tensor) -> float: - abs_max = x.abs().max() - scale = (abs_max.float() / 448.0).clamp(min=1e-30) - return scale +def quantize_fp8_tensorwise_get_scale(x: Tensor) -> Tensor: + # max|x| == max(max x, -min x): one pass over x, no full-tensor abs() copy + min_val, max_val = torch.aminmax(x) + abs_max = torch.maximum(max_val, min_val.neg()) + return (abs_max.float() / 448.0).clamp(min=1e-30) + +def quantize_fp8_tensorwise(x: Tensor) -> tuple[Tensor, Tensor]: + scale = quantize_fp8_tensorwise_get_scale(x) + q = quantize_fp8(x, scale) + return q, scale + +def quantize_fp8_tensorwise_chunked(x: Tensor) -> tuple[Tensor, Tensor]: + scale = quantize_fp8_tensorwise_get_scale(x) + q = torch.empty_like(x, dtype=torch.float8_e4m3fn) + rows = max(1, _QUANTIZE_CHUNK_ELEMENTS // x[0].numel()) + for i in range(0, x.shape[0], rows): + q[i:i + rows] = quantize_fp8(x[i:i + rows], scale) + return q, scale def quantize_fp8_axiswise_get_scale(x: Tensor, dim: int) -> Tensor: abs_max = x.abs().amax(dim=dim, keepdim=True) scale = (abs_max.float() / 448.0).clamp(min=1e-30) return scale -def quantize_fp8_tensorwise(x: Tensor) -> tuple[Tensor, float]: - scale = quantize_fp8_tensorwise_get_scale(x) - q = quantize_fp8(x, scale) - return q, scale - def quantize_fp8_axiswise(x: Tensor, dim: int) -> tuple[Tensor, Tensor]: scale = quantize_fp8_axiswise_get_scale(x, dim) q = quantize_fp8(x, scale) return q, scale +def quantize_axiswise(x: Tensor, dim: int, dtype: torch.dtype) -> tuple[Tensor, Tensor]: + if dtype == torch.int8: + return quantize_int8_axiswise(x, dim) + elif dtype == torch.float8_e4m3fn: + return quantize_fp8_axiswise(x, dim) + else: + raise NotImplementedError(f"{dtype} is not an 8-bit quantization dtype") + def dequantize(q: Tensor, scale: float | Tensor) -> Tensor: return q.float() * scale @@ -102,6 +133,7 @@ def __replace_linear_layers( construct_fn, keep_in_fp32_modules: list[str] | None = None, filters: list[ModuleFilter] | None = None, + fallback_construct_fn = None, copy_parameters: bool = False, name_prefix: str = "", visited_modules: set[int] | None = None, @@ -121,10 +153,12 @@ def __replace_linear_layers( if isinstance(parent_module, (nn.ModuleList, nn.Sequential, nn.ModuleDict)): for key, module in (parent_module.items() if isinstance(parent_module, nn.ModuleDict) else enumerate(parent_module)): if isinstance(module, convert_type): - if filters is not None and len(filters) > 0 and not any(f.matches(name_prefix) for f in filters): + matches = filters is None or len(filters) == 0 or any(f.matches(name_prefix) for f in filters) + fn = construct_fn if matches else fallback_construct_fn + if fn is None: continue - quant_linear = __create_linear_layer(construct_fn, module, copy_parameters) + quant_linear = __create_linear_layer(fn, module, copy_parameters) parent_module[key] = quant_linear del module elif id(module) not in visited_modules: @@ -133,6 +167,7 @@ def __replace_linear_layers( construct_fn=construct_fn, keep_in_fp32_modules=keep_in_fp32_modules, filters=filters, + fallback_construct_fn=fallback_construct_fn, copy_parameters=copy_parameters, name_prefix=f"{name_prefix}.{key}", visited_modules=visited_modules, @@ -145,10 +180,12 @@ def __replace_linear_layers( module = getattr(parent_module, attr_name) if isinstance(module, convert_type): key_name = attr_name if name_prefix == "" else f"{name_prefix}.{attr_name}" - if filters is not None and len(filters) > 0 and not any(f.matches(key_name) for f in filters): + matches = filters is None or len(filters) == 0 or any(f.matches(key_name) for f in filters) + fn = construct_fn if matches else fallback_construct_fn + if fn is None: continue - quant_linear = __create_linear_layer(construct_fn, module, copy_parameters) + quant_linear = __create_linear_layer(fn, module, copy_parameters) setattr(parent_module, attr_name, quant_linear) del module elif isinstance(module, nn.Module) and id(module) not in visited_modules: @@ -157,46 +194,55 @@ def __replace_linear_layers( construct_fn=construct_fn, keep_in_fp32_modules=keep_in_fp32_modules, filters=filters, + fallback_construct_fn=fallback_construct_fn, copy_parameters=copy_parameters, name_prefix=attr_name if name_prefix == "" else f"{name_prefix}.{attr_name}", visited_modules=visited_modules, ) -def replace_linear_with_quantized_layers( - parent_module: nn.Module, - dtype: DataType, - keep_in_fp32_modules: list[str] | None = None, - quantization: QuantizationConfig | None = None, - copy_parameters: bool = False, -): +def __quantized_linear_class_and_kwargs(dtype: DataType): + # only the W8A8 dtypes have a compressed variant, so this covers every layer that can be built compressed + if dtype.is_compressed() and not nvcomp_util.available(): + raise RuntimeError("a compressed weight data type is selected but nvCOMP is not available") + + # deferred imports: the quantized Linear modules pull in heavy backends, so they stay + # out of module scope and are only imported once a quantized dtype is actually requested from modules.module.quantized.LinearFp8 import LinearFp8 from modules.module.quantized.LinearGGUFA8 import LinearGGUFA8 - from modules.module.quantized.LinearSVD import make_svd_linear from modules.module.quantized.LinearW8A8 import LinearW8A8 - kwargs = {} if dtype.quantize_nf4(): - linear_class = LinearNf4 + return LinearNf4, {} elif dtype.quantize_int8(): - linear_class = bnb.nn.Linear8bitLt - kwargs = {'has_fp16_weights': False} + return bnb.nn.Linear8bitLt, {'has_fp16_weights': False} elif dtype.quantize_fp8(): - linear_class = LinearFp8 + return LinearFp8, {} elif dtype.quantize_intW8A8(): - linear_class = LinearW8A8 - kwargs = {'dtype': torch.int8} + return LinearW8A8, {'dtype': torch.int8, 'compress': dtype.is_compressed()} elif dtype.quantize_fpW8A8(): - linear_class=LinearW8A8 - kwargs = {'dtype': torch.float8_e4m3fn} + return LinearW8A8, {'dtype': torch.float8_e4m3fn, 'compress': dtype.is_compressed()} elif dtype == DataType.GGUF_A8_INT: - linear_class=LinearGGUFA8 - kwargs = {'dtype': torch.int8} + return LinearGGUFA8, {'dtype': torch.int8} elif dtype == DataType.GGUF_A8_FLOAT: - linear_class=LinearGGUFA8 - kwargs = {'dtype': torch.float8_e4m3fn} + return LinearGGUFA8, {'dtype': torch.float8_e4m3fn} else: + return None, {} + +def replace_linear_with_quantized_layers( + parent_module: nn.Module, + dtype: DataType, + keep_in_fp32_modules: list[str] | None = None, + quantization: QuantizationConfig | None = None, + copy_parameters: bool = False, +): + from modules.module.quantized.LinearGGUFA8 import LinearGGUFA8 + from modules.module.quantized.LinearSVD import make_svd_linear + + linear_class, kwargs = __quantized_linear_class_and_kwargs(dtype) + if linear_class is None: return + fallback_construct_fn = None if quantization is not None: if quantization.svd_dtype != DataType.NONE: if dtype.is_gguf(): @@ -210,6 +256,11 @@ def replace_linear_with_quantized_layers( ModuleFilter(pattern, use_regex=quantization.layer_filter_regex) for pattern in quantization.layer_filter.split(",") ] + + if not dtype.is_gguf() and quantization.fallback_dtype.is_quantized(): + fallback_linear_class, fallback_kwargs = __quantized_linear_class_and_kwargs(quantization.fallback_dtype) + if fallback_linear_class is not None: + fallback_construct_fn = partial(fallback_linear_class, **fallback_kwargs) else: quant_filters = None @@ -219,6 +270,7 @@ def replace_linear_with_quantized_layers( construct_fn=partial(linear_class, **kwargs), keep_in_fp32_modules=keep_in_fp32_modules, filters=quant_filters, + fallback_construct_fn=fallback_construct_fn, copy_parameters=copy_parameters, convert_type=convert_type, ) @@ -262,29 +314,39 @@ def is_quantized_parameter( return False -def quantize_layers(module: nn.Module, device: torch.device, train_dtype: DataType, config: TrainConfig, compress: bool = False): +def is_quantized_module(module: nn.Module) -> bool: + return any(is_quantized_parameter(module, name) + for name, _ in module.named_parameters(recurse=False)) + + +def quantize_layers(module: nn.Module, device: torch.device, train_dtype: DataType, config: TrainConfig): if module is None: return child_modules = list(module.modules()) - compressible = [m for m in child_modules if isinstance(m, CompressedWeightMixin)] - if compress and not nvcomp_util.available(): - raise RuntimeError("a compressed weight data type is selected but nvCOMP is not available") - for m in compressible: - m.compress = compress - - for _ in multi.master_first(): #avoid cache writing conflicts - for child_module in tqdm(child_modules, desc="Quantizing model weights", total=len(child_modules), delay=5, smoothing=0.1): - if isinstance(child_module, (QuantizedModuleMixin, GGUFLinear)): - child_module.compute_dtype = train_dtype.torch_dtype() - if isinstance(child_module, QuantizedModuleMixin): - child_module.quantize(device=device) - - if multi.is_master() and compress: - uncompressed = sum(m.uncompressed_bytes() for m in compressible) - compressed = sum(m.weight.nbytes for m in compressible) - if uncompressed > 0: - tqdm.write(f"nvCOMP weight compression ({type(module).__name__}): {uncompressed / 2**20:.0f} -> {compressed / 2**20:.0f} MiB ({(1 - compressed / uncompressed) * 100:.0f}% saved)") + for child_module in tqdm(child_modules, desc="Quantizing model weights", total=len(child_modules), delay=5, smoothing=0.1): + if isinstance(child_module, (QuantizedModuleMixin, GGUFLinear)): + child_module.compute_dtype = train_dtype.torch_dtype() + if isinstance(child_module, QuantizedModuleMixin): + child_module.quantize(device=device) + + report_compression(module) + + +def report_compression(module: nn.Module): + # one line per component, summed over its compressed layers. Reads the measured lengths, which outlive the blob, + # so the streamed path -- where the layers are back on meta by now -- reports the same numbers as the resident one. + # A streamed component is cold-materialized once per eviction cycle, so the line is emitted on the first one only. + if not multi.is_master() or getattr(module, "_compression_reported", False): + return + compressible = [m for m in module.modules() + if isinstance(m, CompressedWeightMixin) and m.compressed_bytes() is not None] + if not compressible: + return + module._compression_reported = True + uncompressed = sum(m.uncompressed_bytes() for m in compressible) + compressed = sum(m.compressed_bytes() for m in compressible) + tqdm.write(f"nvCOMP weight compression ({type(module).__name__}): {uncompressed / 2**20:.0f} -> {compressed / 2**20:.0f} MiB ({(1 - compressed / uncompressed) * 100:.0f}% saved)") def get_unquantized_weight(module: nn.Linear, dtype: torch.dtype, device: torch.device) -> Tensor: assert isinstance(module, nn.Linear) @@ -320,6 +382,9 @@ def get_offload_tensors(module: nn.Module) -> list[torch.Tensor]: def get_offload_tensor_bytes(module: nn.Module) -> int: + if isinstance(module, QuantizedLinearMixin) and module.weight.is_meta: + return module.predict_offload_bytes() + tensors = get_offload_tensors(module) return sum(t.element_size() * t.numel() for t in tensors) @@ -329,15 +394,13 @@ def offload_quantized( module: nn.Module, device: torch.device, non_blocking: bool = False, - allocator: Callable[[torch.tensor], torch.tensor] | None = None, + place: Callable[[torch.Tensor, bool], torch.Tensor] | None = None, ): tensors = get_offload_tensors(module) - if allocator is None: + if place is None: for tensor in tensors: tensor.data = tensor.data.to(device=device, non_blocking=non_blocking) else: for tensor in tensors: - new_tensor = allocator(tensor) - new_tensor.copy_(tensor.data, non_blocking=non_blocking) - tensor.data = new_tensor + tensor.data = place(tensor, non_blocking) diff --git a/modules/util/staged_pipeline.py b/modules/util/staged_pipeline.py new file mode 100644 index 000000000..073be3c08 --- /dev/null +++ b/modules/util/staged_pipeline.py @@ -0,0 +1,62 @@ +import inspect +from collections.abc import Callable + +from modules.util.tqdm_util import tqdm + + +def _accepted(process: Callable, available: dict) -> dict: + # Pass a stage only the arguments its signature names (or everything, if it + # takes **kwargs). This keeps each stage's parameter list limited to what it + # actually consumes, even though the context carries more. + params = inspect.signature(process).parameters + if any(p.kind == p.VAR_KEYWORD for p in params.values()): + return available + return {name: value for name, value in available.items() if name in params} + + +def run_staged_pipeline( + stages: list[tuple[str, Callable]], + inputs: dict[str, list], + shared: dict | None = None, +) -> list: + # Run every item through stage 0, then every item through stage 1, and so on + # (stage-major order), instead of running each item through all stages before + # starting the next (item-major order). Each stage is a (label, process) pair; + # `label` names the stage on its progress bar, `process` is a callable run + # once per item. Stage-major order is what bounds how often a model part has to + # be brought on-device. A sampler stage materializes the part it needs and evicts + # the rest, so item-major would cycle text encoder -> transformer -> vae once per + # item, moving all three onto the train device every time; stage-major moves each + # one there once and keeps it for the stage's whole run. What that saves scales + # with how an evicted part is held - a host-to-device copy of weights already in + # RAM, or a re-read of the weights from disk once parts are streamed. + # + # `inputs` is column-oriented: each key maps to the list of that argument's + # per-item values (all lists the same length), transposed here into one context + # dict per item. Each context then accumulates every non-final stage's output, + # so a value produced (or supplied) early is available to any later stage + # without being threaded through the ones in between. The final stage's return + # value is collected into a separate result list - one entry per item - and that + # list is what this function returns, so a pipeline outputs whatever its last + # stage returns. `shared` holds batch-level arguments (e.g. a progress reporter) + # offered to every stage that names them. + shared = shared or {} + names = list(inputs) + count = len(inputs[names[0]]) if names else 0 + contexts = [{name: inputs[name][i] for name in names} for i in range(count)] + results = [None] * count + last_index = len(stages) - 1 + for index, (label, process) in enumerate(stages): + # One bar per stage, counting off the batch's items, so a stage that shows no + # inner progress of its own (text encoding, VAE decoding) is still visible while + # it runs. leave=False keeps the finished stages from piling up in the log under + # the training bars. A stage with its own inner bar (denoising counts diffusion + # steps) nests below this one; tqdm assigns the positions itself. + for i, context in enumerate(tqdm(contexts, desc=label, leave=False)): + output = process(**_accepted(process, {**shared, **context})) + if index == last_index: + results[i] = output + else: + context.update(output) + # with no stages there is nothing to produce; hand back the raw contexts + return results if stages else contexts diff --git a/modules/util/torch_util.py b/modules/util/torch_util.py index 408100bf9..8130074a4 100644 --- a/modules/util/torch_util.py +++ b/modules/util/torch_util.py @@ -1,7 +1,9 @@ +import contextlib import gc +import platform +import time from collections.abc import Callable from contextlib import nullcontext -from typing import Any import torch @@ -14,6 +16,36 @@ torch_version = packaging.version.parse(torch.__version__) +@contextlib.contextmanager +def timed(label: str, enabled: bool = True): + # wall-clock timing around a block; sync the compute device before and after so the measurement includes the + # async device transfer + (re)quantization rather than just the launch overhead. Forces a cuda sync per block, + # so enable only for ad-hoc profiling, not on the hot per-step path. + if not enabled: + yield + return + if torch.cuda.is_available(): + torch.cuda.synchronize() + start = time.perf_counter() + yield + if torch.cuda.is_available(): + torch.cuda.synchronize() + print(f"[timing] {label}: {time.perf_counter() - start:.3f}s") + + +def supports_mem_pool(device: torch.device) -> bool: + return device.type == "cuda" + + +def create_mem_pool(device: torch.device): + # a dedicated MemPool the caller can allocate into; None on devices without MemPool support (cpu/mps) + return torch.cuda.MemPool() if supports_mem_pool(device) else None + + +def mem_pool_context(mem_pool): + # route allocations made in this context into the given MemPool; no-op when it is None + return torch.cuda.use_mem_pool(mem_pool) if mem_pool is not None else nullcontext() + def state_dict_has_prefix(state_dict: dict | None, prefix: str): if not state_dict: @@ -42,73 +74,6 @@ def get_tensor_data( return tensors -def has_grad_fn( - data: torch.Tensor | list | tuple | dict, - include_parameter_indices: list[int] | None = None, -) -> bool: - if isinstance(data, torch.Tensor) and include_parameter_indices is None: - return data.grad_fn is not None - elif isinstance(data, list | tuple): - for i, elem in enumerate(data): - if include_parameter_indices is None or i in include_parameter_indices: - if has_grad_fn(elem): - return True - elif isinstance(data, dict) and include_parameter_indices is None: - for elem in data.values(): - if has_grad_fn(elem): - return True - - return False - -def add_dummy_grad_fn_( - data: torch.Tensor | list | tuple | dict, -) -> Any: - if isinstance(data, torch.Tensor): - if data.grad_fn is not None: - return data - grad_tensor = torch\ - .zeros(size=(0, *data.shape[1:]), requires_grad=True, device=data.device, dtype=data.dtype) - return torch.cat([data, grad_tensor], dim=0) - if isinstance(data, list): - for i, elem in enumerate(data): - if isinstance(elem, torch.Tensor): - if elem.grad_fn is not None: - return data - grad_tensor = torch\ - .zeros(size=(0, *elem.shape[1:]), requires_grad=True, device=elem.device, dtype=elem.dtype) - data[i] = torch.cat([elem, grad_tensor], dim=0) - return data - else: - data[i] = add_dummy_grad_fn_(elem) - if isinstance(data, tuple): - for i, elem in enumerate(data): - if isinstance(elem, torch.Tensor): - if elem.grad_fn is not None: - return data - grad_tensor = torch\ - .zeros(size=(0, *elem.shape[1:]), requires_grad=True, device=elem.device, dtype=elem.dtype) - data = list(data) - data[i] = torch.cat([elem, grad_tensor], dim=0) - data = tuple(data) - return data - else: - data = list(data) - data[i] = add_dummy_grad_fn_(elem) - data = tuple(data) - elif isinstance(data, dict): - for key, elem in data.items(): - if isinstance(elem, torch.Tensor): - if elem.grad_fn is not None: - return data - grad_tensor = torch \ - .zeros(size=(0, *elem.shape[1:]), requires_grad=True, device=elem.device, dtype=elem.dtype) - data[key] = torch.cat([elem, grad_tensor], dim=0) - return data - else: - data[key] = add_dummy_grad_fn_(elem) - - return data - def tensors_to_device_( data: torch.Tensor | list | tuple | dict, device: torch.device, @@ -251,14 +216,24 @@ def pin_tensor_(x): # not implemented for other device types if torch.cuda.is_available(): cudart = torch.cuda.cudart() + num_bytes = x.numel() * x.element_size() err = cudart.cudaHostRegister( x.data_ptr(), - x.numel() * x.element_size(), + num_bytes, 0, ) if err.value != 0: - raise RuntimeError(f"CUDA Error while trying to pin memory. error: {err.value}, ptr: {x.data_ptr()}, size: {x.numel() * x.element_size()}") + hint = "" + if err.value == 1 and num_bytes >= 2**31: + # the kernel's page list for a registration holds one entry per 4 KiB, so at 2 GiB it reaches + # 4 MiB, the largest single kmalloc there is, and the pin is refused whatever the driver or + # GPU. cudaErrorInvalidValue at this size has no other cause, so the attribution is safe. + hint = (f". A single pinned allocation of {num_bytes / 2**30:.2f} GiB failed: linux " + f"{platform.release()} cannot pin 2 GiB or more in one call, a kernel bug present in " + f"6.11 and 6.12 and fixed in 6.13. Update the kernel, or run on a host with 6.13 or " + f"newer. This attempt leaked {num_bytes / 2**30:.2f} GiB of host memory until reboot") + raise RuntimeError(f"CUDA Error while trying to pin memory. error: {err.value}, ptr: {x.data_ptr()}, size: {num_bytes}{hint}") def unpin_tensor_(x): diff --git a/modules/util/tqdm_util.py b/modules/util/tqdm_util.py new file mode 100644 index 000000000..838e38fee --- /dev/null +++ b/modules/util/tqdm_util.py @@ -0,0 +1,49 @@ +from tqdm import tqdm as _tqdm + +#progress bars created through the tqdm below, innermost last +_bars = [] + + +class tqdm(_tqdm): + _status = None + + @classmethod + def get_lock(cls): + #tqdm caches the terminal write lock on the class that first asks for it, so a subclass + #would get one of its own and stop serializing against bars drawn by tqdm itself. + return _tqdm.get_lock() + + @classmethod + def show_status(cls, message: str): + #status of long-running work - a compile, an autotune sweep - goes into the innermost bar's + #postfix rather than on a line of its own, and stands until the next postfix write replaces it. + bar = next((bar for bar in reversed(_bars) if not bar.disable), None) + if bar is None: + cls.write(message) + else: + bar.set_postfix_str(message) + bar._status = message + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + _bars.append(self) + + def _clear_status(self, refresh=True): + #anything written over the status - the training loop's loss - is left standing. + if self._status is not None and self.postfix == self._status: + self.set_postfix_str("", refresh=refresh) + self._status = None + + def update(self, n=1): + #the status describes work that was running while the bar stood still, so the step that + #follows it is the point where it stops being current. + self._clear_status(refresh=False) + return super().update(n) + + def close(self): + self._clear_status() + super().close() + #compared by identity: tqdm's __eq__ is by screen position, so a bar closed late by __del__ + #would drop whichever live bar has taken over its line. + global _bars + _bars = [bar for bar in _bars if bar is not self] diff --git a/modules/util/triton_mm_8bit.py b/modules/util/triton_mm_8bit.py index 754d9ceda..53fd0682d 100644 --- a/modules/util/triton_mm_8bit.py +++ b/modules/util/triton_mm_8bit.py @@ -1,15 +1,13 @@ -#This is a 8bit matmul kernel adapted from the Triton tutorial here: +#8bit matmul kernels adapted from the Triton tutorial here: #https://triton-lang.org/main/getting-started/tutorials/03-matrix-multiplication.html -#It is not optimized and about 10% slower than torch._int_mm and torch._scaled_mm -#However, the torch functions don't work on row-major rhs matrices: -#_scaled_mm fails, _int_mm automatically converts to column-major -# -#Converting to column-major is slow, which is significant because the weights matrix -#of a Linear layer is always column-major during the backward pass. -# -#In these cases, this Triton kernel is *much* faster because it can access the -#row-major weight matrix directly, using strided memory access +#All three mm entry points share the _mm_accumulate compute core (grouped launch order +#for L2 reuse, compile-time layout/divisibility specialization, EVEN_K loop selection) +#and differ only in the epilogue: _mm_kernel stores the raw int32/fp32 product, +#_scaled_mm_kernel folds a per-row dequant scale in and casts to the output dtype, and +#_scaled_lora_mm_kernel additionally fuses a low-rank update into the same tile. + +from modules.util.tqdm_util import tqdm import torch @@ -25,59 +23,150 @@ def _prefer_mxfp8(device: torch.device) -> bool: return torch.version.cuda is not None and torch.cuda.get_device_capability(device)[0] >= 12 -@triton.autotune( - configs=[ - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 128}, num_stages=3,num_warps=8), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64}, num_stages=3,num_warps=8), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 64}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 64}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128}, num_stages=3,num_warps=8), - triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 128}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=5,num_warps=2), - triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 128}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=5,num_warps=2), - - - ], - key=[ - 'QUANTIZED_M', #only tune roughly on M, because M is the transformer sequence length - can vary on data - 'N', - 'K', - 'stride_bk' #use stride of b as key, to autotune again for a strided rhs matrix (backward pass) - ], -) +def announce_autotuning(kernel, name=None): + prefix = f"autotuning {name} " if name else "autotuning " + orig_check_disk_cache = kernel.check_disk_cache + variants = 0 + def check_disk_cache(tuning_key, configs, bench_fn): + def announced_bench(): + nonlocal variants + variants += 1 + tqdm.show_status(f"{prefix}variant #{variants}...") + bench_fn() + return orig_check_disk_cache(tuning_key, configs, announced_bench) + kernel.check_disk_cache = check_disk_cache + +#tiled 8-bit transpose, used to rewrite the backward's B matrix to k-major before the mm. +#The int8 tensor-core op needs its B operand k-major and Ada has no 8-bit ldmatrix.trans, so an +#n-major B - a Linear weight in the backward pass - makes the mm emulate the transpose with byte +#shuffles in its inner loop, ~30% slower on every shape. This copy runs at memory bandwidth +#(~570GB/s) once per weight instead, and the mm then takes the fast k-major path. +_TRANSPOSE_AUTOTUNE_CONFIGS = [ + triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 128}, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64}, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128}, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128}, num_warps=8), +] + +@triton.autotune(configs=_TRANSPOSE_AUTOTUNE_CONFIGS, key=['M', 'N'], cache_results=True) @triton.jit -def _mm_kernel( - a_ptr, b_ptr, c_ptr, +def _transpose_kernel( + src_ptr, dst_ptr, + M, N, + stride_sm, stride_dn, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + tile = tl.load(src_ptr + offs_m[:, None] * stride_sm + offs_n[None, :], + mask=(offs_m[:, None] < M) & (offs_n[None, :] < N)) + tl.store(dst_ptr + offs_n[:, None] * stride_dn + offs_m[None, :], + tl.trans(tile), + mask=(offs_n[:, None] < N) & (offs_m[None, :] < M)) + +def transpose_8bit(src: torch.Tensor) -> torch.Tensor: + #returns src^T as a new contiguous tensor (any 1-byte dtype) + assert src.stride(1) == 1, "src must be contiguous along axis 1" + M, N = src.shape + dst = torch.empty((N, M), device=src.device, dtype=src.dtype) + + def grid(META): + return (triton.cdiv(M, META['BLOCK_SIZE_M']), triton.cdiv(N, META['BLOCK_SIZE_N'])) + _transpose_kernel[grid]( + src, dst, + M, N, + src.stride(0), dst.stride(0), + ) + return dst + +announce_autotuning(_transpose_kernel, name="8-bit transpose") + +#minimum M for the transpose-to-k-major rewrite in the mm wrappers below: the mm saves ~30% +#(~205 -> ~275 TOPS) but the copy costs 2*K*N bytes of traffic regardless of M, so the rewrite +#only pays above a token count. Breakeven is shape-dependent - measured on Ada (4070 Ti SUPER) +#at ~1550 for the widest layers and below 512 for the attention projections - and the wrappers +#see only M, so this takes the widest layer's breakeven and every shape wins above it +_TRANSPOSE_MIN_M = 1536 + + +_AUTOTUNE_KEY = [ + #M is batch*sequence, so unlike N and K it is data-dependent and unbounded: it moves with + #resolution, frame count, batch size and (on models that prune prompt padding) the longest + #caption in the batch. bucketing it per doubling keeps the number of tuning keys logarithmic + #in M instead of linear. proportional resolution is the right shape because the winning config + #is decided by block count against SM count, which is linear in M - so equal ratios of M + #matter equally at every scale, and a fixed stride is too fine at large M and too coarse at small + 'QUANTIZED_M', + 'N', + 'K', + 'stride_bk' #use stride of b as key, to autotune again for a strided rhs matrix (backward pass) +] + +#the LoRA epilogue keys on the rank tiling as well. up's row stride is deliberately not a key even +#though it varies: it is 1 in the forward (up is a transposed view of lora_up) and N in the backward +#(up is lora_down, stored row-major), so keying on it would double the key count on every model. +#the two layouts load differently but the slab is only BLOCK_R x BLOCK_N and is reused by every +#M block in the group, too little traffic next to A and B to move the tile choice +_LORA_AUTOTUNE_KEY = [*_AUTOTUNE_KEY, 'R_TILES', 'BLOCK_R'] + +#configs for the shared _mm_accumulate core: GROUP_SIZE_M is required by its grouped launch +#order. Shared memory per config is stages*(BLOCK_M+BLOCK_N)*BLOCK_K bytes and must stay +#under the ~99KB per-CTA limit (identical on sm86/sm89/sm120, i.e. consumer/workstation +#Ampere through Blackwell); oversized configs would be skipped by the autotuner. +#GROUP_SIZE_M=8 suits the large L2 of Ada/Blackwell; the 16-variants keep fewer B +#panels in flight per wave, for the small L2 (4-6MB) of Ampere - autotuning picks +_AUTOTUNE_CONFIGS = [ + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 16}, num_stages=4, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 16}, num_stages=4, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=3, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 16}, num_stages=3, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=3, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=5, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=5, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=5, num_warps=2), +] + + +#shared compute core of the 8-bit mm kernels: grouped launch order, divisibility hints and +#the EVEN_K-specialized main loop. Returns the raw accumulator plus the output tile offsets; +#each entry kernel below adds its own epilogue and stores. +@triton.jit +def _mm_accumulate( + a_ptr, b_ptr, M, N, K, - stride_am, stride_ak, - stride_bk, stride_bn, - stride_cm, stride_cn, - BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, - QUANTIZED_M, - FLOAT: tl.constexpr, - MXFP8_MMA: tl.constexpr, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, ): - pid_n = tl.program_id(axis=0) - pid_m = tl.program_id(axis=1) + #grouped launch order: consecutive pids walk down GROUP_SIZE_M M-blocks before advancing + #to the next N-block, so the concurrent wave covers a rectangle of blocks and each B panel + #is read from DRAM once and reused by GROUP_SIZE_M CTAs out of L2. A naive row-major grid + #re-reads all of B per M-row instead, which makes the mm DRAM-bound + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m tl.assume(pid_m >= 0) tl.assume(pid_n >= 0) tl.assume(stride_am > 0) - tl.assume(stride_ak > 0) tl.assume(stride_bn > 0) tl.assume(stride_bk > 0) - tl.assume(stride_cm > 0) - tl.assume(stride_cn > 0) offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N @@ -94,27 +183,47 @@ def _mm_kernel( a_scale = tl.full((BLOCK_SIZE_M, BLOCK_SIZE_K // 32), 127, dtype=tl.uint8) b_scale = tl.full((BLOCK_SIZE_N, BLOCK_SIZE_K // 32), 127, dtype=tl.uint8) - for k in range(tl.cdiv(K, BLOCK_SIZE_K)): - a_mask = (offs_am[:, None] < M) & (offs_k[None, :] < K - k*BLOCK_SIZE_K) - b_mask = (offs_bn[None, :] < N) & (offs_k[:, None] < K - k*BLOCK_SIZE_K) - a = tl.load(a_ptrs, mask=a_mask, other=0.0) - b = tl.load(b_ptrs, mask=b_mask, other=0.0) + if EVEN_K: + for _k in range(K // BLOCK_SIZE_K): + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + + if MXFP8_MMA: + accumulator = tl.dot_scaled(a, a_scale, "e4m3", b, b_scale, "e4m3", acc=accumulator) + else: + accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) + + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + else: + for k in range(tl.cdiv(K, BLOCK_SIZE_K)): + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k*BLOCK_SIZE_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k*BLOCK_SIZE_K, other=0.0) - if MXFP8_MMA: - accumulator = tl.dot_scaled(a, a_scale, "e4m3", b, b_scale, "e4m3", acc=accumulator) - else: - accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) + if MXFP8_MMA: + accumulator = tl.dot_scaled(a, a_scale, "e4m3", b, b_scale, "e4m3", acc=accumulator) + else: + accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) - a_ptrs += BLOCK_SIZE_K * stride_ak - b_ptrs += BLOCK_SIZE_K * stride_bk + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) - c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + return accumulator, offs_cm, offs_cn + + +#shared epilogue tail for every mm kernel: c is allocated row-major in _prepare_mm, so only +#stride_cm is passed and the n stride is 1 +@triton.jit +def _store_c(c_ptr, value, offs_cm, offs_cn, M, N, stride_cm): + tl.assume(stride_cm > 0) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + offs_cn[None, :] c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) - tl.store(c_ptrs, accumulator, mask=c_mask) + tl.store(c_ptrs, value, mask=c_mask) -def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + +def _prepare_mm(a: torch.Tensor, b: torch.Tensor, out_dtype: torch.dtype): assert a.shape[1] == b.shape[0], "Incompatible dimensions" assert a.is_contiguous(), "Matrix A must be contiguous" assert a.dtype == b.dtype, "Incompatible dtypes" @@ -124,18 +233,294 @@ def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: M, K = a.shape K, N = b.shape - c = torch.empty((M, N), device=a.device, dtype=torch.float32 if FLOAT else torch.int32) + #the kernel handles exactly two B layouts: k-major (forward, weight.T) and n-major (backward, weight) + B_K_MAJOR = (b.stride(0) == 1) + assert B_K_MAJOR or b.stride(1) == 1, "Matrix B must be contiguous along one axis" + #n-major B runs the mm ~30% slower; for large M a transpose copy pays for itself (see transpose_8bit) + if not B_K_MAJOR and M >= _TRANSPOSE_MIN_M: + b = transpose_8bit(b).t() + B_K_MAJOR = True + c = torch.empty((M, N), device=a.device, dtype=out_dtype) + return b, c, M, N, K, FLOAT + + +def _prepare_scaled_mm(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype): + #_prepare_mm plus the per-row scale reshape (the scaled kernels fold it into the epilogue) + b, c, M, N, K, FLOAT = _prepare_mm(a, b, out_dtype) + scale = scale.reshape(-1).to(torch.float32).contiguous() + assert scale.shape[0] == M, "scale must have one entry per row of a" + return b, c, scale, M, N, K, FLOAT + + +def _prepare_rowcol_scaled_mm(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype): + b, c, scale, M, N, K, FLOAT = _prepare_scaled_mm(a, b, scale, out_dtype) + scale_n = scale_n.reshape(-1).to(torch.float32).contiguous() + assert scale_n.shape[0] == N, "scale_n must have one entry per column of b" + return b, c, scale, scale_n, M, N, K, FLOAT + + +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_AUTOTUNE_KEY, cache_results=True) +@triton.jit +def _mm_kernel( + a_ptr, b_ptr, c_ptr, + M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, +): + accumulator, offs_cm, offs_cn = _mm_accumulate( + a_ptr, b_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, + FLOAT, EVEN_K, MXFP8_MMA, + ) + + _store_c(c_ptr, accumulator, offs_cm, offs_cn, M, N, stride_cm) + +announce_autotuning(_mm_kernel, name="8-bit matmul") + +#Opaque custom ops work around pytorch#164124: torch.compile otherwise absorbs the traced +#@triton.autotune kernels and freezes the config benchmarked for the first shape. Kept opaque, +#these bodies run eagerly, so Triton's autotuner selects per key and the JIT sees real sizes +@torch.library.custom_op("ot_quant::mm_8bit", mutates_args=()) +def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + out_dtype = torch.float32 if a.dtype == torch.float8_e4m3fn else torch.int32 + b, c, M, N, K, FLOAT = _prepare_mm(a, b, out_dtype) + + #1D grid: the kernel derives pid_m/pid_n itself in grouped order for L2 reuse def grid(META): - return (triton.cdiv(N, META['BLOCK_SIZE_N']) , triton.cdiv(M, META['BLOCK_SIZE_M']), ) + return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) _mm_kernel[grid]( a, b, c, M, N, K, - a.stride(0), a.stride(1), - b.stride(0), b.stride(1), - c.stride(0), c.stride(1), - QUANTIZED_M = M // 64, - FLOAT = FLOAT, - MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), ) return c + +@mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + out_dtype = torch.float32 if a.dtype == torch.float8_e4m3fn else torch.int32 + return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) + + +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_AUTOTUNE_KEY, cache_results=True) +@triton.jit +def _scaled_mm_kernel( + a_ptr, b_ptr, c_ptr, scale_ptr, + M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, +): + accumulator, offs_cm, offs_cn = _mm_accumulate( + a_ptr, b_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, + FLOAT, EVEN_K, MXFP8_MMA, + ) + + #per-row scale on axis 0 (M), fold into the epilogue and cast to the output (compute) dtype directly + scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) + result = accumulator.to(tl.float32) * scale[:, None] + result = result.to(c_ptr.dtype.element_ty) + + _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) + +announce_autotuning(_scaled_mm_kernel, name="8-bit scaled matmul") + +@torch.library.custom_op("ot_quant::scaled_mm_8bit", mutates_args=()) +def scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + b, c, scale, M, N, K, FLOAT = _prepare_scaled_mm(a, b, scale, out_dtype) + + def grid(META): + return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) + _scaled_mm_kernel[grid]( + a, b, c, scale, + M, N, K, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), + ) + return c + +@scaled_mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) + + +#_scaled_mm_kernel's epilogue with a second dequant scale on axis 1 (N), for callers whose weight +#is quantized axiswise rather than tensorwise (LinearGGUFA8), where the weight scale is one entry +#per output column and so cannot be folded into the per-row scale +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_AUTOTUNE_KEY, cache_results=True) +@triton.jit +def _rowcol_scaled_mm_kernel( + a_ptr, b_ptr, c_ptr, scale_ptr, scale_n_ptr, + M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, +): + accumulator, offs_cm, offs_cn = _mm_accumulate( + a_ptr, b_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, + FLOAT, EVEN_K, MXFP8_MMA, + ) + + scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) + scale_n = tl.load(scale_n_ptr + offs_cn, mask=offs_cn < N, other=0.0) + result = accumulator.to(tl.float32) * scale[:, None] * scale_n[None, :] + result = result.to(c_ptr.dtype.element_ty) + + _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) + +announce_autotuning(_rowcol_scaled_mm_kernel, name="8-bit row/column scaled matmul") + +@torch.library.custom_op("ot_quant::rowcol_scaled_mm_8bit", mutates_args=()) +def rowcol_scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + b, c, scale, scale_n, M, N, K, FLOAT = _prepare_rowcol_scaled_mm(a, b, scale, scale_n, out_dtype) + + def grid(META): + return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) + _rowcol_scaled_mm_kernel[grid]( + a, b, c, scale, scale_n, + M, N, K, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), + ) + return c + +@rowcol_scaled_mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) + + +#rank-tile width is min(next_pow2(R), cap): the cap bounds the staged tile, and with it the +#epilogue's shared memory, so occupancy stays rank-independent. 64 keeps +#2*BLOCK_R*(BLOCK_M+BLOCK_N) inside the main mm's pool on Ada; 128 overflowed it (backward OOM) +_LORA_BLOCK_R_CAP = 64 + + +def _prepare_lora(a: torch.Tensor, b: torch.Tensor, xd: torch.Tensor, up: torch.Tensor): + R = xd.shape[1] + assert xd.shape[0] == a.shape[0] and up.shape[0] == R and up.shape[1] == b.shape[1], "Incompatible low-rank dimensions" + assert xd.stride(1) == 1, "xd must be contiguous along the rank axis" + assert xd.dtype == up.dtype + return R, min(max(16, triton.next_power_of_2(R)), _LORA_BLOCK_R_CAP) + +@triton.jit +def _add_lora( + result, xd_ptr, up_ptr, offs_cm, offs_cn, M, N, R, + stride_xdm, stride_upr, stride_upn, + BLOCK_R: tl.constexpr, R_TILES: tl.constexpr, +): + for r0 in tl.static_range(R_TILES): + offs_r = r0 * BLOCK_R + tl.arange(0, BLOCK_R) + xd_ptrs = xd_ptr + offs_cm[:, None] * stride_xdm + offs_r[None, :] + up_ptrs = up_ptr + offs_r[:, None] * stride_upr + offs_cn[None, :] * stride_upn + xd = tl.load(xd_ptrs, mask=(offs_cm[:, None] < M) & (offs_r[None, :] < R), other=0.0) + up = tl.load(up_ptrs, mask=(offs_r[:, None] < R) & (offs_cn[None, :] < N), other=0.0) + result += tl.dot(xd, up) + return result + + +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_LORA_AUTOTUNE_KEY, cache_results=True) +@triton.jit +def _scaled_lora_mm_kernel( + a_ptr, b_ptr, c_ptr, scale_ptr, xd_ptr, up_ptr, + M, N, K, R, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_xdm, stride_upr, stride_upn, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, + BLOCK_R: tl.constexpr, R_TILES: tl.constexpr, +): + accumulator, offs_cm, offs_cn = _mm_accumulate( + a_ptr, b_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, + FLOAT, EVEN_K, MXFP8_MMA, + ) + + scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) + result = accumulator.to(tl.float32) * scale[:, None] + + result = _add_lora(result, xd_ptr, up_ptr, offs_cm, offs_cn, M, N, R, + stride_xdm, stride_upr, stride_upn, BLOCK_R, R_TILES) + result = result.to(c_ptr.dtype.element_ty) + + _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) + +announce_autotuning(_scaled_lora_mm_kernel, name="8-bit scaled LoRA matmul") + +@torch.library.custom_op("ot_quant::scaled_lora_mm_8bit", mutates_args=()) +def scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + #returns (a @ b) * scale + xd @ up in out_dtype (rank-tiled epilogue). xd is (M, r), up is + #(r, N) (the lora_up weight transposed, alpha folded in) + R, block_r = _prepare_lora(a, b, xd, up) + b, c, scale, M, N, K, FLOAT = _prepare_scaled_mm(a, b, scale, out_dtype) + + def grid(META): + return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) + _scaled_lora_mm_kernel[grid]( + a, b, c, scale, xd, up, + M, N, K, R, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), xd.stride(0), up.stride(0), up.stride(1), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), + BLOCK_R = block_r, R_TILES = triton.cdiv(R, block_r), + ) + return c + +@scaled_lora_mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) + + +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_LORA_AUTOTUNE_KEY, cache_results=True) +@triton.jit +def _rowcol_scaled_lora_mm_kernel( + a_ptr, b_ptr, c_ptr, scale_ptr, scale_n_ptr, xd_ptr, up_ptr, + M, N, K, R, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_xdm, stride_upr, stride_upn, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, + BLOCK_R: tl.constexpr, R_TILES: tl.constexpr, +): + accumulator, offs_cm, offs_cn = _mm_accumulate( + a_ptr, b_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, + FLOAT, EVEN_K, MXFP8_MMA, + ) + + scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) + scale_n = tl.load(scale_n_ptr + offs_cn, mask=offs_cn < N, other=0.0) + result = accumulator.to(tl.float32) * scale[:, None] * scale_n[None, :] + + result = _add_lora(result, xd_ptr, up_ptr, offs_cm, offs_cn, M, N, R, + stride_xdm, stride_upr, stride_upn, BLOCK_R, R_TILES) + result = result.to(c_ptr.dtype.element_ty) + + _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) + +announce_autotuning(_rowcol_scaled_lora_mm_kernel, name="8-bit row/column scaled LoRA matmul") + +@torch.library.custom_op("ot_quant::rowcol_scaled_lora_mm_8bit", mutates_args=()) +def rowcol_scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + R, block_r = _prepare_lora(a, b, xd, up) + b, c, scale, scale_n, M, N, K, FLOAT = _prepare_rowcol_scaled_mm(a, b, scale, scale_n, out_dtype) + + def grid(META): + return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) + _rowcol_scaled_lora_mm_kernel[grid]( + a, b, c, scale, scale_n, xd, up, + M, N, K, R, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), xd.stride(0), up.stride(0), up.stride(1), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), + BLOCK_R = block_r, R_TILES = triton.cdiv(R, block_r), + ) + return c + +@rowcol_scaled_lora_mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) diff --git a/modules/util/ui/pyside6_util.py b/modules/util/ui/pyside6_util.py index 07d77dce5..eb0f84853 100644 --- a/modules/util/ui/pyside6_util.py +++ b/modules/util/ui/pyside6_util.py @@ -1,4 +1,5 @@ import locale +import os import signal import sys from abc import ABCMeta @@ -20,6 +21,14 @@ def create_application() -> QApplication: # active and Ctrl+C would be ignored. signal.signal(signal.SIGINT, signal.SIG_DFL) + # On desktops without the xdg-desktop-portal Settings interface, Qt spams two + # "qt.qpa.theme.gnome: dbus reply error ... org.freedesktop.portal.Settings" + # lines while probing for the system color scheme. Silence just that category; + # the rules string is read when Qt's logging initializes at QApplication init. + _gnome_theme_rule = "qt.qpa.theme.gnome=false" + existing_rules = os.environ.get("QT_LOGGING_RULES") + os.environ["QT_LOGGING_RULES"] = f"{existing_rules};{_gnome_theme_rule}" if existing_rules else _gnome_theme_rule + app = QApplication(sys.argv) # QApplication initializes the C locale from the environment (setlocale(LC_ALL, "")), which sets LC_NUMERIC # to a locale whose decimal separator may be a comma. C libraries then misparse '.' floats: protobuf's upb diff --git a/requirements-global.txt b/requirements-global.txt index 44c8317ce..1753c8663 100644 --- a/requirements-global.txt +++ b/requirements-global.txt @@ -5,7 +5,7 @@ pillow==12.3.0 imagesize==1.4.1 #for concept statistics tqdm==4.67.1 PyYAML==6.0.2 -huggingface-hub==1.22.0 +huggingface-hub==1.27.0 scipy==1.15.3 matplotlib==3.10.3 av==16.1.0 @@ -20,9 +20,11 @@ safetensors==0.8.0 tensorboard==2.20.0 # diffusion models --e git+https://github.com/huggingface/diffusers.git@1ffa423#egg=diffusers +-e git+https://github.com/huggingface/diffusers.git@175fe6b#egg=diffusers gguf==0.17.1 -transformers==5.5.4 # pinned below 5.6, see https://github.com/Nerogar/OneTrainer/pull/1506 +transformers==5.14.1 # LTX-2.5's Gemma 4 encoder needs >=5.10; 5.15 breaks its RoPE init. CLIP models + # (SD1.x, SDXL, Flux, HunyuanVideo, Würstchen) stay broken on this branch until the + # flattening migration of https://github.com/Nerogar/OneTrainer/pull/1506 is done. sentencepiece==0.2.1 # transitive dependency of transformers for tokenizer loading omegaconf==2.3.0 # needed to load stable diffusion from single ckpt files invisible-watermark==0.2.0 # needed for the SDXL pipeline @@ -32,7 +34,7 @@ pooch==1.8.2 open-clip-torch==2.32.0 # data loader --e git+https://github.com/Nerogar/mgds.git@3a6994a#egg=mgds +-e git+https://github.com/dxqb/MGDS.git@fc1391a#egg=mgds # ltx2-squashed, not yet upstream # optimizers dadaptation==3.2 # dadaptation optimizers diff --git a/resources/sd_model_spec/ltx_2-lora.json b/resources/sd_model_spec/ltx_2-lora.json new file mode 100644 index 000000000..8a18353f4 --- /dev/null +++ b/resources/sd_model_spec/ltx_2-lora.json @@ -0,0 +1,6 @@ +{ + "modelspec.sai_model_spec": "1.0.0", + "modelspec.architecture": "ltx-2/lora", + "modelspec.implementation": "https://github.com/huggingface/diffusers", + "modelspec.title": "LTX 2 LoRA" +} diff --git a/resources/sd_model_spec/ltx_2.json b/resources/sd_model_spec/ltx_2.json new file mode 100644 index 000000000..18a8ef509 --- /dev/null +++ b/resources/sd_model_spec/ltx_2.json @@ -0,0 +1,6 @@ +{ + "modelspec.sai_model_spec": "1.0.0", + "modelspec.architecture": "ltx-2", + "modelspec.implementation": "https://github.com/huggingface/diffusers", + "modelspec.title": "LTX 2" +} diff --git a/scripts/train_remote.py b/scripts/train_remote.py index 37b90931a..11b65ca28 100644 --- a/scripts/train_remote.py +++ b/scripts/train_remote.py @@ -88,6 +88,13 @@ def main(): trainer.start() trainer.train() + except Exception: + # print the traceback before end() runs. end() saves the final model, which takes long enough that a log + # ending on its "Saving ..." line is indistinguishable from a clean stop until the traceback finally + # appears after it. The local UI path already orders it this way. + traceback.print_exc() + raise + finally: if args.command_path: stop_event.set() diff --git a/scripts/util/import_util.py b/scripts/util/import_util.py index 150df4d8a..e269a0505 100644 --- a/scripts/util/import_util.py +++ b/scripts/util/import_util.py @@ -1,7 +1,9 @@ def script_imports(allow_zluda: bool = True): import logging import os + import re import sys + import warnings from pathlib import Path # Filter out the Triton warning on startup. @@ -10,6 +12,43 @@ def script_imports(allow_zluda: bool = True): .getLogger("xformers") \ .addFilter(lambda record: 'A matching Triton is not available' not in record.getMessage()) + # Silence specific non-actionable startup/compile warnings. A logger filter + # targets the exact emitting logger, since a parent logger's filter misses + # records from child loggers. Set OT_DEBUG_WARNINGS to see them all. + if not os.environ.get("OT_DEBUG_WARNINGS"): + # diffusers/transformers chatty logger.warning() lines at import/load time. + logging.getLogger("diffusers.modular_pipelines").addFilter( + lambda record: 'Modular Diffusers is currently an experimental feature' not in record.getMessage() + ) + # The subject of these two is interpolated into the message, so match the whole + # sentence with .* standing in for the runtime value. + logging.getLogger("diffusers.configuration_utils").addFilter( + lambda record: not re.search( + r"The config attributes .* were passed to .*, but are not expected and will be ignored", + record.getMessage(), + ) + ) + logging.getLogger("diffusers.models.modeling_utils").addFilter( + lambda record: 'Attention backends are an experimental feature' not in record.getMessage() + ) + logging.getLogger("transformers.modeling_utils").addFilter( + lambda record: not re.search( + r"`loss_type=.*` was set in the config but it is unrecognized", record.getMessage() + ) + ) + + # A dependency still calls hf_hub_download with the removed local_dir_use_symlinks + # argument; the deprecation warning is not actionable. + warnings.filterwarnings("ignore", message=r".*local_dir_use_symlinks.*") + + # torch.compile emits performance notes when inductor falls back or can't use a + # fast path; harmless and noisy for normal runs. The SMs note is a logger.warning() + # on its exact emitting logger; the complex-operators note is a warnings.warn(). + warnings.filterwarnings("ignore", message=r".*does not support code generation for complex operators.*") + logging.getLogger("torch._inductor.utils").addFilter( + lambda record: 'Not enough SMs to use max_autotune_gemm mode' not in record.getMessage() + ) + # Insert ourselves as the highest-priority library path, so our modules are # always found without any risk of being shadowed by another import path. # 3 .parent calls to navigate from /scripts/util/import_util.py to the main directory diff --git a/training_presets/LTX 2/#ltx2.5 LoRA 16GB.json b/training_presets/LTX 2/#ltx2.5 LoRA 16GB.json new file mode 100644 index 000000000..f2057ee74 --- /dev/null +++ b/training_presets/LTX 2/#ltx2.5 LoRA 16GB.json @@ -0,0 +1,55 @@ +{ + "base_model_name": "Lightricks/LTX-2.5-Diffusers", + "model_type": "LTX_2", + "training_method": "LORA", + "output_model_format": "DIFFUSERS_LORA", + "batch_size": 2, + "learning_rate": 3e-4, + "resolution": "512", + "frames": "121", + "dataloader_threads": 1, + "compile": true, + "train_dtype": "BFLOAT_16", + "output_dtype": "BFLOAT_16", + "transformer": { + "train": true, + "weight_dtype": "INT_W8A8_COMPRESSED", + "offload_fraction": 0.5, + "simplex_offloading": true, + "cache_in_ram": false + }, + "low_noise_transformer": { + "model_name": "Lightricks/LTX-2.5-Diffusers", + "train": false, + "weight_dtype": "INT_W8A8_COMPRESSED", + "gradient_checkpointing": false, + "offload_fraction": 0.3, + "cache_in_ram": false + }, + "text_encoder": { + "train": false, + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.8, + "cache_in_ram": false + }, + "connectors": { + "train": false, + "weight_dtype": "BFLOAT_16", + "gradient_checkpointing": false, + "offload_fraction": 0.8, + "cache_in_ram": false + }, + "vae": { + "weight_dtype": "FLOAT_32" + }, + "layer_filter": ".attn1,.attn2,.ff", + "layer_filter_preset": "video-attn-mlp", + "quantization": { + "layer_filter": "^transformer_blocks\\.(?![01]\\.|4[67]\\.)", + "layer_filter_preset": "lightricks-2.3-fp8", + "layer_filter_regex": true, + "fallback_dtype": "BFLOAT_16" + }, + "timestep_distribution": "LOGIT_NORMAL", + "dynamic_timestep_shifting": false +}