ClearRealityV1-CoreAI / spandrel_source.py
enginil's picture
Upload 9 files
15557f0 verified
Raw History Blame Contribute Delete
3.06 kB
#!/usr/bin/env python3.11
from __future__ import annotations
from copy import deepcopy
from pathlib import Path
from typing import Any
import torch
from spandrel import ModelLoader
def load_spandrel_model(
weights: str | Path,
*,
device: str = "cpu",
dtype: torch.dtype = torch.float32,
):
"""Load exactly the PyTorch module selected by Spandrel."""
weights = Path(weights).expanduser().resolve()
descriptor = ModelLoader().load_from_file(weights)
model = descriptor.model.eval().to(device=device, dtype=dtype)
arch = getattr(descriptor, "architecture", None)
arch_name = getattr(arch, "name", None) or type(arch).__name__
info = {
"architecture": arch_name,
"scale": int(getattr(descriptor, "scale", 4)),
"supports_half": bool(getattr(descriptor, "supports_half", False)),
"input_channels": int(getattr(descriptor, "input_channels", 3)),
"output_channels": int(getattr(descriptor, "output_channels", 3)),
}
return model, descriptor, info
def _get_parent_and_leaf(root: torch.nn.Module, qualified_name: str):
parts = qualified_name.split(".")
parent = root
for part in parts[:-1]:
parent = getattr(parent, part)
return parent, parts[-1]
def freeze_spandrel_reparameterized_convs(model: torch.nn.Module) -> tuple[torch.nn.Module, list[str]]:
"""Freeze Spandrel's eval-time Conv3XC-style reparameterization.
SPAN's Conv3XC.forward() mutates module state during every eval forward:
update_params()
eval_conv(x)
torch.export should see a static inference graph. We therefore:
1. deep-copy the exact Spandrel model;
2. ask Spandrel's OWN update_params() implementation to materialize the
fused eval_conv weights;
3. replace each reparameterizing wrapper module with that exact eval_conv.
No convolution fusion math is reimplemented here.
"""
frozen = deepcopy(model).eval()
candidates: list[str] = []
# Snapshot names first because we replace modules afterwards.
for name, module in frozen.named_modules():
if not name:
continue
if callable(getattr(module, "update_params", None)) and isinstance(
getattr(module, "eval_conv", None), torch.nn.Conv2d
):
candidates.append(name)
# IMPORTANT: no torch.inference_mode() here. Tensors created under
# inference_mode do not track version counters and break torch.export.
with torch.no_grad():
for name in candidates:
module = frozen.get_submodule(name)
module.update_params()
replaced: list[str] = []
# Replace deepest modules first.
for name in sorted(candidates, key=lambda s: s.count("."), reverse=True):
module = frozen.get_submodule(name)
static_conv = deepcopy(module.eval_conv).eval()
parent, leaf = _get_parent_and_leaf(frozen, name)
setattr(parent, leaf, static_conv)
replaced.append(name)
frozen.eval()
return frozen, sorted(replaced)