Spaces:
Sleeping
Sleeping
nicholasLane commited on
Commit ·
a141401
1
Parent(s): 1fe7f79
release model and provide dependencies
Browse files- DISCLAIMER +1 -0
- GRDFNet.py +132 -0
- LICENSE +25 -0
- NEOSR/__init__.py +268 -0
- NEOSR/cachedsr_dataset.py +157 -0
- NEOSR/notes.md +2 -0
- NEOSR/train_grdfnet.toml +78 -0
- README.md +30 -4
- SPANDREL/GRDFNet.py +136 -0
- SPANDREL/__init__.py +111 -0
- SPANDREL/notes.md +4 -0
- app.py +20 -0
- requirements.txt +0 -0
DISCLAIMER
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
This project is maintained independently and is not affiliated with any organization. Commercial use requires prior written permission.
|
GRDFNet.py
ADDED
|
@@ -0,0 +1,132 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from torch import nn
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
def conv3x3(in_channels, out_channels, bias=True):
|
| 6 |
+
return nn.Conv2d(
|
| 7 |
+
in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=bias
|
| 8 |
+
)
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
class LWGRB(nn.Module):
|
| 12 |
+
def __init__(self, channels: int, bias: bool = True, identity=True):
|
| 13 |
+
super().__init__()
|
| 14 |
+
self.identity = identity
|
| 15 |
+
self.conv1 = conv3x3(channels, channels, bias)
|
| 16 |
+
self.act = nn.LeakyReLU(0.1, inplace=True)
|
| 17 |
+
self.conv2 = conv3x3(channels, channels, bias)
|
| 18 |
+
nn.init.zeros_(self.conv2.weight)
|
| 19 |
+
nn.init.zeros_(self.conv2.bias)
|
| 20 |
+
|
| 21 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 22 |
+
r = self.conv2(self.act(self.conv1(x)))
|
| 23 |
+
a = torch.sigmoid(r)
|
| 24 |
+
return x + r * a
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
class LWDRB(nn.Module):
|
| 28 |
+
def __init__(self, c, bias=True, dil=4):
|
| 29 |
+
super().__init__()
|
| 30 |
+
self.conv1 = conv3x3(c, c, bias)
|
| 31 |
+
self.act = nn.LeakyReLU(0.1, inplace=True)
|
| 32 |
+
self.conv2 = nn.Conv2d(c, c, 3, padding=dil, dilation=dil, bias=bias)
|
| 33 |
+
nn.init.zeros_(self.conv2.weight)
|
| 34 |
+
nn.init.zeros_(self.conv2.bias)
|
| 35 |
+
|
| 36 |
+
def forward(self, x):
|
| 37 |
+
r = self.conv2(self.act(self.conv1(x)))
|
| 38 |
+
return x + r
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
class LWGRBShuffle(nn.Module):
|
| 42 |
+
def __init__(self, in_ch, out_ch, scale, bias: bool = True):
|
| 43 |
+
super().__init__()
|
| 44 |
+
self.expand = conv3x3(in_ch, out_ch * scale * scale, bias)
|
| 45 |
+
self.up = nn.PixelShuffle(scale)
|
| 46 |
+
self.refine = LWGRB(out_ch, bias=bias)
|
| 47 |
+
|
| 48 |
+
def forward(self, x):
|
| 49 |
+
x = self.expand(x)
|
| 50 |
+
x = self.up(x)
|
| 51 |
+
return self.refine(x)
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class GRDFNet(nn.Module):
|
| 55 |
+
"""
|
| 56 |
+
GRDFNet (Gated + Residual Dilated Fast Network)
|
| 57 |
+
Configurable block stacking via integer `num_sets`.
|
| 58 |
+
|
| 59 |
+
Structure:
|
| 60 |
+
Stem: 3xLWGRB + 2xLWDRB
|
| 61 |
+
Repeated: num_sets x [LWGRB + LWDRB]
|
| 62 |
+
"""
|
| 63 |
+
|
| 64 |
+
def __init__(
|
| 65 |
+
self,
|
| 66 |
+
num_in_ch: int = 3,
|
| 67 |
+
num_out_ch: int = 3,
|
| 68 |
+
feature_channels: int = 32,
|
| 69 |
+
upscale: int = 1,
|
| 70 |
+
bias: bool = True,
|
| 71 |
+
norm: bool = False,
|
| 72 |
+
img_range: float = 1.0,
|
| 73 |
+
rgb_mean=(0.5, 0.5, 0.5),
|
| 74 |
+
num_sets: int = 3,
|
| 75 |
+
):
|
| 76 |
+
super().__init__()
|
| 77 |
+
|
| 78 |
+
self.in_ch = num_in_ch
|
| 79 |
+
self.out_ch = num_out_ch
|
| 80 |
+
self.c = feature_channels
|
| 81 |
+
self.scale = upscale
|
| 82 |
+
self.img_range = img_range
|
| 83 |
+
self.gamma = nn.Parameter(torch.tensor(0.5))
|
| 84 |
+
self.num_sets = num_sets
|
| 85 |
+
|
| 86 |
+
self.mean = torch.Tensor(rgb_mean).view(1, 3, 1, 1)
|
| 87 |
+
if not norm:
|
| 88 |
+
self.register_buffer("no_norm", torch.zeros(1))
|
| 89 |
+
else:
|
| 90 |
+
self.no_norm = None
|
| 91 |
+
|
| 92 |
+
self.head = conv3x3(self.in_ch, self.c, bias)
|
| 93 |
+
self.body = self._make_body(bias)
|
| 94 |
+
self.tail = conv3x3(self.c, self.out_ch, bias)
|
| 95 |
+
|
| 96 |
+
if self.scale == 1:
|
| 97 |
+
self.upsample0 = nn.Identity()
|
| 98 |
+
else:
|
| 99 |
+
self.upsample0 = LWGRBShuffle(self.out_ch, self.out_ch, self.scale)
|
| 100 |
+
|
| 101 |
+
def _make_body(self, bias: bool):
|
| 102 |
+
blocks = [
|
| 103 |
+
LWGRB(self.c, bias),
|
| 104 |
+
LWGRB(self.c, bias),
|
| 105 |
+
LWGRB(self.c, bias),
|
| 106 |
+
LWDRB(self.c, bias),
|
| 107 |
+
LWDRB(self.c, bias),
|
| 108 |
+
]
|
| 109 |
+
for _ in range(self.num_sets):
|
| 110 |
+
blocks += [LWGRB(self.c, bias), LWDRB(self.c, bias)]
|
| 111 |
+
return nn.Sequential(*blocks)
|
| 112 |
+
|
| 113 |
+
@property
|
| 114 |
+
def is_norm(self) -> bool:
|
| 115 |
+
return getattr(self, "no_norm", None) is None
|
| 116 |
+
|
| 117 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 118 |
+
if self.is_norm:
|
| 119 |
+
self.mean = self.mean.type_as(x)
|
| 120 |
+
x = (x - self.mean) * self.img_range
|
| 121 |
+
|
| 122 |
+
feat = self.head(x)
|
| 123 |
+
feat = self.body(feat)
|
| 124 |
+
out_feat = self.tail(feat)
|
| 125 |
+
out_feat = (1.0 - self.gamma) * x + self.gamma * out_feat
|
| 126 |
+
|
| 127 |
+
out = self.upsample0(out_feat)
|
| 128 |
+
|
| 129 |
+
if self.is_norm:
|
| 130 |
+
out = out / self.img_range + self.mean
|
| 131 |
+
|
| 132 |
+
return out
|
LICENSE
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
MIT License (With Commercial Use Restriction)
|
| 2 |
+
|
| 3 |
+
Copyright (c) 2025 Nicholas Lane
|
| 4 |
+
|
| 5 |
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 6 |
+
of this software and associated documentation files (the "Software"), to deal
|
| 7 |
+
in the Software without restriction, including without limitation the rights
|
| 8 |
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 9 |
+
copies of the Software, and to permit persons to whom the Software is
|
| 10 |
+
furnished to do so, subject to the following conditions:
|
| 11 |
+
|
| 12 |
+
1. The above copyright notice and this permission notice (including the
|
| 13 |
+
commercial use restriction below) shall be included in all copies or
|
| 14 |
+
substantial portions of the Software.
|
| 15 |
+
2. Commercial use of the Software, including any derivative works, is
|
| 16 |
+
prohibited unless prior explicit written permission is granted by the
|
| 17 |
+
copyright holder.
|
| 18 |
+
|
| 19 |
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 20 |
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 21 |
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 22 |
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 23 |
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 24 |
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 25 |
+
SOFTWARE.
|
NEOSR/__init__.py
ADDED
|
@@ -0,0 +1,268 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import importlib
|
| 2 |
+
import os
|
| 3 |
+
import random
|
| 4 |
+
from copy import deepcopy
|
| 5 |
+
from functools import partial
|
| 6 |
+
from pathlib import Path
|
| 7 |
+
from typing import Any
|
| 8 |
+
|
| 9 |
+
import numpy as np
|
| 10 |
+
import torch
|
| 11 |
+
import torch.utils.data
|
| 12 |
+
from torch.utils import data
|
| 13 |
+
from torch.utils.data.sampler import Sampler
|
| 14 |
+
|
| 15 |
+
from neosr.utils import get_root_logger, scandir
|
| 16 |
+
from neosr.utils.dist_util import get_dist_info
|
| 17 |
+
from neosr.utils.registry import DATASET_REGISTRY
|
| 18 |
+
|
| 19 |
+
__all__ = ["build_dataloader", "build_dataset"]
|
| 20 |
+
|
| 21 |
+
# automatically scan and import dataset modules for registry
|
| 22 |
+
# scan all the files under the data folder with '_dataset' in file names
|
| 23 |
+
data_folder = Path(Path(__file__).resolve()).parent
|
| 24 |
+
dataset_filenames = [
|
| 25 |
+
Path(Path(v).name).stem
|
| 26 |
+
for v in scandir(str(data_folder))
|
| 27 |
+
if v.endswith("_dataset.py")
|
| 28 |
+
]
|
| 29 |
+
|
| 30 |
+
def build_dataset(dataset_opt: dict[str, Any]):
|
| 31 |
+
"""Build dataset from options.
|
| 32 |
+
|
| 33 |
+
Args:
|
| 34 |
+
----
|
| 35 |
+
dataset_opt (dict): Configuration for dataset. It must contain:
|
| 36 |
+
type (str): Dataset type.
|
| 37 |
+
|
| 38 |
+
"""
|
| 39 |
+
dataset_opt = deepcopy(dataset_opt)
|
| 40 |
+
dataset_type = dataset_opt["type"]
|
| 41 |
+
dataset_cls = DATASET_REGISTRY.get(dataset_type) # type: ignore[assignment]
|
| 42 |
+
logger = get_root_logger()
|
| 43 |
+
|
| 44 |
+
def _maybe_cache_to_ram() -> None:
|
| 45 |
+
default_cache = str(dataset_type).lower().startswith("cached")
|
| 46 |
+
cache_requested = dataset_opt.get("cache_in_ram", default_cache)
|
| 47 |
+
dataroot_lq = dataset_opt.get("dataroot_lq")
|
| 48 |
+
dataroot_gt = dataset_opt.get("dataroot_gt")
|
| 49 |
+
if not cache_requested:
|
| 50 |
+
logger.debug(
|
| 51 |
+
"Skipping RAM cache preload for dataset '%s' (cache_in_ram disabled).",
|
| 52 |
+
dataset_type,
|
| 53 |
+
)
|
| 54 |
+
return
|
| 55 |
+
if not dataroot_lq or not dataroot_gt:
|
| 56 |
+
logger.warning(
|
| 57 |
+
"cache_in_ram enabled for dataset '%s' but dataroots are missing. Skipping preload.",
|
| 58 |
+
dataset_type,
|
| 59 |
+
)
|
| 60 |
+
return
|
| 61 |
+
|
| 62 |
+
lq_root = Path(dataroot_lq)
|
| 63 |
+
gt_root = Path(dataroot_gt)
|
| 64 |
+
if not lq_root.exists() or not gt_root.exists():
|
| 65 |
+
logger.warning(
|
| 66 |
+
"cache_in_ram enabled for dataset '%s' but dataroots do not exist. "
|
| 67 |
+
"LQ: %s, GT: %s. Skipping preload.",
|
| 68 |
+
dataset_type,
|
| 69 |
+
dataroot_lq,
|
| 70 |
+
dataroot_gt,
|
| 71 |
+
)
|
| 72 |
+
return
|
| 73 |
+
|
| 74 |
+
try:
|
| 75 |
+
from cv2 import IMREAD_UNCHANGED, imread # type: ignore[attr-defined]
|
| 76 |
+
except ImportError as exc: # pragma: no cover - OpenCV should be available
|
| 77 |
+
logger.warning(
|
| 78 |
+
"cache_in_ram enabled for dataset '%s' but OpenCV is unavailable: %s",
|
| 79 |
+
dataset_type,
|
| 80 |
+
exc,
|
| 81 |
+
)
|
| 82 |
+
return
|
| 83 |
+
|
| 84 |
+
def collect_files(root: Path) -> dict[str, Path]:
|
| 85 |
+
return {
|
| 86 |
+
str(path.relative_to(root)).replace("\\", "/"): path
|
| 87 |
+
for path in root.rglob("*")
|
| 88 |
+
if path.is_file()
|
| 89 |
+
}
|
| 90 |
+
|
| 91 |
+
lq_files = collect_files(lq_root)
|
| 92 |
+
gt_files = collect_files(gt_root)
|
| 93 |
+
common_rel_paths = sorted(set(lq_files) & set(gt_files))
|
| 94 |
+
if not common_rel_paths:
|
| 95 |
+
logger.warning(
|
| 96 |
+
"cache_in_ram enabled for dataset '%s' but no matching LQ/GT pairs were found.",
|
| 97 |
+
dataset_type,
|
| 98 |
+
)
|
| 99 |
+
return
|
| 100 |
+
|
| 101 |
+
if len(common_rel_paths) != len(lq_files) or len(common_rel_paths) != len(gt_files):
|
| 102 |
+
logger.info(
|
| 103 |
+
"cache_in_ram: restricting to %d matched pairs (dropped %d LQ, %d GT).",
|
| 104 |
+
len(common_rel_paths),
|
| 105 |
+
len(lq_files) - len(common_rel_paths),
|
| 106 |
+
len(gt_files) - len(common_rel_paths),
|
| 107 |
+
)
|
| 108 |
+
|
| 109 |
+
lq_images: list[np.ndarray] = []
|
| 110 |
+
gt_images: list[np.ndarray] = []
|
| 111 |
+
total_bytes = 0
|
| 112 |
+
for rel_path in common_rel_paths:
|
| 113 |
+
lq_img = imread(str(lq_files[rel_path]), IMREAD_UNCHANGED)
|
| 114 |
+
gt_img = imread(str(gt_files[rel_path]), IMREAD_UNCHANGED)
|
| 115 |
+
if lq_img is None or gt_img is None:
|
| 116 |
+
logger.warning(
|
| 117 |
+
"cache_in_ram: failed to read pair '%s'. Aborting preload.",
|
| 118 |
+
rel_path,
|
| 119 |
+
)
|
| 120 |
+
lq_images.clear()
|
| 121 |
+
gt_images.clear()
|
| 122 |
+
break
|
| 123 |
+
lq_img = np.ascontiguousarray(lq_img).astype(np.float32) / 255.0
|
| 124 |
+
gt_img = np.ascontiguousarray(gt_img).astype(np.float32) / 255.0
|
| 125 |
+
lq_images.append(lq_img)
|
| 126 |
+
gt_images.append(gt_img)
|
| 127 |
+
total_bytes += lq_img.nbytes + gt_img.nbytes
|
| 128 |
+
if not lq_images or not gt_images:
|
| 129 |
+
return
|
| 130 |
+
|
| 131 |
+
dataset_opt["_cached_lq_images"] = lq_images
|
| 132 |
+
dataset_opt["_cached_gt_images"] = gt_images
|
| 133 |
+
dataset_opt["_cached_rel_paths"] = common_rel_paths
|
| 134 |
+
dataset_opt["_cached_lq_paths"] = [str(lq_files[p]) for p in common_rel_paths]
|
| 135 |
+
dataset_opt["_cached_gt_paths"] = [str(gt_files[p]) for p in common_rel_paths]
|
| 136 |
+
dataset_opt.setdefault("scale", dataset_opt.get("scale", 1))
|
| 137 |
+
logger.info(
|
| 138 |
+
"Preloaded %d image pairs for dataset '%s' into RAM (approx %.2f GiB).",
|
| 139 |
+
len(common_rel_paths),
|
| 140 |
+
dataset_type,
|
| 141 |
+
total_bytes / (1024**3),
|
| 142 |
+
)
|
| 143 |
+
|
| 144 |
+
_maybe_cache_to_ram()
|
| 145 |
+
dataset = dataset_cls(dataset_opt) # type: ignore[operator]
|
| 146 |
+
cache_active = getattr(dataset, "cache_in_ram_active", False) or bool(
|
| 147 |
+
dataset_opt.get("_cached_lq_images")
|
| 148 |
+
)
|
| 149 |
+
if cache_active:
|
| 150 |
+
setattr(dataset, "cache_in_ram_active", True)
|
| 151 |
+
logger.info(f"Dataset [{dataset.__class__.__name__}] is built.")
|
| 152 |
+
return dataset
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
def worker_init_fn(worker_id: int, num_workers: int, rank: int, seed: int) -> None:
|
| 156 |
+
# Set the worker seed to num_workers * rank + worker_id + seed
|
| 157 |
+
worker_seed = num_workers * rank + worker_id + seed
|
| 158 |
+
# NOTE: set seed on old generator as a precaution, but
|
| 159 |
+
# it is redundand since we use np.random.Generator
|
| 160 |
+
np.random.seed(worker_seed) # noqa: NPY002
|
| 161 |
+
torch.manual_seed(worker_seed)
|
| 162 |
+
random.seed(worker_seed)
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
def build_dataloader(
|
| 166 |
+
dataset: data.Dataset,
|
| 167 |
+
dataset_opt: dict[str, Any],
|
| 168 |
+
num_gpu: int = 1,
|
| 169 |
+
dist: bool = False,
|
| 170 |
+
sampler: Sampler | None = None,
|
| 171 |
+
seed: int | None = None,
|
| 172 |
+
) -> data.DataLoader:
|
| 173 |
+
"""Build dataloader.
|
| 174 |
+
|
| 175 |
+
Args:
|
| 176 |
+
----
|
| 177 |
+
dataset (torch.utils.data.Dataset): Dataset.
|
| 178 |
+
dataset_opt (dict): Dataset options. It contains the following keys:
|
| 179 |
+
phase (str): 'train' or 'val'.
|
| 180 |
+
num_worker_per_gpu (int): Number of workers for each GPU.
|
| 181 |
+
batch_size (int): Training batch size for each GPU.
|
| 182 |
+
num_gpu (int): Number of GPUs. Used only in the train phase.
|
| 183 |
+
Default: 1.
|
| 184 |
+
dist (bool): Whether in distributed training. Used only in the train
|
| 185 |
+
phase. Default: False.
|
| 186 |
+
sampler (torch.utils.data.sampler): Data sampler. Default: None.
|
| 187 |
+
seed (int | None): Seed. Default: None
|
| 188 |
+
|
| 189 |
+
"""
|
| 190 |
+
phase = dataset_opt["phase"]
|
| 191 |
+
rank, _ = get_dist_info()
|
| 192 |
+
logger = get_root_logger()
|
| 193 |
+
|
| 194 |
+
# train
|
| 195 |
+
if phase == "train":
|
| 196 |
+
cpu_count = os.cpu_count() or 1
|
| 197 |
+
workers_per_gpu_opt = dataset_opt.get("num_worker_per_gpu")
|
| 198 |
+
# Heuristic tuned for high-end hosts: keep workers high enough to saturate I/O.
|
| 199 |
+
if workers_per_gpu_opt is None or workers_per_gpu_opt == "auto":
|
| 200 |
+
gpu_slots = num_gpu if num_gpu > 0 else 1
|
| 201 |
+
workers_per_gpu = max(4, min(32, cpu_count // gpu_slots))
|
| 202 |
+
else:
|
| 203 |
+
workers_per_gpu = int(workers_per_gpu_opt)
|
| 204 |
+
workers_per_gpu = max(workers_per_gpu, 1)
|
| 205 |
+
num_workers = workers_per_gpu
|
| 206 |
+
|
| 207 |
+
if dist: # distributed training
|
| 208 |
+
batch_size = dataset_opt["batch_size"]
|
| 209 |
+
else: # non-distributed training
|
| 210 |
+
multiplier = 1 if num_gpu == 0 else num_gpu
|
| 211 |
+
batch_size = dataset_opt["batch_size"] * multiplier
|
| 212 |
+
num_workers *= multiplier
|
| 213 |
+
|
| 214 |
+
if "prefetch_factor" in dataset_opt:
|
| 215 |
+
prefetch_factor = dataset_opt["prefetch_factor"]
|
| 216 |
+
else:
|
| 217 |
+
# High worker counts benefit from deeper prefetch queues.
|
| 218 |
+
prefetch_factor = max(4, min(16, workers_per_gpu))
|
| 219 |
+
|
| 220 |
+
cache_in_ram_active = getattr(dataset, "cache_in_ram_active", False)
|
| 221 |
+
if os.name == "nt" and cache_in_ram_active and num_workers > 0:
|
| 222 |
+
logger.warning(
|
| 223 |
+
"Windows detected with RAM-cached dataset '%s'; forcing num_workers from %d to 0 "
|
| 224 |
+
"to avoid spawn pickling overhead.",
|
| 225 |
+
dataset.__class__.__name__,
|
| 226 |
+
num_workers,
|
| 227 |
+
)
|
| 228 |
+
num_workers = 0
|
| 229 |
+
|
| 230 |
+
dataloader_args = {
|
| 231 |
+
"dataset": dataset,
|
| 232 |
+
"batch_size": batch_size,
|
| 233 |
+
"shuffle": False,
|
| 234 |
+
"num_workers": num_workers,
|
| 235 |
+
"sampler": sampler,
|
| 236 |
+
"drop_last": True,
|
| 237 |
+
}
|
| 238 |
+
if num_workers > 0:
|
| 239 |
+
dataloader_args["prefetch_factor"] = prefetch_factor
|
| 240 |
+
if sampler is None:
|
| 241 |
+
dataloader_args["shuffle"] = True
|
| 242 |
+
dataloader_args["worker_init_fn"] = (
|
| 243 |
+
partial(worker_init_fn, num_workers=num_workers, rank=rank, seed=seed)
|
| 244 |
+
if seed is not None
|
| 245 |
+
else None
|
| 246 |
+
)
|
| 247 |
+
|
| 248 |
+
# val
|
| 249 |
+
elif phase in {"val", "test"}:
|
| 250 |
+
dataloader_args = {
|
| 251 |
+
"dataset": dataset,
|
| 252 |
+
"batch_size": 1,
|
| 253 |
+
"shuffle": False,
|
| 254 |
+
"num_workers": 0,
|
| 255 |
+
}
|
| 256 |
+
else:
|
| 257 |
+
msg = f"Wrong dataset phase: {phase}. Supported ones are 'train', 'val' and 'test'."
|
| 258 |
+
raise ValueError(msg)
|
| 259 |
+
|
| 260 |
+
dataloader_args["pin_memory"] = dataset_opt.get("pin_memory", True)
|
| 261 |
+
if "persistent_workers" in dataset_opt:
|
| 262 |
+
dataloader_args["persistent_workers"] = dataset_opt["persistent_workers"]
|
| 263 |
+
else:
|
| 264 |
+
dataloader_args["persistent_workers"] = dataloader_args.get("num_workers", 0) > 0
|
| 265 |
+
if dataloader_args["persistent_workers"] and dataloader_args["num_workers"] == 0:
|
| 266 |
+
dataloader_args["persistent_workers"] = False
|
| 267 |
+
|
| 268 |
+
return data.DataLoader(**dataloader_args)
|
NEOSR/cachedsr_dataset.py
ADDED
|
@@ -0,0 +1,157 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pathlib import Path
|
| 2 |
+
from typing import Any
|
| 3 |
+
|
| 4 |
+
import cv2
|
| 5 |
+
import numpy as np
|
| 6 |
+
from torch.utils import data
|
| 7 |
+
from torchvision.transforms.functional import normalize
|
| 8 |
+
|
| 9 |
+
from neosr.data.transforms import basic_augment, paired_random_crop
|
| 10 |
+
from neosr.utils import get_root_logger, img2tensor
|
| 11 |
+
from neosr.utils.registry import DATASET_REGISTRY
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
@DATASET_REGISTRY.register()
|
| 15 |
+
class CachedSRDataset(data.Dataset):
|
| 16 |
+
"""Paired SR dataset that keeps image pairs in host RAM for fast sampling."""
|
| 17 |
+
|
| 18 |
+
def __init__(self, opt: dict[str, Any]) -> None:
|
| 19 |
+
super().__init__()
|
| 20 |
+
self.opt = opt
|
| 21 |
+
self.logger = get_root_logger()
|
| 22 |
+
self.scale = opt.get("scale", 1)
|
| 23 |
+
self.patch_size = opt.get("patch_size")
|
| 24 |
+
self.phase = opt.get("phase", "train")
|
| 25 |
+
self.mean = opt.get("mean")
|
| 26 |
+
self.std = opt.get("std")
|
| 27 |
+
self.color = opt.get("color", None) != "y"
|
| 28 |
+
self.use_hflip = opt.get("use_hflip", True)
|
| 29 |
+
self.use_rot = opt.get("use_rot", True)
|
| 30 |
+
|
| 31 |
+
cached_lq = opt.get("_cached_lq_images")
|
| 32 |
+
cached_gt = opt.get("_cached_gt_images")
|
| 33 |
+
cached_lq_paths = opt.get("_cached_lq_paths")
|
| 34 |
+
cached_gt_paths = opt.get("_cached_gt_paths")
|
| 35 |
+
cached_rel = opt.get("_cached_rel_paths")
|
| 36 |
+
|
| 37 |
+
self.cache_in_ram_active = False
|
| 38 |
+
if cached_lq and cached_gt:
|
| 39 |
+
self.lq_images = cached_lq
|
| 40 |
+
self.gt_images = cached_gt
|
| 41 |
+
self.lq_paths = cached_lq_paths or [
|
| 42 |
+
f"{opt.get('dataroot_lq', 'lq')}::{idx}" for idx in range(len(self.lq_images))
|
| 43 |
+
]
|
| 44 |
+
self.gt_paths = cached_gt_paths or [
|
| 45 |
+
f"{opt.get('dataroot_gt', 'gt')}::{idx}" for idx in range(len(self.gt_images))
|
| 46 |
+
]
|
| 47 |
+
self.rel_paths = cached_rel or [str(idx) for idx in range(len(self.lq_images))]
|
| 48 |
+
self.cache_in_ram_active = True
|
| 49 |
+
self.logger.info(
|
| 50 |
+
"CachedSRDataset reusing %d preloaded pairs from build_dataset cache.",
|
| 51 |
+
len(self.lq_images),
|
| 52 |
+
)
|
| 53 |
+
else:
|
| 54 |
+
self.logger.info("CachedSRDataset preloading images locally (no shared cache provided).")
|
| 55 |
+
dataroot_lq = opt.get("dataroot_lq")
|
| 56 |
+
dataroot_gt = opt.get("dataroot_gt")
|
| 57 |
+
if not dataroot_lq or not dataroot_gt:
|
| 58 |
+
msg = "CachedSRDataset requires dataroot_lq and dataroot_gt when cache is absent."
|
| 59 |
+
raise ValueError(msg)
|
| 60 |
+
self._load_from_disk(dataroot_lq, dataroot_gt)
|
| 61 |
+
self.cache_in_ram_active = True
|
| 62 |
+
|
| 63 |
+
assert len(self.lq_images) == len(self.gt_images), "Unmatched LQ/GT counts."
|
| 64 |
+
self.length = len(self.lq_images)
|
| 65 |
+
|
| 66 |
+
def _load_from_disk(self, dataroot_lq: str, dataroot_gt: str) -> None:
|
| 67 |
+
lq_root = Path(dataroot_lq)
|
| 68 |
+
gt_root = Path(dataroot_gt)
|
| 69 |
+
if not lq_root.exists() or not gt_root.exists():
|
| 70 |
+
msg = f"CachedSRDataset dataroots do not exist. LQ: {dataroot_lq}, GT: {dataroot_gt}"
|
| 71 |
+
raise FileNotFoundError(msg)
|
| 72 |
+
|
| 73 |
+
def collect(root: Path) -> dict[str, Path]:
|
| 74 |
+
return {
|
| 75 |
+
str(path.relative_to(root)).replace("\\", "/"): path
|
| 76 |
+
for path in root.rglob("*")
|
| 77 |
+
if path.is_file()
|
| 78 |
+
}
|
| 79 |
+
|
| 80 |
+
lq_files = collect(lq_root)
|
| 81 |
+
gt_files = collect(gt_root)
|
| 82 |
+
keys = sorted(set(lq_files) & set(gt_files))
|
| 83 |
+
if not keys:
|
| 84 |
+
msg = "CachedSRDataset could not find any matched LQ/GT pairs."
|
| 85 |
+
raise RuntimeError(msg)
|
| 86 |
+
dropped_lq = len(lq_files) - len(keys)
|
| 87 |
+
dropped_gt = len(gt_files) - len(keys)
|
| 88 |
+
if dropped_lq or dropped_gt:
|
| 89 |
+
self.logger.warning(
|
| 90 |
+
"CachedSRDataset dropping unmatched pairs (LQ: -%d, GT: -%d).",
|
| 91 |
+
dropped_lq,
|
| 92 |
+
dropped_gt,
|
| 93 |
+
)
|
| 94 |
+
|
| 95 |
+
self.lq_images = []
|
| 96 |
+
self.gt_images = []
|
| 97 |
+
self.lq_paths = []
|
| 98 |
+
self.gt_paths = []
|
| 99 |
+
self.rel_paths = keys
|
| 100 |
+
|
| 101 |
+
for rel in keys:
|
| 102 |
+
lq_img = cv2.imread(str(lq_files[rel]), cv2.IMREAD_UNCHANGED)
|
| 103 |
+
gt_img = cv2.imread(str(gt_files[rel]), cv2.IMREAD_UNCHANGED)
|
| 104 |
+
if lq_img is None or gt_img is None:
|
| 105 |
+
msg = f"Failed to read cached pair {rel}."
|
| 106 |
+
raise RuntimeError(msg)
|
| 107 |
+
lq_img = np.ascontiguousarray(lq_img).astype(np.float32) / 255.0
|
| 108 |
+
gt_img = np.ascontiguousarray(gt_img).astype(np.float32) / 255.0
|
| 109 |
+
self.lq_images.append(lq_img)
|
| 110 |
+
self.gt_images.append(gt_img)
|
| 111 |
+
self.lq_paths.append(str(lq_files[rel]))
|
| 112 |
+
self.gt_paths.append(str(gt_files[rel]))
|
| 113 |
+
|
| 114 |
+
def __len__(self) -> int:
|
| 115 |
+
return self.length
|
| 116 |
+
|
| 117 |
+
def __getitem__(self, index: int) -> dict[str, Any]:
|
| 118 |
+
lq_img = self.lq_images[index]
|
| 119 |
+
gt_img = self.gt_images[index]
|
| 120 |
+
gt_path = self.gt_paths[index] if index < len(self.gt_paths) else None
|
| 121 |
+
lq_path = self.lq_paths[index] if index < len(self.lq_paths) else None
|
| 122 |
+
|
| 123 |
+
# Training branch: random crop + augment
|
| 124 |
+
if self.phase == "train" and self.patch_size:
|
| 125 |
+
gt_img, lq_img = paired_random_crop(
|
| 126 |
+
gt_img, lq_img, self.patch_size, self.scale, gt_path or self.rel_paths[index]
|
| 127 |
+
)
|
| 128 |
+
gt_img = np.ascontiguousarray(gt_img)
|
| 129 |
+
lq_img = np.ascontiguousarray(lq_img)
|
| 130 |
+
gt_img, lq_img = basic_augment(
|
| 131 |
+
[gt_img, lq_img],
|
| 132 |
+
hflip=self.use_hflip,
|
| 133 |
+
rotation=self.use_rot,
|
| 134 |
+
) # type: ignore[reportAssignmentType]
|
| 135 |
+
gt_img = np.ascontiguousarray(gt_img)
|
| 136 |
+
lq_img = np.ascontiguousarray(lq_img)
|
| 137 |
+
else:
|
| 138 |
+
# Align GT to LQ size when evaluating
|
| 139 |
+
h_lq, w_lq = lq_img.shape[:2]
|
| 140 |
+
gt_img = gt_img[: h_lq * self.scale, : w_lq * self.scale, ...]
|
| 141 |
+
gt_img = np.ascontiguousarray(gt_img)
|
| 142 |
+
lq_img = np.ascontiguousarray(lq_img)
|
| 143 |
+
|
| 144 |
+
gt_tensor, lq_tensor = img2tensor(
|
| 145 |
+
[gt_img, lq_img], bgr2rgb=True, float32=True, color=self.color
|
| 146 |
+
)
|
| 147 |
+
|
| 148 |
+
if self.mean is not None or self.std is not None:
|
| 149 |
+
normalize(lq_tensor, self.mean, self.std, inplace=True) # type: ignore[arg-type]
|
| 150 |
+
normalize(gt_tensor, self.mean, self.std, inplace=True) # type: ignore[arg-type]
|
| 151 |
+
|
| 152 |
+
return {
|
| 153 |
+
"lq": lq_tensor,
|
| 154 |
+
"gt": gt_tensor,
|
| 155 |
+
"lq_path": lq_path or self.rel_paths[index],
|
| 156 |
+
"gt_path": gt_path or self.rel_paths[index],
|
| 157 |
+
}
|
NEOSR/notes.md
ADDED
|
@@ -0,0 +1,2 @@
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Add cachedsr_dataset and update `__init__.py` in NEOSR for faster training if sytsem RAM allows.
|
| 2 |
+
Minimal training TOML provided for non-GAN training.
|
NEOSR/train_grdfnet.toml
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
|
| 2 |
+
name = 'NAME-Here'
|
| 3 |
+
model_type = 'image'
|
| 4 |
+
scale = 2
|
| 5 |
+
use_amp = true
|
| 6 |
+
bfloat16 = true
|
| 7 |
+
fast_matmul = true
|
| 8 |
+
[datasets.train]
|
| 9 |
+
type = 'CachedSRDataset'
|
| 10 |
+
dataroot_gt = 'GT IMAGES'
|
| 11 |
+
dataroot_lq = 'LQ IMAGES'
|
| 12 |
+
patch_size = 96
|
| 13 |
+
batch_size = 32
|
| 14 |
+
accumulate = 1
|
| 15 |
+
augmentation = [ 'resizemix', 'cutblur'] # 'mixup', 'cutmix', 'resizemix', 'cutblur'
|
| 16 |
+
aug_prob = [ 0.25, 0.15 ]
|
| 17 |
+
use_hflip = true
|
| 18 |
+
use_rot = true
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
[datasets.val]
|
| 22 |
+
name = 'val'
|
| 23 |
+
type = 'paired'
|
| 24 |
+
dataroot_gt = 'GT IMAGES'
|
| 25 |
+
dataroot_lq = 'LQ IMAGES'
|
| 26 |
+
|
| 27 |
+
[val]
|
| 28 |
+
val_freq = 5000
|
| 29 |
+
[val.metrics.psnr]
|
| 30 |
+
type = 'calculate_psnr'
|
| 31 |
+
|
| 32 |
+
[train]
|
| 33 |
+
ema = 0.995
|
| 34 |
+
match_lq_colors = false
|
| 35 |
+
|
| 36 |
+
[network_g]
|
| 37 |
+
type = 'GRDFNet'
|
| 38 |
+
num_sets = 3
|
| 39 |
+
feature_channels = 32
|
| 40 |
+
|
| 41 |
+
[train.scheduler]
|
| 42 |
+
type = 'cosineannealing'
|
| 43 |
+
T_max = 150000
|
| 44 |
+
eta_min = 1e-9
|
| 45 |
+
|
| 46 |
+
[train.optim_g]
|
| 47 |
+
type = 'adam'
|
| 48 |
+
lr = 1e-4
|
| 49 |
+
betas = [0.9, 0.999]
|
| 50 |
+
weight_decay = 5e-7
|
| 51 |
+
|
| 52 |
+
[train.pixel_opt]
|
| 53 |
+
type = 'HuberLoss'
|
| 54 |
+
loss_weight = 1.0
|
| 55 |
+
reduction = 'mean'
|
| 56 |
+
|
| 57 |
+
[train.gw_opt]
|
| 58 |
+
type = "gw_loss"
|
| 59 |
+
loss_weight = 0.15
|
| 60 |
+
criterion = "chc_loss"
|
| 61 |
+
corner = true
|
| 62 |
+
|
| 63 |
+
[train.consistency_opt]
|
| 64 |
+
type = 'consistency_loss'
|
| 65 |
+
loss_weight = 0.25
|
| 66 |
+
criterion = 'chc' # 'l1'
|
| 67 |
+
blur = true
|
| 68 |
+
cosim = true
|
| 69 |
+
saturation = 1.0
|
| 70 |
+
brightness = 1.0
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
[logger]
|
| 74 |
+
total_iter = 150000
|
| 75 |
+
save_checkpoint_freq = 5000
|
| 76 |
+
use_tb_logger = true
|
| 77 |
+
#save_tb_img = true
|
| 78 |
+
print_freq = 500
|
README.md
CHANGED
|
@@ -1,14 +1,40 @@
|
|
| 1 |
---
|
| 2 |
title: GRDFNet
|
| 3 |
-
emoji: 🚀
|
| 4 |
colorFrom: red
|
| 5 |
colorTo: yellow
|
| 6 |
sdk: gradio
|
| 7 |
sdk_version: 5.44.1
|
| 8 |
app_file: app.py
|
| 9 |
-
|
| 10 |
-
license: other
|
| 11 |
short_description: A lightweight image restoration architecture
|
|
|
|
| 12 |
---
|
| 13 |
|
| 14 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
---
|
| 2 |
title: GRDFNet
|
|
|
|
| 3 |
colorFrom: red
|
| 4 |
colorTo: yellow
|
| 5 |
sdk: gradio
|
| 6 |
sdk_version: 5.44.1
|
| 7 |
app_file: app.py
|
| 8 |
+
license: mit
|
|
|
|
| 9 |
short_description: A lightweight image restoration architecture
|
| 10 |
+
pinned: true
|
| 11 |
---
|
| 12 |
|
| 13 |
+
# GRDFNet
|
| 14 |
+
|
| 15 |
+
GRDFNet is a lightweight image restoration network that combines gated and dilated residual blocks to deliver strong perceptual quality with modest compute requirements.
|
| 16 |
+
|
| 17 |
+
## Recommended Configurations
|
| 18 |
+
|
| 19 |
+
- `num_sets = 3`, `feature_channels = 32`: strong quality while staying fast for most desktop workloads.
|
| 20 |
+
- `num_sets = 6`, `feature_channels = 48`: highest quality configuration; expect roughly a 4x slowdown versus the 32-channel model.
|
| 21 |
+
- `num_sets = 3`, `feature_channels = 24`: suggested for lightly compressed video inference; typically 50~75% faster than the 32-channel variant when deployed with TensorRT.
|
| 22 |
+
|
| 23 |
+
## Performance Snapshot
|
| 24 |
+
|
| 25 |
+
Example TensorRT run on an NVIDIA RTX 4080 Super (16 GB):
|
| 26 |
+
```
|
| 27 |
+
DEBUG: TensorRT initialized. Setting shape.
|
| 28 |
+
DEBUG: Shape set. Getting output shape.
|
| 29 |
+
[INFO] Input: 1280x720 -> 1280x720 -> ModelOut: 1280x720 @ 30000/1001 fps
|
| 30 |
+
DEBUG: Before NVENC initialization.
|
| 31 |
+
[prof] frames=209 avg=208.2 fps
|
| 32 |
+
[prof] frames=431 avg=215.1 fps
|
| 33 |
+
[prof] frames=654 avg=217.3 fps
|
| 34 |
+
[INFO] Processed 709 frames in 3.280s -> 216.2 FPS
|
| 35 |
+
```
|
| 36 |
+
|
| 37 |
+
## Resources
|
| 38 |
+
|
| 39 |
+
- Model weights: https://huggingface.co/nicholasLane/GRDFNet
|
| 40 |
+
- Hosted demo: https://huggingface.co/spaces/nicholasLane/GRDFNet
|
SPANDREL/GRDFNet.py
ADDED
|
@@ -0,0 +1,136 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
from torch import nn as nn
|
| 5 |
+
|
| 6 |
+
from spandrel.util import store_hyperparameters
|
| 7 |
+
|
| 8 |
+
def conv3x3(in_channels, out_channels, bias=True):
|
| 9 |
+
return nn.Conv2d(
|
| 10 |
+
in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=bias
|
| 11 |
+
)
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class LWGRB(nn.Module):
|
| 15 |
+
def __init__(self, channels: int, bias: bool = True, identity=True):
|
| 16 |
+
super().__init__()
|
| 17 |
+
self.identity = identity
|
| 18 |
+
self.conv1 = conv3x3(channels, channels, bias)
|
| 19 |
+
self.act = nn.LeakyReLU(0.1, inplace=True)
|
| 20 |
+
self.conv2 = conv3x3(channels, channels, bias)
|
| 21 |
+
nn.init.zeros_(self.conv2.weight)
|
| 22 |
+
nn.init.zeros_(self.conv2.bias)
|
| 23 |
+
|
| 24 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 25 |
+
r = self.conv2(self.act(self.conv1(x)))
|
| 26 |
+
a = torch.sigmoid(r)
|
| 27 |
+
return x + r * a
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
class LWDRB(nn.Module):
|
| 31 |
+
def __init__(self, c, bias=True, dil=4):
|
| 32 |
+
super().__init__()
|
| 33 |
+
self.conv1 = conv3x3(c, c, bias)
|
| 34 |
+
self.act = nn.LeakyReLU(0.1, inplace=True)
|
| 35 |
+
self.conv2 = nn.Conv2d(c, c, 3, padding=dil, dilation=dil, bias=bias)
|
| 36 |
+
nn.init.zeros_(self.conv2.weight)
|
| 37 |
+
nn.init.zeros_(self.conv2.bias)
|
| 38 |
+
|
| 39 |
+
def forward(self, x):
|
| 40 |
+
r = self.conv2(self.act(self.conv1(x)))
|
| 41 |
+
return x + r
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
class LWGRBShuffle(nn.Module):
|
| 45 |
+
def __init__(self, in_ch, out_ch, scale, bias: bool = True):
|
| 46 |
+
super().__init__()
|
| 47 |
+
self.expand = conv3x3(in_ch, out_ch * scale * scale, bias)
|
| 48 |
+
self.up = nn.PixelShuffle(scale)
|
| 49 |
+
self.refine = LWGRB(out_ch, bias=bias)
|
| 50 |
+
|
| 51 |
+
def forward(self, x):
|
| 52 |
+
x = self.expand(x)
|
| 53 |
+
x = self.up(x)
|
| 54 |
+
return self.refine(x)
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
@store_hyperparameters()
|
| 58 |
+
class GRDFNet(nn.Module):
|
| 59 |
+
"""
|
| 60 |
+
GRDFNet (Gated + Residual Dilated Fast Network)
|
| 61 |
+
Configurable block stacking via integer `num_sets`.
|
| 62 |
+
|
| 63 |
+
Structure:
|
| 64 |
+
Stem: 3xLWGRB + 2xLWDRB
|
| 65 |
+
Repeated: num_sets x [LWGRB + LWDRB]
|
| 66 |
+
"""
|
| 67 |
+
|
| 68 |
+
def __init__(
|
| 69 |
+
self,
|
| 70 |
+
num_in_ch: int = 3,
|
| 71 |
+
num_out_ch: int = 3,
|
| 72 |
+
feature_channels: int = 32,
|
| 73 |
+
upscale: int = 1,
|
| 74 |
+
bias: bool = True,
|
| 75 |
+
norm: bool = False,
|
| 76 |
+
img_range: float = 1.0,
|
| 77 |
+
rgb_mean=(0.5, 0.5, 0.5),
|
| 78 |
+
num_sets: int = 3,
|
| 79 |
+
):
|
| 80 |
+
super().__init__()
|
| 81 |
+
|
| 82 |
+
self.in_ch = num_in_ch
|
| 83 |
+
self.out_ch = num_out_ch
|
| 84 |
+
self.c = feature_channels
|
| 85 |
+
self.scale = upscale
|
| 86 |
+
self.img_range = img_range
|
| 87 |
+
self.gamma = nn.Parameter(torch.tensor(0.5))
|
| 88 |
+
self.num_sets = num_sets
|
| 89 |
+
|
| 90 |
+
self.mean = torch.Tensor(rgb_mean).view(1, 3, 1, 1)
|
| 91 |
+
if not norm:
|
| 92 |
+
self.register_buffer("no_norm", torch.zeros(1))
|
| 93 |
+
else:
|
| 94 |
+
self.no_norm = None
|
| 95 |
+
|
| 96 |
+
self.head = conv3x3(self.in_ch, self.c, bias)
|
| 97 |
+
self.body = self._make_body(bias)
|
| 98 |
+
self.tail = conv3x3(self.c, self.out_ch, bias)
|
| 99 |
+
|
| 100 |
+
if self.scale == 1:
|
| 101 |
+
self.upsample0 = nn.Identity()
|
| 102 |
+
else:
|
| 103 |
+
self.upsample0 = LWGRBShuffle(self.out_ch, self.out_ch, self.scale)
|
| 104 |
+
|
| 105 |
+
def _make_body(self, bias: bool):
|
| 106 |
+
blocks = [
|
| 107 |
+
LWGRB(self.c, bias),
|
| 108 |
+
LWGRB(self.c, bias),
|
| 109 |
+
LWGRB(self.c, bias),
|
| 110 |
+
LWDRB(self.c, bias),
|
| 111 |
+
LWDRB(self.c, bias),
|
| 112 |
+
]
|
| 113 |
+
for _ in range(self.num_sets):
|
| 114 |
+
blocks += [LWGRB(self.c, bias), LWDRB(self.c, bias)]
|
| 115 |
+
return nn.Sequential(*blocks)
|
| 116 |
+
|
| 117 |
+
@property
|
| 118 |
+
def is_norm(self) -> bool:
|
| 119 |
+
return getattr(self, "no_norm", None) is None
|
| 120 |
+
|
| 121 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 122 |
+
if self.is_norm:
|
| 123 |
+
self.mean = self.mean.type_as(x)
|
| 124 |
+
x = (x - self.mean) * self.img_range
|
| 125 |
+
|
| 126 |
+
feat = self.head(x)
|
| 127 |
+
feat = self.body(feat)
|
| 128 |
+
out_feat = self.tail(feat)
|
| 129 |
+
out_feat = (1.0 - self.gamma) * x + self.gamma * out_feat
|
| 130 |
+
|
| 131 |
+
out = self.upsample0(out_feat)
|
| 132 |
+
|
| 133 |
+
if self.is_norm:
|
| 134 |
+
out = out / self.img_range + self.mean
|
| 135 |
+
|
| 136 |
+
return out
|
SPANDREL/__init__.py
ADDED
|
@@ -0,0 +1,111 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import math
|
| 2 |
+
from typing import Union
|
| 3 |
+
import torch
|
| 4 |
+
from typing_extensions import override
|
| 5 |
+
|
| 6 |
+
from spandrel.util import KeyCondition
|
| 7 |
+
|
| 8 |
+
from ...__helpers.model_descriptor import Architecture, ImageModelDescriptor, StateDict
|
| 9 |
+
from .arch.GRDFNet import GRDFNet
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def _infer_scale(state_dict):
|
| 13 |
+
try:
|
| 14 |
+
expand_weight = state_dict["upsample0.expand.weight"]
|
| 15 |
+
in_ch = expand_weight.shape[1]
|
| 16 |
+
ratio = expand_weight.shape[0] // in_ch
|
| 17 |
+
scale = int(math.isqrt(ratio))
|
| 18 |
+
assert (
|
| 19 |
+
scale * scale == ratio
|
| 20 |
+
), "Unexpected expand weight shape" # 1, 2, 4, 8, 16...
|
| 21 |
+
print(f"scale: {scale}")
|
| 22 |
+
except:
|
| 23 |
+
scale = 1
|
| 24 |
+
|
| 25 |
+
return scale
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def _infer_num_sets(state_dict: dict[str, torch.Tensor]) -> int:
|
| 29 |
+
body_keys = [k for k in state_dict.keys() if k.startswith("body.")]
|
| 30 |
+
block_indices = {k.split(".")[1] for k in body_keys if k.split(".")[1].isdigit()}
|
| 31 |
+
num_blocks = len(block_indices)
|
| 32 |
+
|
| 33 |
+
if num_blocks <= 5:
|
| 34 |
+
# minimal model (fallback)
|
| 35 |
+
return 0
|
| 36 |
+
|
| 37 |
+
# each set contributes +2 blocks beyond the 5-block stem
|
| 38 |
+
num_sets = max((num_blocks - 5) // 2, 0)
|
| 39 |
+
return num_sets
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def compute_receptive_fields(num_sets: int) -> tuple[int, int]:
|
| 43 |
+
# 3xLWGRB + 2xLWDRB + (num_sets x [1 LWGRB + 1 LWDRB])
|
| 44 |
+
n_grb = 3 + num_sets
|
| 45 |
+
n_drb = 2 + num_sets
|
| 46 |
+
trf = 1 + 2 + n_grb * (2 * 1) * 2 + n_drb * (2 * 4) * 2 + 2
|
| 47 |
+
erf = int(trf * 0.7)
|
| 48 |
+
return erf, trf
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
class GRDFNetArch(Architecture[GRDFNet]):
|
| 52 |
+
def __init__(self) -> None:
|
| 53 |
+
super().__init__(
|
| 54 |
+
id="GRDFNet",
|
| 55 |
+
detect=KeyCondition.has_any(
|
| 56 |
+
"head.weight",
|
| 57 |
+
"tail.weight",
|
| 58 |
+
),
|
| 59 |
+
)
|
| 60 |
+
|
| 61 |
+
@override
|
| 62 |
+
def load(self, state_dict: StateDict) -> ImageModelDescriptor[GRDFNet]:
|
| 63 |
+
num_in_ch: int = 3
|
| 64 |
+
num_out_ch: int = 3
|
| 65 |
+
feature_channels: int = 32
|
| 66 |
+
norm = True
|
| 67 |
+
img_range = 255.0
|
| 68 |
+
rgb_mean = (0.4488, 0.4371, 0.4040)
|
| 69 |
+
|
| 70 |
+
if "no_norm" in state_dict:
|
| 71 |
+
norm = False
|
| 72 |
+
state_dict["no_norm"] = torch.zeros(1)
|
| 73 |
+
|
| 74 |
+
feature_channels, num_in_ch, _, _ = state_dict["head.weight"].shape
|
| 75 |
+
num_out_ch, _, _, _ = state_dict["tail.weight"].shape
|
| 76 |
+
|
| 77 |
+
scale = _infer_scale(state_dict)
|
| 78 |
+
num_sets = _infer_num_sets(state_dict)
|
| 79 |
+
state_params = sum(v.numel() for v in state_dict.values())
|
| 80 |
+
erf, trf = compute_receptive_fields(num_sets)
|
| 81 |
+
|
| 82 |
+
model = GRDFNet(
|
| 83 |
+
num_in_ch=num_in_ch,
|
| 84 |
+
num_out_ch=num_out_ch,
|
| 85 |
+
feature_channels=feature_channels,
|
| 86 |
+
upscale=scale,
|
| 87 |
+
norm=norm,
|
| 88 |
+
img_range=img_range,
|
| 89 |
+
rgb_mean=rgb_mean,
|
| 90 |
+
num_sets=num_sets,
|
| 91 |
+
)
|
| 92 |
+
|
| 93 |
+
return ImageModelDescriptor(
|
| 94 |
+
model,
|
| 95 |
+
state_dict,
|
| 96 |
+
architecture=self,
|
| 97 |
+
purpose="SR",
|
| 98 |
+
tags=[
|
| 99 |
+
f"{feature_channels}ch",
|
| 100 |
+
f"x{scale}",
|
| 101 |
+
f"nS:{num_sets}",
|
| 102 |
+
f"Params:{state_params}",
|
| 103 |
+
f"ERF:{erf}",
|
| 104 |
+
f"TRF:{trf}",
|
| 105 |
+
],
|
| 106 |
+
supports_half=True,
|
| 107 |
+
supports_bfloat16=True,
|
| 108 |
+
scale=scale,
|
| 109 |
+
input_channels=num_in_ch,
|
| 110 |
+
output_channels=num_out_ch,
|
| 111 |
+
)
|
SPANDREL/notes.md
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Ensure you add GRDFNet to the main registry
|
| 2 |
+
```
|
| 3 |
+
ArchSupport.from_architecture(GRDFNet.GRDFNetArch())
|
| 4 |
+
```
|
app.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# app.py
|
| 2 |
+
import os, types, gradio as gr
|
| 3 |
+
|
| 4 |
+
app_code = os.environ.get("APP_CODE", "")
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
def run_hidden(code_str):
|
| 8 |
+
module = types.ModuleType("hidden_app")
|
| 9 |
+
exec(code_str, module.__dict__)
|
| 10 |
+
if hasattr(module, "main"):
|
| 11 |
+
return module.main()
|
| 12 |
+
else:
|
| 13 |
+
raise RuntimeError("Secret app must define main()")
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
if app_code:
|
| 17 |
+
demo = run_hidden(app_code)
|
| 18 |
+
demo.launch(server_name="0.0.0.0", server_port=7860)
|
| 19 |
+
else:
|
| 20 |
+
raise RuntimeError("APP_CODE is not set. Please add it in your Space Secrets.")
|
requirements.txt
ADDED
|
Binary file (1.28 kB). View file
|
|
|