Initial Commit
This commit is contained in:
@@ -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
|
||||
}
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1013 KiB |
Executable
+3
@@ -0,0 +1,3 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
Executable
+2
@@ -0,0 +1,2 @@
|
||||
from .hourglass.image_transformer_v2 import ImageTransformerDenoiserModelV2
|
||||
from .swinir.swinir import SwinIR
|
||||
Executable
Executable
+112
@@ -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)
|
||||
Executable
+60
@@ -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)
|
||||
Executable
+58
@@ -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
|
||||
Executable
+766
@@ -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
|
||||
Executable
Executable
+904
@@ -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
|
||||
Executable
+318
@@ -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
|
||||
@@ -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",
|
||||
}
|
||||
Executable
+172
@@ -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/"])
|
||||
Executable
Executable
+143
@@ -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,
|
||||
}
|
||||
Reference in New Issue
Block a user