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/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..126e0c305 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,20 +52,33 @@ 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. - if include_unconditional_transformer: - unconditional_transformer = self._load_diffusers_sub_module( + # 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; in + # streaming mode that returns a materialize closure, otherwise a plain module. + if include_unconditional_transformer and stream_from_disk: + 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=True, + ) + elif include_unconditional_transformer: + model.unconditional_transformer = self._load_diffusers_sub_module( Ideogram4Transformer2DModel, weight_dtypes.unconditional_transformer, weight_dtypes.train_dtype, @@ -71,49 +87,34 @@ def __load_diffusers( quantization, ) 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 +135,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 +151,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/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..7c542bd00 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,119 @@ 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( - LlamaForCausalLM, - weight_dtypes.text_encoder_4, - weight_dtypes.train_dtype, - text_encoder_4_model_name, - ) + # 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. + # _load_transformers_sub_module returns a (module, materialize_fn) pair only when streaming; a bare + # module otherwise. + if stream_from_disk: + 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=True, + ) + else: + model.text_encoder_4 = self._load_transformers_sub_module( + LlamaForCausalLM, + weight_dtypes.text_encoder_4, + weight_dtypes.train_dtype, + text_encoder_4_model_name, + ) 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 +267,7 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] @@ -277,6 +277,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 +290,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..5f548e63c 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 = name.replace(".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, @@ -157,7 +403,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 +441,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,6 +450,7 @@ def _load_transformers_sub_module( train_dtype: DataType, pretrained_model_name_or_path: str, subfolder: str = "", + stream_from_disk: bool = False, ): user_agent = { "file_type": "model", @@ -213,19 +467,110 @@ 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 + return 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", + ) + + 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, [], 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,6 +579,7 @@ def _load_diffusers_sub_module( pretrained_model_name_or_path: str, subfolder: str | None = None, quantization: QuantizationConfig | None = None, + stream_from_disk: bool = False, ): user_agent = { "file_type": "model", @@ -250,19 +596,67 @@ 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 + return 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", + ) + + 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 +677,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 +728,117 @@ 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, + ): + # 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. + # Always returns a (transformer, materialize_fn) pair -- materialize_fn None when not streamed -- so callers + # pass stream_from_disk through. + if transformer_model_name: + single_file_kwargs = {} + if config is not None: + single_file_kwargs["config"] = config + single_file_kwargs["subfolder"] = "transformer" + + 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 + elif stream_from_disk: + # stream from disk: 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. + # train_dtype is applied per-materialize, not here. + return self._load_diffusers_sub_module( + module_type, + weight_dtypes.transformer, + weight_dtypes.train_dtype, + base_model_name, + "transformer", + quantization, + stream_from_disk=True, + ) + else: + transformer = self._load_diffusers_sub_module( + module_type, + weight_dtypes.transformer, + weight_dtypes.train_dtype, + base_model_name, + "transformer", + quantization, + ) + return transformer, None + + 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. Always 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. + if stream_from_disk: + return self._load_transformers_sub_module( + module_type, + dtype, + train_dtype, + base_model_name, + subfolder, + stream_from_disk=True, + ) + else: + text_encoder = self._load_transformers_sub_module( + module_type, + dtype, + train_dtype, + base_model_name, + subfolder, + ) + return text_encoder, None + + 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..216e0186f 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,78 @@ 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, + model.vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.fallback_train_dtype, + base_model_name, + vae_model_name, + ) + + # the SDXL UNet has no single-file transformer helper and lives in the "unet" subfolder, so it streams via + # _load_diffusers_sub_module directly, which returns a (module, materialize_fn) pair only when streaming and a + # bare module otherwise (train_dtype is applied per-materialize when streaming, so pass None there). + if stream_from_disk: + model.unet, model.materialize_fn["unet"] = self._load_diffusers_sub_module( + UNet2DConditionModel, + weight_dtypes.unet, + None, + base_model_name, + "unet", + quantization, + stream_from_disk=True, ) else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.fallback_train_dtype, + model.unet = self._load_diffusers_sub_module( + UNet2DConditionModel, + weight_dtypes.unet, + weight_dtypes.train_dtype, base_model_name, - "vae", + "unet", + quantization, ) - unet = self._load_diffusers_sub_module( - UNet2DConditionModel, - weight_dtypes.unet, - weight_dtypes.train_dtype, - base_model_name, - "unet", - quantization, - ) - - 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 +250,7 @@ def load( model_names: ModelNames, weight_dtypes: ModelWeightDtypes, quantization: QuantizationConfig, + stream_from_disk: bool = False, ): stacktraces = [] @@ -255,17 +258,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/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/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/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/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..3fcb3b777 100644 --- a/modules/module/quantized/LinearSVD.py +++ b/modules/module/quantized/LinearSVD.py @@ -1,3 +1,4 @@ +import os from abc import abstractmethod from contextlib import suppress @@ -51,6 +52,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 +87,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] diff --git a/modules/module/quantized/LinearW8A8.py b/modules/module/quantized/LinearW8A8.py index 8bbb90e5a..6da9bdc16 100644 --- a/modules/module/quantized/LinearW8A8.py +++ b/modules/module/quantized/LinearW8A8.py @@ -82,33 +82,50 @@ class LinearW8A8( QuantizedLinearMixin, CompressedWeightMixin, ): - 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. + 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 @@ -125,14 +142,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 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/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/GenericTrainer.py b/modules/trainer/GenericTrainer.py index a6547afbe..d406c4c1e 100644 --- a/modules/trainer/GenericTrainer.py +++ b/modules/trainer/GenericTrainer.py @@ -137,6 +137,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 @@ -664,7 +666,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( diff --git a/modules/ui/BaseModelTabView.py b/modules/ui/BaseModelTabView.py index 3b48a46f0..6b61dfb34 100644 --- a/modules/ui/BaseModelTabView.py +++ b/modules/ui/BaseModelTabView.py @@ -128,6 +128,13 @@ def __create_base_dtype_components(self, frame, row: int, ui_state) -> int: row += 1 + # stream from disk + self.components.label(frame, row, 0, "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, 1, ui_state, "stream_from_disk") + + row += 1 + return row def __create_base_components( @@ -151,15 +158,6 @@ def __create_base_components( 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 +176,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 +204,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.") diff --git a/modules/ui/BaseTrainingTabView.py b/modules/ui/BaseTrainingTabView.py index 29a54bb1e..58a252f7d 100644 --- a/modules/ui/BaseTrainingTabView.py +++ b/modules/ui/BaseTrainingTabView.py @@ -455,12 +455,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, diff --git a/modules/ui/SampleWindowController.py b/modules/ui/SampleWindowController.py index b388bca5a..54e0444ec 100644 --- a/modules/ui/SampleWindowController.py +++ b/modules/ui/SampleWindowController.py @@ -90,6 +90,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 diff --git a/modules/util/LayerOffloadConductor.py b/modules/util/LayerOffloadConductor.py index f3e737c0e..06245aec4 100644 --- a/modules/util/LayerOffloadConductor.py +++ b/modules/util/LayerOffloadConductor.py @@ -1,13 +1,25 @@ import math import random +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, @@ -20,6 +32,8 @@ import torch from torch import nn +from tqdm import tqdm + MESSAGES = [] @@ -29,17 +43,23 @@ def log(msg: str = ''): # MESSAGES.append(msg) -def clone_tensor_allocator(tensor: torch.Tensor) -> torch.Tensor: +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 +92,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 +117,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 +129,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 +139,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 +188,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 +250,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 +272,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 +308,99 @@ 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. It is filled at materialize by an explicit GPU->CPU copy, and survives evict only while + # cache_in_ram is on -- with it off the buffer is freed with the weights and refilled from a fresh stream. + 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 there is none; (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 exact + # packed-weight 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. + self.deactivate_cache() + self.__buffer = None + self.__fill_offset = 0 + self.__slots = {} + + class StaticActivationAllocator: __device: torch.device __allocate_statically: bool @@ -254,7 +428,7 @@ def __init__( 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 @@ -290,7 +464,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) @@ -548,7 +722,7 @@ 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] @@ -562,17 +736,24 @@ class LayerOffloadConductor: __is_forward_pass: bool __keep_graph: 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__() @@ -600,8 +781,13 @@ 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 = [] @@ -615,16 +801,29 @@ def __init__( self.__is_forward_pass = False self.__keep_graph = 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 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() @@ -632,63 +831,175 @@ def evict(self): 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: + # the buffer is empty: 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. + 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: + # the buffer survived the evict, which freed only the GPU ring the weights 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 + 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 - torch_gc() + self.__materialized = True + return True def add_layer(self, layer: nn.Module, included_offload_param_indices: list[int] = None): if included_offload_param_indices is None: @@ -698,13 +1009,15 @@ def add_layer(self, layer: nn.Module, included_offload_param_indices: list[int] 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): log("starting forward") - if not self.__is_active: + if not self.__materialized: return if self.__async_transfer: @@ -719,7 +1032,7 @@ def before_layer(self, layer_index: int, call_index: int, activations: Any) -> A log() log(f"before layer {layer_index}, {call_index}") - if not self.__is_active: + if not self.__materialized: return activations self.__call_index_layer_index_map[call_index] = layer_index @@ -780,7 +1093,7 @@ def before_layer(self, layer_index: int, call_index: int, activations: Any) -> A 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: + if not self.__materialized: return # record stream @@ -797,12 +1110,46 @@ def after_layer(self, layer_index: int, call_index: int, activations: Any): event = SyncEvent(self.__train_stream.record_event(), f"train on {self.__train_device}") self.__layer_train_event_map[layer_index] = event + 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 + + 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,10 +1157,37 @@ 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() @@ -870,7 +1244,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 +1269,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) diff --git a/modules/util/checkpointing_util.py b/modules/util/checkpointing_util.py index 2df40065e..058b06ed6 100644 --- a/modules/util/checkpointing_util.py +++ b/modules/util/checkpointing_util.py @@ -245,13 +245,37 @@ 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") + + # simplex is a layer-offloading mode, so with layer offloading off there is no conductor to give the buffer to + # and the switch does nothing -- a stale setting rather than a bad one, so read it as inactive instead of + # rejecting the combinations below against constraints that would never apply. + 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 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 diff --git a/modules/util/config/TrainConfig.py b/modules/util/config/TrainConfig.py index ba11de9d7..55e8f6256 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 @@ -891,6 +898,18 @@ 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, @@ -1046,6 +1065,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)) 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..aa9098c7b 100644 --- a/modules/util/enum/ModelType.py +++ b/modules/util/enum/ModelType.py @@ -219,6 +219,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, @@ -316,6 +319,41 @@ def supported_output_formats(self, training_method: TrainingMethod) -> list[Mode ModelType.IDEOGRAM_4: ("transformer", "text_encoder", "unconditional_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",), +} + class PeftType(Enum): LORA = 'LORA' diff --git a/modules/util/quantization_util.py b/modules/util/quantization_util.py index d7229687d..b61c028c3 100644 --- a/modules/util/quantization_util.py +++ b/modules/util/quantization_util.py @@ -102,6 +102,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 +122,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 +136,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 +149,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 +163,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 +225,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 +239,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 +283,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 +351,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 +363,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/torch_util.py b/modules/util/torch_util.py index 408100bf9..2d2a9c1c6 100644 --- a/modules/util/torch_util.py +++ b/modules/util/torch_util.py @@ -1,4 +1,6 @@ +import contextlib import gc +import time from collections.abc import Callable from contextlib import nullcontext from typing import Any @@ -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: