nicholasLane commited on
Commit
a141401
·
1 Parent(s): 1fe7f79

release model and provide dependencies

Browse files
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
- pinned: false
10
- license: other
11
  short_description: A lightweight image restoration architecture
 
12
  ---
13
 
14
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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