Added legacy paths

This commit is contained in:
Jaret Burkett
2026-08-28 07:55:19 -06:00
parent c3bc8b0b4e
commit 702254688d
45 changed files with 840 additions and 223 deletions

View File

@@ -371,6 +371,43 @@ Decisions:
unload/load the difference
- [ ] Legacy `stable_diffusion_model.py` archs: grandfather or port last
## 100% mixin coverage (2026-08-28)
Every component of every non-legacy arch now loads through OstrisModelMixin —
DiTs, text encoders, VAEs, vision encoders, connectors/vocoders, and the
custom cases that previously bypassed it:
- New wrapper classes: llama, gemma3 + gemma4, mistral3 (×2), qwen3 base,
qwen3-vl base + text-only, wan VAE, diffusers flux2 KL VAE, CLIP vision
(first vision encoder), cosmos DiT, anima text conditioner, the full ltx2
set (transformer, video/audio VAEs, connectors, vocoders).
- Existing-class swaps: omnigen2 (mllm + VAE), boogu (TE + VAE), klein TE,
hidream llama, ideogram4 TE, prx TE.
- Custom restructures: minimax video/audio VAEs and TE, MageVAE
(deferred-load ctor), the shared flux2_kl AutoEncoder (small-decoder sniff
as a class hook; flux2 + ideogram4 route through it), ace_step's bundle
(per-component class loaders), anima (modular-pipeline load replaced with
component-wise v2 loads + update_components), hidream_o1's Qwen3VL DiT,
z_image_l2p rebased onto the v2 class, ltx2's converter builds v2 classes.
- Mixin: kwargs flow through the single-file chain (component ctor args like
MageVAE's sample_posterior).
Verified with real weights: ltx2.3 + ltx2.5 (full v2 family), anima,
ace_step_15 (new harness entry), minimax VAEs, plus the standing harness
coverage. The legacy monolith archs joined the system per the inference-engine
goal ("send a job with any base_model, reload/unload components on the fly"):
six new wrapper classes complete the component vocabulary (UNet2DCondition,
SD3/PixArt ×2/AuraFlow/Lumina2 transformers, Gemma2), and
`adopt_component` (in-place class swap, OstrisLinear-style) rebinds
pipeline-loaded components onto their wrappers at the monolith's single
post-load funnel — covering every legacy arch without touching its fifteen
load branches, with all pipeline references staying valid. Verified: sd1
(UNet/KLVAE/CLIPTextEncoder all mixin instances + generation) and sdxl
(dual-CLIP adoption + 1024² generation); harness gained sd1/sdxl entries,
a legacy scheduler fallback, and the sampler-name pass-through. Full
monolith decomposition (per-arch v2 loading with comfy candidates) remains
future work, but every resident component is now poolable by the engine.
## Testing
- [x] `testing/test_model_loading.py`: per-arch load + one small sample through
@@ -418,8 +455,12 @@ Decisions:
flux2_klein_9b, prx_pixel, zeta_chroma, zimage_l2p, both qwen edit
archs, f-lite. Fixes found by the run: ltx2.5's fp32 scale_shift
tables promoted hidden states into bf16 linears under the diffusers
class (ComfyUI casts per-op) — the DiT/connectors now cast to compute
dtype after quantized attach (ConvRot storage immune); harness configs
class — fixed ComfyUI-style (toolkit/util/mixed_precision.py):
per-op input casting hooks on every weighted module + the stored-fp32
tensors pinned against parent .to(dtype) casts (device moves pass
through), so the tables stay genuinely fp32 at sample time. The
mechanism is general — any mixed-precision comfy checkpoint on a
diffusers-class arch can use it; harness configs
for e1 resolution and zeta/l2p extras_name_or_path corrected to match
the UI defaults.
- [ ] Each newly migrated model adds its test in the same PR as its migration.

View File

@@ -42,6 +42,24 @@ class BasicModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMix
pass
def adopt_component(module: torch.nn.Module, wrapper_cls: type) -> torch.nn.Module:
"""Rebind an already-loaded component instance onto its v2 wrapper class
(in-place class swap, like the OstrisLinear conversion). Used when a
component arrives from machinery that builds the base class directly
(e.g. a diffusers pipeline load), so every resident component is a mixin
instance the inference engine can pool and manage."""
if isinstance(module, wrapper_cls):
return module
base = wrapper_cls.__mro__[1]
if not isinstance(module, base):
raise TypeError(
f"cannot adopt {type(module).__name__} into {wrapper_cls.__name__} "
f"(expected a {base.__name__})"
)
module.__class__ = wrapper_cls
return module
class OstrisModelMixin:
# ---- per-model configuration, override in subclasses ----
# subfolder that holds this model inside a diffusers style checkpoint
@@ -195,6 +213,7 @@ class OstrisModelMixin:
config_path=config_path,
config=config,
subfolder=subfolder,
**kwargs,
)
else:
if os.path.isdir(name_or_path):
@@ -259,6 +278,7 @@ class OstrisModelMixin:
config_path: Optional[str] = None,
config=None,
subfolder: Optional[str] = None,
**kwargs,
):
state_dict = load_file(file_path)
return cls.load_from_state_dict(
@@ -267,6 +287,7 @@ class OstrisModelMixin:
config_path=config_path,
config=config,
subfolder=subfolder,
**kwargs,
)
@classmethod
@@ -284,6 +305,7 @@ class OstrisModelMixin:
config_path: Optional[str] = None,
config=None,
subfolder: Optional[str] = None,
**kwargs,
):
"""Build the model and load an already-read single-file state dict
(the tail of the single-file path; also callable directly for

View File

@@ -0,0 +1,13 @@
from diffusers import AuraFlowTransformer2DModel as DiffusersAuraFlowTransformer2DModel
from .._mixin import OstrisModelMixin
class AuraFlowTransformer2DModel(
DiffusersAuraFlowTransformer2DModel, OstrisModelMixin
):
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["joint_transformer_blocks", "single_transformer_blocks"]

View File

@@ -0,0 +1,15 @@
from diffusers.models import (
CosmosTransformer3DModel as DiffusersCosmosTransformer3DModel,
)
from .._mixin import OstrisModelMixin
class CosmosTransformer3DModel(DiffusersCosmosTransformer3DModel, OstrisModelMixin):
"""Cosmos video/image DiT (anima)."""
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["transformer_blocks"]

View File

@@ -0,0 +1,34 @@
from diffusers.models.transformers import (
LTX2VideoTransformer3DModel as DiffusersLTX2VideoTransformer3DModel,
)
from diffusers.pipelines.ltx2 import (
LTX2TextConnectors as DiffusersLTX2TextConnectors,
)
from diffusers.pipelines.ltx2 import LTX2Vocoder as DiffusersLTX2Vocoder
from diffusers.pipelines.ltx2 import (
LTX2VocoderWithBWE as DiffusersLTX2VocoderWithBWE,
)
from .._mixin import OstrisModelMixin
class LTX2VideoTransformer3DModel(
DiffusersLTX2VideoTransformer3DModel, OstrisModelMixin
):
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["transformer_blocks"]
class LTX2TextConnectors(DiffusersLTX2TextConnectors, OstrisModelMixin):
aitk_subfolder = "connectors"
class LTX2Vocoder(DiffusersLTX2Vocoder, OstrisModelMixin):
aitk_subfolder = "vocoder"
class LTX2VocoderWithBWE(DiffusersLTX2VocoderWithBWE, OstrisModelMixin):
aitk_subfolder = "vocoder"

View File

@@ -0,0 +1,13 @@
from diffusers import Lumina2Transformer2DModel as DiffusersLumina2Transformer2DModel
from .._mixin import OstrisModelMixin
class Lumina2Transformer2DModel(
DiffusersLumina2Transformer2DModel, OstrisModelMixin
):
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["layers"]

View File

@@ -0,0 +1,24 @@
from diffusers import PixArtTransformer2DModel as DiffusersPixArtTransformer2DModel
from diffusers import Transformer2DModel as DiffusersTransformer2DModel
from .._mixin import OstrisModelMixin
class PixArtTransformer2DModel(
DiffusersPixArtTransformer2DModel, OstrisModelMixin
):
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["transformer_blocks"]
class Transformer2DModel(DiffusersTransformer2DModel, OstrisModelMixin):
"""The generic diffusers DiT the pixart sigma path loads."""
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["transformer_blocks"]

View File

@@ -0,0 +1,11 @@
from diffusers import SD3Transformer2DModel as DiffusersSD3Transformer2DModel
from .._mixin import OstrisModelMixin
class SD3Transformer2DModel(DiffusersSD3Transformer2DModel, OstrisModelMixin):
aitk_subfolder = "transformer"
@classmethod
def get_transformer_block_names(cls):
return ["transformer_blocks"]

View File

@@ -0,0 +1,13 @@
from diffusers import UNet2DConditionModel as DiffusersUNet2DConditionModel
from .._mixin import OstrisModelMixin
class UNet2DConditionModel(DiffusersUNet2DConditionModel, OstrisModelMixin):
"""The SD1/SD2/SDXL-family UNet."""
aitk_subfolder = "unet"
@classmethod
def get_transformer_block_names(cls):
return ["down_blocks", "up_blocks"]

View File

@@ -0,0 +1,9 @@
from diffusers import AnimaTextConditioner as DiffusersAnimaTextConditioner
from .._mixin import OstrisModelMixin
class AnimaTextConditioner(DiffusersAnimaTextConditioner, OstrisModelMixin):
"""Anima's learned text conditioner (rides next to the Qwen3 encoder)."""
aitk_subfolder = "text_conditioner"

View File

@@ -0,0 +1,14 @@
from transformers import Gemma2Model
from .._mixin import OstrisTransformersMixin
class Gemma2ModelEncoder(Gemma2Model, OstrisTransformersMixin):
"""Gemma2 base model (lumina2's text encoder — what AutoModel resolves)."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["layers"]

View File

@@ -0,0 +1,32 @@
from transformers import Gemma3ForConditionalGeneration
from .._mixin import OstrisTransformersMixin
class Gemma3TextEncoder(Gemma3ForConditionalGeneration, OstrisTransformersMixin):
"""Gemma3 conditioning stack (ltx2 / ltx2.3)."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
# both layouts seen across transformers versions; missing paths skip
return ["model.language_model.layers", "language_model.model.layers"]
try:
from transformers.models.gemma4.modeling_gemma4 import Gemma4TextModel
class Gemma4TextEncoder(Gemma4TextModel, OstrisTransformersMixin):
"""Gemma4 text decoder (ltx2.5's conditioning stack)."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["layers"]
except ImportError:
Gemma4TextEncoder = None

View File

@@ -0,0 +1,14 @@
from transformers import LlamaForCausalLM
from .._mixin import OstrisTransformersMixin
class LlamaTextEncoder(LlamaForCausalLM, OstrisTransformersMixin):
"""Llama causal-LM text encoder (hidream's text_encoder_4)."""
aitk_subfolder = "text_encoder_4"
aitk_tokenizer_subfolder = "tokenizer_4"
@classmethod
def get_transformer_block_names(cls):
return ["model.layers"]

View File

@@ -0,0 +1,25 @@
from transformers import Mistral3ForConditionalGeneration, Mistral3Model
from .._mixin import OstrisTransformersMixin
class Mistral3TextEncoder(Mistral3ForConditionalGeneration, OstrisTransformersMixin):
"""Mistral3 conditioning stack (flux2)."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["model.language_model.layers", "language_model.model.layers"]
class Mistral3ModelEncoder(Mistral3Model, OstrisTransformersMixin):
"""The inner Mistral3 base model (ernie_image's text encoder)."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["language_model.layers"]

View File

@@ -1,4 +1,4 @@
from transformers import Qwen3ForCausalLM
from transformers import Qwen3ForCausalLM, Qwen3Model
from .._mixin import OstrisTransformersMixin
@@ -15,3 +15,14 @@ class Qwen3TextEncoder(Qwen3ForCausalLM, OstrisTransformersMixin):
@classmethod
def get_transformer_block_names(cls):
return ["model.layers"]
class Qwen3ModelEncoder(Qwen3Model, OstrisTransformersMixin):
"""The inner Qwen3 base model (anima's text encoder)."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["layers"]

View File

@@ -1,6 +1,10 @@
import torch
import torch.nn.functional as F
from transformers import Qwen3VLForConditionalGeneration
from transformers import (
Qwen3VLForConditionalGeneration,
Qwen3VLModel,
Qwen3VLTextModel,
)
from .._mixin import OstrisTransformersMixin
@@ -51,3 +55,35 @@ class Qwen3VLTextEncoder(Qwen3VLForConditionalGeneration, OstrisTransformersMixi
"""Keep the vision tower (reference images ride into the embeddings)
but swap its Conv3d patch_embed for an equivalent GEMM."""
return patch_qwen_vl_patch_embed(self)
class Qwen3VLModelEncoder(Qwen3VLModel, OstrisTransformersMixin):
"""The inner Qwen3-VL base model (what AutoModel resolves): its
last_hidden_state is the instruction feature boogu_image / ideogram4
consume."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["language_model.layers"]
def drop_vision_tower(self):
if getattr(self, "visual", None) is not None:
self.visual = None
return self
def patch_vision_patch_embed(self) -> int:
return patch_qwen_vl_patch_embed(self)
class Qwen3VLTextOnlyEncoder(Qwen3VLTextModel, OstrisTransformersMixin):
"""The Qwen3-VL text tower alone (prx_pixel)."""
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
@classmethod
def get_transformer_block_names(cls):
return ["layers"]

View File

@@ -0,0 +1,10 @@
from diffusers import AutoencoderKLFlux2
from .._mixin import OstrisModelMixin
class Flux2KLVAE(AutoencoderKLFlux2, OstrisModelMixin):
"""The diffusers Flux2 KL VAE (ernie_image; distinct from the hand-rolled
BFL AutoEncoder in vae/flux2_kl.py that flux2/ideogram4 use)."""
aitk_subfolder = "vae"

View File

@@ -20,6 +20,8 @@ import torch.utils.checkpoint as ckpt
from einops import rearrange
from torch import Tensor, nn
from .._mixin import OstrisModelMixin
@dataclass
class AutoEncoderParams:
@@ -354,7 +356,14 @@ class Decoder(nn.Module):
return h
class AutoEncoder(nn.Module):
class AutoEncoder(nn.Module, OstrisModelMixin):
@classmethod
def aitk_config_from_state_dict(cls, state_dict):
# small decoder builds report 96ch at the first decoder block
if state_dict["decoder.up.0.block.0.conv1.bias"].shape[0] == 96:
return AutoEncoderSmallDecoderParams()
return AutoEncoderParams()
def __init__(self, params: AutoEncoderParams):
super().__init__()
self.params = params

View File

@@ -0,0 +1,16 @@
from diffusers.models.autoencoders import (
AutoencoderKLLTX2Audio as DiffusersAutoencoderKLLTX2Audio,
)
from diffusers.models.autoencoders import (
AutoencoderKLLTX2Video as DiffusersAutoencoderKLLTX2Video,
)
from .._mixin import OstrisModelMixin
class LTX2VideoVAE(DiffusersAutoencoderKLLTX2Video, OstrisModelMixin):
aitk_subfolder = "vae"
class LTX2AudioVAE(DiffusersAutoencoderKLLTX2Audio, OstrisModelMixin):
aitk_subfolder = "audio_vae"

View File

@@ -0,0 +1,9 @@
from diffusers import AutoencoderKLWan
from .._mixin import OstrisModelMixin
class WanVAE(AutoencoderKLWan, OstrisModelMixin):
"""The wan 2.1/2.2 causal video VAE."""
aitk_subfolder = "vae"

View File

@@ -0,0 +1,13 @@
from transformers import CLIPVisionModel
from .._mixin import OstrisTransformersMixin
class CLIPVisionEncoder(CLIPVisionModel, OstrisTransformersMixin):
"""CLIP vision tower (wan21 i2v image conditioning)."""
aitk_subfolder = "image_encoder"
@classmethod
def get_transformer_block_names(cls):
return ["vision_model.encoder.layers"]

View File

@@ -44,6 +44,7 @@ from typing import Any, Callable, Dict, List, Optional, Union
from toolkit.models.wan21.wan_lora_convert import convert_to_diffusers, convert_to_original
from toolkit.util.quantize import quantize_model
from toolkit.models.v2.text_encoders.umt5 import UMT5TextEncoder
from toolkit.models.v2.vae.wan import WanVAE
from toolkit.metadata import get_meta_for_safetensors
# for generation only?
@@ -431,11 +432,9 @@ class Wan21(BaseModel):
if self._wan_vae_path is not None:
# load the vae from individual repo
vae = AutoencoderKLWan.from_pretrained(
self._wan_vae_path, torch_dtype=dtype).to(dtype=dtype)
vae = WanVAE.load_model(self._wan_vae_path, dtype=dtype, subfolder="")
else:
vae = AutoencoderKLWan.from_pretrained(
vae_path, subfolder="vae", torch_dtype=dtype).to(dtype=dtype)
vae = WanVAE.load_model(vae_path, dtype=dtype)
flush()
self.print_and_status_update("Making pipe")

View File

@@ -1,4 +1,5 @@
# WIP, coming soon ish
from toolkit.models.v2.vision_encoders.clip_vision import CLIPVisionEncoder
from functools import partial
import torch
import yaml
@@ -314,10 +315,8 @@ class Wan21I2V(Wan21):
self.model_config.extras_name_or_path ,
subfolder="image_processor"
)
self.image_encoder = CLIPVisionModel.from_pretrained(
self.model_config.extras_name_or_path,
subfolder="image_encoder",
torch_dtype=dtype,
self.image_encoder = CLIPVisionEncoder.load_model(
self.model_config.extras_name_or_path, dtype=dtype
)
except Exception as e:
# load from name_or_path
@@ -325,10 +324,8 @@ class Wan21I2V(Wan21):
self.model_config.name_or_path_original,
subfolder="image_processor"
)
self.image_encoder = CLIPVisionModel.from_pretrained(
self.model_config.name_or_path_original,
subfolder="image_encoder",
torch_dtype=dtype,
self.image_encoder = CLIPVisionEncoder.load_model(
self.model_config.name_or_path_original, dtype=dtype
)
self.image_encoder.to(self.device_torch, dtype=dtype)
self.image_encoder.eval()

View File

@@ -101,6 +101,46 @@ DO_NOT_TRAIN_WEIGHTS = [
DeviceStatePreset = Literal['cache_latents', 'generate']
# diffusers-class name -> v2 wrapper for the legacy archs' components. The
# adoption is an in-place class swap, so pipeline-held references stay valid.
_V2_ADOPTION_MAP = {
"UNet2DConditionModel": ("toolkit.models.v2.diffusion_models.unet", "UNet2DConditionModel"),
"AutoencoderKL": ("toolkit.models.v2.vae.autoencoder_kl", "KLVAE"),
"CLIPTextModel": ("toolkit.models.v2.text_encoders.clip", "CLIPTextEncoder"),
"CLIPTextModelWithProjection": ("toolkit.models.v2.text_encoders.clip", "CLIPTextEncoderWithProjection"),
"T5EncoderModel": ("toolkit.models.v2.text_encoders.t5", "T5TextEncoder"),
"UMT5EncoderModel": ("toolkit.models.v2.text_encoders.umt5", "UMT5TextEncoder"),
"SD3Transformer2DModel": ("toolkit.models.v2.diffusion_models.sd3", "SD3Transformer2DModel"),
"PixArtTransformer2DModel": ("toolkit.models.v2.diffusion_models.pixart", "PixArtTransformer2DModel"),
"Transformer2DModel": ("toolkit.models.v2.diffusion_models.pixart", "Transformer2DModel"),
"AuraFlowTransformer2DModel": ("toolkit.models.v2.diffusion_models.auraflow", "AuraFlowTransformer2DModel"),
"FluxTransformer2DModel": ("toolkit.models.v2.diffusion_models.flux", "FluxTransformer2DModel"),
"Lumina2Transformer2DModel": ("toolkit.models.v2.diffusion_models.lumina2", "Lumina2Transformer2DModel"),
"Gemma2Model": ("toolkit.models.v2.text_encoders.gemma2", "Gemma2ModelEncoder"),
}
def _adopt_v2(module):
"""Rebind a loaded legacy component onto its v2 wrapper class so every
resident component is an OstrisModelMixin instance the inference engine
can pool and hot-swap. No-op for unknown or already-adopted classes."""
import importlib
from toolkit.models.v2._mixin import OstrisModelMixin, adopt_component
if module is None or isinstance(module, OstrisModelMixin):
return module
entry = _V2_ADOPTION_MAP.get(type(module).__name__)
if entry is None:
return module
try:
wrapper = getattr(importlib.import_module(entry[0]), entry[1])
return adopt_component(module, wrapper)
except (ImportError, TypeError):
return module
class BlankNetwork:
def __init__(self):
@@ -1048,6 +1088,16 @@ class StableDiffusion:
self.unet.requires_grad_(False)
self.unet.eval()
# every resident component joins the v2 mixin system (in-place class
# adoption for components the pipeline loaders built directly)
_adopt_v2(self.unet)
_adopt_v2(self.vae)
if isinstance(text_encoder, list):
for te in text_encoder:
_adopt_v2(te)
elif text_encoder is not None:
_adopt_v2(text_encoder)
# load any loras we have
if self.model_config.lora_path is not None and not self.is_flux and not self.is_lumina2:
pipe.load_lora_weights(self.model_config.lora_path, adapter_name="lora1")

View File

@@ -0,0 +1,93 @@
"""ComfyUI-style handling for mixed-precision checkpoints (e.g. fp32
scale_shift tables next to bf16 linears in one file).
Two pieces, used together on a loaded module tree:
- attach_per_op_casting(root): every weighted module casts its floating-point
inputs to its own weight dtype at forward (what comfy's manual-cast ops do),
so activations promoted to fp32 by a stored-fp32 tensor drop back to the
layer's dtype at the next op instead of erroring in a bf16 matmul.
- pin_stored_fp32(root): tensors that are fp32 after load stay fp32 through
any parent-level dtype cast (`module.to(dtype=...)`, `.half()`, ...) —
device moves still apply. Without this, a holder's blanket `.to(dtype)`
would silently downcast the deliberately-fp32 pieces.
"""
import types
import torch
from torch import nn
from toolkit.util.ostris_quant import OstrisLinear
def _cast_args(args, kwargs, dtype):
def cast(t):
if torch.is_tensor(t) and t.is_floating_point() and t.dtype != dtype:
return t.to(dtype)
return t
return tuple(cast(a) for a in args), {k: cast(v) for k, v in kwargs.items()}
def attach_per_op_casting(root: nn.Module) -> int:
"""Register forward-pre-hooks casting inputs to each weighted module's own
weight dtype. Covers Linear/Conv/Norm-style modules that own a floating
``weight`` parameter, and OstrisLinear via its stored orig dtype (its
``weight`` property would materialize the dequantized tensor). Returns the
number of modules hooked."""
hooked = 0
for module in root.modules():
if isinstance(module, OstrisLinear):
def hook(mod, args, kwargs):
return _cast_args(args, kwargs, mod.ostris_orig_dtype)
else:
weight = module._parameters.get("weight", None)
if weight is None or not weight.is_floating_point():
continue
def hook(mod, args, kwargs):
return _cast_args(args, kwargs, mod._parameters["weight"].dtype)
module.register_forward_pre_hook(hook, with_kwargs=True)
hooked += 1
return hooked
def _tag_fp32(root: nn.Module):
for t in list(root.parameters()) + list(root.buffers()):
if t.is_floating_point() and t.dtype == torch.float32:
t._pin_dtype = True
if isinstance(t, nn.Parameter):
t.data._pin_dtype = True
def pin_stored_fp32(root: nn.Module):
"""Mark every currently-fp32 param/buffer in ``root`` and wrap the root's
``_apply`` so parent dtype casts skip them (device moves still apply).
The wrapper probes the cast function with a scalar to detect whether it
changes float dtypes, so ``.to(device)`` passes through untouched."""
_tag_fp32(root)
orig_apply = root._apply
def _apply(self, fn, *args, **kwargs):
probe = fn(torch.zeros((), dtype=torch.float32))
if probe.dtype == torch.float32:
return orig_apply(fn, *args, **kwargs)
def fn_pinned(t):
if getattr(t, "_pin_dtype", False):
out = t if t.device == probe.device else t.to(probe.device)
out._pin_dtype = True
return out
return fn(t)
result = orig_apply(fn_pinned, *args, **kwargs)
# _apply may rewrap tensors (dropping attributes); the invariant is
# simple — the pinned set IS the fp32 set — so re-tag
_tag_fp32(root)
return result
root._apply = types.MethodType(_apply, root)
return root