Initial Commit

This commit is contained in:
2kpr
2024-10-10 23:07:41 +00:00
parent 3544c936b5
commit 7c01417f15
16 changed files with 2878 additions and 0 deletions
+143
View File
@@ -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
View File
@@ -0,0 +1,3 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+2
View File
@@ -0,0 +1,2 @@
from .hourglass.image_transformer_v2 import ImageTransformerDenoiserModelV2
from .swinir.swinir import SwinIR
View File
+112
View File
@@ -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)
+60
View File
@@ -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)
+58
View File
@@ -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
+766
View File
@@ -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
View File
+904
View File
@@ -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
+318
View File
@@ -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
Executable
+197
View File
@@ -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",
}
+172
View File
@@ -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/"])
View File
+143
View File
@@ -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,
}