#!/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)