Download spandrel_source.py from enginil/ClearRealityV1-CoreAI: direct link, hf CLI and curl.
- Browser
- Download file 3.06 kB
-
https://huggingface.co/enginil/ClearRealityV1-CoreAI/resolve/main/spandrel_source.py
- Command line
-
hf download hf://enginil/ClearRealityV1-CoreAI/spandrel_source.py
-
curl -L -o spandrel_source.py https://huggingface.co/enginil/ClearRealityV1-CoreAI/resolve/main/spandrel_source.py
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) | |