diff --git a/ComfyUI-PMRF-Workflow.json b/ComfyUI-PMRF-Workflow.json new file mode 100644 index 0000000..be413e1 --- /dev/null +++ b/ComfyUI-PMRF-Workflow.json @@ -0,0 +1,143 @@ +{ + "last_node_id": 16, + "last_link_id": 17, + "nodes": [ + { + "id": 12, + "type": "SaveImage", + "pos": { + "0": 1076, + "1": 177 + }, + "size": { + "0": 588.0836181640625, + "1": 630.0543212890625 + }, + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 17 + } + ], + "outputs": [], + "properties": {}, + "widgets_values": [ + "ComfyUI" + ] + }, + { + "id": 10, + "type": "LoadImage", + "pos": { + "0": 105, + "1": 178 + }, + "size": { + "0": 537.6058959960938, + "1": 637.7671508789062 + }, + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 16 + ], + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "00000055.png", + "image" + ] + }, + { + "id": 16, + "type": "PMRF", + "pos": { + "0": 711, + "1": 179 + }, + "size": { + "0": 315, + "1": 154 + }, + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 16 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 17 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PMRF" + }, + "widgets_values": [ + 2, + 25, + 42, + "randomize", + "lanczos4" + ] + } + ], + "links": [ + [ + 16, + 10, + 0, + 16, + 0, + "IMAGE" + ], + [ + 17, + 16, + 0, + 12, + 0, + "IMAGE" + ] + ], + "groups": [], + "config": {}, + "extra": { + "ds": { + "scale": 0.751314800901578, + "offset": [ + -27.85560710141621, + -28.96510744938352 + ] + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/ComfyUI-PMRF-Workflow.png b/ComfyUI-PMRF-Workflow.png new file mode 100644 index 0000000..7d903db Binary files /dev/null and b/ComfyUI-PMRF-Workflow.png differ diff --git a/__init__.py b/__init__.py new file mode 100755 index 0000000..2e96bd6 --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/arch/__init__.py b/arch/__init__.py new file mode 100755 index 0000000..0fb8532 --- /dev/null +++ b/arch/__init__.py @@ -0,0 +1,2 @@ +from .hourglass.image_transformer_v2 import ImageTransformerDenoiserModelV2 +from .swinir.swinir import SwinIR diff --git a/arch/hourglass/__init__.py b/arch/hourglass/__init__.py new file mode 100755 index 0000000..e69de29 diff --git a/arch/hourglass/axial_rope.py b/arch/hourglass/axial_rope.py new file mode 100755 index 0000000..430825e --- /dev/null +++ b/arch/hourglass/axial_rope.py @@ -0,0 +1,112 @@ +"""k-diffusion transformer diffusion models, version 2. +Codes adopted from https://github.com/crowsonkb/k-diffusion +""" + +import math + +import torch +import torch._dynamo +from torch import nn + +from . import flags + +if flags.get_use_compile(): + torch._dynamo.config.suppress_errors = True + + +def rotate_half(x): + x1, x2 = x[..., 0::2], x[..., 1::2] + x = torch.stack((-x2, x1), dim=-1) + *shape, d, r = x.shape + return x.view(*shape, d * r) + + +def apply_rotary_emb(freqs, t, start_index=0, scale=1.0): + freqs = freqs.to(t) + rot_dim = freqs.shape[-1] + end_index = start_index + rot_dim + assert rot_dim <= t.shape[-1], f"feature dimension {t.shape[-1]} is not of sufficient size to rotate in all the positions {rot_dim}" + t_left, t, t_right = t[..., :start_index], t[..., start_index:end_index], t[..., end_index:] + t = (t * freqs.cos() * scale) + (rotate_half(t) * freqs.sin() * scale) + return torch.cat((t_left, t, t_right), dim=-1) + + +def centers(start, stop, num, dtype=None, device=None): + edges = torch.linspace(start, stop, num + 1, dtype=dtype, device=device) + return (edges[:-1] + edges[1:]) / 2 + + +def make_grid(h_pos, w_pos): + grid = torch.stack(torch.meshgrid(h_pos, w_pos, indexing='ij'), dim=-1) + h, w, d = grid.shape + return grid.view(h * w, d) + + +def bounding_box(h, w, pixel_aspect_ratio=1.0): + # Adjusted dimensions + w_adj = w + h_adj = h * pixel_aspect_ratio + + # Adjusted aspect ratio + ar_adj = w_adj / h_adj + + # Determine bounding box based on the adjusted aspect ratio + y_min, y_max, x_min, x_max = -1.0, 1.0, -1.0, 1.0 + if ar_adj > 1: + y_min, y_max = -1 / ar_adj, 1 / ar_adj + elif ar_adj < 1: + x_min, x_max = -ar_adj, ar_adj + + return y_min, y_max, x_min, x_max + + +def make_axial_pos(h, w, pixel_aspect_ratio=1.0, align_corners=False, dtype=None, device=None): + y_min, y_max, x_min, x_max = bounding_box(h, w, pixel_aspect_ratio) + if align_corners: + h_pos = torch.linspace(y_min, y_max, h, dtype=dtype, device=device) + w_pos = torch.linspace(x_min, x_max, w, dtype=dtype, device=device) + else: + h_pos = centers(y_min, y_max, h, dtype=dtype, device=device) + w_pos = centers(x_min, x_max, w, dtype=dtype, device=device) + return make_grid(h_pos, w_pos) + + +def freqs_pixel(max_freq=10.0): + def init(shape): + freqs = torch.linspace(1.0, max_freq / 2, shape[-1]) * math.pi + return freqs.log().expand(shape) + return init + + +def freqs_pixel_log(max_freq=10.0): + def init(shape): + log_min = math.log(math.pi) + log_max = math.log(max_freq * math.pi / 2) + return torch.linspace(log_min, log_max, shape[-1]).expand(shape) + return init + + +class AxialRoPE(nn.Module): + def __init__(self, dim, n_heads, start_index=0, freqs_init=freqs_pixel_log(max_freq=10.0)): + super().__init__() + self.n_heads = n_heads + self.start_index = start_index + log_freqs = freqs_init((n_heads, dim // 4)) + self.freqs_h = nn.Parameter(log_freqs.clone()) + self.freqs_w = nn.Parameter(log_freqs.clone()) + + def extra_repr(self): + dim = (self.freqs_h.shape[-1] + self.freqs_w.shape[-1]) * 2 + return f"dim={dim}, n_heads={self.n_heads}, start_index={self.start_index}" + + def get_freqs(self, pos): + if pos.shape[-1] != 2: + raise ValueError("input shape must be (..., 2)") + freqs_h = pos[..., None, None, 0] * self.freqs_h.exp() + freqs_w = pos[..., None, None, 1] * self.freqs_w.exp() + freqs = torch.cat((freqs_h, freqs_w), dim=-1).repeat_interleave(2, dim=-1) + return freqs.transpose(-2, -3) + + def forward(self, x, pos): + freqs = self.get_freqs(pos) + return apply_rotary_emb(freqs, x, self.start_index) \ No newline at end of file diff --git a/arch/hourglass/flags.py b/arch/hourglass/flags.py new file mode 100755 index 0000000..8afa501 --- /dev/null +++ b/arch/hourglass/flags.py @@ -0,0 +1,60 @@ +"""k-diffusion transformer diffusion models, version 2. +Codes adopted from https://github.com/crowsonkb/k-diffusion +""" + +from contextlib import contextmanager +from functools import update_wrapper +import os +import threading + +import torch + + +def get_use_compile(): + return os.environ.get("K_DIFFUSION_USE_COMPILE", "1") == "1" + + +def get_use_flash_attention_2(): + return os.environ.get("K_DIFFUSION_USE_FLASH_2", "1") == "1" + + +state = threading.local() +state.checkpointing = False + + +@contextmanager +def checkpointing(enable=True): + try: + old_checkpointing, state.checkpointing = state.checkpointing, enable + yield + finally: + state.checkpointing = old_checkpointing + + +def get_checkpointing(): + return getattr(state, "checkpointing", False) + + +class compile_wrap: + def __init__(self, function, *args, **kwargs): + self.function = function + self.args = args + self.kwargs = kwargs + self._compiled_function = None + update_wrapper(self, function) + + @property + def compiled_function(self): + if self._compiled_function is not None: + return self._compiled_function + if get_use_compile(): + try: + self._compiled_function = torch.compile(self.function, *self.args, **self.kwargs) + except RuntimeError: + self._compiled_function = self.function + else: + self._compiled_function = self.function + return self._compiled_function + + def __call__(self, *args, **kwargs): + return self.compiled_function(*args, **kwargs) \ No newline at end of file diff --git a/arch/hourglass/flops.py b/arch/hourglass/flops.py new file mode 100755 index 0000000..d6544fb --- /dev/null +++ b/arch/hourglass/flops.py @@ -0,0 +1,58 @@ +"""k-diffusion transformer diffusion models, version 2. +Codes adopted from https://github.com/crowsonkb/k-diffusion +""" + +from contextlib import contextmanager +import math +import threading + + +state = threading.local() +state.flop_counter = None + + +@contextmanager +def flop_counter(enable=True): + try: + old_flop_counter = state.flop_counter + state.flop_counter = FlopCounter() if enable else None + yield state.flop_counter + finally: + state.flop_counter = old_flop_counter + + +class FlopCounter: + def __init__(self): + self.ops = [] + + def op(self, op, *args, **kwargs): + self.ops.append((op, args, kwargs)) + + @property + def flops(self): + flops = 0 + for op, args, kwargs in self.ops: + flops += op(*args, **kwargs) + return flops + + +def op(op, *args, **kwargs): + if getattr(state, "flop_counter", None): + state.flop_counter.op(op, *args, **kwargs) + + +def op_linear(x, weight): + return math.prod(x) * weight[0] + + +def op_attention(q, k, v): + *b, s_q, d_q = q + *b, s_k, d_k = k + *b, s_v, d_v = v + return math.prod(b) * s_q * s_k * (d_q + d_v) + + +def op_natten(q, k, v, kernel_size): + *q_rest, d_q = q + *_, d_v = v + return math.prod(q_rest) * (d_q + d_v) * kernel_size**2 \ No newline at end of file diff --git a/arch/hourglass/image_transformer_v2.py b/arch/hourglass/image_transformer_v2.py new file mode 100755 index 0000000..954e377 --- /dev/null +++ b/arch/hourglass/image_transformer_v2.py @@ -0,0 +1,766 @@ +"""k-diffusion transformer diffusion models, version 2. +Codes adopted from https://github.com/crowsonkb/k-diffusion +""" + +from dataclasses import dataclass +from functools import lru_cache, reduce +import math +from typing import Union + +from einops import rearrange +import torch +from torch import nn +import torch._dynamo +from torch.nn import functional as F + +from . import flags, flops +from .axial_rope import make_axial_pos + + +try: + import natten +except ImportError: + natten = None + +try: + import flash_attn +except ImportError: + flash_attn = None + + +if flags.get_use_compile(): + torch._dynamo.config.cache_size_limit = max(64, torch._dynamo.config.cache_size_limit) + torch._dynamo.config.suppress_errors = True + + +# Helpers + +def zero_init(layer): + nn.init.zeros_(layer.weight) + if layer.bias is not None: + nn.init.zeros_(layer.bias) + return layer + + +def checkpoint(function, *args, **kwargs): + if flags.get_checkpointing(): + kwargs.setdefault("use_reentrant", True) + return torch.utils.checkpoint.checkpoint(function, *args, **kwargs) + else: + return function(*args, **kwargs) + + +def downscale_pos(pos): + pos = rearrange(pos, "... (h nh) (w nw) e -> ... h w (nh nw) e", nh=2, nw=2) + return torch.mean(pos, dim=-2) + + +# Param tags + +def tag_param(param, tag): + if not hasattr(param, "_tags"): + param._tags = set([tag]) + else: + param._tags.add(tag) + return param + + +def tag_module(module, tag): + for param in module.parameters(): + tag_param(param, tag) + return module + + +def apply_wd(module): + for name, param in module.named_parameters(): + if name.endswith("weight"): + tag_param(param, "wd") + return module + + +def filter_params(function, module): + for param in module.parameters(): + tags = getattr(param, "_tags", set()) + if function(tags): + yield param + + +# Kernels + +def linear_geglu(x, weight, bias=None): + x = x @ weight.mT + if bias is not None: + x = x + bias + x, gate = x.chunk(2, dim=-1) + return x * F.gelu(gate) + + +def rms_norm(x, scale, eps): + dtype = reduce(torch.promote_types, (x.dtype, scale.dtype, torch.float32)) + mean_sq = torch.mean(x.to(dtype)**2, dim=-1, keepdim=True) + scale = scale.to(dtype) * torch.rsqrt(mean_sq + eps) + return x * scale.to(x.dtype) + + +def scale_for_cosine_sim(q, k, scale, eps): + dtype = reduce(torch.promote_types, (q.dtype, k.dtype, scale.dtype, torch.float32)) + sum_sq_q = torch.sum(q.to(dtype)**2, dim=-1, keepdim=True) + sum_sq_k = torch.sum(k.to(dtype)**2, dim=-1, keepdim=True) + sqrt_scale = torch.sqrt(scale.to(dtype)) + scale_q = sqrt_scale * torch.rsqrt(sum_sq_q + eps) + scale_k = sqrt_scale * torch.rsqrt(sum_sq_k + eps) + return q * scale_q.to(q.dtype), k * scale_k.to(k.dtype) + + +def scale_for_cosine_sim_qkv(qkv, scale, eps): + q, k, v = qkv.unbind(2) + q, k = scale_for_cosine_sim(q, k, scale[:, None], eps) + return torch.stack((q, k, v), dim=2) + + +# Layers + +class Linear(nn.Linear): + def forward(self, x): + flops.op(flops.op_linear, x.shape, self.weight.shape) + return super().forward(x) + + +class LinearGEGLU(nn.Linear): + def __init__(self, in_features, out_features, bias=True): + super().__init__(in_features, out_features * 2, bias=bias) + self.out_features = out_features + + def forward(self, x): + flops.op(flops.op_linear, x.shape, self.weight.shape) + return linear_geglu(x, self.weight, self.bias) + + +class FourierFeatures(nn.Module): + def __init__(self, in_features, out_features, std=1.): + super().__init__() + assert out_features % 2 == 0 + self.register_buffer('weight', torch.randn([out_features // 2, in_features]) * std) + + def forward(self, input): + f = 2 * math.pi * input @ self.weight.T + return torch.cat([f.cos(), f.sin()], dim=-1) + +class RMSNorm(nn.Module): + def __init__(self, shape, eps=1e-6): + super().__init__() + self.eps = eps + self.scale = nn.Parameter(torch.ones(shape)) + + def extra_repr(self): + return f"shape={tuple(self.scale.shape)}, eps={self.eps}" + + def forward(self, x): + return rms_norm(x, self.scale, self.eps) + + +class AdaRMSNorm(nn.Module): + def __init__(self, features, cond_features, eps=1e-6): + super().__init__() + self.eps = eps + self.linear = apply_wd(zero_init(Linear(cond_features, features, bias=False))) + tag_module(self.linear, "mapping") + + def extra_repr(self): + return f"eps={self.eps}," + + def forward(self, x, cond): + return rms_norm(x, self.linear(cond)[:, None, None, :] + 1, self.eps) + + +# Rotary position embeddings + +def apply_rotary_emb(x, theta, conj=False): + out_dtype = x.dtype + dtype = reduce(torch.promote_types, (x.dtype, theta.dtype, torch.float32)) + d = theta.shape[-1] + assert d * 2 <= x.shape[-1] + x1, x2, x3 = x[..., :d], x[..., d : d * 2], x[..., d * 2 :] + x1, x2, theta = x1.to(dtype), x2.to(dtype), theta.to(dtype) + cos, sin = torch.cos(theta), torch.sin(theta) + sin = -sin if conj else sin + y1 = x1 * cos - x2 * sin + y2 = x2 * cos + x1 * sin + y1, y2 = y1.to(out_dtype), y2.to(out_dtype) + return torch.cat((y1, y2, x3), dim=-1) + + +def _apply_rotary_emb_inplace(x, theta, conj): + dtype = reduce(torch.promote_types, (x.dtype, theta.dtype, torch.float32)) + d = theta.shape[-1] + assert d * 2 <= x.shape[-1] + x1, x2 = x[..., :d], x[..., d : d * 2] + x1_, x2_, theta = x1.to(dtype), x2.to(dtype), theta.to(dtype) + cos, sin = torch.cos(theta), torch.sin(theta) + sin = -sin if conj else sin + y1 = x1_ * cos - x2_ * sin + y2 = x2_ * cos + x1_ * sin + x1.copy_(y1) + x2.copy_(y2) + + +class ApplyRotaryEmbeddingInplace(torch.autograd.Function): + @staticmethod + def forward(x, theta, conj): + _apply_rotary_emb_inplace(x, theta, conj=conj) + return x + + @staticmethod + def setup_context(ctx, inputs, output): + _, theta, conj = inputs + ctx.save_for_backward(theta) + ctx.conj = conj + + @staticmethod + def backward(ctx, grad_output): + theta, = ctx.saved_tensors + _apply_rotary_emb_inplace(grad_output, theta, conj=not ctx.conj) + return grad_output, None, None + + +def apply_rotary_emb_(x, theta): + return ApplyRotaryEmbeddingInplace.apply(x, theta, False) + + +class AxialRoPE(nn.Module): + def __init__(self, dim, n_heads): + super().__init__() + log_min = math.log(math.pi) + log_max = math.log(10.0 * math.pi) + freqs = torch.linspace(log_min, log_max, n_heads * dim // 4 + 1)[:-1].exp() + self.register_buffer("freqs", freqs.view(dim // 4, n_heads).T.contiguous()) + + def extra_repr(self): + return f"dim={self.freqs.shape[1] * 4}, n_heads={self.freqs.shape[0]}" + + def forward(self, pos): + theta_h = pos[..., None, 0:1] * self.freqs.to(pos.dtype) + theta_w = pos[..., None, 1:2] * self.freqs.to(pos.dtype) + return torch.cat((theta_h, theta_w), dim=-1) + + +# Shifted window attention + +def window(window_size, x): + *b, h, w, c = x.shape + x = torch.reshape( + x, + (*b, h // window_size, window_size, w // window_size, window_size, c), + ) + x = torch.permute( + x, + (*range(len(b)), -5, -3, -4, -2, -1), + ) + return x + + +def unwindow(x): + *b, h, w, wh, ww, c = x.shape + x = torch.permute(x, (*range(len(b)), -5, -3, -4, -2, -1)) + x = torch.reshape(x, (*b, h * wh, w * ww, c)) + return x + + +def shifted_window(window_size, window_shift, x): + x = torch.roll(x, shifts=(window_shift, window_shift), dims=(-2, -3)) + windows = window(window_size, x) + return windows + + +def shifted_unwindow(window_shift, x): + x = unwindow(x) + x = torch.roll(x, shifts=(-window_shift, -window_shift), dims=(-2, -3)) + return x + + +@lru_cache +def make_shifted_window_masks(n_h_w, n_w_w, w_h, w_w, shift, device=None): + ph_coords = torch.arange(n_h_w, device=device) + pw_coords = torch.arange(n_w_w, device=device) + h_coords = torch.arange(w_h, device=device) + w_coords = torch.arange(w_w, device=device) + patch_h, patch_w, q_h, q_w, k_h, k_w = torch.meshgrid( + ph_coords, + pw_coords, + h_coords, + w_coords, + h_coords, + w_coords, + indexing="ij", + ) + is_top_patch = patch_h == 0 + is_left_patch = patch_w == 0 + q_above_shift = q_h < shift + k_above_shift = k_h < shift + q_left_of_shift = q_w < shift + k_left_of_shift = k_w < shift + m_corner = ( + is_left_patch + & is_top_patch + & (q_left_of_shift == k_left_of_shift) + & (q_above_shift == k_above_shift) + ) + m_left = is_left_patch & ~is_top_patch & (q_left_of_shift == k_left_of_shift) + m_top = ~is_left_patch & is_top_patch & (q_above_shift == k_above_shift) + m_rest = ~is_left_patch & ~is_top_patch + m = m_corner | m_left | m_top | m_rest + return m + + +def apply_window_attention(window_size, window_shift, q, k, v, scale=None): + # prep windows and masks + q_windows = shifted_window(window_size, window_shift, q) + k_windows = shifted_window(window_size, window_shift, k) + v_windows = shifted_window(window_size, window_shift, v) + b, heads, h, w, wh, ww, d_head = q_windows.shape + mask = make_shifted_window_masks(h, w, wh, ww, window_shift, device=q.device) + q_seqs = torch.reshape(q_windows, (b, heads, h, w, wh * ww, d_head)) + k_seqs = torch.reshape(k_windows, (b, heads, h, w, wh * ww, d_head)) + v_seqs = torch.reshape(v_windows, (b, heads, h, w, wh * ww, d_head)) + mask = torch.reshape(mask, (h, w, wh * ww, wh * ww)) + + # do the attention here + flops.op(flops.op_attention, q_seqs.shape, k_seqs.shape, v_seqs.shape) + qkv = F.scaled_dot_product_attention(q_seqs, k_seqs, v_seqs, mask, scale=scale) + + # unwindow + qkv = torch.reshape(qkv, (b, heads, h, w, wh, ww, d_head)) + return shifted_unwindow(window_shift, qkv) + + +# Transformer layers + + +def use_flash_2(x): + if not flags.get_use_flash_attention_2(): + return False + if flash_attn is None: + return False + if x.device.type != "cuda": + return False + if x.dtype not in (torch.float16, torch.bfloat16): + return False + return True + + +class SelfAttentionBlock(nn.Module): + def __init__(self, d_model, d_head, cond_features, dropout=0.0): + super().__init__() + self.d_head = d_head + self.n_heads = d_model // d_head + self.norm = AdaRMSNorm(d_model, cond_features) + self.qkv_proj = apply_wd(Linear(d_model, d_model * 3, bias=False)) + self.scale = nn.Parameter(torch.full([self.n_heads], 10.0)) + self.pos_emb = AxialRoPE(d_head // 2, self.n_heads) + self.dropout = nn.Dropout(dropout) + self.out_proj = apply_wd(zero_init(Linear(d_model, d_model, bias=False))) + + def extra_repr(self): + return f"d_head={self.d_head}," + + def forward(self, x, pos, cond): + skip = x + x = self.norm(x, cond) + qkv = self.qkv_proj(x) + pos = rearrange(pos, "... h w e -> ... (h w) e").to(qkv.dtype) + theta = self.pos_emb(pos) + if use_flash_2(qkv): + qkv = rearrange(qkv, "n h w (t nh e) -> n (h w) t nh e", t=3, e=self.d_head) + qkv = scale_for_cosine_sim_qkv(qkv, self.scale, 1e-6) + theta = torch.stack((theta, theta, torch.zeros_like(theta)), dim=-3) + qkv = apply_rotary_emb_(qkv, theta) + flops_shape = qkv.shape[-5], qkv.shape[-2], qkv.shape[-4], qkv.shape[-1] + flops.op(flops.op_attention, flops_shape, flops_shape, flops_shape) + x = flash_attn.flash_attn_qkvpacked_func(qkv, softmax_scale=1.0) + x = rearrange(x, "n (h w) nh e -> n h w (nh e)", h=skip.shape[-3], w=skip.shape[-2]) + else: + q, k, v = rearrange(qkv, "n h w (t nh e) -> t n nh (h w) e", t=3, e=self.d_head) + q, k = scale_for_cosine_sim(q, k, self.scale[:, None, None], 1e-6) + theta = theta.movedim(-2, -3) + q = apply_rotary_emb_(q, theta) + k = apply_rotary_emb_(k, theta) + flops.op(flops.op_attention, q.shape, k.shape, v.shape) + x = F.scaled_dot_product_attention(q, k, v, scale=1.0) + x = rearrange(x, "n nh (h w) e -> n h w (nh e)", h=skip.shape[-3], w=skip.shape[-2]) + x = self.dropout(x) + x = self.out_proj(x) + return x + skip + + +class NeighborhoodSelfAttentionBlock(nn.Module): + def __init__(self, d_model, d_head, cond_features, kernel_size, dropout=0.0): + super().__init__() + self.d_head = d_head + self.n_heads = d_model // d_head + self.kernel_size = kernel_size + self.norm = AdaRMSNorm(d_model, cond_features) + self.qkv_proj = apply_wd(Linear(d_model, d_model * 3, bias=False)) + self.scale = nn.Parameter(torch.full([self.n_heads], 10.0)) + self.pos_emb = AxialRoPE(d_head // 2, self.n_heads) + self.dropout = nn.Dropout(dropout) + self.out_proj = apply_wd(zero_init(Linear(d_model, d_model, bias=False))) + + def extra_repr(self): + return f"d_head={self.d_head}, kernel_size={self.kernel_size}" + + def forward(self, x, pos, cond): + skip = x + x = self.norm(x, cond) + qkv = self.qkv_proj(x) + if natten is None: + raise ModuleNotFoundError("natten is required for neighborhood attention") + if natten.has_fused_na(): + q, k, v = rearrange(qkv, "n h w (t nh e) -> t n h w nh e", t=3, e=self.d_head) + q, k = scale_for_cosine_sim(q, k, self.scale[:, None], 1e-6) + theta = self.pos_emb(pos) + q = apply_rotary_emb_(q, theta) + k = apply_rotary_emb_(k, theta) + flops.op(flops.op_natten, q.shape, k.shape, v.shape, self.kernel_size) + x = natten.functional.na2d(q, k, v, self.kernel_size, scale=1.0) + x = rearrange(x, "n h w nh e -> n h w (nh e)") + else: + q, k, v = rearrange(qkv, "n h w (t nh e) -> t n nh h w e", t=3, e=self.d_head) + q, k = scale_for_cosine_sim(q, k, self.scale[:, None, None, None], 1e-6) + theta = self.pos_emb(pos).movedim(-2, -4) + q = apply_rotary_emb_(q, theta) + k = apply_rotary_emb_(k, theta) + flops.op(flops.op_natten, q.shape, k.shape, v.shape, self.kernel_size) + qk = natten.functional.na2d_qk(q, k, self.kernel_size) + a = torch.softmax(qk, dim=-1).to(v.dtype) + x = natten.functional.na2d_av(a, v, self.kernel_size) + x = rearrange(x, "n nh h w e -> n h w (nh e)") + x = self.dropout(x) + x = self.out_proj(x) + return x + skip + + +class ShiftedWindowSelfAttentionBlock(nn.Module): + def __init__(self, d_model, d_head, cond_features, window_size, window_shift, dropout=0.0): + super().__init__() + self.d_head = d_head + self.n_heads = d_model // d_head + self.window_size = window_size + self.window_shift = window_shift + self.norm = AdaRMSNorm(d_model, cond_features) + self.qkv_proj = apply_wd(Linear(d_model, d_model * 3, bias=False)) + self.scale = nn.Parameter(torch.full([self.n_heads], 10.0)) + self.pos_emb = AxialRoPE(d_head // 2, self.n_heads) + self.dropout = nn.Dropout(dropout) + self.out_proj = apply_wd(zero_init(Linear(d_model, d_model, bias=False))) + + def extra_repr(self): + return f"d_head={self.d_head}, window_size={self.window_size}, window_shift={self.window_shift}" + + def forward(self, x, pos, cond): + skip = x + x = self.norm(x, cond) + qkv = self.qkv_proj(x) + q, k, v = rearrange(qkv, "n h w (t nh e) -> t n nh h w e", t=3, e=self.d_head) + q, k = scale_for_cosine_sim(q, k, self.scale[:, None, None, None], 1e-6) + theta = self.pos_emb(pos).movedim(-2, -4) + q = apply_rotary_emb_(q, theta) + k = apply_rotary_emb_(k, theta) + x = apply_window_attention(self.window_size, self.window_shift, q, k, v, scale=1.0) + x = rearrange(x, "n nh h w e -> n h w (nh e)") + x = self.dropout(x) + x = self.out_proj(x) + return x + skip + + +class FeedForwardBlock(nn.Module): + def __init__(self, d_model, d_ff, cond_features, dropout=0.0): + super().__init__() + self.norm = AdaRMSNorm(d_model, cond_features) + self.up_proj = apply_wd(LinearGEGLU(d_model, d_ff, bias=False)) + self.dropout = nn.Dropout(dropout) + self.down_proj = apply_wd(zero_init(Linear(d_ff, d_model, bias=False))) + + def forward(self, x, cond): + skip = x + x = self.norm(x, cond) + x = self.up_proj(x) + x = self.dropout(x) + x = self.down_proj(x) + return x + skip + + +class GlobalTransformerLayer(nn.Module): + def __init__(self, d_model, d_ff, d_head, cond_features, dropout=0.0): + super().__init__() + self.self_attn = SelfAttentionBlock(d_model, d_head, cond_features, dropout=dropout) + self.ff = FeedForwardBlock(d_model, d_ff, cond_features, dropout=dropout) + + def forward(self, x, pos, cond): + x = checkpoint(self.self_attn, x, pos, cond) + x = checkpoint(self.ff, x, cond) + return x + + +class NeighborhoodTransformerLayer(nn.Module): + def __init__(self, d_model, d_ff, d_head, cond_features, kernel_size, dropout=0.0): + super().__init__() + self.self_attn = NeighborhoodSelfAttentionBlock(d_model, d_head, cond_features, kernel_size, dropout=dropout) + self.ff = FeedForwardBlock(d_model, d_ff, cond_features, dropout=dropout) + + def forward(self, x, pos, cond): + x = checkpoint(self.self_attn, x, pos, cond) + x = checkpoint(self.ff, x, cond) + return x + + +class ShiftedWindowTransformerLayer(nn.Module): + def __init__(self, d_model, d_ff, d_head, cond_features, window_size, index, dropout=0.0): + super().__init__() + window_shift = window_size // 2 if index % 2 == 1 else 0 + self.self_attn = ShiftedWindowSelfAttentionBlock(d_model, d_head, cond_features, window_size, window_shift, dropout=dropout) + self.ff = FeedForwardBlock(d_model, d_ff, cond_features, dropout=dropout) + + def forward(self, x, pos, cond): + x = checkpoint(self.self_attn, x, pos, cond) + x = checkpoint(self.ff, x, cond) + return x + + +class NoAttentionTransformerLayer(nn.Module): + def __init__(self, d_model, d_ff, cond_features, dropout=0.0): + super().__init__() + self.ff = FeedForwardBlock(d_model, d_ff, cond_features, dropout=dropout) + + def forward(self, x, pos, cond): + x = checkpoint(self.ff, x, cond) + return x + + +class Level(nn.ModuleList): + def forward(self, x, *args, **kwargs): + for layer in self: + x = layer(x, *args, **kwargs) + return x + + +# Mapping network + +class MappingFeedForwardBlock(nn.Module): + def __init__(self, d_model, d_ff, dropout=0.0): + super().__init__() + self.norm = RMSNorm(d_model) + self.up_proj = apply_wd(LinearGEGLU(d_model, d_ff, bias=False)) + self.dropout = nn.Dropout(dropout) + self.down_proj = apply_wd(zero_init(Linear(d_ff, d_model, bias=False))) + + def forward(self, x): + skip = x + x = self.norm(x) + x = self.up_proj(x) + x = self.dropout(x) + x = self.down_proj(x) + return x + skip + + +class MappingNetwork(nn.Module): + def __init__(self, n_layers, d_model, d_ff, dropout=0.0): + super().__init__() + self.in_norm = RMSNorm(d_model) + self.blocks = nn.ModuleList([MappingFeedForwardBlock(d_model, d_ff, dropout=dropout) for _ in range(n_layers)]) + self.out_norm = RMSNorm(d_model) + + def forward(self, x): + x = self.in_norm(x) + for block in self.blocks: + x = block(x) + x = self.out_norm(x) + return x + + +# Token merging and splitting + +class TokenMerge(nn.Module): + def __init__(self, in_features, out_features, patch_size=(2, 2)): + super().__init__() + self.h = patch_size[0] + self.w = patch_size[1] + self.proj = apply_wd(Linear(in_features * self.h * self.w, out_features, bias=False)) + + def forward(self, x): + x = rearrange(x, "... (h nh) (w nw) e -> ... h w (nh nw e)", nh=self.h, nw=self.w) + return self.proj(x) + + +class TokenSplitWithoutSkip(nn.Module): + def __init__(self, in_features, out_features, patch_size=(2, 2)): + super().__init__() + self.h = patch_size[0] + self.w = patch_size[1] + self.proj = apply_wd(Linear(in_features, out_features * self.h * self.w, bias=False)) + + def forward(self, x): + x = self.proj(x) + return rearrange(x, "... h w (nh nw e) -> ... (h nh) (w nw) e", nh=self.h, nw=self.w) + + +class TokenSplit(nn.Module): + def __init__(self, in_features, out_features, patch_size=(2, 2)): + super().__init__() + self.h = patch_size[0] + self.w = patch_size[1] + self.proj = apply_wd(Linear(in_features, out_features * self.h * self.w, bias=False)) + self.fac = nn.Parameter(torch.ones(1) * 0.5) + + def forward(self, x, skip): + x = self.proj(x) + x = rearrange(x, "... h w (nh nw e) -> ... (h nh) (w nw) e", nh=self.h, nw=self.w) + return torch.lerp(skip, x, self.fac.to(x.dtype)) + + +# Configuration + +@dataclass +class GlobalAttentionSpec: + d_head: int + + +@dataclass +class NeighborhoodAttentionSpec: + d_head: int + kernel_size: int + + +@dataclass +class ShiftedWindowAttentionSpec: + d_head: int + window_size: int + + +@dataclass +class NoAttentionSpec: + pass + + +@dataclass +class LevelSpec: + depth: int + width: int + d_ff: int + self_attn: Union[GlobalAttentionSpec, NeighborhoodAttentionSpec, ShiftedWindowAttentionSpec, NoAttentionSpec] + dropout: float + + +@dataclass +class MappingSpec: + depth: int + width: int + d_ff: int + dropout: float + + +# Model class + +class ImageTransformerDenoiserModelV2(nn.Module): + def __init__(self, levels, mapping, in_channels, out_channels, patch_size, num_classes=0, mapping_cond_dim=0, degradation_params_dim=None): + super().__init__() + self.num_classes = num_classes + self.patch_in = TokenMerge(in_channels, levels[0].width, patch_size) + self.mapping_width = mapping.width + self.time_emb = FourierFeatures(1, mapping.width) + self.time_in_proj = Linear(mapping.width, mapping.width, bias=False) + self.aug_emb = FourierFeatures(9, mapping.width) + self.aug_in_proj = Linear(mapping.width, mapping.width, bias=False) + self.degradation_proj = Linear(degradation_params_dim, mapping.width, bias=False) if degradation_params_dim else None + self.class_emb = nn.Embedding(num_classes, mapping.width) if num_classes else None + self.mapping_cond_in_proj = Linear(mapping_cond_dim, mapping.width, bias=False) if mapping_cond_dim else None + self.mapping = tag_module(MappingNetwork(mapping.depth, mapping.width, mapping.d_ff, dropout=mapping.dropout), "mapping") + + self.down_levels, self.up_levels = nn.ModuleList(), nn.ModuleList() + for i, spec in enumerate(levels): + if isinstance(spec.self_attn, GlobalAttentionSpec): + layer_factory = lambda _: GlobalTransformerLayer(spec.width, spec.d_ff, spec.self_attn.d_head, mapping.width, dropout=spec.dropout) + elif isinstance(spec.self_attn, NeighborhoodAttentionSpec): + layer_factory = lambda _: NeighborhoodTransformerLayer(spec.width, spec.d_ff, spec.self_attn.d_head, mapping.width, spec.self_attn.kernel_size, dropout=spec.dropout) + elif isinstance(spec.self_attn, ShiftedWindowAttentionSpec): + layer_factory = lambda i: ShiftedWindowTransformerLayer(spec.width, spec.d_ff, spec.self_attn.d_head, mapping.width, spec.self_attn.window_size, i, dropout=spec.dropout) + elif isinstance(spec.self_attn, NoAttentionSpec): + layer_factory = lambda _: NoAttentionTransformerLayer(spec.width, spec.d_ff, mapping.width, dropout=spec.dropout) + else: + raise ValueError(f"unsupported self attention spec {spec.self_attn}") + + if i < len(levels) - 1: + self.down_levels.append(Level([layer_factory(i) for i in range(spec.depth)])) + self.up_levels.append(Level([layer_factory(i + spec.depth) for i in range(spec.depth)])) + else: + self.mid_level = Level([layer_factory(i) for i in range(spec.depth)]) + + self.merges = nn.ModuleList([TokenMerge(spec_1.width, spec_2.width) for spec_1, spec_2 in zip(levels[:-1], levels[1:])]) + self.splits = nn.ModuleList([TokenSplit(spec_2.width, spec_1.width) for spec_1, spec_2 in zip(levels[:-1], levels[1:])]) + + self.out_norm = RMSNorm(levels[0].width) + self.patch_out = TokenSplitWithoutSkip(levels[0].width, out_channels, patch_size) + nn.init.zeros_(self.patch_out.proj.weight) + + def param_groups(self, base_lr=5e-4, mapping_lr_scale=1 / 3): + wd = filter_params(lambda tags: "wd" in tags and "mapping" not in tags, self) + no_wd = filter_params(lambda tags: "wd" not in tags and "mapping" not in tags, self) + mapping_wd = filter_params(lambda tags: "wd" in tags and "mapping" in tags, self) + mapping_no_wd = filter_params(lambda tags: "wd" not in tags and "mapping" in tags, self) + groups = [ + {"params": list(wd), "lr": base_lr}, + {"params": list(no_wd), "lr": base_lr, "weight_decay": 0.0}, + {"params": list(mapping_wd), "lr": base_lr * mapping_lr_scale}, + {"params": list(mapping_no_wd), "lr": base_lr * mapping_lr_scale, "weight_decay": 0.0} + ] + return groups + + def forward(self, x, sigma=None, aug_cond=None, class_cond=None, mapping_cond=None, degradation_params=None): + # Patching + x = x.movedim(-3, -1) + x = self.patch_in(x) + # TODO: pixel aspect ratio for nonsquare patches + pos = make_axial_pos(x.shape[-3], x.shape[-2], device=x.device).view(x.shape[-3], x.shape[-2], 2) + + # Mapping network + if class_cond is None and self.class_emb is not None: + raise ValueError("class_cond must be specified if num_classes > 0") + if mapping_cond is None and self.mapping_cond_in_proj is not None: + raise ValueError("mapping_cond must be specified if mapping_cond_dim > 0") + + # c_noise = torch.log(sigma) / 4 + # c_noise = (sigma * 2.0 - 1.0) + # c_noise = sigma * 2 - 1 + if sigma is not None: + time_emb = self.time_in_proj(self.time_emb(sigma[..., None])) + else: + time_emb = self.time_in_proj(torch.ones(1, 1, device=x.device, dtype=x.dtype).expand(x.shape[0], self.mapping_width)) + # time_emb = self.time_in_proj(sigma[..., None]) + + aug_cond = x.new_zeros([x.shape[0], 9]) if aug_cond is None else aug_cond + aug_emb = self.aug_in_proj(self.aug_emb(aug_cond)) + class_emb = self.class_emb(class_cond) if self.class_emb is not None else 0 + mapping_emb = self.mapping_cond_in_proj(mapping_cond) if self.mapping_cond_in_proj is not None else 0 + degradation_emb = self.degradation_proj(degradation_params) if degradation_params is not None else 0 + cond = self.mapping(time_emb + aug_emb + class_emb + mapping_emb + degradation_emb) + + # Hourglass transformer + skips, poses = [], [] + for down_level, merge in zip(self.down_levels, self.merges): + x = down_level(x, pos, cond) + skips.append(x) + poses.append(pos) + x = merge(x) + pos = downscale_pos(pos) + + x = self.mid_level(x, pos, cond) + + for up_level, split, skip, pos in reversed(list(zip(self.up_levels, self.splits, skips, poses))): + x = split(x, skip) + x = up_level(x, pos, cond) + + # Unpatching + x = self.out_norm(x) + x = self.patch_out(x) + x = x.movedim(-1, -3) + + return x \ No newline at end of file diff --git a/arch/swinir/__init__.py b/arch/swinir/__init__.py new file mode 100755 index 0000000..e69de29 diff --git a/arch/swinir/swinir.py b/arch/swinir/swinir.py new file mode 100755 index 0000000..2f18836 --- /dev/null +++ b/arch/swinir/swinir.py @@ -0,0 +1,904 @@ +# ----------------------------------------------------------------------------------- +# SwinIR: Image Restoration Using Swin Transformer, https://arxiv.org/abs/2108.10257 +# Originally Written by Ze Liu, Modified by Jingyun Liang. +# ----------------------------------------------------------------------------------- +# Borrowed from DifFace (https://github.com/zsyOAOA/DifFace/blob/master/models/swinir.py) + +import math +from typing import Set + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.utils.checkpoint as checkpoint +from timm.models.layers import DropPath, to_2tuple, trunc_normal_ + + +class Mlp(nn.Module): + def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.): + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + self.fc1 = nn.Linear(in_features, hidden_features) + self.act = act_layer() + self.fc2 = nn.Linear(hidden_features, out_features) + self.drop = nn.Dropout(drop) + + def forward(self, x): + x = self.fc1(x) + x = self.act(x) + x = self.drop(x) + x = self.fc2(x) + x = self.drop(x) + return x + + +def window_partition(x, window_size): + """ + Args: + x: (B, H, W, C) + window_size (int): window size + + Returns: + windows: (num_windows*B, window_size, window_size, C) + """ + B, H, W, C = x.shape + x = x.view(B, H // window_size, window_size, W // window_size, window_size, C) + windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C) + return windows + + +def window_reverse(windows, window_size, H, W): + """ + Args: + windows: (num_windows*B, window_size, window_size, C) + window_size (int): Window size + H (int): Height of image + W (int): Width of image + + Returns: + x: (B, H, W, C) + """ + B = int(windows.shape[0] / (H * W / window_size / window_size)) + x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1) + x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1) + return x + + +class WindowAttention(nn.Module): + r""" Window based multi-head self attention (W-MSA) module with relative position bias. + It supports both of shifted and non-shifted window. + + Args: + dim (int): Number of input channels. + window_size (tuple[int]): The height and width of the window. + num_heads (int): Number of attention heads. + qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True + qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set + attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0 + proj_drop (float, optional): Dropout ratio of output. Default: 0.0 + """ + + def __init__(self, dim, window_size, num_heads, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0.): + + super().__init__() + self.dim = dim + self.window_size = window_size # Wh, Ww + self.num_heads = num_heads + head_dim = dim // num_heads + self.scale = qk_scale or head_dim ** -0.5 + + # define a parameter table of relative position bias + self.relative_position_bias_table = nn.Parameter( + torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)) # 2*Wh-1 * 2*Ww-1, nH + + # get pair-wise relative position index for each token inside the window + coords_h = torch.arange(self.window_size[0]) + coords_w = torch.arange(self.window_size[1]) + # coords = torch.stack(torch.meshgrid([coords_h, coords_w])) # 2, Wh, Ww + # Fix: Pass indexing="ij" to avoid warning + coords = torch.stack(torch.meshgrid([coords_h, coords_w], indexing="ij")) # 2, Wh, Ww + coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww + relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww + relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2 + relative_coords[:, :, 0] += self.window_size[0] - 1 # shift to start from 0 + relative_coords[:, :, 1] += self.window_size[1] - 1 + relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1 + relative_position_index = relative_coords.sum(-1) # Wh*Ww, Wh*Ww + self.register_buffer("relative_position_index", relative_position_index) + + self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) + self.attn_drop = nn.Dropout(attn_drop) + self.proj = nn.Linear(dim, dim) + + self.proj_drop = nn.Dropout(proj_drop) + + trunc_normal_(self.relative_position_bias_table, std=.02) + self.softmax = nn.Softmax(dim=-1) + + def forward(self, x, mask=None): + """ + Args: + x: input features with shape of (num_windows*B, N, C) + mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None + """ + B_, N, C = x.shape + qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) + q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple) + + q = q * self.scale + attn = (q @ k.transpose(-2, -1)) + + relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view( + self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1) # Wh*Ww,Wh*Ww,nH + relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww + attn = attn + relative_position_bias.unsqueeze(0) + + if mask is not None: + nW = mask.shape[0] + attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0) + attn = attn.view(-1, self.num_heads, N, N) + attn = self.softmax(attn) + else: + attn = self.softmax(attn) + + attn = self.attn_drop(attn) + + x = (attn @ v).transpose(1, 2).reshape(B_, N, C) + x = self.proj(x) + x = self.proj_drop(x) + return x + + def extra_repr(self) -> str: + return f'dim={self.dim}, window_size={self.window_size}, num_heads={self.num_heads}' + + def flops(self, N): + # calculate flops for 1 window with token length of N + flops = 0 + # qkv = self.qkv(x) + flops += N * self.dim * 3 * self.dim + # attn = (q @ k.transpose(-2, -1)) + flops += self.num_heads * N * (self.dim // self.num_heads) * N + # x = (attn @ v) + flops += self.num_heads * N * N * (self.dim // self.num_heads) + # x = self.proj(x) + flops += N * self.dim * self.dim + return flops + + +class SwinTransformerBlock(nn.Module): + r""" Swin Transformer Block. + + Args: + dim (int): Number of input channels. + input_resolution (tuple[int]): Input resulotion. + num_heads (int): Number of attention heads. + window_size (int): Window size. + shift_size (int): Shift size for SW-MSA. + mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. + qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True + qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set. + drop (float, optional): Dropout rate. Default: 0.0 + attn_drop (float, optional): Attention dropout rate. Default: 0.0 + drop_path (float, optional): Stochastic depth rate. Default: 0.0 + act_layer (nn.Module, optional): Activation layer. Default: nn.GELU + norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm + """ + + def __init__(self, dim, input_resolution, num_heads, window_size=7, shift_size=0, + mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0., + act_layer=nn.GELU, norm_layer=nn.LayerNorm): + super().__init__() + self.dim = dim + self.input_resolution = input_resolution + self.num_heads = num_heads + self.window_size = window_size + self.shift_size = shift_size + self.mlp_ratio = mlp_ratio + if min(self.input_resolution) <= self.window_size: + # if window size is larger than input resolution, we don't partition windows + self.shift_size = 0 + self.window_size = min(self.input_resolution) + assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size" + + self.norm1 = norm_layer(dim) + self.attn = WindowAttention( + dim, window_size=to_2tuple(self.window_size), num_heads=num_heads, + qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop) + + self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity() + self.norm2 = norm_layer(dim) + mlp_hidden_dim = int(dim * mlp_ratio) + self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop) + + if self.shift_size > 0: + attn_mask = self.calculate_mask(self.input_resolution) + else: + attn_mask = None + + self.register_buffer("attn_mask", attn_mask) + + def calculate_mask(self, x_size): + # calculate attention mask for SW-MSA + H, W = x_size + img_mask = torch.zeros((1, H, W, 1)) # 1 H W 1 + h_slices = (slice(0, -self.window_size), + slice(-self.window_size, -self.shift_size), + slice(-self.shift_size, None)) + w_slices = (slice(0, -self.window_size), + slice(-self.window_size, -self.shift_size), + slice(-self.shift_size, None)) + cnt = 0 + for h in h_slices: + for w in w_slices: + img_mask[:, h, w, :] = cnt + cnt += 1 + + mask_windows = window_partition(img_mask, self.window_size) # nW, window_size, window_size, 1 + mask_windows = mask_windows.view(-1, self.window_size * self.window_size) + attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2) + attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0)) + + return attn_mask + + def forward(self, x, x_size): + H, W = x_size + B, L, C = x.shape + # assert L == H * W, "input feature has wrong size" + + shortcut = x + x = self.norm1(x) + x = x.view(B, H, W, C) + + # cyclic shift + if self.shift_size > 0: + shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2)) + else: + shifted_x = x + + # partition windows + x_windows = window_partition(shifted_x, self.window_size) # nW*B, window_size, window_size, C + x_windows = x_windows.view(-1, self.window_size * self.window_size, C) # nW*B, window_size*window_size, C + + # W-MSA/SW-MSA (to be compatible for testing on images whose shapes are the multiple of window size + if self.input_resolution == x_size: + attn_windows = self.attn(x_windows, mask=self.attn_mask) # nW*B, window_size*window_size, C + else: + attn_windows = self.attn(x_windows, mask=self.calculate_mask(x_size).to(x.device)) + + # merge windows + attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C) + shifted_x = window_reverse(attn_windows, self.window_size, H, W) # B H' W' C + + # reverse cyclic shift + if self.shift_size > 0: + x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2)) + else: + x = shifted_x + x = x.view(B, H * W, C) + + # FFN + x = shortcut + self.drop_path(x) + x = x + self.drop_path(self.mlp(self.norm2(x))) + + return x + + def extra_repr(self) -> str: + return f"dim={self.dim}, input_resolution={self.input_resolution}, num_heads={self.num_heads}, " \ + f"window_size={self.window_size}, shift_size={self.shift_size}, mlp_ratio={self.mlp_ratio}" + + def flops(self): + flops = 0 + H, W = self.input_resolution + # norm1 + flops += self.dim * H * W + # W-MSA/SW-MSA + nW = H * W / self.window_size / self.window_size + flops += nW * self.attn.flops(self.window_size * self.window_size) + # mlp + flops += 2 * H * W * self.dim * self.dim * self.mlp_ratio + # norm2 + flops += self.dim * H * W + return flops + + +class PatchMerging(nn.Module): + r""" Patch Merging Layer. + + Args: + input_resolution (tuple[int]): Resolution of input feature. + dim (int): Number of input channels. + norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm + """ + + def __init__(self, input_resolution, dim, norm_layer=nn.LayerNorm): + super().__init__() + self.input_resolution = input_resolution + self.dim = dim + self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False) + self.norm = norm_layer(4 * dim) + + def forward(self, x): + """ + x: B, H*W, C + """ + H, W = self.input_resolution + B, L, C = x.shape + assert L == H * W, "input feature has wrong size" + assert H % 2 == 0 and W % 2 == 0, f"x size ({H}*{W}) are not even." + + x = x.view(B, H, W, C) + + x0 = x[:, 0::2, 0::2, :] # B H/2 W/2 C + x1 = x[:, 1::2, 0::2, :] # B H/2 W/2 C + x2 = x[:, 0::2, 1::2, :] # B H/2 W/2 C + x3 = x[:, 1::2, 1::2, :] # B H/2 W/2 C + x = torch.cat([x0, x1, x2, x3], -1) # B H/2 W/2 4*C + x = x.view(B, -1, 4 * C) # B H/2*W/2 4*C + + x = self.norm(x) + x = self.reduction(x) + + return x + + def extra_repr(self) -> str: + return f"input_resolution={self.input_resolution}, dim={self.dim}" + + def flops(self): + H, W = self.input_resolution + flops = H * W * self.dim + flops += (H // 2) * (W // 2) * 4 * self.dim * 2 * self.dim + return flops + + +class BasicLayer(nn.Module): + """ A basic Swin Transformer layer for one stage. + + Args: + dim (int): Number of input channels. + input_resolution (tuple[int]): Input resolution. + depth (int): Number of blocks. + num_heads (int): Number of attention heads. + window_size (int): Local window size. + mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. + qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True + qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set. + drop (float, optional): Dropout rate. Default: 0.0 + attn_drop (float, optional): Attention dropout rate. Default: 0.0 + drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0 + norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm + downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None + use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False. + """ + + def __init__(self, dim, input_resolution, depth, num_heads, window_size, + mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., + drop_path=0., norm_layer=nn.LayerNorm, downsample=None, use_checkpoint=False): + + super().__init__() + self.dim = dim + self.input_resolution = input_resolution + self.depth = depth + self.use_checkpoint = use_checkpoint + + # build blocks + self.blocks = nn.ModuleList([ + SwinTransformerBlock(dim=dim, input_resolution=input_resolution, + num_heads=num_heads, window_size=window_size, + shift_size=0 if (i % 2 == 0) else window_size // 2, + mlp_ratio=mlp_ratio, + qkv_bias=qkv_bias, qk_scale=qk_scale, + drop=drop, attn_drop=attn_drop, + drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path, + norm_layer=norm_layer) + for i in range(depth)]) + + # patch merging layer + if downsample is not None: + self.downsample = downsample(input_resolution, dim=dim, norm_layer=norm_layer) + else: + self.downsample = None + + def forward(self, x, x_size): + for blk in self.blocks: + if self.use_checkpoint: + x = checkpoint.checkpoint(blk, x, x_size) + else: + x = blk(x, x_size) + if self.downsample is not None: + x = self.downsample(x) + return x + + def extra_repr(self) -> str: + return f"dim={self.dim}, input_resolution={self.input_resolution}, depth={self.depth}" + + def flops(self): + flops = 0 + for blk in self.blocks: + flops += blk.flops() + if self.downsample is not None: + flops += self.downsample.flops() + return flops + + +class RSTB(nn.Module): + """Residual Swin Transformer Block (RSTB). + + Args: + dim (int): Number of input channels. + input_resolution (tuple[int]): Input resolution. + depth (int): Number of blocks. + num_heads (int): Number of attention heads. + window_size (int): Local window size. + mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. + qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True + qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set. + drop (float, optional): Dropout rate. Default: 0.0 + attn_drop (float, optional): Attention dropout rate. Default: 0.0 + drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0 + norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm + downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None + use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False. + img_size: Input image size. + patch_size: Patch size. + resi_connection: The convolutional block before residual connection. + """ + + def __init__(self, dim, input_resolution, depth, num_heads, window_size, + mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., + drop_path=0., norm_layer=nn.LayerNorm, downsample=None, use_checkpoint=False, + img_size=224, patch_size=4, resi_connection='1conv'): + super(RSTB, self).__init__() + + self.dim = dim + self.input_resolution = input_resolution + + self.residual_group = BasicLayer(dim=dim, + input_resolution=input_resolution, + depth=depth, + num_heads=num_heads, + window_size=window_size, + mlp_ratio=mlp_ratio, + qkv_bias=qkv_bias, qk_scale=qk_scale, + drop=drop, attn_drop=attn_drop, + drop_path=drop_path, + norm_layer=norm_layer, + downsample=downsample, + use_checkpoint=use_checkpoint) + + if resi_connection == '1conv': + self.conv = nn.Conv2d(dim, dim, 3, 1, 1) + elif resi_connection == '3conv': + # to save parameters and memory + self.conv = nn.Sequential(nn.Conv2d(dim, dim // 4, 3, 1, 1), nn.LeakyReLU(negative_slope=0.2, inplace=True), + nn.Conv2d(dim // 4, dim // 4, 1, 1, 0), + nn.LeakyReLU(negative_slope=0.2, inplace=True), + nn.Conv2d(dim // 4, dim, 3, 1, 1)) + + self.patch_embed = PatchEmbed( + img_size=img_size, patch_size=patch_size, in_chans=0, embed_dim=dim, + norm_layer=None) + + self.patch_unembed = PatchUnEmbed( + img_size=img_size, patch_size=patch_size, in_chans=0, embed_dim=dim, + norm_layer=None) + + def forward(self, x, x_size): + return self.patch_embed(self.conv(self.patch_unembed(self.residual_group(x, x_size), x_size))) + x + + def flops(self): + flops = 0 + flops += self.residual_group.flops() + H, W = self.input_resolution + flops += H * W * self.dim * self.dim * 9 + flops += self.patch_embed.flops() + flops += self.patch_unembed.flops() + + return flops + + +class PatchEmbed(nn.Module): + r""" Image to Patch Embedding + + Args: + img_size (int): Image size. Default: 224. + patch_size (int): Patch token size. Default: 4. + in_chans (int): Number of input image channels. Default: 3. + embed_dim (int): Number of linear projection output channels. Default: 96. + norm_layer (nn.Module, optional): Normalization layer. Default: None + """ + + def __init__(self, img_size=224, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None): + super().__init__() + img_size = to_2tuple(img_size) + patch_size = to_2tuple(patch_size) + patches_resolution = [img_size[0] // patch_size[0], img_size[1] // patch_size[1]] + self.img_size = img_size + self.patch_size = patch_size + self.patches_resolution = patches_resolution + self.num_patches = patches_resolution[0] * patches_resolution[1] + + self.in_chans = in_chans + self.embed_dim = embed_dim + + if norm_layer is not None: + self.norm = norm_layer(embed_dim) + else: + self.norm = None + + def forward(self, x): + x = x.flatten(2).transpose(1, 2) # B Ph*Pw C + if self.norm is not None: + x = self.norm(x) + return x + + def flops(self): + flops = 0 + H, W = self.img_size + if self.norm is not None: + flops += H * W * self.embed_dim + return flops + + +class PatchUnEmbed(nn.Module): + r""" Image to Patch Unembedding + + Args: + img_size (int): Image size. Default: 224. + patch_size (int): Patch token size. Default: 4. + in_chans (int): Number of input image channels. Default: 3. + embed_dim (int): Number of linear projection output channels. Default: 96. + norm_layer (nn.Module, optional): Normalization layer. Default: None + """ + + def __init__(self, img_size=224, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None): + super().__init__() + img_size = to_2tuple(img_size) + patch_size = to_2tuple(patch_size) + patches_resolution = [img_size[0] // patch_size[0], img_size[1] // patch_size[1]] + self.img_size = img_size + self.patch_size = patch_size + self.patches_resolution = patches_resolution + self.num_patches = patches_resolution[0] * patches_resolution[1] + + self.in_chans = in_chans + self.embed_dim = embed_dim + + def forward(self, x, x_size): + B, HW, C = x.shape + x = x.transpose(1, 2).view(B, self.embed_dim, x_size[0], x_size[1]) # B Ph*Pw C + return x + + def flops(self): + flops = 0 + return flops + + +class Upsample(nn.Sequential): + """Upsample module. + + Args: + scale (int): Scale factor. Supported scales: 2^n and 3. + num_feat (int): Channel number of intermediate features. + """ + + def __init__(self, scale, num_feat): + m = [] + if (scale & (scale - 1)) == 0: # scale = 2^n + for _ in range(int(math.log(scale, 2))): + m.append(nn.Conv2d(num_feat, 4 * num_feat, 3, 1, 1)) + m.append(nn.PixelShuffle(2)) + elif scale == 3: + m.append(nn.Conv2d(num_feat, 9 * num_feat, 3, 1, 1)) + m.append(nn.PixelShuffle(3)) + else: + raise ValueError(f'scale {scale} is not supported. ' 'Supported scales: 2^n and 3.') + super(Upsample, self).__init__(*m) + + +class UpsampleOneStep(nn.Sequential): + """UpsampleOneStep module (the difference with Upsample is that it always only has 1conv + 1pixelshuffle) + Used in lightweight SR to save parameters. + + Args: + scale (int): Scale factor. Supported scales: 2^n and 3. + num_feat (int): Channel number of intermediate features. + + """ + + def __init__(self, scale, num_feat, num_out_ch, input_resolution=None): + self.num_feat = num_feat + self.input_resolution = input_resolution + m = [] + m.append(nn.Conv2d(num_feat, (scale ** 2) * num_out_ch, 3, 1, 1)) + m.append(nn.PixelShuffle(scale)) + super(UpsampleOneStep, self).__init__(*m) + + def flops(self): + H, W = self.input_resolution + flops = H * W * self.num_feat * 3 * 9 + return flops + + +class SwinIR(nn.Module): + r""" SwinIR + A PyTorch impl of : `SwinIR: Image Restoration Using Swin Transformer`, based on Swin Transformer. + + Args: + img_size (int | tuple(int)): Input image size. Default 64 + patch_size (int | tuple(int)): Patch size. Default: 1 + in_chans (int): Number of input image channels. Default: 3 + embed_dim (int): Patch embedding dimension. Default: 96 + depths (tuple(int)): Depth of each Swin Transformer layer. + num_heads (tuple(int)): Number of attention heads in different layers. + window_size (int): Window size. Default: 7 + mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4 + qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True + qk_scale (float): Override default qk scale of head_dim ** -0.5 if set. Default: None + drop_rate (float): Dropout rate. Default: 0 + attn_drop_rate (float): Attention dropout rate. Default: 0 + drop_path_rate (float): Stochastic depth rate. Default: 0.1 + norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm. + ape (bool): If True, add absolute position embedding to the patch embedding. Default: False + patch_norm (bool): If True, add normalization after patch embedding. Default: True + use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False + sf: Upscale factor. 2/3/4/8 for image SR, 1 for denoising and compress artifact reduction + img_range: Image range. 1. or 255. + upsampler: The reconstruction reconstruction module. 'pixelshuffle'/'pixelshuffledirect'/'nearest+conv'/None + resi_connection: The convolutional block before residual connection. '1conv'/'3conv' + """ + + def __init__( + self, + img_size=64, + patch_size=1, + in_chans=3, + num_out_ch=3, + embed_dim=96, + depths=[6, 6, 6, 6], + num_heads=[6, 6, 6, 6], + window_size=7, + mlp_ratio=4., + qkv_bias=True, + qk_scale=None, + drop_rate=0., + attn_drop_rate=0., + drop_path_rate=0.1, + norm_layer=nn.LayerNorm, + ape=False, + patch_norm=True, + use_checkpoint=False, + sf=4, + img_range=1., + upsampler='', + resi_connection='1conv', + unshuffle=False, + unshuffle_scale=None, + hq_key: str = "jpg", + lq_key: str = "hint", + learning_rate: float = None, + weight_decay: float = None + ) -> "SwinIR": + super(SwinIR, self).__init__() + num_in_ch = in_chans * (unshuffle_scale ** 2) if unshuffle else in_chans + num_feat = 64 + self.img_range = img_range + if in_chans == 3: + rgb_mean = (0.4488, 0.4371, 0.4040) + self.mean = torch.Tensor(rgb_mean).view(1, 3, 1, 1) + else: + self.mean = torch.zeros(1, 1, 1, 1) + self.upscale = sf + self.upsampler = upsampler + self.window_size = window_size + self.unshuffle_scale = unshuffle_scale + self.unshuffle = unshuffle + + ##################################################################################################### + ################################### 1, shallow feature extraction ################################### + if unshuffle: + assert unshuffle_scale is not None + self.conv_first = nn.Sequential( + nn.PixelUnshuffle(sf), + nn.Conv2d(num_in_ch, embed_dim, 3, 1, 1), + ) + else: + self.conv_first = nn.Conv2d(num_in_ch, embed_dim, 3, 1, 1) + + ##################################################################################################### + ################################### 2, deep feature extraction ###################################### + self.num_layers = len(depths) + self.embed_dim = embed_dim + self.ape = ape + self.patch_norm = patch_norm + self.num_features = embed_dim + self.mlp_ratio = mlp_ratio + + # split image into non-overlapping patches + self.patch_embed = PatchEmbed( + img_size=img_size, patch_size=patch_size, in_chans=embed_dim, embed_dim=embed_dim, + norm_layer=norm_layer if self.patch_norm else None + ) + num_patches = self.patch_embed.num_patches + patches_resolution = self.patch_embed.patches_resolution + self.patches_resolution = patches_resolution + + # merge non-overlapping patches into image + self.patch_unembed = PatchUnEmbed( + img_size=img_size, patch_size=patch_size, in_chans=embed_dim, embed_dim=embed_dim, + norm_layer=norm_layer if self.patch_norm else None + ) + + # absolute position embedding + if self.ape: + self.absolute_pos_embed = nn.Parameter(torch.zeros(1, num_patches, embed_dim)) + trunc_normal_(self.absolute_pos_embed, std=.02) + + self.pos_drop = nn.Dropout(p=drop_rate) + + # stochastic depth + dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule + + # build Residual Swin Transformer blocks (RSTB) + self.layers = nn.ModuleList() + for i_layer in range(self.num_layers): + layer = RSTB( + dim=embed_dim, + input_resolution=(patches_resolution[0], patches_resolution[1]), + depth=depths[i_layer], + num_heads=num_heads[i_layer], + window_size=window_size, + mlp_ratio=self.mlp_ratio, + qkv_bias=qkv_bias, qk_scale=qk_scale, + drop=drop_rate, attn_drop=attn_drop_rate, + drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])], # no impact on SR results + norm_layer=norm_layer, + downsample=None, + use_checkpoint=use_checkpoint, + img_size=img_size, + patch_size=patch_size, + resi_connection=resi_connection + ) + self.layers.append(layer) + self.norm = norm_layer(self.num_features) + + # build the last conv layer in deep feature extraction + if resi_connection == '1conv': + self.conv_after_body = nn.Conv2d(embed_dim, embed_dim, 3, 1, 1) + elif resi_connection == '3conv': + # to save parameters and memory + self.conv_after_body = nn.Sequential( + nn.Conv2d(embed_dim, embed_dim // 4, 3, 1, 1), + nn.LeakyReLU(negative_slope=0.2, inplace=True), + nn.Conv2d(embed_dim // 4, embed_dim // 4, 1, 1, 0), + nn.LeakyReLU(negative_slope=0.2, inplace=True), + nn.Conv2d(embed_dim // 4, embed_dim, 3, 1, 1) + ) + + ##################################################################################################### + ################################ 3, high quality image reconstruction ################################ + if self.upsampler == 'pixelshuffle': + # for classical SR + self.conv_before_upsample = nn.Sequential( + nn.Conv2d(embed_dim, num_feat, 3, 1, 1), + nn.LeakyReLU(inplace=True) + ) + self.upsample = Upsample(sf, num_feat) + self.conv_last = nn.Conv2d(num_feat, num_out_ch, 3, 1, 1) + elif self.upsampler == 'pixelshuffledirect': + # for lightweight SR (to save parameters) + self.upsample = UpsampleOneStep( + sf, embed_dim, num_out_ch, + (patches_resolution[0], patches_resolution[1]) + ) + elif self.upsampler == 'nearest+conv': + # for real-world SR (less artifacts) + self.conv_before_upsample = nn.Sequential( + nn.Conv2d(embed_dim, num_feat, 3, 1, 1), + nn.LeakyReLU(inplace=True) + ) + self.conv_up1 = nn.Conv2d(num_feat, num_feat, 3, 1, 1) + if self.upscale == 4: + self.conv_up2 = nn.Conv2d(num_feat, num_feat, 3, 1, 1) + elif self.upscale == 8: + self.conv_up2 = nn.Conv2d(num_feat, num_feat, 3, 1, 1) + self.conv_up3 = nn.Conv2d(num_feat, num_feat, 3, 1, 1) + self.conv_hr = nn.Conv2d(num_feat, num_feat, 3, 1, 1) + self.conv_last = nn.Conv2d(num_feat, num_out_ch, 3, 1, 1) + self.lrelu = nn.LeakyReLU(negative_slope=0.2, inplace=True) + else: + # for image denoising and JPEG compression artifact reduction + self.conv_last = nn.Conv2d(embed_dim, num_out_ch, 3, 1, 1) + + self.apply(self._init_weights) + + def _init_weights(self, m: nn.Module) -> None: + if isinstance(m, nn.Linear): + trunc_normal_(m.weight, std=.02) + if isinstance(m, nn.Linear) and m.bias is not None: + nn.init.constant_(m.bias, 0) + elif isinstance(m, nn.LayerNorm): + nn.init.constant_(m.bias, 0) + nn.init.constant_(m.weight, 1.0) + + # TODO: What's this ? + @torch.jit.ignore + def no_weight_decay(self) -> Set[str]: + return {'absolute_pos_embed'} + + @torch.jit.ignore + def no_weight_decay_keywords(self) -> Set[str]: + return {'relative_position_bias_table'} + + def check_image_size(self, x: torch.Tensor) -> torch.Tensor: + _, _, h, w = x.size() + mod_pad_h = (self.window_size - h % self.window_size) % self.window_size + mod_pad_w = (self.window_size - w % self.window_size) % self.window_size + x = F.pad(x, (0, mod_pad_w, 0, mod_pad_h), 'reflect') + return x + + def forward_features(self, x: torch.Tensor) -> torch.Tensor: + x_size = (x.shape[2], x.shape[3]) + x = self.patch_embed(x) + if self.ape: + x = x + self.absolute_pos_embed + x = self.pos_drop(x) + + for layer in self.layers: + x = layer(x, x_size) + + x = self.norm(x) # B L C + x = self.patch_unembed(x, x_size) + + return x + + def forward(self, x: torch.Tensor) -> torch.Tensor: + H, W = x.shape[2:] + x = self.check_image_size(x) + + self.mean = self.mean.type_as(x) + x = (x - self.mean) * self.img_range + + if self.upsampler == 'pixelshuffle': + # for classical SR + x = self.conv_first(x) + x = self.conv_after_body(self.forward_features(x)) + x + x = self.conv_before_upsample(x) + x = self.conv_last(self.upsample(x)) + elif self.upsampler == 'pixelshuffledirect': + # for lightweight SR + x = self.conv_first(x) + x = self.conv_after_body(self.forward_features(x)) + x + x = self.upsample(x) + elif self.upsampler == 'nearest+conv': + # for real-world SR + x = self.conv_first(x) + x = self.conv_after_body(self.forward_features(x)) + x + x = self.conv_before_upsample(x) + x = self.lrelu(self.conv_up1(torch.nn.functional.interpolate(x, scale_factor=2, mode='nearest'))) + if self.upscale == 4: + x = self.lrelu(self.conv_up2(torch.nn.functional.interpolate(x, scale_factor=2, mode='nearest'))) + elif self.upscale == 8: + x = self.lrelu(self.conv_up2(torch.nn.functional.interpolate(x, scale_factor=2, mode='nearest'))) + x = self.lrelu(self.conv_up3(torch.nn.functional.interpolate(x, scale_factor=2, mode='nearest'))) + x = self.conv_last(self.lrelu(self.conv_hr(x))) + else: + # for image denoising and JPEG compression artifact reduction + x_first = self.conv_first(x) + res = self.conv_after_body(self.forward_features(x_first)) + x_first + x = x + self.conv_last(res) + + x = x / self.img_range + self.mean + + return x[:, :, :H * self.upscale, :W * self.upscale] + + def flops(self) -> int: + flops = 0 + H, W = self.patches_resolution + flops += H * W * 3 * self.embed_dim * 9 + flops += self.patch_embed.flops() + for i, layer in enumerate(self.layers): + flops += layer.flops() + flops += H * W * 3 * self.embed_dim * self.embed_dim + flops += self.upsample.flops() + return flops diff --git a/lightning_models/mmse_rectified_flow.py b/lightning_models/mmse_rectified_flow.py new file mode 100755 index 0000000..7db7ee3 --- /dev/null +++ b/lightning_models/mmse_rectified_flow.py @@ -0,0 +1,318 @@ +import os +from contextlib import contextmanager, nullcontext + +import torch +#import wandb +from pytorch_lightning import LightningModule +from torch.nn.functional import mse_loss +from torch.nn.functional import sigmoid +from torch.optim import AdamW +from torch_ema import ExponentialMovingAverage as EMA +from torchmetrics.image.fid import FrechetInceptionDistance +from torchmetrics.image.inception import InceptionScore +from torchvision.transforms.functional import to_pil_image +from torchvision.utils import save_image +from ..utils.create_arch import create_arch +from huggingface_hub import PyTorchModelHubMixin + + + +class MMSERectifiedFlow(LightningModule, + PyTorchModelHubMixin, + pipeline_tag="image-to-image", + license="mit", + ): + def __init__(self, + stage, + arch, + conditional=False, + mmse_model_ckpt_path=None, + mmse_model_arch=None, + lr=5e-4, + weight_decay=1e-3, + betas=(0.9, 0.95), + mmse_noise_std=0.1, + num_flow_steps=50, + ema_decay=0.9999, + eps=0.0, + t_schedule='stratified_uniform', + *args, + **kwargs + ): + super().__init__() + self.save_hyperparameters(logger=False) + + if stage == 'flow': + if conditional: + condition_channels = 3 + else: + condition_channels = 0 + if mmse_model_arch is None and 'colorization' in kwargs and kwargs['colorization']: + condition_channels //= 3 + self.model = create_arch(arch, condition_channels) + self.mmse_model = create_arch(mmse_model_arch, 0) if mmse_model_arch is not None else None + if mmse_model_ckpt_path is not None: + ckpt = torch.load(mmse_model_ckpt_path, map_location="cpu") + if mmse_model_arch is None: + mmse_model_arch = ckpt['hyper_parameters']['arch'] + self.mmse_model = create_arch(mmse_model_arch, 0) + if 'ema' in ckpt: + # ema_decay doesn't affect anything here, because we are doing load_state_dict + mmse_ema = EMA(self.mmse_model.parameters(), decay=ema_decay) + mmse_ema.load_state_dict(ckpt['ema']) + mmse_ema.copy_to() + elif 'params_ema' in ckpt: + self.mmse_model.load_state_dict(ckpt['params_ema']) + else: + state_dict = ckpt['state_dict'] + state_dict = {layer_name.replace('model.', ''): weights for layer_name, weights in + state_dict.items()} + state_dict = {layer_name.replace('module.', ''): weights for layer_name, weights in + state_dict.items()} + self.mmse_model.load_state_dict(state_dict) + for param in self.mmse_model.parameters(): + param.requires_grad = False + self.mmse_model.eval() + else: + assert stage == 'mmse' or stage == 'naive_flow' + assert not conditional + self.model = create_arch(arch, 0) + self.mmse_model = None + if 'flow' in stage: + self.fid = FrechetInceptionDistance(reset_real_features=True, normalize=True) + self.inception_score = InceptionScore(normalize=True) + + self.ema = EMA(self.model.parameters(), decay=ema_decay) if self.ema_wanted else None + self.test_results_path = None + + @property + def ema_wanted(self): + return self.hparams.ema_decay != -1 + + def on_save_checkpoint(self, checkpoint: dict) -> None: + if self.ema_wanted: + checkpoint['ema'] = self.ema.state_dict() + return super().on_save_checkpoint(checkpoint) + + def on_load_checkpoint(self, checkpoint: dict) -> None: + if self.ema_wanted: + self.ema.load_state_dict(checkpoint['ema']) + return super().on_load_checkpoint(checkpoint) + + def on_before_zero_grad(self, optimizer) -> None: + if self.ema_wanted: + self.ema.update(self.model.parameters()) + return super().on_before_zero_grad(optimizer) + + def to(self, *args, **kwargs): + if self.ema_wanted: + self.ema.to(*args, **kwargs) + return super().to(*args, **kwargs) + + # This will use the contextmanager of ema, to copy the EMA weights to the flow model during validation, and then restore them for training. + @contextmanager + def maybe_ema(self): + ema = self.ema + ctx = nullcontext if ema is None else ema.average_parameters + yield ctx + + def forward_mmse(self, y): + return self.model(y).clip(0, 1) + + def forward_flow(self, x_t, t, y=None): + if self.hparams.conditional: + if self.mmse_model is not None: + with torch.no_grad(): + self.mmse_model.eval() + condition = self.mmse_model(y).clip(0, 1) + else: + condition = y + x_t = torch.cat((x_t, condition), dim=1) + return self.model(x_t, t) + + def forward(self, x_t, t, y): + if 'flow' in self.hparams.stage: + return self.forward_flow(x_t, t, y) + else: + return self.forward_mmse(y) + + @torch.no_grad() + def create_source_distribution_samples(self, x, y, non_noisy_z0): + with torch.no_grad(): + if self.hparams.conditional: + source_dist_samples = torch.randn_like(x) + else: + if self.hparams.stage == 'flow': + if non_noisy_z0 is None: + self.mmse_model.eval() + non_noisy_z0 = self.mmse_model(y).clip(0, 1) + source_dist_samples = non_noisy_z0 + torch.randn_like(non_noisy_z0) * self.hparams.mmse_noise_std + else: + assert self.hparams.stage == 'naive_flow' + if non_noisy_z0 is not None: + source_dist_samples = non_noisy_z0 + else: + source_dist_samples = y + if source_dist_samples.shape[1] != x.shape[1]: + assert source_dist_samples.shape[1] == 1 # Colorization + source_dist_samples = source_dist_samples.expand(-1, x.shape[1], -1, -1) + if self.hparams.mmse_noise_std is not None: + source_dist_samples = source_dist_samples + torch.randn_like(source_dist_samples) * self.hparams.mmse_noise_std + return source_dist_samples + + @staticmethod + def stratified_uniform(bs, group=0, groups=1, dtype=None, device=None): + if groups <= 0: + raise ValueError(f"groups must be positive, got {groups}") + if group < 0 or group >= groups: + raise ValueError(f"group must be in [0, {groups})") + n = bs * groups + offsets = torch.arange(group, n, groups, dtype=dtype, device=device) + u = torch.rand(bs, dtype=dtype, device=device) + return ((offsets + u) / n).view(bs, 1, 1, 1) + + def generate_random_t(self, bs, dtype=None): + if self.hparams.t_schedule == 'logit-normal': + return sigmoid(torch.randn(bs, 1, 1, 1, device=self.device)) * (1.0 - self.hparams.eps) + self.hparams.eps + elif self.hparams.t_schedule == 'uniform': + return torch.rand(bs, 1, 1, 1, device=self.device) * (1.0 - self.hparams.eps) + self.hparams.eps + elif self.hparams.t_schedule == 'stratified_uniform': + return self.stratified_uniform(bs, self.trainer.global_rank, self.trainer.world_size, dtype=dtype, + device=self.device) * (1.0 - self.hparams.eps) + self.hparams.eps + else: + raise NotImplementedError() + + def training_step(self, batch, batch_idx): + x = batch['x'] + y = batch['y'] + non_noisy_z0 = batch['non_noisy_z0'] if 'non_noisy_z0' in batch else None + if 'flow' in self.hparams.stage: + with torch.no_grad(): + t = self.generate_random_t(x.shape[0], dtype=x.dtype) + source_dist_samples = self.create_source_distribution_samples(x, y, non_noisy_z0) + x_t = t * x + (1.0 - t) * source_dist_samples + v_t = self(x_t, t.squeeze(), y) + loss = mse_loss(v_t, x - source_dist_samples) + else: + xhat = self(x_t=None, t=None, y=y) + loss = mse_loss(xhat, x) + self.log("train/loss", loss) + return loss + + @torch.no_grad() + def generate_reconstructions(self, x, y, non_noisy_z0, num_flow_steps, result_device): + with self.maybe_ema(): + if 'flow' in self.hparams.stage: + source_dist_samples = self.create_source_distribution_samples(x, y, non_noisy_z0) + + dt = (1.0 / num_flow_steps) * (1.0 - self.hparams.eps) + x_t_next = source_dist_samples.clone() + x_t_seq = [x_t_next] + t_one = torch.ones(x.shape[0], device=self.device) + for i in range(num_flow_steps): + num_t = (i / num_flow_steps) * (1.0 - self.hparams.eps) + self.hparams.eps + v_t_next = self(x_t=x_t_next, t=t_one * num_t, y=y).to(x_t_next.dtype) + x_t_next = x_t_next.clone() + v_t_next * dt + x_t_seq.append(x_t_next.to(result_device)) + + xhat = x_t_seq[-1].clip(0, 1).to(torch.float32) + source_dist_samples = source_dist_samples.to(result_device) + else: + xhat = self(x_t=None, t=None, y=y).to(torch.float32) + x_t_seq = None + source_dist_samples = None + return xhat.to(result_device), x_t_seq, source_dist_samples + + def validation_step(self, batch, batch_idx): + x = batch['x'] + y = batch['y'] + non_noisy_z0 = batch['non_noisy_z0'] if 'non_noisy_z0' in batch else None + xhat, x_t_seq, source_dist_samples = self.generate_reconstructions(x, y, non_noisy_z0, self.hparams.num_flow_steps, + self.device) + x = x.to(torch.float32) + y = y.to(torch.float32) + self.log_dict({"val_metrics/mse": ((x - xhat) ** 2).mean()}, on_step=False, on_epoch=True, sync_dist=True, + batch_size=x.shape[0]) + + if 'flow' in self.hparams.stage: + self.fid.update(x, real=True) + self.fid.update(xhat, real=False) + self.inception_score.update(xhat) + + if batch_idx == 0: + ''' + wandb_logger = self.logger.experiment + wandb_logger.log({'val_images/x': [wandb.Image(to_pil_image(create_grid(x)))], + 'val_images/y': [wandb.Image(to_pil_image(create_grid(y.clip(0, 1))))], + 'val_images/xhat': [wandb.Image(to_pil_image(create_grid(xhat)))], }) + if 'flow' in self.hparams.stage: + wandb_logger.log({'val_images/x_t_seq': [wandb.Image(to_pil_image(create_grid( + torch.cat([elem[0].unsqueeze(0).to(torch.float32) for elem in x_t_seq], dim=0).clip(0, 1), + num_images=len(x_t_seq))))], 'val_images/source_distribution_samples': [ + wandb.Image(to_pil_image(create_grid(source_dist_samples.clip(0, 1).to(torch.float32))))]}) + if self.mmse_model is not None: + xhat_mmse = self.mmse_model(y).clip(0, 1) + wandb_logger.log({'val_images/xhat_mmse': [ + wandb.Image(to_pil_image(create_grid(xhat_mmse.to(torch.float32))))]}) + ''' + + def on_validation_epoch_end(self): + if 'flow' in self.hparams.stage: + inception_score_mean, inception_score_std = self.inception_score.compute() + self.log_dict( + {'val_metrics/fid': self.fid.compute(), + 'val_metrics/inception_score_mean': inception_score_mean, + 'val_metrics/inception_score_std': inception_score_std}, + on_epoch=True, on_step=False, sync_dist=True, + batch_size=1) + self.fid.reset() + self.inception_score.reset() + + def test_step(self, batch, batch_idx): + assert self.test_results_path is not None, "Please set test_results_path before testing." + assert os.path.isdir(self.test_results_path), 'Please make sure the test_result_path dir exists.' + + def save_image_batch(images, folder, image_file_names): + os.makedirs(folder, exist_ok=True) + for i, img in enumerate(images): + save_image(images[i].clip(0, 1), os.path.join(folder, image_file_names[i])) + + os.makedirs(self.test_results_path, exist_ok=True) + x = batch['x'] + y = batch['y'] + non_noisy_z0 = batch['non_noisy_z0'] if 'non_noisy_z0' in batch else None + y_path = os.path.join(self.test_results_path, 'y') + save_image_batch(y, y_path, batch['img_file_name']) + + if 'flow' in self.hparams.stage: + source_dist_samples_to_save = None + + for num_flow_steps in self.num_test_flow_steps: + xhat, x_t_seq, source_dist_samples = self.generate_reconstructions(x, y, non_noisy_z0, num_flow_steps, + torch.device("cpu")) + xhat_path = os.path.join(self.test_results_path, f"num_flow_steps={num_flow_steps}", 'xhat') + save_image_batch(xhat, xhat_path, batch['img_file_name']) + if source_dist_samples_to_save is None: + source_dist_samples_to_save = source_dist_samples + + source_distribution_samples_path = os.path.join(self.test_results_path, 'source_distribution_samples') + save_image_batch(source_dist_samples_to_save, source_distribution_samples_path, batch['img_file_name']) + if self.mmse_model is not None: + mmse_estimates = self.mmse_model(y).clip(0, 1) + mmse_samples_path = os.path.join(self.test_results_path, 'mmse_samples') + save_image_batch(mmse_estimates, mmse_samples_path, batch['img_file_name']) + + + else: + xhat, _, _ = self.generate_reconstructions(x, y, non_noisy_z0, None, torch.device('cpu')) + xhat_path = os.path.join(self.test_results_path, 'xhat') + save_image_batch(xhat, xhat_path, batch['img_file_name']) + + def configure_optimizers(self): + # Add here a learning rate scheduler if you wish to do so. + optimizer = AdamW(self.model.parameters(), + betas=self.hparams.betas, + eps=1e-8, + lr=self.hparams.lr, + weight_decay=self.hparams.weight_decay) + return optimizer diff --git a/nodes.py b/nodes.py new file mode 100755 index 0000000..237ac16 --- /dev/null +++ b/nodes.py @@ -0,0 +1,197 @@ +# Code from https://huggingface.co/spaces/ohayonguy/PMRF/blob/main/app.py +# Some of the implementations below are adopted from +# https://huggingface.co/spaces/sczhou/CodeFormer and https://huggingface.co/spaces/wzhouxiff/RestoreFormerPlusPlus +import cv2 +from tqdm import tqdm +import torch +from basicsr.archs.rrdbnet_arch import RRDBNet +from basicsr.utils import img2tensor, tensor2img +from facexlib.utils.face_restoration_helper import FaceRestoreHelper +from realesrgan.utils import RealESRGANer +from .lightning_models.mmse_rectified_flow import MMSERectifiedFlow +import torchvision +import numpy as np +from PIL import Image +import os +import folder_paths + +device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + +def set_realesrgan(): + use_half = False + if torch.cuda.is_available(): # set False in CPU/MPS mode + no_half_gpu_list = ["1650", "1660"] # set False for GPUs that don't support f16 + if not True in [ + gpu in torch.cuda.get_device_name(0) for gpu in no_half_gpu_list + ]: + use_half = True + + model = RRDBNet( + num_in_ch=3, num_out_ch=3, num_feat=64, num_block=23, num_grow_ch=32, scale=2, + ) + upscale_models_path = os.path.join(folder_paths.models_dir, "upscale_models") + realesrgan_path = os.path.join(upscale_models_path, "RealESRGAN_x2plus.pth") + upsampler = RealESRGANer( + scale=2, + model_path=realesrgan_path, + model=model, + tile=400, + tile_pad=40, + pre_pad=0, + half=use_half, + ) + return upsampler + +pmrf_path = os.path.join(folder_paths.models_dir, "pmrf") +upsampler = set_realesrgan() +pmrf = MMSERectifiedFlow.from_pretrained(pmrf_path).to(device=device) + +def generate_reconstructions(pmrf_model, x, y, non_noisy_z0, num_flow_steps, device): + source_dist_samples = pmrf_model.create_source_distribution_samples( + x, y, non_noisy_z0 + ) + dt = (1.0 / num_flow_steps) * (1.0 - pmrf_model.hparams.eps) + x_t_next = source_dist_samples.clone() + t_one = torch.ones(x.shape[0], device=device) + for i in tqdm(range(num_flow_steps)): + num_t = (i / num_flow_steps) * ( + 1.0 - pmrf_model.hparams.eps + ) + pmrf_model.hparams.eps + v_t_next = pmrf_model(x_t=x_t_next, t=t_one * num_t, y=y).to(x_t_next.dtype) + x_t_next = x_t_next.clone() + v_t_next * dt + return x_t_next.clip(0, 1) + +def resize(img, size, interpolation): + # From https://github.com/sczhou/CodeFormer/blob/master/facelib/utils/face_restoration_helper.py + h, w = img.shape[0:2] + scale = float(size) / float(min(h, w)) + h, w = int(round(h * scale)), int(round(w * scale)) + return cv2.resize(img, (w, h), interpolation=interpolation) + +@torch.inference_mode() +def enhance_face(img, face_helper, num_flow_steps, scale=2, interpolation=cv2.INTER_LANCZOS4): + face_helper.clean_all() + face_helper.read_image(img) + face_helper.input_img = resize(face_helper.input_img, 640, interpolation) + face_helper.get_face_landmarks_5(only_center_face=False, eye_dist_threshold=5) + face_helper.align_warp_face() + + # face restoration + for i, cropped_face in tqdm(enumerate(face_helper.cropped_faces)): + cropped_face_t = img2tensor(cropped_face / 255.0, bgr2rgb=True, float32=True) + cropped_face_t = cropped_face_t.unsqueeze(0).to(device) + + output = generate_reconstructions( + pmrf, + torch.zeros_like(cropped_face_t), + cropped_face_t, + None, + num_flow_steps, + device, + ) + restored_face = tensor2img( + output.to(torch.float32).squeeze(0), rgb2bgr=True, min_max=(0, 1) + ) + restored_face = restored_face.astype("uint8") + face_helper.add_restored_face(restored_face) + + # upsample the background + # Now only support RealESRGAN for upsampling background + bg_img = upsampler.enhance(img, outscale=scale)[0] + face_helper.get_inverse_affine(None) + # paste each restored face to the input image + restored_img = face_helper.paste_faces_to_input_image(upsample_img=bg_img) + return face_helper.cropped_faces, face_helper.restored_faces, restored_img + +@torch.inference_mode() +def inference( + imgs, + scale, + num_flow_steps, + seed, + interpolation, +): + torch.manual_seed(seed) + if interpolation == "lanczos4": + interpolation = cv2.INTER_LANCZOS4 + elif interpolation == "nearest": + interpolation = cv2.INTER_NEAREST + elif interpolation == "linear": + interpolation = cv2.INTER_LINEAR + elif interpolation == "cubic": + interpolation = cv2.INTER_CUBIC + elif interpolation == "area": + interpolation = cv2.INTER_AREA + elif interpolation == "linear_exact": + interpolation = cv2.INTER_LINEAR_EXACT + elif interpolation == "nearest_exact": + interpolation = cv2.INTER_NEAREST_EXACT + imgs_output = [] + for img in imgs: + img = img.permute(2, 0, 1) + img = torchvision.transforms.functional.to_pil_image(img.clamp(0, 1)).convert("RGB") + img = np.array(img) + img = img[:, :, ::-1].copy() + h, w = img.shape[0:2] + size = min(h, w) + + face_scale = scale*(size/640) + face_scale = face_scale if face_scale < scale else scale + face_scale = face_scale if face_scale > 1.0 else 1.0 + + face_helper = FaceRestoreHelper( + face_scale, + face_size=512, + crop_ratio=(1, 1), + det_model="retinaface_resnet50", + save_ext="png", + use_parse=True, + device=device, + model_rootpath=None, + ) + + cropped_face, restored_faces, restored_img = enhance_face( + img, + face_helper, + num_flow_steps=num_flow_steps, + scale=face_scale, + interpolation=interpolation, + ) + + output = restored_img + output = cv2.cvtColor(output, cv2.COLOR_BGR2RGB) + output = resize(output, size*scale, interpolation) + + torch.cuda.empty_cache() + output = torchvision.transforms.functional.pil_to_tensor(Image.fromarray(output)).to(torch.float32) / 255.0 + output = output.permute(1, 2, 0) + imgs_output.append(output[None,]) + return (torch.cat(tuple(imgs_output), dim=0),) + +class PMRF: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "images": ("IMAGE",), + "scale": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 40.0, "step": 0.1}), + "num_steps": ("INT", {"default": 25, "min": 1, "max": 400, "step": 1}), + "seed": ("INT", {"default": 123, "min": 0, "max": 2**32, "step": 1}), + "interpolation": (["lanczos4", "nearest", "linear", "cubic", "area", "linear_exact", "nearest_exact"],) + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images", ) + FUNCTION = "pmrf" + CATEGORY = "PMRF" + + def pmrf(self, images, scale, num_steps, seed, interpolation): + return inference(images, scale, num_steps, seed, interpolation) + +NODE_CLASS_MAPPINGS = { + "PMRF": PMRF, +} +NODE_DISPLAY_NAME_MAPPINGS = { + "PMRF": "PMRF", +} diff --git a/prestartup_script.py b/prestartup_script.py new file mode 100755 index 0000000..35be46b --- /dev/null +++ b/prestartup_script.py @@ -0,0 +1,172 @@ +import pkg_resources +import subprocess +import sys +import huggingface_hub +import importlib.util +import importlib.metadata +import folder_paths +import os +import pathlib +from packaging.version import Version +import time + +pmrf_path = os.path.join(folder_paths.models_dir, "pmrf") +pmrf_model_path = os.path.join(pmrf_path, "model.safetensors") +pmrf_model_json_path = os.path.join(pmrf_path, "config.json") +if not (os.path.exists(pmrf_model_path) and os.path.exists(pmrf_model_json_path)): + print("Downloading PMRF model from ohayonguy/PMRF_blind_face_image_restoration...") + if not os.path.exists(pmrf_path): + os.makedirs(pmrf_path) + huggingface_hub.snapshot_download( + repo_id="ohayonguy/PMRF_blind_face_image_restoration", + local_dir=pmrf_path, + ) +upscale_models_path = os.path.join(folder_paths.models_dir, "upscale_models") +models = ["RealESRGAN_x2plus.pth", "RealESRGAN_x4plus.pth"] +for model in models: + realesrgan_path = os.path.join(upscale_models_path, model) + if not os.path.exists(realesrgan_path): + print(f"Downloading {model} model from 2kpr/Real-ESRGAN...") + huggingface_hub.snapshot_download( + repo_id="2kpr/Real-ESRGAN", + allow_patterns=model, + local_dir=upscale_models_path, + ) + +packages = [ + {"name": "realesrgan", "version": "0.2.5"}, + {"name": "torchvision", "version": "0.19.0"}, + {"name": "torch_fidelity", "version": "0.3.0"}, + {"name": "torch_ema", "version": "0.3"}, + {"name": "pytorch_lightning", "version": "2.4.0"}, + {"name": "timm", "version": "1.0.7"}, +] + +for package in packages: + if importlib.util.find_spec(package["name"]): + #print(f'Found package {package["name"]}') + #print(f'Version: {package["version"]}') + #print(f'Version: {importlib.metadata.version(package["name"])}') + if Version(package["version"]) > Version(importlib.metadata.version(package["name"])): + print(f'Updating {package["name"]} for PMRF...') + subprocess.check_call([sys.executable, "-m", "pip", "install", f'{package["name"]}>={package["version"]}', "--upgrade"]) + else: + print(f'Installing {package["name"]} for PMRF...') + subprocess.check_call([sys.executable, "-m", "pip", "install", f'{package["name"]}>={package["version"]}', "--upgrade"]) + +if importlib.util.find_spec("basicsr"): + path = pathlib.Path(importlib.util.find_spec("basicsr").origin).parent.joinpath("data/degradations.py") + if os.path.exists(path): + with open(path, "r", encoding="utf-8") as f: + content = f.read() + if "from torchvision.transforms.functional_tensor import rgb_to_grayscale" in content: + print(f"Patching basicsr with fix from https://github.com/XPixelGroup/BasicSR/pull/650 for PMRF...") + content = content.replace( + "from torchvision.transforms.functional_tensor import rgb_to_grayscale", + "from torchvision.transforms.functional import rgb_to_grayscale", + ) + with open(path, "w", encoding="utf-8") as f: + f.write(content) + +if not importlib.util.find_spec("natten"): + print(f'Installing natten for PMRF...') + cuda_version = "" + torch_version = "" + print("Searching for CUDA and Torch versions for installing atten needed by PMRF...") + for p in pkg_resources.working_set: + if p.project_name.startswith("nvidia-cuda-runtime"): + if p.version.startswith("12.4"): + cuda_version = "cu124" + print("- Found CUDA 12.4") + elif p.version.startswith("12.1"): + cuda_version = "cu121" + print("- Found CUDA 12.1") + elif p.version.startswith("11.8"): + cuda_version = "cu118" + print("- Found CUDA 11.8") + elif p.project_name == "torch": + if p.version.startswith("2.4"): + torch_version = "torch240" + print("- Found Torch 2.4") + elif p.version.startswith("2.3"): + torch_version = "torch230" + print("- Found Torch 2.3") + elif p.version.startswith("2.2"): + torch_version = "torch220" + print("- Found Torch 2.2") + elif p.version.startswith("2.1"): + torch_version = "torch210" + print("- Found Torch 2.1") + if cuda_version == "": + py_path = os.path.join(folder_paths.temp_directory, "torchcudaversion.py") + if not os.path.exists(py_path): + if not os.path.exists(folder_paths.temp_directory): + os.makedirs(folder_paths.temp_directory) + with open(py_path, "w", encoding="utf-8") as f: + f.write("import torch\nprint(torch.version.cuda)") + cuda_version = subprocess.check_output([sys.executable, f"{py_path}"]).decode().strip() + if cuda_version == "12.4": + cuda_version = "cu124" + print("- Found CUDA 12.4") + elif cuda_version == "12.1": + cuda_version = "cu121" + print("- Found CUDA 12.1") + elif cuda_version == "11.8": + cuda_version = "cu118" + print("- Found CUDA 11.8") + if cuda_version == "": + print("************************************") + print("Error: Can't find CUDA runtime version, can't install natten") + print(" PMRF will not work until natten is installed, see https://github.com/SHI-Labs/NATTEN for help in installing natten.") + print("************************************") + time.sleep(4) + elif torch_version == "": + print("************************************") + print("Error: Can't find torch version, can't install natten") + print(" PMRF will not work until natten is installed, see https://github.com/SHI-Labs/NATTEN for help in installing natten.") + print("************************************") + time.sleep(4) + elif cuda_version == "cu124" and torch_version != "torch240": + print("************************************") + print("Error: Can't install natten, which is needed by PMRF since CUDA runtime version is 12.4 but torch is not version 2.4") + print(" PMRF will not work until natten is installed, see https://github.com/SHI-Labs/NATTEN for help in installing natten.") + print("************************************") + time.sleep(4) + elif os.name == "nt" and cuda_version != "cu124": + print("************************************") + print("Error: Can't install natten on windows if CUDA runtime version is not 12.4 unless you build natten yourself, see https://github.com/SHI-Labs/NATTEN/blob/main/docs/install.md#build-with-msvc") + print(" PMRF will not work until natten is installed, see https://github.com/SHI-Labs/NATTEN for help in installing natten.") + print("************************************") + time.sleep(4) + elif os.name == "nt" and torch_version != "torch240": + print("************************************") + print("Error: Can't install natten on windows if torch version is not 2.4 unless you build natten yourself, see https://github.com/SHI-Labs/NATTEN/blob/main/docs/install.md#build-with-msvc") + print(" PMRF will not work until natten is installed, see https://github.com/SHI-Labs/NATTEN for help in installing natten.") + print("************************************") + time.sleep(4) + elif os.name == "nt" and (sys.version_info[1] < 10 or sys.version_info[1] > 12): + print("************************************") + print("Error: Can't install natten on windows if python version isn't 3.10, 3.11, or 3.12, unless you build natten yourself, see https://github.com/SHI-Labs/NATTEN/blob/main/docs/install.md#build-with-msvc") + print(" PMRF will not work until natten is installed, see https://github.com/SHI-Labs/NATTEN for help in installing natten.") + print("************************************") + time.sleep(4) + elif os.name == "nt": + if sys.version_info[1] == 10: + whl = "natten-0.17.2.dev0-py310-none-win_amd64.whl" + elif sys.version_info[1] == 11: + whl = "natten-0.17.2.dev0-py311-none-win_amd64.whl" + elif sys.version_info[1] == 12: + whl = "natten-0.17.2.dev0-py312-none-win_amd64.whl" + whl_path = os.path.join(folder_paths.temp_directory, whl) + if not os.path.exists(whl_path): + if not os.path.exists(folder_paths.temp_directory): + os.makedirs(folder_paths.temp_directory) + print(f"Downloading {whl} from 2kpr/NATTEN-Windows...") + huggingface_hub.snapshot_download( + repo_id="2kpr/NATTEN-Windows", + allow_patterns=whl, + local_dir=folder_paths.temp_directory, + ) + subprocess.check_call([sys.executable, "-m", "pip", "install", f"{whl_path}"]) + else: + subprocess.check_call([sys.executable, "-m", "pip", "install", f"natten==0.17.1+{torch_version}{cuda_version}", "-f", "https://shi-labs.com/natten/wheels/"]) \ No newline at end of file diff --git a/utils/__init__.py b/utils/__init__.py new file mode 100755 index 0000000..e69de29 diff --git a/utils/create_arch.py b/utils/create_arch.py new file mode 100755 index 0000000..3cd2a8e --- /dev/null +++ b/utils/create_arch.py @@ -0,0 +1,143 @@ +from ..arch.hourglass import image_transformer_v2 as itv2 +from ..arch.hourglass.image_transformer_v2 import ImageTransformerDenoiserModelV2 +from ..arch.swinir.swinir import SwinIR + + +def create_arch(arch, condition_channels=0): + # arch should be, e.g., swinir_XL, or hdit_XL + arch_name, arch_size = arch.split('_') + arch_config = arch_configs[arch_name][arch_size].copy() + arch_config['in_channels'] += condition_channels + return arch_name_to_object[arch_name](**arch_config) + + +arch_configs = { + 'hdit': { + "ImageNet256Sp4": { + 'in_channels': 3, + 'out_channels': 3, + 'widths': [256, 512, 1024], + 'depths': [2, 2, 8], + 'patch_size': [4, 4], + 'self_attns': [ + {"type": "neighborhood", "d_head": 64, "kernel_size": 7}, + {"type": "neighborhood", "d_head": 64, "kernel_size": 7}, + {"type": "global", "d_head": 64} + ], + 'mapping_depth': 2, + 'mapping_width': 768, + 'dropout_rate': [0, 0, 0], + 'mapping_dropout_rate': 0.0 + }, + "XL2": { + 'in_channels': 3, + 'out_channels': 3, + 'widths': [384, 768], + 'depths': [2, 11], + 'patch_size': [4, 4], + 'self_attns': [ + {"type": "neighborhood", "d_head": 64, "kernel_size": 7}, + {"type": "global", "d_head": 64} + ], + 'mapping_depth': 2, + 'mapping_width': 768, + 'dropout_rate': [0, 0], + 'mapping_dropout_rate': 0.0 + } + + }, + 'swinir': { + "M": { + 'in_channels': 3, + 'out_channels': 3, + 'embed_dim': 120, + 'depths': [6, 6, 6, 6, 6], + 'num_heads': [6, 6, 6, 6, 6], + 'resi_connection': '1conv', + 'sf': 8 + + }, + "L": { + 'in_channels': 3, + 'out_channels': 3, + 'embed_dim': 180, + 'depths': [6, 6, 6, 6, 6, 6, 6, 6], + 'num_heads': [6, 6, 6, 6, 6, 6, 6, 6], + 'resi_connection': '1conv', + 'sf': 8 + }, + }, +} + + +def create_swinir_model(in_channels, out_channels, embed_dim, depths, num_heads, resi_connection, + sf): + return SwinIR( + img_size=64, + patch_size=1, + in_chans=in_channels, + num_out_ch=out_channels, + embed_dim=embed_dim, + depths=depths, + num_heads=num_heads, + window_size=8, + mlp_ratio=2, + sf=sf, + img_range=1.0, + upsampler="nearest+conv", + resi_connection=resi_connection, + unshuffle=True, + unshuffle_scale=8 + ) + + +def create_hdit_model(widths, + depths, + self_attns, + dropout_rate, + mapping_depth, + mapping_width, + mapping_dropout_rate, + in_channels, + out_channels, + patch_size + ): + assert len(widths) == len(depths) + assert len(widths) == len(self_attns) + assert len(widths) == len(dropout_rate) + mapping_d_ff = mapping_width * 3 + d_ffs = [] + for width in widths: + d_ffs.append(width * 3) + + levels = [] + for depth, width, d_ff, self_attn, dropout in zip(depths, widths, d_ffs, self_attns, dropout_rate): + if self_attn['type'] == 'global': + self_attn = itv2.GlobalAttentionSpec(self_attn.get('d_head', 64)) + elif self_attn['type'] == 'neighborhood': + self_attn = itv2.NeighborhoodAttentionSpec(self_attn.get('d_head', 64), self_attn.get('kernel_size', 7)) + elif self_attn['type'] == 'shifted-window': + self_attn = itv2.ShiftedWindowAttentionSpec(self_attn.get('d_head', 64), self_attn['window_size']) + elif self_attn['type'] == 'none': + self_attn = itv2.NoAttentionSpec() + else: + raise ValueError(f'unsupported self attention type {self_attn["type"]}') + levels.append(itv2.LevelSpec(depth, width, d_ff, self_attn, dropout)) + mapping = itv2.MappingSpec(mapping_depth, mapping_width, mapping_d_ff, mapping_dropout_rate) + model = ImageTransformerDenoiserModelV2( + levels=levels, + mapping=mapping, + in_channels=in_channels, + out_channels=out_channels, + patch_size=patch_size, + num_classes=0, + mapping_cond_dim=0, + ) + + return model + + +arch_name_to_object = { + 'hdit': create_hdit_model, + 'swinir': create_swinir_model, +}