Initial commit

This commit is contained in:
BobRandomNumber
2025-08-04 18:50:22 -04:00
committed by GitHub
parent 7dc06c1741
commit f35fe344e9
26 changed files with 9709 additions and 0 deletions
+112
View File
@@ -0,0 +1,112 @@
# ComfyUI-TLBVFI
A LLM coded node pack for ComfyUI that provides video frame interpolation using the **TLB-VFI** model.
This is a wrapper for the [TLB-VFI: Temporal-Aware Latent Brownian Bridge Diffusion for Video Frame Interpolation](https://github.com/ZonglinL/TLBVFI) project, allowing for integration into ComfyUI.
## Features
- **High-Quality Interpolation**: Leverages a powerful latent diffusion model to generate smooth and detailed in-between frames.
- **Configurable Interpolation Steps**: Easily double, quadruple, or octuple your frame rate by adjusting the `times_to_interpolate` setting.
---
## ⚙️ Installation
Please follow these steps carefully to ensure the node is set up correctly.
### Step 1: Install the Custom Node
If you are using the [ComfyUI-Manager](https://github.com/Comfy-Org/ComfyUI-Manager), you can install this node from there.
Alternatively, you can install it manually by cloning this repository into your `ComfyUI/custom_nodes/` directory...
# Navigate to your ComfyUI custom_nodes directory
```bash
cd ComfyUI/custom_nodes/
# Clone this repository
git clone https://github.com/BobRandomNumber/ComfyUI-TLBVFI.git
```
### Step 2: Install Dependencies
```bash
# Navigate into the newly created custom node directory
cd ComfyUI/custom_nodes/ComfyUI-TLBVFI/
pip install -r requirements.txt
```
### Step 3: Download the Pre-trained Model
Only one model file is required to run the interpolation.
- **Full Model:** `vimeo_unet.pth`
Download the file from the official Hugging Face repository:
- **[ucfzl/TLBVFI on Hugging Face](https://huggingface.co/ucfzl/TLBVFI/tree/main)**
### Step 4: Place Model in the `interpolation` Folder
This node looks for models in the `ComfyUI/models/interpolation/` directory.
1. Place the downloaded `vimeo_unet.pth` file into this folder.
For better organization, you are welcome to create a subdirectory. The node will find the model automatically.
**Example Folder Structure:**
```
ComfyUI/
└── models/
└── interpolation/
└── tlbvfi_models/
└── vimeo_unet.pth
```
> **For advanced users:** If you prefer to store models elsewhere, you can add a path to your `extra_model_paths.yaml` file and assign it the type `interpolation`.
### Step 5: Restart ComfyUI
After completing all the steps, **restart ComfyUI**.
---
## 🚀 Usage
1. In ComfyUI, add the **TLBVFI Frame Interpolation** node. You can find it by right-clicking and searching, or under the `frame_interpolation/TLBVFI` category.
2. Connect a batch of loaded images (e.g., from a `Load Video` or `Load Image Batch` node) to the `images` input.
3. Select the correct model from the dropdown menu:
* **`model_name`**: Choose `vimeo_unet.pth` (or `tlbvfi_models/vimeo_unet.pth` if you used a subfolder).
4. Adjust **`times_to_interpolate`** to control how many new frames are generated between each pair of original frames:
* `1`: **Doubles** the frame count (1 new frame).
* `2`: **Quadruples** the frame count (3 new frames).
* `3`: **8x** the frame count (7 new frames).
5. Connect the output `IMAGE` to a `Save Image` or `Preview Image` node to see your interpolated sequence.
---
## 🧠 How It Works
This node uses a two-stage **latent diffusion** process:
1. **VQGAN (Autoencoder)**: First, the VQGAN model takes your full-resolution input frames and compresses them into a small, efficient "latent space."
2. **UNet (Diffusion Model)**: The core interpolation logic happens in this latent space. The UNet takes the compressed representations of the start and end frames and generates the latent representation for the frame in between.
3. **VQGAN (Decoder)**: Finally, the VQGAN's decoder takes the newly generated latent and reconstructs it back into a full-resolution, detailed image.
This approach is highly efficient and allows for the generation of high-quality, temporally consistent frames.
---
## 🙏 Acknowledgements and Citation
This node is a wrapper implementation for ComfyUI. All credit for the model architecture, training, and research goes to the original authors of TLB-VFI. If you use this model in your research, please cite their work.
- **Original GitHub Repository**: [https://github.com/ZonglinL/TLBVFI](https://github.com/ZonglinL/TLBVFI)
- **Project Page**: [https://zonglinl.github.io/tlbvfi_page/](https://zonglinl.github.io/tlbvfi_page/)
```bibtex
@article{lyu2025tlbvfitemporalawarelatentbrownian,
title={TLB-VFI: Temporal-Aware Latent Brownian Bridge Diffusion for Video Frame Interpolation},
author={Zonglin Lyu and Chen Chen},
year={2025},
eprint={2507.04984},
archivePrefix={arXiv},
primaryClass={cs.CV},
}
```
+1
View File
@@ -0,0 +1 @@
+647
View File
@@ -0,0 +1,647 @@
import os
import sys
import torch
import torch.nn as nn
from torch.nn import init
import torch.nn.functional as F
from torch.nn.parallel import DistributedDataParallel
import functools
import copy
from functools import partial, reduce
import numpy as np
import itertools
import math
from collections import OrderedDict
from timm.models.layers import DropPath, to_2tuple, trunc_normal_
sys.path.append('../..')
from VFI.archs.warplayer import warp
from VFI.archs.transformer_layers import TFModel
def make_layer(block, n_layers):
layers = []
for _ in range(n_layers):
layers.append(block())
return nn.Sequential(*layers)
class ResidualBlock(nn.Module):
def __init__(self, nf, kernel_size=3, stride=1, padding=1, dilation=1, act='relu'):
super().__init__()
self.conv1 = nn.Conv2d(nf, nf, kernel_size=kernel_size, stride=stride, padding=padding, dilation=dilation)
self.conv2 = nn.Conv2d(nf, nf, kernel_size=kernel_size, stride=stride, padding=padding, dilation=dilation)
if act == 'relu':
self.act = nn.ReLU(inplace=True)
else:
self.act = nn.LeakyReLU(0.1, inplace=True)
def forward(self, x):
out = self.conv2(self.act(self.conv1(x)))
return out + x
def deconv(in_planes, out_planes, kernel_size=4, stride=2, padding=1):
return nn.Sequential(
torch.nn.ConvTranspose2d(in_channels=in_planes, out_channels=out_planes, kernel_size=4, stride=2, padding=1),
nn.PReLU(out_planes)
)
def conv_wo_act(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation, bias=True),
)
def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation, bias=True),
nn.PReLU(out_planes)
)
class Conv2(nn.Module):
def __init__(self, in_planes, out_planes, stride=2):
super().__init__()
self.conv1 = conv(in_planes, out_planes, 3, stride, 1)
self.conv2 = conv(out_planes, out_planes, 3, 1, 1)
def forward(self, x):
x = self.conv1(x)
x = self.conv2(x)
return x
class IFBlock(nn.Module):
def __init__(self, in_planes, scale=1, c=64):
super().__init__()
self.scale = scale
self.conv0 = nn.Sequential(
conv(in_planes, c//2, 3, 2, 1),
conv(c//2, c, 3, 2, 1),
)
self.convblock = nn.Sequential(
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
)
self.conv1 = nn.ConvTranspose2d(c, 4, 4, 2, 1)
def forward(self, x):
if self.scale != 1:
x = F.interpolate(x, scale_factor= 1. / self.scale, mode="bilinear", align_corners=False)
x = self.conv0(x)
x = self.convblock(x) + x
x = self.conv1(x)
flow = x
if self.scale != 1:
flow = F.interpolate(flow, scale_factor= self.scale, mode="bilinear", align_corners=False)
return flow
class IFNet(nn.Module):
def __init__(self, args=None):
super().__init__()
self.block0 = IFBlock(6, scale=4, c=240)
self.block1 = IFBlock(10, scale=2, c=150)
self.block2 = IFBlock(10, scale=1, c=90)
def forward(self, x):
flow0 = self.block0(x)
F1 = flow0
F1_large = F.interpolate(F1, scale_factor=2.0, mode="bilinear", align_corners=False) * 2.0
warped_img0 = warp(x[:, :3], F1_large[:, :2])
warped_img1 = warp(x[:, 3:], F1_large[:, 2:4])
flow1 = self.block1(torch.cat((warped_img0, warped_img1, F1_large), 1))
F2 = (flow0 + flow1)
F2_large = F.interpolate(F2, scale_factor=2.0, mode="bilinear", align_corners=False) * 2.0
warped_img0 = warp(x[:, :3], F2_large[:, :2])
warped_img1 = warp(x[:, 3:], F2_large[:, 2:4])
flow2 = self.block2(torch.cat((warped_img0, warped_img1, F2_large), 1))
F3 = (flow0 + flow1 + flow2)
return F3, [F1, F2, F3]
class FlowRefineNetA(nn.Module):
def __init__(self, context_dim, c=16, r=1, n_iters=4):
super().__init__()
corr_dim = c
flow_dim = c
motion_dim = c
hidden_dim = c
self.n_iters = n_iters
self.r = r
self.n_pts = (r * 2 + 1) ** 2
self.occl_convs = nn.Sequential(nn.Conv2d(2 * context_dim, hidden_dim, 1, 1, 0),
nn.PReLU(hidden_dim),
nn.Conv2d(hidden_dim, hidden_dim, 1, 1, 0),
nn.PReLU(hidden_dim),
nn.Conv2d(hidden_dim, 1, 1, 1, 0),
nn.Sigmoid())
self.corr_convs = nn.Sequential(nn.Conv2d(self.n_pts, hidden_dim, 1, 1, 0),
nn.PReLU(hidden_dim),
nn.Conv2d(hidden_dim, corr_dim, 1, 1, 0),
nn.PReLU(corr_dim))
self.flow_convs = nn.Sequential(nn.Conv2d(2, hidden_dim, 3, 1, 1),
nn.PReLU(hidden_dim),
nn.Conv2d(hidden_dim, flow_dim, 3, 1, 1),
nn.PReLU(flow_dim))
self.motion_convs = nn.Sequential(nn.Conv2d(corr_dim + flow_dim, motion_dim, 3, 1, 1),
nn.PReLU(motion_dim))
self.gru = nn.Sequential(nn.Conv2d(motion_dim + context_dim * 2 + 2, hidden_dim, 3, 1, 1),
nn.PReLU(hidden_dim),
nn.Conv2d(hidden_dim, flow_dim, 3, 1, 1),
nn.PReLU(flow_dim), )
self.flow_head = nn.Sequential(nn.Conv2d(flow_dim, hidden_dim, 3, 1, 1),
nn.PReLU(hidden_dim),
nn.Conv2d(hidden_dim, 2, 3, 1, 1))
def L2normalize(self, x, dim=1):
eps = 1e-12
norm = x ** 2
norm = norm.sum(dim=dim, keepdim=True) + eps
norm = norm ** (0.5)
return (x/norm)
def forward_once(self, x0, x1, flow0, flow1):
B, C, H, W = x0.size()
x0_unfold = F.unfold(x0, kernel_size=(self.r * 2 + 1), padding=1).view(B, C * self.n_pts, H,
W) # (B, C*n_pts, H, W)
x1_unfold = F.unfold(x1, kernel_size=(self.r * 2 + 1), padding=1).view(B, C * self.n_pts, H,
W) # (B, C*n_pts, H, W)
contents0 = warp(x0_unfold, flow0)
contents1 = warp(x1_unfold, flow1)
contents0 = contents0.view(B, C, self.n_pts, H, W)
contents1 = contents1.view(B, C, self.n_pts, H, W)
fea0 = contents0[:, :, self.n_pts // 2, :, :]
fea1 = contents1[:, :, self.n_pts // 2, :, :]
# get context feature
occl = self.occl_convs(torch.cat([fea0, fea1], dim=1))
fea = fea0 * occl + fea1 * (1 - occl)
# get correlation features
fea_view = fea.permute(0, 2, 3, 1).contiguous().view(B * H * W, 1, C)
contents0 = contents0.permute(0, 3, 4, 2, 1).contiguous().view(B * H * W, self.n_pts, C)
contents1 = contents1.permute(0, 3, 4, 2, 1).contiguous().view(B * H * W, self.n_pts, C)
fea_view = self.L2normalize(fea_view, dim=-1)
contents0 = self.L2normalize(contents0, dim=-1)
contents1 = self.L2normalize(contents1, dim=-1)
corr0 = torch.einsum('bic,bjc->bij', fea_view, contents0) # (B*H*W, 1, n_pts)
corr1 = torch.einsum('bic,bjc->bij', fea_view, contents1)
# corr0 = corr0 / torch.sqrt(torch.tensor(C).float())
# corr1 = corr1 / torch.sqrt(torch.tensor(C).float())
corr0 = corr0.view(B, H, W, self.n_pts).permute(0, 3, 1, 2).contiguous() # (B, n_pts, H, W)
corr1 = corr1.view(B, H, W, self.n_pts).permute(0, 3, 1, 2).contiguous()
corr0 = self.corr_convs(corr0) # (B, corr_dim, H, W)
corr1 = self.corr_convs(corr1)
# get flow features
flow0_fea = self.flow_convs(flow0)
flow1_fea = self.flow_convs(flow1)
# merge correlation and flow features, get motion features
motion0 = self.motion_convs(torch.cat([corr0, flow0_fea], dim=1))
motion1 = self.motion_convs(torch.cat([corr1, flow1_fea], dim=1))
# update flows
inp0 = torch.cat([fea, fea0, motion0, flow0], dim=1)
delta_flow0 = self.flow_head(self.gru(inp0))
flow0 = flow0 + delta_flow0
inp1 = torch.cat([fea, fea1, motion1, flow1], dim=1)
delta_flow1 = self.flow_head(self.gru(inp1))
flow1 = flow1 + delta_flow1
return flow0, flow1
def forward(self, x0, x1, flow0, flow1):
for i in range(self.n_iters):
flow0, flow1 = self.forward_once(x0, x1, flow0, flow1)
return torch.cat([flow0, flow1], dim=1)
class FlowRefineNet_Multis_our(nn.Module):
def __init__(self, c=24, n_iters=1):
super().__init__()
self.rf_block1 = FlowRefineNetA(context_dim= c, c= c, r=1, n_iters=n_iters)
self.rf_block2 = FlowRefineNetA(context_dim= c, c= c, r=1, n_iters=n_iters)
self.rf_block3 = FlowRefineNetA(context_dim=2 * c, c=2 * c, r=1, n_iters=n_iters)
self.rf_block4 = FlowRefineNetA(context_dim=2 * c, c=2 * c, r=1, n_iters=n_iters)
def forward(self, feats, flow):
s_1,s_2,s_3,s_4 = feats
bs = s_1.size(0)//2
# update flow from small scale
flow = F.interpolate(flow, scale_factor=0.25, mode="bilinear", align_corners=False) * 0.25 # 1/8
flow = self.rf_block4(s_4[:bs], s_4[bs:], flow[:, :2], flow[:, 2:4]) # 1/8
flow = F.interpolate(flow, scale_factor=2., mode="bilinear", align_corners=False) * 2.
flow = self.rf_block3(s_3[:bs], s_3[bs:], flow[:, :2], flow[:, 2:4]) # 1/4
flow = F.interpolate(flow, scale_factor=2., mode="bilinear", align_corners=False) * 2.
flow = self.rf_block2(s_2[:bs], s_2[bs:], flow[:, :2], flow[:, 2:4]) # 1/2
flow = F.interpolate(flow, scale_factor=2., mode="bilinear", align_corners=False) * 2.
flow = self.rf_block1(s_1[:bs], s_1[bs:], flow[:, :2], flow[:, 2:4]) # 1
# warp features by the updated flow
c0 = [s_1[:bs], s_2[:bs], s_3[:bs], s_4[:bs]]
c1 = [s_1[bs:], s_2[bs:], s_3[bs:], s_4[bs:]]
out0 = self.warp_fea(c0, flow[:, :2])
out1 = self.warp_fea(c1, flow[:, 2:4])
return flow, out0, out1
def warp_fea(self, feas, flow):
outs = []
for i, fea in enumerate(feas):
out = warp(fea, flow)
outs.append(out)
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False) * 0.5
return outs
def get_context(self, feats, flow):
s_1,s_2,s_3,s_4 = feats
bs = s_1.size(0)//2
# warp features by the updated flow
c0 = [s_1[:bs], s_2[:bs], s_3[:bs], s_4[:bs]]
c1 = [s_1[bs:], s_2[bs:], s_3[bs:], s_4[bs:]]
out0 = self.warp_fea(c0, flow[:, :2])
out1 = self.warp_fea(c1, flow[:, 2:4])
return flow, out0, out1
class FlowRefineNet_Multis(nn.Module):
def __init__(self, c=24, n_iters=1):
super().__init__()
self.conv1 = Conv2(3, c, 1)
self.conv2 = Conv2(c, 2 * c)
self.conv3 = Conv2(2 * c, 4 * c)
self.conv4 = Conv2(4 * c, 8 * c)
self.rf_block1 = FlowRefineNetA(context_dim=c, c=c, r=1, n_iters=n_iters)
self.rf_block2 = FlowRefineNetA(context_dim=2 * c, c=2 * c, r=1, n_iters=n_iters)
self.rf_block3 = FlowRefineNetA(context_dim=4 * c, c=4 * c, r=1, n_iters=n_iters)
self.rf_block4 = FlowRefineNetA(context_dim=8 * c, c=8 * c, r=1, n_iters=n_iters)
def get_context(self, x0, x1, flow):
bs = x0.size(0)
inp = torch.cat([x0, x1], dim=0)
s_1 = self.conv1(inp) # 1
s_2 = self.conv2(s_1) # 1/2
s_3 = self.conv3(s_2) # 1/4
s_4 = self.conv4(s_3) # 1/8
# warp features by the updated flow
c0 = [s_1[:bs], s_2[:bs], s_3[:bs], s_4[:bs]]
c1 = [s_1[bs:], s_2[bs:], s_3[bs:], s_4[bs:]]
out0 = self.warp_fea(c0, flow[:, :2])
out1 = self.warp_fea(c1, flow[:, 2:4])
return flow, out0, out1
def forward(self, x0, x1, flow):
bs = x0.size(0)
inp = torch.cat([x0, x1], dim=0)
s_1 = self.conv1(inp) # 1
s_2 = self.conv2(s_1) # 1/2
s_3 = self.conv3(s_2) # 1/4
s_4 = self.conv4(s_3) # 1/8
# update flow from small scale
flow = F.interpolate(flow, scale_factor=0.25, mode="bilinear", align_corners=False) * 0.25 # 1/8
flow = self.rf_block4(s_4[:bs], s_4[bs:], flow[:, :2], flow[:, 2:4]) # 1/8
flow = F.interpolate(flow, scale_factor=2., mode="bilinear", align_corners=False) * 2.
flow = self.rf_block3(s_3[:bs], s_3[bs:], flow[:, :2], flow[:, 2:4]) # 1/4
flow = F.interpolate(flow, scale_factor=2., mode="bilinear", align_corners=False) * 2.
flow = self.rf_block2(s_2[:bs], s_2[bs:], flow[:, :2], flow[:, 2:4]) # 1/2
flow = F.interpolate(flow, scale_factor=2., mode="bilinear", align_corners=False) * 2.
flow = self.rf_block1(s_1[:bs], s_1[bs:], flow[:, :2], flow[:, 2:4]) # 1
# warp features by the updated flow
c0 = [s_1[:bs], s_2[:bs], s_3[:bs], s_4[:bs]]
c1 = [s_1[bs:], s_2[bs:], s_3[bs:], s_4[bs:]]
out0 = self.warp_fea(c0, flow[:, :2])
out1 = self.warp_fea(c1, flow[:, 2:4])
return flow, out0, out1
def warp_fea(self, feas, flow):
outs = []
for i, fea in enumerate(feas):
out = warp(fea, flow)
outs.append(out)
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False) * 0.5
return outs
def warp_batch_fea(self, feas, flow,bs):
c0 = []
c1 = []
for i, fea in enumerate(feas):
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False) * 0.5
out0 = warp(fea[:bs], flow[:,:2])
out1 = warp(fea[bs:], flow[:,2:])
c0.append(out0)
c1.append(out1)
return c0,c1
class VFIformer(nn.Module):
def __init__(self):
super().__init__()
self.phase = 'test'
self.device = 'cuda'
c = 24
height = 192
width = 192
window_size = 8
embed_dim = 160
self.flownet = IFNet()
self.refinenet = FlowRefineNet_Multis(c=c, n_iters=1)
self.fuse_block = nn.Sequential(nn.Conv2d(12, 2*c, 3, 1, 1),
nn.LeakyReLU(negative_slope=0.2, inplace=True),
nn.Conv2d(2*c, 2*c, 3, 1, 1),
nn.LeakyReLU(negative_slope=0.2, inplace=True),)
self.transformer = TFModel(img_size=(height, width), in_chans=2*c, out_chans=4, fuse_c=c,
window_size=window_size, img_range=1.,
depths=[[3, 3], [3, 3], [3, 3], [1, 1]],
embed_dim=embed_dim, num_heads=[[2, 2], [2, 2], [2, 2], [2, 2]], mlp_ratio=2,
resi_connection='1conv',
use_crossattn=[[[False, False, False, False], [True, True, True, True]], \
[[False, False, False, False], [True, True, True, True]], \
[[False, False, False, False], [True, True, True, True]], \
[[False, False, False, False], [False, False, False, False]]])
self.apply(self._init_weights)
def load_networks(self, net_name = 'VFIformer', resume = '/scratch/zl3958/VLPR/net_220.pth', strict=True):
load_path = resume
network = getattr(self, net_name)
if isinstance(network, nn.DataParallel) or isinstance(network, DistributedDataParallel):
network = network.module
load_net = torch.load(load_path, map_location=torch.device(self.device))
load_net_clean = OrderedDict() # remove unnecessary 'module.'
for k, v in load_net.items():
if k.startswith('module.'):
load_net_clean[k[7:]] = v
else:
load_net_clean[k] = v
if 'optimizer' or 'scheduler' in net_name:
network.load_state_dict(load_net_clean)
else:
network.load_state_dict(load_net_clean, strict=strict)
print('load pretrained VFIformer')
def _init_weights(self, m):
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)
def get_flow(self, img0, img1):
imgs = torch.cat((img0, img1), 1)
flow, flow_list = self.flownet(imgs)
flow, c0, c1 = self.refinenet(img0, img1, flow)
return flow
def forward(self, img0, img1, flow_pre=None):
B, _, H, W = img0.size()
imgs = torch.cat((img0, img1), 1)
if flow_pre is not None:
flow = flow_pre
_, c0, c1 = self.refinenet.get_context(img0, img1, flow)
else:
flow, flow_list = self.flownet(imgs)
flow, c0, c1 = self.refinenet(img0, img1, flow)
warped_img0 = warp(img0, flow[:, :2])
warped_img1 = warp(img1, flow[:, 2:])
x = self.fuse_block(torch.cat([img0, img1, warped_img0, warped_img1], dim=1))
refine_output = self.transformer(x, c0, c1)
res = torch.sigmoid(refine_output[:, :3]) * 2 - 1
mask = torch.sigmoid(refine_output[:, 3:4])
merged_img = warped_img0 * mask + warped_img1 * (1 - mask)
pred = merged_img + res
pred = torch.clamp(pred, 0, 1)
if self.phase == 'train':
return pred, flow_list
else:
return pred, flow
#-------------------------------------
# light-weight version
#-------------------------------------
class FlowRefineNet_Multis_Simple(nn.Module):
def __init__(self, c=24, n_iters=1):
super(FlowRefineNet_Multis_Simple, self).__init__()
self.conv1 = Conv2(3, c, 1)
self.conv2 = Conv2(c, 2 * c)
self.conv3 = Conv2(2 * c, 4 * c)
self.conv4 = Conv2(4 * c, 8 * c)
def forward(self, x0, x1, flow):
bs = x0.size(0)
inp = torch.cat([x0, x1], dim=0)
s_1 = self.conv1(inp) # 1
s_2 = self.conv2(s_1) # 1/2
s_3 = self.conv3(s_2) # 1/4
s_4 = self.conv4(s_3) # 1/8
flow = F.interpolate(flow, scale_factor=2., mode="bilinear", align_corners=False) * 2.
# warp features by the updated flow
c0 = [s_1[:bs], s_2[:bs], s_3[:bs], s_4[:bs]]
c1 = [s_1[bs:], s_2[bs:], s_3[bs:], s_4[bs:]]
out0 = self.warp_fea(c0, flow[:, :2])
out1 = self.warp_fea(c1, flow[:, 2:4])
return flow, out0, out1
def warp_fea(self, feas, flow):
outs = []
for i, fea in enumerate(feas):
out = warp(fea, flow)
outs.append(out)
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False) * 0.5
return outs
class VFIformerSmall(nn.Module):
def __init__(self, args):
super(VFIformerSmall, self).__init__()
self.phase = args.phase
self.device = args.device
c = 24
height = args.crop_size
width = args.crop_size
window_size = 4
embed_dim = 136
self.flownet = IFNet()
self.refinenet = FlowRefineNet_Multis_Simple(c=c, n_iters=1)
self.fuse_block = nn.Sequential(nn.Conv2d(12, 2*c, 3, 1, 1),
nn.LeakyReLU(negative_slope=0.2, inplace=True),
nn.Conv2d(2*c, 2*c, 3, 1, 1),
nn.LeakyReLU(negative_slope=0.2, inplace=True),)
self.transformer = TFModel(img_size=(height, width), in_chans=2*c, out_chans=4, fuse_c=c,
window_size=window_size, img_range=1.,
depths=[[3, 3], [3, 3], [3, 3], [1, 1]],
embed_dim=embed_dim, num_heads=[[2, 2], [2, 2], [2, 2], [2, 2]], mlp_ratio=2,
resi_connection='1conv',
use_crossattn=[[[False, False, False, False], [True, True, True, True]], \
[[False, False, False, False], [True, True, True, True]], \
[[False, False, False, False], [True, True, True, True]], \
[[False, False, False, False], [False, False, False, False]]])
self.apply(self._init_weights)
if args.resume_flownet:
self.load_networks('flownet', args.resume_flownet)
print('------ flownet loaded --------')
def load_networks(self, net_name, resume, strict=True):
load_path = resume
network = getattr(self, net_name)
if isinstance(network, nn.DataParallel) or isinstance(network, DistributedDataParallel):
network = network.module
load_net = torch.load(load_path, map_location=torch.device(self.device))
load_net_clean = OrderedDict() # remove unnecessary 'module.'
for k, v in load_net.items():
if k.startswith('module.'):
load_net_clean[k[7:]] = v
else:
load_net_clean[k] = v
if 'optimizer' or 'scheduler' in net_name:
network.load_state_dict(load_net_clean)
else:
network.load_state_dict(load_net_clean, strict=strict)
def _init_weights(self, m):
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)
def get_flow(self, img0, img1):
imgs = torch.cat((img0, img1), 1)
flow, flow_list = self.flownet(imgs)
flow, c0, c1 = self.refinenet(img0, img1, flow)
return flow
def forward(self, img0, img1, flow_pre=None):
B, _, H, W = img0.size()
imgs = torch.cat((img0, img1), 1)
if flow_pre is not None:
flow = flow_pre
_, c0, c1 = self.refinenet(img0, img1, flow)
else:
flow, flow_list = self.flownet(imgs)
flow, c0, c1 = self.refinenet(img0, img1, flow)
warped_img0 = warp(img0, flow[:, :2])
warped_img1 = warp(img1, flow[:, 2:])
x = self.fuse_block(torch.cat([img0, img1, warped_img0, warped_img1], dim=1))
refine_output = self.transformer(x, c0, c1)
res = torch.sigmoid(refine_output[:, :3]) * 2 - 1
mask = torch.sigmoid(refine_output[:, 3:4])
merged_img = warped_img0 * mask + warped_img1 * (1 - mask)
pred = merged_img + res
pred = torch.clamp(pred, 0, 1)
if self.phase == 'train':
return pred, flow_list
else:
return pred, flow
if __name__ == "__main__":
# try:
# from models.archs.dcn.deform_conv import ModulatedDeformConvPack as DCN
# except ImportError:
# raise ImportError('Failed to import DCNv2 module.')
import argparse
parser = argparse.ArgumentParser(description='test')
parser.add_argument('--phase', default='train', type=str)
parser.add_argument('--device', default='cuda', type=str)
parser.add_argument('--crop_size', default=192, type=int)
args = parser.parse_args()
device = 'cuda'
net = Swin_Fuse_CrossScaleV2_MaskV5_Normal_WoRefine_ConvBaseline(args).to(device)
print('----- generator parameters: %f -----' % (sum(param.numel() for param in net.parameters()) / (10**6)))
w = 192
img0 = torch.randn((2, 3, w, w)).to(device)
img1 = torch.randn((2, 3, w, w)).to(device)
out = net(img0, img1)
print(out[0].size())
File diff suppressed because it is too large Load Diff
+30
View File
@@ -0,0 +1,30 @@
import torch
import torch.nn as nn
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
backwarp_tenGrid = {}
def warp(tenInput, tenFlow):
k = (str(tenFlow.device), str(tenFlow.size()))
if k not in backwarp_tenGrid:
tenHorizontal = torch.linspace(-1.0, 1.0, tenFlow.shape[3], device=device).view(
1, 1, 1, tenFlow.shape[3]).expand(tenFlow.shape[0], -1, tenFlow.shape[2], -1)
tenVertical = torch.linspace(-1.0, 1.0, tenFlow.shape[2], device=device).view(
1, 1, tenFlow.shape[2], 1).expand(tenFlow.shape[0], -1, -1, tenFlow.shape[3])
backwarp_tenGrid[k] = torch.cat(
[tenHorizontal, tenVertical], 1).to(device)
tenFlow = torch.cat([tenFlow[:, 0:1, :, :] / ((tenInput.shape[3] - 1.0) / 2.0),
tenFlow[:, 1:2, :, :] / ((tenInput.shape[2] - 1.0) / 2.0)], 1)
g = (backwarp_tenGrid[k] + tenFlow).permute(0, 2, 3, 1)
return torch.nn.functional.grid_sample(input=tenInput, grid=g, mode='bilinear', padding_mode='border', align_corners=True)
def flow_reversal(flow):
# flow: (B, 2, H, W)
B, _, H, W = flow.size()
flow_r = warp(flow, flow.clone())
flow_r = -1 * flow_r
return flow_r
+139
View File
@@ -0,0 +1,139 @@
# Latent Brownian Bridge Diffusion Model Template(Latent Space)
runner: "BBDMRunner"
training:
n_epochs: 400
n_steps: 1000000
save_interval: 1
sample_interval: 1
validation_interval: 1
accumulate_grad_batches: 1
testing:
clip_denoised: False
sample_num: 1
data:
dataset_name: 'DAVIS' ## this folder stores training logs
dataset_type: 'Interpolation'
dataset_config:
dataset_path: /home/zo258499/data ## path to your data directory
image_size: 256
channels: 3
to_normal: True
flip: True
cat: False
aug_noise: False
aug_cut: False
eval: 'DAVIS' ## options {"UCF", "MidB", "DAVIS","FILM"}
mode: 'easy' ## options{"easy","medium","hard","extreme"}
train:
batch_size: 48
shuffle: True
val:
batch_size: 32
shuffle: False
test:
batch_size: 1
# shuffle: False
model:
model_name: "LBBDM-f32" # part of result path
model_type: "LBBDM" # specify a module
latent_before_quant_conv: False
normalize_latent: False
only_load_latent_mean_std: False
# model_load_path: # model checkpoint path
# optim_sche_load_path: # optimizer scheduler checkpoint path
EMA:
use_ema: True
ema_decay: 0.995
update_ema_interval: 8 # step
start_ema_step: 3000
CondStageParams:
n_stages: 4
in_channels: 3
out_channels: 3
VQGAN:
params:
ckpt_path: "results/VQGAN/vimeo_new.ckpt"
embed_dim: 3
n_embed: 8192
ddconfig:
double_z: False
z_channels: 3
resolution: 256
in_channels: 3
out_ch: 3
ch: 64
ch_mult: !!python/list
- 1
- 2
- 2
- 2
- 4
num_res_blocks: 1
cond_type: max_cross_attn
attn_type: max
attn_resolutions: [16]
dropout: 0.0
load_VFI: #'net_220.pth'
num_head_channels: -1
lossconfig:
target: torch.nn.Identity
cond_stage_config: __is_first_stage__
BB:
optimizer:
weight_decay: 0.
optimizer: 'Adam'
lr: 1.e-4
beta1: 0.9
lr_scheduler:
factor: 0.5
patience: 3000
threshold: 0.0001
cooldown: 3000
min_lr: 5.e-7
params:
mt_type: 'linear' # options {'linear', 'sin'}
objective: 'BB' # options {'grad', 'noise', 'ysubx','BB'}
loss_type: 'l2' # options {'l1', 'l2'}
skip_sample: True
sample_type: 'linear' # options {"linear", "sin"}
sample_step: 10
num_timesteps: 1000 # timesteps
eta: 1.0 # DDIM reverse process eta
max_var: 1.0 # maximum variance
UNetParams:
image_size: 8
in_channels: 6
model_channels: 32
out_channels: 3
num_res_blocks: 1
attention_resolutions: !!python/tuple
- 8
- 4
- 2
channel_mult: !!python/tuple
- 1
- 1
- 1
conv_resample: True
dims: 3
num_heads: 1
use_scale_shift_norm: True
resblock_updown: True
use_max_self_attn: False # replace all full self-attention with MaxViT
context_dim:
dropout: 0
condition_key: "first_stage" # options {"SpatialRescaler", "first_stage", "nocond"}
View File
+721
View File
@@ -0,0 +1,721 @@
import torch
import cupy
import re
class Stream:
ptr = torch.cuda.current_stream().cuda_stream
# end
kernel_DSepconv_updateOutput = '''
extern "C" __global__ void kernel_DSepconv_updateOutput(
const int n,
const float* input,
const float* vertical,
const float* horizontal,
const float* offset_x,
const float* offset_y,
const float* mask,
float* output
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
float dblOutput = 0.0;
const int intSample = ( intIndex / SIZE_3(output) / SIZE_2(output) / SIZE_1(output) ) % SIZE_0(output);
const int intDepth = ( intIndex / SIZE_3(output) / SIZE_2(output) ) % SIZE_1(output);
const int intY = ( intIndex / SIZE_3(output) ) % SIZE_2(output);
const int intX = ( intIndex ) % SIZE_3(output);
for (int intFilterY = 0; intFilterY < SIZE_1(vertical); intFilterY += 1) {
for (int intFilterX = 0; intFilterX < SIZE_1(horizontal); intFilterX += 1) {
float delta_x = OFFSET_4(offset_y, intSample, intFilterY*SIZE_1(vertical) + intFilterX, intY, intX);
float delta_y = OFFSET_4(offset_x, intSample, intFilterY*SIZE_1(vertical) + intFilterX, intY, intX);
float position_x = delta_x + intX + intFilterX - (SIZE_1(horizontal) - 1) / 2 + 1;
float position_y = delta_y + intY + intFilterY - (SIZE_1(vertical) - 1) / 2 + 1;
if (position_x < 0)
position_x = 0;
if (position_x > SIZE_3(input) - 1)
position_x = SIZE_3(input) - 1;
if (position_y < 0)
position_y = 0;
if (position_y > SIZE_2(input) - 1)
position_y = SIZE_2(input) - 1;
int left = floor(delta_x + intX + intFilterX - (SIZE_1(horizontal) - 1) / 2 + 1);
int right = left + 1;
if (left < 0)
left = 0;
if (left > SIZE_3(input) - 1)
left = SIZE_3(input) - 1;
if (right < 0)
right = 0;
if (right > SIZE_3(input) - 1)
right = SIZE_3(input) - 1;
int top = floor(delta_y + intY + intFilterY - (SIZE_1(vertical) - 1) / 2 + 1);
int bottom = top + 1;
if (top < 0)
top = 0;
if (top > SIZE_2(input) - 1)
top = SIZE_2(input) - 1;
if (bottom < 0)
bottom = 0;
if (bottom > SIZE_2(input) - 1)
bottom = SIZE_2(input) - 1;
float floatValue = VALUE_4(input, intSample, intDepth, top, left) * (1 + (left - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, intDepth, top, right) * (1 - (right - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, intDepth, bottom, left) * (1 + (left - position_x)) * (1 - (bottom - position_y)) +
VALUE_4(input, intSample, intDepth, bottom, right) * (1 - (right - position_x)) * (1 - (bottom - position_y));
dblOutput += floatValue * VALUE_4(vertical, intSample, intFilterY, intY, intX) * VALUE_4(horizontal, intSample, intFilterX, intY, intX) * VALUE_4(mask, intSample, SIZE_1(vertical)*intFilterY + intFilterX, intY, intX);
}
}
output[intIndex] = dblOutput;
} }
'''
kernel_DSepconv_updateGradVertical = '''
extern "C" __global__ void kernel_DSepconv_updateGradVertical(
const int n,
const float* gradLoss,
const float* input,
const float* horizontal,
const float* offset_x,
const float* offset_y,
const float* mask,
float* gradVertical
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
float floatOutput = 0.0;
const int intSample = ( intIndex / SIZE_3(gradVertical) / SIZE_2(gradVertical) / SIZE_1(gradVertical) ) % SIZE_0(gradVertical);
const int intFilterY = ( intIndex / SIZE_3(gradVertical) / SIZE_2(gradVertical) ) % SIZE_1(gradVertical);
const int intY = ( intIndex / SIZE_3(gradVertical) ) % SIZE_2(gradVertical);
const int intX = ( intIndex ) % SIZE_3(gradVertical);
for (int intFilterX = 0; intFilterX < SIZE_1(horizontal); intFilterX += 1){
int intDepth = intFilterY * SIZE_1(horizontal) + intFilterX;
float delta_x = OFFSET_4(offset_y, intSample, intDepth, intY, intX);
float delta_y = OFFSET_4(offset_x, intSample, intDepth, intY, intX);
float position_x = delta_x + intX + intFilterX - (SIZE_1(horizontal) - 1) / 2 + 1;
float position_y = delta_y + intY + intFilterY - (SIZE_1(horizontal) - 1) / 2 + 1;
if (position_x < 0)
position_x = 0;
if (position_x > SIZE_3(input) - 1)
position_x = SIZE_3(input) - 1;
if (position_y < 0)
position_y = 0;
if (position_y > SIZE_2(input) - 1)
position_y = SIZE_2(input) - 1;
int left = floor(delta_x + intX + intFilterX - (SIZE_1(horizontal) - 1) / 2 + 1);
int right = left + 1;
if (left < 0)
left = 0;
if (left > SIZE_3(input) - 1)
left = SIZE_3(input) - 1;
if (right < 0)
right = 0;
if (right > SIZE_3(input) - 1)
right = SIZE_3(input) - 1;
int top = floor(delta_y + intY + intFilterY - (SIZE_1(horizontal) - 1) / 2 + 1);
int bottom = top + 1;
if (top < 0)
top = 0;
if (top > SIZE_2(input) - 1)
top = SIZE_2(input) - 1;
if (bottom < 0)
bottom = 0;
if (bottom > SIZE_2(input) - 1)
bottom = SIZE_2(input) - 1;
float floatSampled0 = VALUE_4(input, intSample, 0, top, left) * (1 + (left - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 0, top, right) * (1 - (right - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 0, bottom, left) * (1 + (left - position_x)) * (1 - (bottom - position_y)) +
VALUE_4(input, intSample, 0, bottom, right) * (1 - (right - position_x)) * (1 - (bottom - position_y));
float floatSampled1 = VALUE_4(input, intSample, 1, top, left) * (1 + (left - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 1, top, right) * (1 - (right - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 1, bottom, left) * (1 + (left - position_x)) * (1 - (bottom - position_y)) +
VALUE_4(input, intSample, 1, bottom, right) * (1 - (right - position_x)) * (1 - (bottom - position_y));
float floatSampled2 = VALUE_4(input, intSample, 2, top, left) * (1 + (left - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 2, top, right) * (1 - (right - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 2, bottom, left) * (1 + (left - position_x)) * (1 - (bottom - position_y)) +
VALUE_4(input, intSample, 2, bottom, right) * (1 - (right - position_x)) * (1 - (bottom - position_y));
floatOutput += VALUE_4(gradLoss, intSample, 0, intY, intX) * floatSampled0 * VALUE_4(horizontal, intSample, intFilterX, intY, intX) * VALUE_4(mask, intSample, intDepth, intY, intX) +
VALUE_4(gradLoss, intSample, 1, intY, intX) * floatSampled1 * VALUE_4(horizontal, intSample, intFilterX, intY, intX) * VALUE_4(mask, intSample, intDepth, intY, intX) +
VALUE_4(gradLoss, intSample, 2, intY, intX) * floatSampled2 * VALUE_4(horizontal, intSample, intFilterX, intY, intX) * VALUE_4(mask, intSample, intDepth, intY, intX);
}
gradVertical[intIndex] = floatOutput;
} }
'''
kernel_DSepconv_updateGradHorizontal = '''
extern "C" __global__ void kernel_DSepconv_updateGradHorizontal(
const int n,
const float* gradLoss,
const float* input,
const float* vertical,
const float* offset_x,
const float* offset_y,
const float* mask,
float* gradHorizontal
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
float floatOutput = 0.0;
const int intSample = ( intIndex / SIZE_3(gradHorizontal) / SIZE_2(gradHorizontal) / SIZE_1(gradHorizontal) ) % SIZE_0(gradHorizontal);
const int intFilterX = ( intIndex / SIZE_3(gradHorizontal) / SIZE_2(gradHorizontal) ) % SIZE_1(gradHorizontal);
const int intY = ( intIndex / SIZE_3(gradHorizontal) ) % SIZE_2(gradHorizontal);
const int intX = ( intIndex ) % SIZE_3(gradHorizontal);
for (int intFilterY = 0; intFilterY < SIZE_1(vertical); intFilterY += 1){
int intDepth = intFilterY * SIZE_1(vertical) + intFilterX;
float delta_x = OFFSET_4(offset_y, intSample, intDepth, intY, intX);
float delta_y = OFFSET_4(offset_x, intSample, intDepth, intY, intX);
float position_x = delta_x + intX + intFilterX - (SIZE_1(vertical) - 1) / 2 + 1;
float position_y = delta_y + intY + intFilterY - (SIZE_1(vertical) - 1) / 2 + 1;
if (position_x < 0)
position_x = 0;
if (position_x > SIZE_3(input) - 1)
position_x = SIZE_3(input) - 1;
if (position_y < 0)
position_y = 0;
if (position_y > SIZE_2(input) - 1)
position_y = SIZE_2(input) - 1;
int left = floor(delta_x + intX + intFilterX - (SIZE_1(vertical) - 1) / 2 + 1);
int right = left + 1;
if (left < 0)
left = 0;
if (left > SIZE_3(input) - 1)
left = SIZE_3(input) - 1;
if (right < 0)
right = 0;
if (right > SIZE_3(input) - 1)
right = SIZE_3(input) - 1;
int top = floor(delta_y + intY + intFilterY - (SIZE_1(vertical) - 1) / 2 + 1);
int bottom = top + 1;
if (top < 0)
top = 0;
if (top > SIZE_2(input) - 1)
top = SIZE_2(input) - 1;
if (bottom < 0)
bottom = 0;
if (bottom > SIZE_2(input) - 1)
bottom = SIZE_2(input) - 1;
float floatSampled0 = VALUE_4(input, intSample, 0, top, left) * (1 + (left - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 0, top, right) * (1 - (right - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 0, bottom, left) * (1 + (left - position_x)) * (1 - (bottom - position_y)) +
VALUE_4(input, intSample, 0, bottom, right) * (1 - (right - position_x)) * (1 - (bottom - position_y));
float floatSampled1 = VALUE_4(input, intSample, 1, top, left) * (1 + (left - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 1, top, right) * (1 - (right - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 1, bottom, left) * (1 + (left - position_x)) * (1 - (bottom - position_y)) +
VALUE_4(input, intSample, 1, bottom, right) * (1 - (right - position_x)) * (1 - (bottom - position_y));
float floatSampled2 = VALUE_4(input, intSample, 2, top, left) * (1 + (left - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 2, top, right) * (1 - (right - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, 2, bottom, left) * (1 + (left - position_x)) * (1 - (bottom - position_y)) +
VALUE_4(input, intSample, 2, bottom, right) * (1 - (right - position_x)) * (1 - (bottom - position_y));
floatOutput += VALUE_4(gradLoss, intSample, 0, intY, intX) * floatSampled0 * VALUE_4(vertical, intSample, intFilterY, intY, intX) * VALUE_4(mask, intSample, intDepth, intY, intX) +
VALUE_4(gradLoss, intSample, 1, intY, intX) * floatSampled1 * VALUE_4(vertical, intSample, intFilterY, intY, intX) * VALUE_4(mask, intSample, intDepth, intY, intX) +
VALUE_4(gradLoss, intSample, 2, intY, intX) * floatSampled2 * VALUE_4(vertical, intSample, intFilterY, intY, intX) * VALUE_4(mask, intSample, intDepth, intY, intX);
}
gradHorizontal[intIndex] = floatOutput;
} }
'''
kernel_DSepconv_updateGradMask = '''
extern "C" __global__ void kernel_DSepconv_updateGradMask(
const int n,
const float* gradLoss,
const float* input,
const float* vertical,
const float* horizontal,
const float* offset_x,
const float* offset_y,
float* gradMask
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
float floatOutput = 0.0;
const int intSample = ( intIndex / SIZE_3(gradMask) / SIZE_2(gradMask) / SIZE_1(gradMask) ) % SIZE_0(gradMask);
const int intDepth = ( intIndex / SIZE_3(gradMask) / SIZE_2(gradMask) ) % SIZE_1(gradMask);
const int intY = ( intIndex / SIZE_3(gradMask) ) % SIZE_2(gradMask);
const int intX = ( intIndex ) % SIZE_3(gradMask);
int intFilterY = intDepth / SIZE_1(vertical);
int intFilterX = intDepth % SIZE_1(vertical);
float delta_x = OFFSET_4(offset_y, intSample, intDepth, intY, intX);
float delta_y = OFFSET_4(offset_x, intSample, intDepth, intY, intX);
float position_x = delta_x + intX + intFilterX - (SIZE_1(vertical) - 1) / 2 + 1;
float position_y = delta_y + intY + intFilterY - (SIZE_1(vertical) - 1) / 2 + 1;
if (position_x < 0)
position_x = 0;
if (position_x > SIZE_3(input) - 1)
position_x = SIZE_3(input) - 1;
if (position_y < 0)
position_y = 0;
if (position_y > SIZE_2(input) - 1)
position_y = SIZE_2(input) - 1;
int left = floor(delta_x + intX + intFilterX - (SIZE_1(vertical) - 1) / 2 + 1);
int right = left + 1;
if (left < 0)
left = 0;
if (left > SIZE_3(input) - 1)
left = SIZE_3(input) - 1;
if (right < 0)
right = 0;
if (right > SIZE_3(input) - 1)
right = SIZE_3(input) - 1;
int top = floor(delta_y + intY + intFilterY - (SIZE_1(vertical) - 1) / 2 + 1);
int bottom = top + 1;
if (top < 0)
top = 0;
if (top > SIZE_2(input) - 1)
top = SIZE_2(input) - 1;
if (bottom < 0)
bottom = 0;
if (bottom > SIZE_2(input) - 1)
bottom = SIZE_2(input) - 1;
for (int intChannel = 0; intChannel < 3; intChannel++){
floatOutput += VALUE_4(gradLoss, intSample, intChannel, intY, intX) * (
VALUE_4(input, intSample, intChannel, top, left) * (1 + (left - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, intChannel, top, right) * (1 - (right - position_x)) * (1 + (top - position_y)) +
VALUE_4(input, intSample, intChannel, bottom, left) * (1 + (left - position_x)) * (1 - (bottom - position_y)) +
VALUE_4(input, intSample, intChannel, bottom, right) * (1 - (right - position_x)) * (1 - (bottom - position_y))
) * VALUE_4(vertical, intSample, intFilterY, intY, intX) * VALUE_4(horizontal, intSample, intFilterX, intY, intX);
}
gradMask[intIndex] = floatOutput;
} }
'''
kernel_DSepconv_updateGradOffsetX = '''
extern "C" __global__ void kernel_DSepconv_updateGradOffsetX(
const int n,
const float* gradLoss,
const float* input,
const float* vertical,
const float* horizontal,
const float* offset_x,
const float* offset_y,
const float* mask,
float* gradOffsetX
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
float floatOutput = 0.0;
const int intSample = ( intIndex / SIZE_3(gradOffsetX) / SIZE_2(gradOffsetX) / SIZE_1(gradOffsetX) ) % SIZE_0(gradOffsetX);
const int intDepth = ( intIndex / SIZE_3(gradOffsetX) / SIZE_2(gradOffsetX) ) % SIZE_1(gradOffsetX);
const int intY = ( intIndex / SIZE_3(gradOffsetX) ) % SIZE_2(gradOffsetX);
const int intX = ( intIndex ) % SIZE_3(gradOffsetX);
int intFilterY = intDepth / SIZE_1(vertical);
int intFilterX = intDepth % SIZE_1(vertical);
float delta_x = OFFSET_4(offset_y, intSample, intDepth, intY, intX);
float delta_y = OFFSET_4(offset_x, intSample, intDepth, intY, intX);
float position_x = delta_x + intX + intFilterX - (SIZE_1(vertical) - 1) / 2 + 1;
float position_y = delta_y + intY + intFilterY - (SIZE_1(vertical) - 1) / 2 + 1;
if (position_x < 0)
position_x = 0;
if (position_x > SIZE_3(input) - 1)
position_x = SIZE_3(input) - 1;
if (position_y < 0)
position_y = 0;
if (position_y > SIZE_2(input) - 1)
position_y = SIZE_2(input) - 1;
int left = floor(delta_x + intX + intFilterX - (SIZE_1(vertical) - 1) / 2 + 1);
int right = left + 1;
if (left < 0)
left = 0;
if (left > SIZE_3(input) - 1)
left = SIZE_3(input) - 1;
if (right < 0)
right = 0;
if (right > SIZE_3(input) - 1)
right = SIZE_3(input) - 1;
int top = floor(delta_y + intY + intFilterY - (SIZE_1(vertical) - 1) / 2 + 1);
int bottom = top + 1;
if (top < 0)
top = 0;
if (top > SIZE_2(input) - 1)
top = SIZE_2(input) - 1;
if (bottom < 0)
bottom = 0;
if (bottom > SIZE_2(input) - 1)
bottom = SIZE_2(input) - 1;
for (int intChannel = 0; intChannel < 3; intChannel++){
floatOutput += VALUE_4(gradLoss, intSample, intChannel, intY, intX) * (
- VALUE_4(input, intSample, intChannel, top, left) * (1 + (left - position_x))
- VALUE_4(input, intSample, intChannel, top, right) * (1 - (right - position_x))
+ VALUE_4(input, intSample, intChannel, bottom, left) * (1 + (left - position_x))
+ VALUE_4(input, intSample, intChannel, bottom, right) * (1 - (right - position_x))
)
* VALUE_4(vertical, intSample, intFilterY, intY, intX) * VALUE_4(horizontal, intSample, intFilterX, intY, intX)
* VALUE_4(mask, intSample, intDepth, intY, intX);
}
gradOffsetX[intIndex] = floatOutput;
} }
'''
kernel_DSepconv_updateGradOffsetY = '''
extern "C" __global__ void kernel_DSepconv_updateGradOffsetY(
const int n,
const float* gradLoss,
const float* input,
const float* vertical,
const float* horizontal,
const float* offset_x,
const float* offset_y,
const float* mask,
float* gradOffsetY
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
float floatOutput = 0.0;
const int intSample = ( intIndex / SIZE_3(gradOffsetX) / SIZE_2(gradOffsetX) / SIZE_1(gradOffsetX) ) % SIZE_0(gradOffsetX);
const int intDepth = ( intIndex / SIZE_3(gradOffsetX) / SIZE_2(gradOffsetX) ) % SIZE_1(gradOffsetX);
const int intY = ( intIndex / SIZE_3(gradOffsetX) ) % SIZE_2(gradOffsetX);
const int intX = ( intIndex ) % SIZE_3(gradOffsetX);
int intFilterY = intDepth / SIZE_1(vertical);
int intFilterX = intDepth % SIZE_1(vertical);
float delta_x = OFFSET_4(offset_y, intSample, intDepth, intY, intX);
float delta_y = OFFSET_4(offset_x, intSample, intDepth, intY, intX);
float position_x = delta_x + intX + intFilterX - (SIZE_1(vertical) - 1) / 2 + 1;
float position_y = delta_y + intY + intFilterY - (SIZE_1(vertical) - 1) / 2 + 1;
if (position_x < 0)
position_x = 0;
if (position_x > SIZE_3(input) - 1)
position_x = SIZE_3(input) - 1;
if (position_y < 0)
position_y = 0;
if (position_y > SIZE_2(input) - 1)
position_y = SIZE_2(input) - 1;
int left = floor(delta_x + intX + intFilterX - (SIZE_1(vertical) - 1) / 2 + 1);
int right = left + 1;
if (left < 0)
left = 0;
if (left > SIZE_3(input) - 1)
left = SIZE_3(input) - 1;
if (right < 0)
right = 0;
if (right > SIZE_3(input) - 1)
right = SIZE_3(input) - 1;
int top = floor(delta_y + intY + intFilterY - (SIZE_1(vertical) - 1) / 2 + 1);
int bottom = top + 1;
if (top < 0)
top = 0;
if (top > SIZE_2(input) - 1)
top = SIZE_2(input) - 1;
if (bottom < 0)
bottom = 0;
if (bottom > SIZE_2(input) - 1)
bottom = SIZE_2(input) - 1;
for (int intChannel = 0; intChannel < 3; intChannel++){
floatOutput += VALUE_4(gradLoss, intSample, intChannel, intY, intX) * (
- VALUE_4(input, intSample, intChannel, top, left) * (1 + (top - position_y))
+ VALUE_4(input, intSample, intChannel, top, right) * (1 + (top - position_y))
- VALUE_4(input, intSample, intChannel, bottom, left) * (1 - (bottom - position_y))
+ VALUE_4(input, intSample, intChannel, bottom, right) * (1 - (bottom - position_y))
)
* VALUE_4(vertical, intSample, intFilterY, intY, intX) * VALUE_4(horizontal, intSample, intFilterX, intY, intX)
* VALUE_4(mask, intSample, intDepth, intY, intX);
}
gradOffsetY[intIndex] = floatOutput;
} }
'''
def cupy_kernel(strFunction, objectVariables):
strKernel = globals()[strFunction]
while True:
objectMatch = re.search('(SIZE_)([0-4])(\()([^\)]*)(\))', strKernel)
if objectMatch is None:
break
# end
intArg = int(objectMatch.group(2))
strTensor = objectMatch.group(4)
intSizes = objectVariables[strTensor].size()
strKernel = strKernel.replace(objectMatch.group(), str(intSizes[intArg]))
# end
while True:
objectMatch = re.search('(VALUE_)([0-4])(\()([^\)]+)(\))', strKernel)
if objectMatch is None:
break
# end
intArgs = int(objectMatch.group(2))
strArgs = objectMatch.group(4).split(',')
strTensor = strArgs[0]
intStrides = objectVariables[strTensor].stride()
strIndex = ['((' + strArgs[intArg + 1].replace('{', '(').replace('}', ')').strip() + ')*' + str(
intStrides[intArg]) + ')' for intArg in range(intArgs)]
strKernel = strKernel.replace(objectMatch.group(0), strTensor + '[' + str.join('+', strIndex) + ']')
# end
while True:
objectMatch = re.search('(OFFSET_)([0-4])(\()([^\)]+)(\))', strKernel)
if objectMatch is None:
break
# end
intArgs = int(objectMatch.group(2))
strArgs = objectMatch.group(4).split(',')
strTensor = strArgs[0]
intStrides = objectVariables[strTensor].stride()
strIndex = ['((' + strArgs[intArg + 1].replace('{', '(').replace('}', ')').strip() + ')*' + str(
intStrides[intArg]) + ')' for intArg in range(intArgs)]
strKernel = strKernel.replace(objectMatch.group(0), strTensor + '[' + str.join('+', strIndex) + ']')
# end
return strKernel
# end
@cupy.memoize(for_each_device=True)
def cupy_launch(strFunction, strKernel):
# return cupy.cuda.compile_with_cache(strKernel).get_function(strFunction)
return cupy.RawKernel(strKernel, strFunction)
# end
class _FunctionDSepconv(torch.autograd.Function):
@staticmethod
def forward(self, input, vertical, horizontal, offset_x, offset_y, mask):
self.save_for_backward(input, vertical, horizontal, offset_x, offset_y, mask)
intSample = input.size(0)
intInputDepth = input.size(1)
intInputHeight = input.size(2)
intInputWidth = input.size(3)
intFilterSize = min(vertical.size(1), horizontal.size(1))
intOutputHeight = min(vertical.size(2), horizontal.size(2))
intOutputWidth = min(vertical.size(3), horizontal.size(3))
assert (intInputHeight == intOutputHeight + intFilterSize - 1)
assert (intInputWidth == intOutputWidth + intFilterSize - 1)
assert (input.is_contiguous() == True)
assert (vertical.is_contiguous() == True)
assert (horizontal.is_contiguous() == True)
assert (offset_x.is_contiguous() == True)
assert (offset_y.is_contiguous() == True)
assert (mask.is_contiguous() == True)
output = input.new_zeros([intSample, intInputDepth, intOutputHeight, intOutputWidth])
if input.is_cuda == True:
n = output.nelement()
cupy_launch('kernel_DSepconv_updateOutput', cupy_kernel('kernel_DSepconv_updateOutput', {
'input': input,
'vertical': vertical,
'horizontal': horizontal,
'offset_x': offset_x,
'offset_y': offset_y,
'mask': mask,
'output': output
}))(
grid=tuple([int((n + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[n, input.data_ptr(), vertical.data_ptr(), horizontal.data_ptr(), offset_x.data_ptr(), offset_y.data_ptr(),
mask.data_ptr(), output.data_ptr()],
stream=Stream
)
elif input.is_cuda == False:
raise NotImplementedError()
# end
return output
# end
@staticmethod
def backward(self, gradOutput):
input, vertical, horizontal, offset_x, offset_y, mask = self.saved_tensors
intSample = input.size(0)
intInputDepth = input.size(1)
intInputHeight = input.size(2)
intInputWidth = input.size(3)
intFilterSize = min(vertical.size(1), horizontal.size(1))
intOutputHeight = min(vertical.size(2), horizontal.size(2))
intOutputWidth = min(vertical.size(3), horizontal.size(3))
assert (intInputHeight == intOutputHeight + intFilterSize - 1)
assert (intInputWidth == intOutputWidth + intFilterSize - 1)
assert (gradOutput.is_contiguous() == True)
gradInput = input.new_zeros([intSample, intInputDepth, intInputHeight, intInputWidth]) if \
self.needs_input_grad[0] == True else None
gradVertical = input.new_zeros([intSample, intFilterSize, intOutputHeight, intOutputWidth]) if \
self.needs_input_grad[1] == True else None
gradHorizontal = input.new_zeros([intSample, intFilterSize, intOutputHeight, intOutputWidth]) if \
self.needs_input_grad[2] == True else None
gradOffsetX = input.new_zeros([intSample, intFilterSize * intFilterSize, intOutputHeight, intOutputWidth]) if \
self.needs_input_grad[3] == True else None
gradOffsetY = input.new_zeros([intSample, intFilterSize * intFilterSize, intOutputHeight, intOutputWidth]) if \
self.needs_input_grad[4] == True else None
gradMask = input.new_zeros([intSample, intFilterSize * intFilterSize, intOutputHeight, intOutputWidth]) if \
self.needs_input_grad[5] == True else None
if input.is_cuda == True:
nv = gradVertical.nelement()
cupy_launch('kernel_DSepconv_updateGradVertical', cupy_kernel('kernel_DSepconv_updateGradVertical', {
'gradLoss': gradOutput,
'input': input,
'horizontal': horizontal,
'offset_x': offset_x,
'offset_y': offset_y,
'mask': mask,
'gradVertical': gradVertical
}))(
grid=tuple([int((nv + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[nv, gradOutput.data_ptr(), input.data_ptr(), horizontal.data_ptr(), offset_x.data_ptr(),
offset_y.data_ptr(), mask.data_ptr(), gradVertical.data_ptr()],
stream=Stream
)
nh = gradHorizontal.nelement()
cupy_launch('kernel_DSepconv_updateGradHorizontal', cupy_kernel('kernel_DSepconv_updateGradHorizontal', {
'gradLoss': gradOutput,
'input': input,
'vertical': vertical,
'offset_x': offset_x,
'offset_y': offset_y,
'mask': mask,
'gradHorizontal': gradHorizontal
}))(
grid=tuple([int((nh + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[nh, gradOutput.data_ptr(), input.data_ptr(), vertical.data_ptr(), offset_x.data_ptr(),
offset_y.data_ptr(), mask.data_ptr(), gradHorizontal.data_ptr()],
stream=Stream
)
nx = gradOffsetX.nelement()
cupy_launch('kernel_DSepconv_updateGradOffsetX', cupy_kernel('kernel_DSepconv_updateGradOffsetX', {
'gradLoss': gradOutput,
'input': input,
'vertical': vertical,
'horizontal': horizontal,
'offset_x': offset_x,
'offset_y': offset_y,
'mask': mask,
'gradOffsetX': gradOffsetX
}))(
grid=tuple([int((nx + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[nx, gradOutput.data_ptr(), input.data_ptr(), vertical.data_ptr(), horizontal.data_ptr(), offset_x.data_ptr(),
offset_y.data_ptr(), mask.data_ptr(), gradOffsetX.data_ptr()],
stream=Stream
)
ny = gradOffsetY.nelement()
cupy_launch('kernel_DSepconv_updateGradOffsetY', cupy_kernel('kernel_DSepconv_updateGradOffsetY', {
'gradLoss': gradOutput,
'input': input,
'vertical': vertical,
'horizontal': horizontal,
'offset_x': offset_x,
'offset_y': offset_y,
'mask': mask,
'gradOffsetX': gradOffsetY
}))(
grid=tuple([int((ny + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[ny, gradOutput.data_ptr(), input.data_ptr(), vertical.data_ptr(), horizontal.data_ptr(),
offset_x.data_ptr(),
offset_y.data_ptr(), mask.data_ptr(), gradOffsetY.data_ptr()],
stream=Stream
)
nm = gradMask.nelement()
cupy_launch('kernel_DSepconv_updateGradMask', cupy_kernel('kernel_DSepconv_updateGradMask', {
'gradLoss': gradOutput,
'input': input,
'vertical': vertical,
'horizontal': horizontal,
'offset_x': offset_x,
'offset_y': offset_y,
'gradMask': gradMask
}))(
grid=tuple([int((nm + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[nm, gradOutput.data_ptr(), input.data_ptr(), vertical.data_ptr(), horizontal.data_ptr(),
offset_x.data_ptr(),
offset_y.data_ptr(), gradMask.data_ptr()],
stream=Stream
)
elif input.is_cuda == False:
raise NotImplementedError()
# end
return gradInput, gradVertical, gradHorizontal, gradOffsetX, gradOffsetY, gradMask
# end
# end
def FunctionDSepconv(tensorInput, tensorVertical, tensorHorizontal, tensorOffsetX, tensorOffsetY, tensorMask):
return _FunctionDSepconv.apply(tensorInput, tensorVertical, tensorHorizontal, tensorOffsetX, tensorOffsetY, tensorMask)
# end
class ModuleDSepconv(torch.nn.Module):
def __init__(self):
super(ModuleDSepconv, self).__init__()
# end
def forward(self, tensorInput, tensorVertical, tensorHorizontal, tensorOffsetX, tensorOffsetY, tensorMask):
return _FunctionDSepconv.apply(tensorInput, tensorVertical, tensorHorizontal, tensorOffsetX, tensorOffsetY, tensorMask)
# end
# end
# float floatValue = VALUE_4(input, intSample, intDepth, top, left) * (1 - (delta_x - floor(delta_x))) * (1 - (delta_y - floor(delta_y))) +
# VALUE_4(input, intSample, intDepth, top, right) * (delta_x - floor(delta_x)) * (1 - (delta_y - floor(delta_y))) +
# VALUE_4(input, intSample, intDepth, bottom, left) * (1 - (delta_x - floor(delta_x))) * (delta_y - floor(delta_y)) +
# VALUE_4(input, intSample, intDepth, bottom, right) * (delta_x - floor(delta_x)) * (delta_y - floor(delta_y));
@@ -0,0 +1,332 @@
import pdb
import torch
import torch.nn as nn
import torch.nn.functional as F
from functools import partial
from tqdm.autonotebook import tqdm
import numpy as np
from model.utils import extract, default
from model.BrownianBridge.base.modules.diffusionmodules.openaimodel import UNetModel
from model.BrownianBridge.base.modules.encoders.modules import SpatialRescaler
class BrownianBridgeModel(nn.Module):
def __init__(self, model_config):
super().__init__()
self.model_config = model_config
# model hyperparameters
model_params = model_config.BB.params
self.num_timesteps = model_params.num_timesteps
self.mt_type = model_params.mt_type
self.max_var = model_params.max_var if model_params.__contains__("max_var") else 1
self.eta = model_params.eta if model_params.__contains__("eta") else 1
self.skip_sample = model_params.skip_sample
self.sample_type = model_params.sample_type
self.sample_step = model_params.sample_step
self.steps = None
self.register_schedule()
self.next_frame = False
# loss and objective
self.loss_type = model_params.loss_type
self.objective = model_params.objective
# UNet
self.image_size = model_params.UNetParams.image_size
self.channels = model_params.UNetParams.in_channels
self.condition_key = model_params.UNetParams.condition_key
self.denoise_fn = UNetModel(**vars(model_params.UNetParams))
def register_schedule(self):
T = self.num_timesteps
if self.mt_type == "linear":
m_min, m_max = 0.001, 0.999
m_t = np.linspace(m_min, m_max, T)
elif self.mt_type == "sin":
m_t = 1.0075 ** np.linspace(0, T, T)
m_t = m_t / m_t[-1]
m_t[-1] = 0.999
else:
raise NotImplementedError
m_tminus = np.append(0, m_t[:-1]) ## left shifted mt
variance_t = 2. * (m_t - m_t ** 2) * self.max_var ## delta_t in the paper (variance of BB)
variance_tminus = np.append(0., variance_t[:-1]) ## left shifted delta_t
variance_t_tminus = variance_t - variance_tminus * ((1. - m_t) / (1. - m_tminus)) ** 2 ## delta t|t-1
posterior_variance_t = variance_t_tminus * variance_tminus / variance_t ## delta t in the reverse process
to_torch = partial(torch.tensor, dtype=torch.float32)
self.register_buffer('m_t', to_torch(m_t))
self.register_buffer('m_tminus', to_torch(m_tminus))
self.register_buffer('variance_t', to_torch(variance_t))
self.register_buffer('variance_tminus', to_torch(variance_tminus))
self.register_buffer('variance_t_tminus', to_torch(variance_t_tminus))
self.register_buffer('posterior_variance_t', to_torch(posterior_variance_t))
if self.skip_sample:
if self.sample_type == 'linear':
midsteps = torch.arange(self.num_timesteps - 1, 1,
step=-((self.num_timesteps - 1) / (self.sample_step - 2))).long()
self.steps = torch.cat((midsteps, torch.Tensor([1, 0]).long()), dim=0)
elif self.sample_type == 'cosine':
steps = np.linspace(start=0, stop=self.num_timesteps, num=self.sample_step + 1)
steps = (np.cos(steps / self.num_timesteps * np.pi) + 1.) / 2. * self.num_timesteps
self.steps = torch.from_numpy(steps)
else:
self.steps = torch.arange(self.num_timesteps-1, -1, -1)
def apply(self, weight_init):
self.denoise_fn.apply(weight_init)
return self
def get_parameters(self):
return self.denoise_fn.parameters()
def forward(self, x,y, context = None):
## context is default to None
## x is the target to sample (interpolated frame)
if self.condition_key == "nocond":
context = None
else:
context = y if context is None else context
b, c, f, h, w, device, img_size, = *x.shape, x.device, self.image_size
assert h == img_size and w == img_size, f'height and width of image must be {img_size}'
t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
return self.p_losses(x, y, context, t)
def compute_loss(self,x,y,loss_weights = 1):
diff = x - y
if self.loss_type == 'l1':
diff = diff.abs()
else:
diff = diff.pow(2.)
diff = diff*loss_weights
return diff.mean()
def bi_p_losses(self, x, y, z, context_y,context_z, t, noise=None):
"""
model loss
:param x: encoded x current frame
:param y: encoded y (previous frame)
:param z: encoded z (next frame)
:param t: timestep
:param noise: Standard Gaussian Noise
:return: loss
"""
b, c, h, w = x.shape
noise = default(noise, lambda: torch.randn_like(x))
loss_weights = 1
var_t = extract(self.variance_t, t, x.shape)
tmp = var_t
snr = torch.sqrt(1/tmp)
loss_weights = snr.clamp_(max = 5)
if self.next_frame:
## if we do next frame prediction, we diffuse from z to x to y
x_t_1, objective_1 = self.q_sample(x, y, t, noise)
x_t_2,objective_2 = self.q_sample(z, x, t, noise)
else:
## otherwise, we diffuse from x to y and x to z
x_t_1,objective_1 = self.q_sample(x, y, t, noise) ## from x to y
x_t_2,objective_2 = self.q_sample(x, z, t, noise) ## from x to z
if self.next_frame:
## next frame prediction will only condition on previous frame instead of bidirectional
objective_recon_1 = None
objective_recon_2 = self.denoise_fn(x_t_2, x_t_1, timesteps=t, context=context_y)
else:
if np.random.rand()>0.5:
objective_recon_2 = self.denoise_fn(x_t_2, x_t_1, timesteps=self.num_timesteps + (t+1), context=context_y)
## when we predicting noise from the next frame, change the context to next frame
else:
objective_recon_2 = self.denoise_fn(x_t_1, x_t_2, timesteps=self.num_timesteps - (t+1), context=context_z)
objective_2 = objective_1
if self.next_frame:
recloss = self.compute_loss(objective_2,objective_recon_2,loss_weights)
else:
recloss = self.compute_loss(objective_2,objective_recon_2,loss_weights) #+ self.compute_loss(objective_1,objective_recon_1,loss_weights)
'''
if self.loss_type == 'l1':
if self.next_frame:
#recloss = (objective_2 - objective_recon_2).abs().mean()
recloss = self.compute_loss(objective_2,objective_recon_2,loss_weights)
else:
recloss = (objective_1 - objective_recon_1).abs().mean() + (objective_2 - objective_recon_2).abs().mean()
elif self.loss_type == 'l2':
if self.next_frame:
recloss = F.mse_loss(objective_2, objective_recon_2)
else:
recloss = F.mse_loss(objective_1, objective_recon_1) + F.mse_loss(objective_2, objective_recon_2)
else:
raise NotImplementedError()
'''
if self.next_frame:
x0_recon = self.predict_x0_from_objective(x_t_2, z, t, objective_recon_2)
else:
#x0_recon = self.predict_x0_from_objective(x_t_1, y, t, objective_recon_1)
x0_recon = self.predict_x0_from_objective(x_t_2, y, t, objective_recon_2)
"""
x0_recon_next = self.predict_x0_from_objective(x_t_2, y, t, objective_recon_2)
consistent_loss = self.compute_loss(x0_recon,x0_recon_next,loss_weights)
"""
log_dict = {
"loss": recloss,
"x0_recon": x0_recon
}
return recloss, log_dict
def p_losses(self, x0, y, context, t, noise=None):
"""
model loss
:param x0: encoded x_ori, E(x_ori) = x0
:param y: encoded y_ori, E(y_ori) = y
:param y_ori: original source domain image
:param t: timestep
:param noise: Standard Gaussian Noise
:return: loss
"""
b, c, f, h, w = x0.shape
noise = default(noise, lambda: torch.randn_like(x0))
x_t, objective = self.q_sample(x0, y, t, noise)
objective_recon = self.denoise_fn(x_t, cond = y, timesteps=t, context=context)
if self.loss_type == 'l1':
recloss = (objective - objective_recon).abs().mean()
elif self.loss_type == 'l2':
recloss = F.mse_loss(objective, objective_recon)
else:
raise NotImplementedError()
x0_recon = self.predict_x0_from_objective(x_t, y, t, objective_recon)
log_dict = {
"loss": recloss,
"x0_recon": x0_recon
}
return recloss, log_dict
def q_sample(self, x0, y, t, noise=None):
noise = default(noise, lambda: torch.randn_like(x0))
m_t = extract(self.m_t, t, x0.shape)
var_t = extract(self.variance_t, t, x0.shape)
sigma_t = torch.sqrt(var_t)
x_t = (1. - m_t) * x0 + m_t * y + sigma_t * noise
if self.objective == 'grad':
#objective = m_t * (y - x0) + sigma_t * noise
objective = x_t - x0
elif self.objective == 'noise':
objective = noise
elif self.objective == 'ysubx':
objective = y - x0
elif self.objective == 'BB':
objective = x_t - x0
else:
raise NotImplementedError()
return (
x_t,
objective
)
def predict_x0_from_objective(self, x_t, y, t, objective_recon):
if self.objective == 'grad':
x0_recon = x_t - objective_recon
elif self.objective == 'noise':
m_t = extract(self.m_t, t, x_t.shape)
var_t = extract(self.variance_t, t, x_t.shape)
sigma_t = torch.sqrt(var_t)
x0_recon = (x_t - m_t * y - sigma_t * objective_recon) / (1. - m_t)
elif self.objective == 'ysubx':
x0_recon = y - objective_recon
elif self.objective == 'BB':
x0_recon = -objective_recon + x_t ## if predicting xt - x0
else:
raise NotImplementedError
return x0_recon
@torch.no_grad()
def q_sample_loop(self, x0, y):
imgs = [x0]
for i in tqdm(range(self.num_timesteps), desc='q sampling loop', total=self.num_timesteps):
t = torch.full((y.shape[0],), i, device=x0.device, dtype=torch.long)
img, _ = self.q_sample(x0, y, t)
imgs.append(img)
return imgs
@torch.no_grad()
def p_sample(self, x_t, y, context, i, clip_denoised=False):
b, *_, device = *x_t.shape, x_t.device
if self.steps[i] == 0:
t = torch.full((x_t.shape[0],), self.steps[i], device=x_t.device, dtype=torch.long)
objective_recon = self.denoise_fn(x_t, timesteps=t, context=context)
x0_recon = self.predict_x0_from_objective(x_t, y, t, objective_recon=objective_recon)
if clip_denoised:
x0_recon.clamp_(-1., 1.)
return x0_recon, x0_recon
else:
t = torch.full((x_t.shape[0],), self.steps[i], device=x_t.device, dtype=torch.long)
n_t = torch.full((x_t.shape[0],), self.steps[i+1], device=x_t.device, dtype=torch.long)
objective_recon = self.denoise_fn(x_t, timesteps=t, cond = None, context=context)
x0_recon = self.predict_x0_from_objective(x_t, y, t, objective_recon=objective_recon)
if clip_denoised:
x0_recon.clamp_(-1., 1.)
m_t = extract(self.m_t, t, x_t.shape)
m_nt = extract(self.m_t, n_t, x_t.shape)
var_t = extract(self.variance_t, t, x_t.shape)
var_nt = extract(self.variance_t, n_t, x_t.shape)
sigma2_t = (var_t - var_nt * (1. - m_t) ** 2 / (1. - m_nt) ** 2) * var_nt / var_t
sigma_t = torch.sqrt(sigma2_t) * self.eta
noise = torch.randn_like(x_t)
x_tminus_mean = (1. - m_nt) * x0_recon + m_nt * y + torch.sqrt((var_nt - sigma2_t) / var_t) * \
(x_t - (1. - m_t) * x0_recon - m_t * y)
return x_tminus_mean + sigma_t * noise, x0_recon
@torch.no_grad()
def p_sample_loop(self, y, context=None, clip_denoised=True, sample_mid_step=False):
if self.condition_key == "nocond":
context = None
else:
context = y if context is None else context
if sample_mid_step:
imgs, one_step_imgs = [y], []
for i in tqdm(range(len(self.steps)), desc=f'sampling loop time step', total=len(self.steps)):
img, x0_recon = self.p_sample(x_t=imgs[-1], y=y, context=context, i=i, clip_denoised=clip_denoised)
imgs.append(img)
one_step_imgs.append(x0_recon)
return imgs, one_step_imgs
else:
img = y
for i in tqdm(range(len(self.steps)), desc=f'sampling loop time step', total=len(self.steps)):
img, _ = self.p_sample(x_t=img, y=y, context=context, i=i, clip_denoised=clip_denoised)
return img
@torch.no_grad()
def sample(self, y,z, context_y=None, context_z=None, clip_denoised=True, sample_mid_step=False):
## y: previous frame
## z: next frame in interpolation, bad frame in inpainting, current frame in next frame prediction
## context will be concatenated to x, also cross_attended
return self.p_sample_loop(y, z, context_y,context_z, clip_denoised, sample_mid_step)
@@ -0,0 +1,130 @@
import itertools
import pdb
import random
import torch
import torch.nn as nn
from tqdm.autonotebook import tqdm
from einops import rearrange,repeat
from model.BrownianBridge.BrownianBridgeModel import BrownianBridgeModel
from model.BrownianBridge.base.modules.encoders.modules import SpatialRescaler
from model.VQGAN.vqgan import VQFlowNetInterface
def disabled_train(self, mode=True):
"""Overwrite model.train with this function to make sure train/eval mode
does not change anymore."""
return self
class LatentBrownianBridgeModel(BrownianBridgeModel):
def __init__(self, model_config):
super().__init__(model_config)
self.vqgan = VQFlowNetInterface(**vars(model_config.VQGAN.params)).eval()
self.vqgan.train = disabled_train
for param in self.vqgan.parameters():
param.requires_grad = False
# Condition Stage Model
if self.condition_key == 'nocond':
self.cond_stage_model = None
elif self.condition_key == 'first_stage':
self.cond_stage_model = self.vqgan ## VQGAN quantization
elif self.condition_key == 'SpatialRescaler':
self.cond_stage_model = SpatialRescaler(**vars(model_config.CondStageParams)) ## interpolation
else:
raise NotImplementedError
def get_ema_net(self):
return self
def get_parameters(self):
if self.condition_key == 'SpatialRescaler':
print("get parameters to optimize: SpatialRescaler, UNet")
params = itertools.chain(self.denoise_fn.parameters(), self.cond_stage_model.parameters())
else:
print("get parameters to optimize: UNet")
params = self.denoise_fn.parameters()
return params
def apply(self, weights_init):
super().apply(weights_init)
if self.cond_stage_model is not None:
self.cond_stage_model.apply(weights_init)
return self
def forward(self, x, y, z, context=None):
with torch.no_grad():
gt,_ = self.encode(torch.cat([y,x,z],dim = 0).detach())
gt = torch.stack(torch.chunk(gt,3),2)
latent,_ = self.encode(torch.cat([y,torch.zeros_like(y),z],dim = 0))
latent = torch.stack(torch.chunk(latent,3),2)
return super().forward(gt, latent, context)
def get_cond_stage_context(self, x_cond):
if self.cond_stage_model is not None:
if self.condition_key == 'first_stage':
context = self.encode(x_cond,cond = True)[0].detach()
else:
context = self.cond_stage_model(x_cond)
else:
context = None
return context
@torch.no_grad()
def encode(self, x, cond=True, normalize=None):
model = self.vqgan
if cond:
x_latent,phi_list = model.encode(x,ret_feature=cond)
return x_latent,phi_list
else:
x_latent = model.encode(x,ret_feature=cond)
return x_latent
@torch.no_grad()
def decode(self, x_latent, prev_img,next_img,phi_list,scale = 0.5):
model = self.vqgan ## latent B C F H W
x_latent = x_latent.permute(0,2,1,3,4) ## B C F H W --> B F C H W
x_latent = rearrange(x_latent,'b f c h w -> (b f) c h w') ## BF C H W
out = model.decode(x_latent,prev_img,next_img,phi_list,scale = scale)
return out
@torch.no_grad()
def latent_p_sample_loop(self, latent, y, context, clip_denoised=True, sample_mid_step=False, disable_progress=False):
imgs,one_step_imgs = [y],[]
# Added disable=disable_progress to the tqdm call
for i in tqdm(range(len(self.steps)), desc=f'sampling loop time step', total=len(self.steps), disable=disable_progress):
img, x0_recon = self.p_sample(x_t = imgs[-1],y=y,context = context,i = i)
imgs.append(img)
one_step_imgs.append(x0_recon)
return imgs,one_step_imgs
@torch.no_grad()
def sample(self, y, z, clip_denoised=False, sample_mid_step=False,scale = 0.5, disable_progress=False):
x = torch.zeros_like(y)
latent,phi_list = self.encode(torch.cat([y,x,z],0))
latent = torch.stack(torch.chunk(latent,3),2)
context = latent
# Pass disable_progress to the sampling loop
imgs,one_step_imgs = self.latent_p_sample_loop(latent = latent,
y = latent,
context = context,
clip_denoised=clip_denoised,
sample_mid_step=sample_mid_step,
disable_progress=disable_progress)
with torch.no_grad():
out = self.decode(imgs[-1].detach(), y,z,phi_list,scale = scale)
return out
@torch.no_grad()
def sample_vqgan(self, x):
x_rec, _ = self.vqgan(x)
return x_rec
def get_flow(self, img0, img1,feats):
return self.vqgan.get_flow(self,img0,img1,feats)
@@ -0,0 +1,349 @@
import pdb
from inspect import isfunction
import math
import torch
import torch.nn.functional as F
from torch import nn, einsum
from einops import rearrange, repeat
from model.BrownianBridge.base.modules.diffusionmodules.util import checkpoint
def exists(val):
return val is not None
def uniq(arr):
return{el: True for el in arr}.keys()
def default(val, d):
if exists(val):
return val
return d() if isfunction(d) else d
def max_neg_value(t):
return -torch.finfo(t.dtype).max
def init_(tensor):
dim = tensor.shape[-1]
std = 1 / math.sqrt(dim)
tensor.uniform_(-std, std)
return tensor
# feedforward
class GEGLU(nn.Module):
def __init__(self, dim_in, dim_out):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out * 2)
def forward(self, x):
x, gate = self.proj(x).chunk(2, dim=-1)
return x * F.gelu(gate)
class FeedForward(nn.Module):
def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0.):
super().__init__()
inner_dim = int(dim * mult)
dim_out = default(dim_out, dim)
project_in = nn.Sequential(
nn.Linear(dim, inner_dim),
nn.GELU()
) if not glu else GEGLU(dim, inner_dim)
self.net = nn.Sequential(
project_in,
nn.Dropout(dropout),
nn.Linear(inner_dim, dim_out)
)
def forward(self, x):
return self.net(x)
def zero_module(module):
"""
Zero out the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().zero_()
return module
def Normalize(in_channels):
return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
class LinearAttention(nn.Module):
def __init__(self, dim, heads=4, dim_head=32):
super().__init__()
self.heads = heads
hidden_dim = dim_head * heads
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
self.to_out = nn.Conv2d(hidden_dim, dim, 1)
def forward(self, x):
b, c, h, w = x.shape
qkv = self.to_qkv(x)
q, k, v = rearrange(qkv, 'b (qkv heads c) h w -> qkv b heads c (h w)', heads = self.heads, qkv=3)
k = k.softmax(dim=-1)
context = torch.einsum('bhdn,bhen->bhde', k, v)
out = torch.einsum('bhde,bhdn->bhen', context, q)
out = rearrange(out, 'b heads c (h w) -> b (heads c) h w', heads=self.heads, h=h, w=w)
return self.to_out(out)
class SpatialSelfAttention(nn.Module):
def __init__(self, in_channels):
super().__init__()
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.k = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.v = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.proj_out = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
def forward(self, x):
h_ = x
h_ = self.norm(h_)
q = self.q(h_)
k = self.k(h_)
v = self.v(h_)
# compute attention
b,c,h,w = q.shape
q = rearrange(q, 'b c h w -> b (h w) c')
k = rearrange(k, 'b c h w -> b c (h w)')
w_ = torch.einsum('bij,bjk->bik', q, k)
w_ = w_ * (int(c)**(-0.5))
w_ = torch.nn.functional.softmax(w_, dim=2)
# attend to values
v = rearrange(v, 'b c h w -> b c (h w)')
w_ = rearrange(w_, 'b i j -> b j i')
h_ = torch.einsum('bij,bjk->bik', v, w_)
h_ = rearrange(h_, 'b c (h w) -> b c h w', h=h)
h_ = self.proj_out(h_)
return x+h_
class CrossAttention(nn.Module):
def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0.):
super().__init__()
inner_dim = dim_head * heads
context_dim = default(context_dim, query_dim)
self.scale = dim_head ** -0.5
self.heads = heads
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
self.to_out = nn.Sequential(
nn.Linear(inner_dim, query_dim),
nn.Dropout(dropout)
)
def forward(self, x, context=None, mask=None):
h = self.heads
q = self.to_q(x)
if context is not None:
context = rearrange(context, 'b c h w -> b (h w) c')
context = default(context, x)
k = self.to_k(context)
v = self.to_v(context)
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v))
sim = einsum('b i d, b j d -> b i j', q, k) * self.scale
if exists(mask):
mask = rearrange(mask, 'b ... -> b (...)')
max_neg_value = -torch.finfo(sim.dtype).max
mask = repeat(mask, 'b j -> (b h) () j', h=h)
sim.masked_fill_(~mask, max_neg_value)
# attention, what we cannot get enough of
attn = sim.softmax(dim=-1)
out = einsum('b i j, b j d -> b i d', attn, v)
out = rearrange(out, '(b h) n d -> b n (h d)', h=h)
return self.to_out(out)
class BasicTransformerBlock(nn.Module):
def __init__(self, dim, n_heads, d_head, dropout=0., context_dim=None, gated_ff=True, checkpoint=True):
super().__init__()
self.attn1 = CrossAttention(query_dim=dim, heads=n_heads, dim_head=d_head, dropout=dropout) # is a self-attention
self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff)
self.attn2 = CrossAttention(query_dim=dim, context_dim=context_dim,
heads=n_heads, dim_head=d_head, dropout=dropout) # is self-attn if context is none
self.norm1 = nn.LayerNorm(dim)
self.norm2 = nn.LayerNorm(dim)
self.norm3 = nn.LayerNorm(dim)
self.checkpoint = checkpoint
def forward(self, x, context=None):
return checkpoint(self._forward, (x, context), self.parameters(), self.checkpoint)
def _forward(self, x, context=None):
x = self.attn1(self.norm1(x)) + x
x = self.attn2(self.norm2(x), context=context) + x
x = self.ff(self.norm3(x)) + x
return x
class SpatialTransformer(nn.Module):
"""
Transformer block for image-like data.
First, project the input (aka embedding)
and reshape to b, t, d.
Then apply standard transformer action.
Finally, reshape to image
"""
def __init__(self, in_channels, n_heads, d_head,
depth=1, dropout=0., context_dim=None):
super().__init__()
self.in_channels = in_channels
inner_dim = n_heads * d_head
self.norm = Normalize(in_channels)
self.proj_in = nn.Conv2d(in_channels,
inner_dim,
kernel_size=1,
stride=1,
padding=0)
self.transformer_blocks = nn.ModuleList(
[BasicTransformerBlock(inner_dim, n_heads, d_head, dropout=dropout, context_dim=context_dim)
for d in range(depth)]
)
self.proj_out = zero_module(nn.Conv2d(inner_dim,
in_channels,
kernel_size=1,
stride=1,
padding=0))
def forward(self, x, context=None):
# note: if no context is given, cross-attention defaults to self-attention
b, c, h, w = x.shape
x_in = x
x = self.norm(x)
x = self.proj_in(x)
x = rearrange(x, 'b c h w -> b (h w) c')
for block in self.transformer_blocks:
x = block(x, context=context)
x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w)
x = self.proj_out(x)
return x + x_in
class SpatialCrossAttentionWithPosEmb(nn.Module):
'''
Cross-attention block for image-like data.
First image reshape to b, t, d.
Perform self-attention if context is None, else cross-attention.
The dims of the input and output of the block are the same (arg query_dim).
'''
def __init__(self, in_channels=None, heads=8, dim_head=64, dropout=0.):
super().__init__()
inner_dim = dim_head * heads
self.scale = dim_head ** -0.5
self.heads = heads
self.proj_in = nn.Conv2d(in_channels,
inner_dim,
kernel_size=1,
stride=1,
padding=0)
self.to_q = nn.Linear(inner_dim, inner_dim, bias=False)
self.to_k = nn.Linear(inner_dim, inner_dim, bias=False)
self.to_v = nn.Linear(inner_dim, inner_dim, bias=False)
self.to_out = nn.Sequential(
nn.Linear(inner_dim, inner_dim),
nn.Dropout(dropout)
)
self.proj_out = zero_module(nn.Conv2d(inner_dim,
in_channels,
kernel_size=1,
stride=1,
padding=0))
self.norm = nn.LayerNorm(inner_dim)
def forward(self, x, context=None):
b, c, h, w = x.shape
x_in = x
context = default(context, x)
x = self.proj_in(x) # (b,d,h,w)
context = self.proj_in(context) # (b,d,h,w)
# positional embedding
pe = posemb_sincos_2d(x) # (n,d)
# re-arrange image data to b, n, d.
x = rearrange(x, 'b c h w -> b (h w) c')
if (len(context.shape) == 4):
context = rearrange(context, 'b c h w -> b (h w) c')
# add pos emb
x += pe
if context.shape[1] != x.shape[1]:
context[:,:h*w] += pe
context[:,h*w:] += pe
else:
context += pe
heads = self.heads
x = self.norm(x)
context = self.norm(context)
q = self.to_q(x)
k = self.to_k(context)
v = self.to_v(context)
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=heads), (q, k, v))
sim = einsum('b i d, b j d -> b i j', q, k) * self.scale
# attention, what we cannot get enough of
attn = sim.softmax(dim=-1)
out = einsum('b i j, b j d -> b i d', attn, v)
out = rearrange(out, '(b h) n d -> b n (h d)', h=heads)
out = self.to_out(out)
# restore image shape
out = rearrange(out, 'b (h w) c -> b c h w', h=h, w=w)
return x_in + out
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,994 @@
import pdb
from abc import abstractmethod
from functools import partial
import math
from typing import Iterable
import numpy as np
import torch as th
import torch.nn as nn
import torch.nn.functional as F
from model.BrownianBridge.base.modules.diffusionmodules.util import (
checkpoint,
conv_nd,
linear,
avg_pool_nd,
zero_module,
normalization,
timestep_embedding,
)
from model.BrownianBridge.base.modules.attention import SpatialTransformer
from model.BrownianBridge.base.modules.maxvit import SpatialTransformerWithMax, MaxAttentionBlock
# dummy replace
def convert_module_to_f16(x):
pass
def convert_module_to_f32(x):
pass
## go
class AttentionPool2d(nn.Module):
"""
Adapted from CLIP: https://github.com/openai/CLIP/blob/main/clip/model.py
"""
def __init__(
self,
spacial_dim: int,
embed_dim: int,
num_heads_channels: int,
output_dim: int = None,
):
super().__init__()
self.positional_embedding = nn.Parameter(th.randn(embed_dim, spacial_dim ** 2 + 1) / embed_dim ** 0.5)
self.qkv_proj = conv_nd(1, embed_dim, 3 * embed_dim, 1)
self.c_proj = conv_nd(1, embed_dim, output_dim or embed_dim, 1)
self.num_heads = embed_dim // num_heads_channels
self.attention = QKVAttention(self.num_heads)
def forward(self, x):
b, c, *_spatial = x.shape
x = x.reshape(b, c, -1) # NC(HW)
x = th.cat([x.mean(dim=-1, keepdim=True), x], dim=-1) # NC(HW+1)
x = x + self.positional_embedding[None, :, :].to(x.dtype) # NC(HW+1)
x = self.qkv_proj(x)
x = self.attention(x)
x = self.c_proj(x)
return x[:, :, 0]
class TimestepBlock(nn.Module):
"""
Any module where forward() takes timestep embeddings as a second argument.
"""
@abstractmethod
def forward(self, x, emb):
"""
Apply the module to `x` given `emb` timestep embeddings.
"""
class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
"""
A sequential module that passes timestep embeddings to the children that
support it as an extra input.
"""
def forward(self, x, emb, context=None):
# pdb.set_trace()
for layer in self:
if isinstance(layer, TimestepBlock):
x = layer(x, emb)
elif isinstance(layer, SpatialTransformer):
x = layer(x, context)
else:
x = layer(x)
return x
class Upsample(nn.Module):
"""
An upsampling layer with an optional convolution.
:param channels: channels in the inputs and outputs.
:param use_conv: a bool determining if a convolution is applied.
:param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then
upsampling occurs in the inner-two dimensions.
"""
def __init__(self, channels, use_conv, dims=2, out_channels=None, padding=1):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.use_conv = use_conv
self.dims = dims
if use_conv:
self.conv = conv_nd(dims, self.channels, self.out_channels, 3, padding=padding)
def forward(self, x):
assert x.shape[1] == self.channels
if self.dims == 3:
x = F.interpolate(
x, (x.shape[2], x.shape[3] * 2, x.shape[4] * 2), mode="nearest"
)
else:
x = F.interpolate(x, scale_factor=2, mode="nearest")
if self.use_conv:
x = self.conv(x)
return x
class TransposedUpsample(nn.Module):
'Learned 2x upsampling without padding'
def __init__(self, channels, out_channels=None, ks=5):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.up = nn.ConvTranspose2d(self.channels,self.out_channels,kernel_size=ks,stride=2)
def forward(self,x):
return self.up(x)
class Downsample(nn.Module):
"""
A downsampling layer with an optional convolution.
:param channels: channels in the inputs and outputs.
:param use_conv: a bool determining if a convolution is applied.
:param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then
downsampling occurs in the inner-two dimensions.
"""
def __init__(self, channels, use_conv, dims=2, out_channels=None,padding=1):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.use_conv = use_conv
self.dims = dims
stride = 2 if dims != 3 else (1, 2, 2)
if use_conv:
self.op = conv_nd(
dims, self.channels, self.out_channels, 3, stride=stride, padding=padding
)
else:
assert self.channels == self.out_channels
self.op = avg_pool_nd(dims, kernel_size=stride, stride=stride)
def forward(self, x):
assert x.shape[1] == self.channels
return self.op(x)
class ResBlock(TimestepBlock):
"""
A residual block that can optionally change the number of channels.
:param channels: the number of input channels.
:param emb_channels: the number of timestep embedding channels.
:param dropout: the rate of dropout.
:param out_channels: if specified, the number of out channels.
:param use_conv: if True and out_channels is specified, use a spatial
convolution instead of a smaller 1x1 convolution to change the
channels in the skip connection.
:param dims: determines if the signal is 1D, 2D, or 3D.
:param use_checkpoint: if True, use gradient checkpointing on this module.
:param up: if True, use this block for upsampling.
:param down: if True, use this block for downsampling.
"""
def __init__(
self,
channels,
emb_channels,
dropout,
out_channels=None,
use_conv=False,
use_scale_shift_norm=False,
dims=2,
use_checkpoint=False,
up=False,
down=False,
):
super().__init__()
self.channels = channels
self.emb_channels = emb_channels
self.dropout = dropout
self.out_channels = out_channels or channels
self.use_conv = use_conv
self.use_checkpoint = use_checkpoint
self.use_scale_shift_norm = use_scale_shift_norm
self.in_layers = nn.Sequential(
normalization(channels),
nn.SiLU(),
conv_nd(dims, channels, self.out_channels, 3, padding=1),
)
self.updown = up or down
if up:
self.h_upd = Upsample(channels, False, dims)
self.x_upd = Upsample(channels, False, dims)
elif down:
self.h_upd = Downsample(channels, False, dims)
self.x_upd = Downsample(channels, False, dims)
else:
self.h_upd = self.x_upd = nn.Identity()
self.emb_layers = nn.Sequential(
nn.SiLU(),
linear(
emb_channels,
2 * self.out_channels if use_scale_shift_norm else self.out_channels,
),
)
self.out_layers = nn.Sequential(
normalization(self.out_channels),
nn.SiLU(),
nn.Dropout(p=dropout),
zero_module(
conv_nd(dims, self.out_channels, self.out_channels, 3, padding=1)
),
)
if self.out_channels == channels:
self.skip_connection = nn.Identity()
elif use_conv:
self.skip_connection = conv_nd(
dims, channels, self.out_channels, 3, padding=1
)
else:
self.skip_connection = conv_nd(dims, channels, self.out_channels, 1)
def forward(self, x, emb):
"""
Apply the block to a Tensor, conditioned on a timestep embedding.
:param x: an [N x C x ...] Tensor of features.
:param emb: an [N x emb_channels] Tensor of timestep embeddings.
:return: an [N x C x ...] Tensor of outputs.
"""
return checkpoint(
self._forward, (x, emb), self.parameters(), self.use_checkpoint
)
def _forward(self, x, emb):
if self.updown:
in_rest, in_conv = self.in_layers[:-1], self.in_layers[-1]
h = in_rest(x)
h = self.h_upd(h)
x = self.x_upd(x)
h = in_conv(h)
else:
h = self.in_layers(x)
emb_out = self.emb_layers(emb).type(h.dtype)
while len(emb_out.shape) < len(h.shape):
emb_out = emb_out[..., None]
if self.use_scale_shift_norm:
out_norm, out_rest = self.out_layers[0], self.out_layers[1:]
scale, shift = th.chunk(emb_out, 2, dim=1)
h = out_norm(h) * (1 + scale) + shift
h = out_rest(h)
else:
h = h + emb_out
h = self.out_layers(h)
return self.skip_connection(x) + h
class AttentionBlock(nn.Module):
"""
An attention block that allows spatial positions to attend to each other.
Originally ported from here, but adapted to the N-d case.
https://github.com/hojonathanho/diffusion/blob/1e0dceb3b3495bbe19116a5e1b3596cd0706c543/diffusion_tf/models/unet.py#L66.
"""
def __init__(
self,
channels,
num_heads=1,
num_head_channels=-1,
use_checkpoint=False,
use_new_attention_order=False,
):
super().__init__()
self.channels = channels
if num_head_channels == -1:
self.num_heads = num_heads
else:
assert (
channels % num_head_channels == 0
), f"q,k,v channels {channels} is not divisible by num_head_channels {num_head_channels}"
self.num_heads = channels // num_head_channels
self.use_checkpoint = use_checkpoint
self.norm = normalization(channels)
self.qkv = conv_nd(1, channels, channels * 3, 1)
if use_new_attention_order:
# split qkv before split heads
self.attention = QKVAttention(self.num_heads)
else:
# split heads before split qkv
self.attention = QKVAttentionLegacy(self.num_heads)
self.proj_out = zero_module(conv_nd(1, channels, channels, 1))
def forward(self, x):
return checkpoint(self._forward, (x,), self.parameters(), True) # TODO: check checkpoint usage, is True # TODO: fix the .half call!!!
#return pt_checkpoint(self._forward, x) # pytorch
def _forward(self, x):
b, c, *spatial = x.shape
x = x.reshape(b, c, -1)
qkv = self.qkv(self.norm(x))
h = self.attention(qkv)
h = self.proj_out(h)
return (x + h).reshape(b, c, *spatial)
def count_flops_attn(model, _x, y):
"""
A counter for the `thop` package to count the operations in an
attention operation.
Meant to be used like:
macs, params = thop.profile(
model,
inputs=(inputs, timestamps),
custom_ops={QKVAttention: QKVAttention.count_flops},
)
"""
b, c, *spatial = y[0].shape
num_spatial = int(np.prod(spatial))
# We perform two matmuls with the same number of ops.
# The first computes the weight matrix, the second computes
# the combination of the value vectors.
matmul_ops = 2 * b * (num_spatial ** 2) * c
model.total_ops += th.DoubleTensor([matmul_ops])
class QKVAttentionLegacy(nn.Module):
"""
A module which performs QKV attention. Matches legacy QKVAttention + input/ouput heads shaping
"""
def __init__(self, n_heads):
super().__init__()
self.n_heads = n_heads
def forward(self, qkv):
"""
Apply QKV attention.
:param qkv: an [N x (H * 3 * C) x T] tensor of Qs, Ks, and Vs.
:return: an [N x (H * C) x T] tensor after attention.
"""
bs, width, length = qkv.shape
assert width % (3 * self.n_heads) == 0
ch = width // (3 * self.n_heads)
q, k, v = qkv.reshape(bs * self.n_heads, ch * 3, length).split(ch, dim=1)
scale = 1 / math.sqrt(math.sqrt(ch))
weight = th.einsum(
"bct,bcs->bts", q * scale, k * scale
) # More stable with f16 than dividing afterwards
weight = th.softmax(weight.float(), dim=-1).type(weight.dtype)
a = th.einsum("bts,bcs->bct", weight, v)
return a.reshape(bs, -1, length)
@staticmethod
def count_flops(model, _x, y):
return count_flops_attn(model, _x, y)
class QKVAttention(nn.Module):
"""
A module which performs QKV attention and splits in a different order.
"""
def __init__(self, n_heads):
super().__init__()
self.n_heads = n_heads
def forward(self, qkv):
"""
Apply QKV attention.
:param qkv: an [N x (3 * H * C) x T] tensor of Qs, Ks, and Vs.
:return: an [N x (H * C) x T] tensor after attention.
"""
bs, width, length = qkv.shape
assert width % (3 * self.n_heads) == 0
ch = width // (3 * self.n_heads)
q, k, v = qkv.chunk(3, dim=1)
scale = 1 / math.sqrt(math.sqrt(ch))
weight = th.einsum(
"bct,bcs->bts",
(q * scale).view(bs * self.n_heads, ch, length),
(k * scale).view(bs * self.n_heads, ch, length),
) # More stable with f16 than dividing afterwards
weight = th.softmax(weight.float(), dim=-1).type(weight.dtype)
a = th.einsum("bts,bcs->bct", weight, v.reshape(bs * self.n_heads, ch, length))
return a.reshape(bs, -1, length)
@staticmethod
def count_flops(model, _x, y):
return count_flops_attn(model, _x, y)
class UNetModel(nn.Module):
"""
The full UNet model with attention and timestep embedding.
:param in_channels: channels in the input Tensor.
:param model_channels: base channel count for the model.
:param out_channels: channels in the output Tensor.
:param num_res_blocks: number of residual blocks per downsample.
:param attention_resolutions: a collection of downsample rates at which
attention will take place. May be a set, list, or tuple.
For example, if this contains 4, then at 4x downsampling, attention
will be used.
:param dropout: the dropout probability.
:param channel_mult: channel multiplier for each level of the UNet.
:param conv_resample: if True, use learned convolutions for upsampling and
downsampling.
:param dims: determines if the signal is 1D, 2D, or 3D.
:param num_classes: if specified (as an int), then this model will be
class-conditional with `num_classes` classes.
:param use_checkpoint: use gradient checkpointing to reduce memory usage.
:param num_heads: the number of attention heads in each attention layer.
:param num_heads_channels: if specified, ignore num_heads and instead use
a fixed channel width per attention head.
:param num_heads_upsample: works with num_heads to set a different number
of heads for upsampling. Deprecated.
:param use_scale_shift_norm: use a FiLM-like conditioning mechanism.
:param resblock_updown: use residual blocks for up/downsampling.
:param use_new_attention_order: use a different attention pattern for potentially
increased efficiency.
"""
def __init__(
self,
image_size,
in_channels,
model_channels,
out_channels,
num_res_blocks,
attention_resolutions,
dropout=0,
channel_mult=(1, 2, 4, 8),
conv_resample=True,
dims=2,
num_classes=None,
use_checkpoint=False,
use_fp16=False,
num_heads=-1,
num_head_channels=-1,
num_heads_upsample=-1,
use_scale_shift_norm=False,
resblock_updown=False,
use_new_attention_order=False,
use_max_self_attn = False,
use_max_spatial_transfomer = False,
use_spatial_transformer=False, # custom transformer support
transformer_depth=1, # custom transformer support
context_dim=None, # custom transformer support
n_embed=None, # custom support for prediction of discrete ids into codebook of first stage vq model
legacy=True,
condition_key="concat",
):
super().__init__()
if use_spatial_transformer:
assert context_dim is not None, 'Fool!! You forgot to include the dimension of your cross-attention conditioning...'
if context_dim is not None:
assert use_spatial_transformer, 'Fool!! You forgot to use the spatial transformer for your cross-attention conditioning...'
from omegaconf.listconfig import ListConfig
if type(context_dim) == ListConfig:
context_dim = list(context_dim)
if num_heads_upsample == -1:
num_heads_upsample = num_heads
if num_heads == -1:
assert num_head_channels != -1, 'Either num_heads or num_head_channels has to be set'
if num_head_channels == -1:
assert num_heads != -1, 'Either num_heads or num_head_channels has to be set'
self.image_size = image_size
self.in_channels = in_channels
self.model_channels = model_channels
self.out_channels = out_channels
self.num_res_blocks = num_res_blocks
self.attention_resolutions = attention_resolutions
self.dropout = dropout
self.channel_mult = channel_mult
self.conv_resample = conv_resample
self.num_classes = num_classes
self.use_checkpoint = use_checkpoint
self.dtype = th.float16 if use_fp16 else th.float32
self.num_heads = num_heads
self.num_head_channels = num_head_channels
self.num_heads_upsample = num_heads_upsample
self.predict_codebook_ids = n_embed is not None
self.condition_key = condition_key
time_embed_dim = model_channels * 4
self.time_embed = nn.Sequential(
linear(model_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
)
if self.num_classes is not None:
self.label_emb = nn.Embedding(num_classes, time_embed_dim)
self.input_blocks = nn.ModuleList(
[
TimestepEmbedSequential(
conv_nd(dims, in_channels, model_channels, 3, padding=1)
)
]
)
self._feature_size = model_channels
input_block_chans = [model_channels]
ch = model_channels
ds = 1
max_self_attn_ws = min(self.image_size // 4, 8)
for level, mult in enumerate(channel_mult):
for _ in range(num_res_blocks):
layers = [
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=mult * model_channels,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = mult * model_channels
if ds in attention_resolutions:
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
if legacy:
#num_heads = 1
dim_head = ch // num_heads if use_spatial_transformer else num_head_channels
layers.append(
AttentionBlock(
ch,
use_checkpoint=use_checkpoint,
num_heads=num_heads,
num_head_channels=dim_head,
use_new_attention_order=use_new_attention_order,
) if not use_max_self_attn and not use_spatial_transformer else MaxAttentionBlock(
ch, num_heads, dim_head, window_size=max_self_attn_ws
) if not use_spatial_transformer else SpatialTransformer(
ch, num_heads, dim_head, depth=transformer_depth, context_dim=context_dim
) if not use_max_spatial_transfomer else SpatialTransformerWithMax(
ch, num_heads, dim_head, context_dim=context_dim
)
)
self.input_blocks.append(TimestepEmbedSequential(*layers))
self._feature_size += ch
input_block_chans.append(ch)
if level != len(channel_mult) - 1:
out_ch = ch
self.input_blocks.append(
TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
down=True,
)
if resblock_updown
else Downsample(
ch, conv_resample, dims=dims, out_channels=out_ch
)
)
)
ch = out_ch
input_block_chans.append(ch)
ds *= 2
self._feature_size += ch
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
if legacy:
#num_heads = 1
dim_head = ch // num_heads if use_spatial_transformer else num_head_channels
self.middle_block = TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
AttentionBlock(
ch,
use_checkpoint=use_checkpoint,
num_heads=num_heads,
num_head_channels=dim_head,
use_new_attention_order=use_new_attention_order,
) if not use_max_self_attn and not use_spatial_transformer else MaxAttentionBlock(
ch, num_heads, dim_head, window_size=max_self_attn_ws
) if not use_spatial_transformer else SpatialTransformer(
ch, num_heads, dim_head, depth=transformer_depth, context_dim=context_dim
) if not use_max_spatial_transfomer else SpatialTransformerWithMax(
ch, num_heads, dim_head, context_dim=context_dim
),
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
)
self._feature_size += ch
self.output_blocks = nn.ModuleList([])
for level, mult in list(enumerate(channel_mult))[::-1]:
for i in range(num_res_blocks + 1):
ich = input_block_chans.pop()
layers = [
ResBlock(
ch + ich,
time_embed_dim,
dropout,
out_channels=model_channels * mult,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = model_channels * mult
if ds in attention_resolutions:
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
if legacy:
#num_heads = 1
dim_head = ch // num_heads if use_spatial_transformer else num_head_channels
layers.append(
AttentionBlock(
ch,
use_checkpoint=use_checkpoint,
num_heads=num_heads_upsample,
num_head_channels=dim_head,
use_new_attention_order=use_new_attention_order,
) if not use_max_self_attn and not use_spatial_transformer else MaxAttentionBlock(
ch, num_heads, dim_head, window_size=max_self_attn_ws
) if not use_spatial_transformer else SpatialTransformer(
ch, num_heads, dim_head, depth=transformer_depth, context_dim=context_dim
) if not use_max_spatial_transfomer else SpatialTransformerWithMax(
ch, num_heads, dim_head, context_dim=context_dim
)
)
if level and i == num_res_blocks:
out_ch = ch
layers.append(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
up=True,
)
if resblock_updown
else Upsample(ch, conv_resample, dims=dims, out_channels=out_ch)
)
ds //= 2
self.output_blocks.append(TimestepEmbedSequential(*layers))
self._feature_size += ch
self.out = nn.Sequential(
normalization(ch),
nn.SiLU(),
zero_module(conv_nd(dims, model_channels, out_channels, 3, padding=1)),
)
if self.predict_codebook_ids:
self.id_predictor = nn.Sequential(
normalization(ch),
conv_nd(dims, model_channels, n_embed, 1),
#nn.LogSoftmax(dim=1) # change to cross_entropy and produce non-normalized logits
)
def get_parameter_number(model):
total_num = sum(p.numel() for p in model.parameters())
trainable_num = sum(p.numel() for p in model.parameters() if p.requires_grad)
print("Total Number of parameter: %.2fM" % (total_num / 1e6))
print("Trainable Number of parameter: %.2fM" % (total_num / 1e6))
def convert_to_fp16(self):
"""
Convert the torso of the model to float16.
"""
self.input_blocks.apply(convert_module_to_f16)
self.middle_block.apply(convert_module_to_f16)
self.output_blocks.apply(convert_module_to_f16)
def convert_to_fp32(self):
"""
Convert the torso of the model to float32.
"""
self.input_blocks.apply(convert_module_to_f32)
self.middle_block.apply(convert_module_to_f32)
self.output_blocks.apply(convert_module_to_f32)
def forward(self, x, cond = None, timesteps=None, context=None, y=None,**kwargs):
"""
Apply the model to an input batch.
:param x: an [N x C x ...] Tensor of inputs.
:param cond: an N C ... Tensor, condition of another path
:param timesteps: a 1-D batch of timesteps.
:param context: conditioning plugged in via crossattn
:param y: an [N] Tensor of labels, if class-conditional.
:return: an [N x C x ...] Tensor of outputs.
"""
assert (y is not None) == (
self.num_classes is not None
), "must specify y if and only if the model is class-conditional"
hs = []
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False)
emb = self.time_embed(t_emb)
if self.num_classes is not None:
assert y.shape == (x.shape[0],)
emb = emb + self.label_emb(y)
if self.condition_key != 'nocond':
x = th.cat([x, context], dim=1)
#x = th.cat([x,cond],dim = 1) ## cat with previous path
h = x.type(self.dtype)
for module in self.input_blocks:
h = module(h, emb, context)
hs.append(h)
h = self.middle_block(h, emb, context)
for module in self.output_blocks:
hspop = hs.pop()
h = th.cat([h, hspop], dim=1)
h = module(h, emb, context)
h = h.type(x.dtype)
if self.predict_codebook_ids:
return self.id_predictor(h)
else:
return self.out(h)
class EncoderUNetModel(nn.Module):
"""
The half UNet model with attention and timestep embedding.
For usage, see UNet.
"""
def __init__(
self,
image_size,
in_channels,
model_channels,
out_channels,
num_res_blocks,
attention_resolutions,
dropout=0,
channel_mult=(1, 2, 4, 8),
conv_resample=True,
dims=2,
use_checkpoint=False,
use_fp16=False,
num_heads=1,
num_head_channels=-1,
num_heads_upsample=-1,
use_scale_shift_norm=False,
resblock_updown=False,
use_new_attention_order=False,
pool="adaptive",
*args,
**kwargs
):
super().__init__()
if num_heads_upsample == -1:
num_heads_upsample = num_heads
self.in_channels = in_channels
self.model_channels = model_channels
self.out_channels = out_channels
self.num_res_blocks = num_res_blocks
self.attention_resolutions = attention_resolutions
self.dropout = dropout
self.channel_mult = channel_mult
self.conv_resample = conv_resample
self.use_checkpoint = use_checkpoint
self.dtype = th.float16 if use_fp16 else th.float32
self.num_heads = num_heads
self.num_head_channels = num_head_channels
self.num_heads_upsample = num_heads_upsample
time_embed_dim = model_channels * 4
self.time_embed = nn.Sequential(
linear(model_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
)
self.input_blocks = nn.ModuleList(
[
TimestepEmbedSequential(
conv_nd(dims, in_channels, model_channels, 3, padding=1)
)
]
)
self._feature_size = model_channels
input_block_chans = [model_channels]
ch = model_channels
ds = 1
for level, mult in enumerate(channel_mult):
for _ in range(num_res_blocks):
layers = [
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=mult * model_channels,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = mult * model_channels
if ds in attention_resolutions:
layers.append(
AttentionBlock(
ch,
use_checkpoint=use_checkpoint,
num_heads=num_heads,
num_head_channels=num_head_channels,
use_new_attention_order=use_new_attention_order,
)
)
self.input_blocks.append(TimestepEmbedSequential(*layers))
self._feature_size += ch
input_block_chans.append(ch)
if level != len(channel_mult) - 1:
out_ch = ch
self.input_blocks.append(
TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
down=True,
)
if resblock_updown
else Downsample(
ch, conv_resample, dims=dims, out_channels=out_ch
)
)
)
ch = out_ch
input_block_chans.append(ch)
ds *= 2
self._feature_size += ch
self.middle_block = TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
AttentionBlock(
ch,
use_checkpoint=use_checkpoint,
num_heads=num_heads,
num_head_channels=num_head_channels,
use_new_attention_order=use_new_attention_order,
),
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
)
self._feature_size += ch
self.pool = pool
if pool == "adaptive":
self.out = nn.Sequential(
normalization(ch),
nn.SiLU(),
nn.AdaptiveAvgPool2d((1, 1)),
zero_module(conv_nd(dims, ch, out_channels, 1)),
nn.Flatten(),
)
elif pool == "attention":
assert num_head_channels != -1
self.out = nn.Sequential(
normalization(ch),
nn.SiLU(),
AttentionPool2d(
(image_size // ds), ch, num_head_channels, out_channels
),
)
elif pool == "spatial":
self.out = nn.Sequential(
nn.Linear(self._feature_size, 2048),
nn.ReLU(),
nn.Linear(2048, self.out_channels),
)
elif pool == "spatial_v2":
self.out = nn.Sequential(
nn.Linear(self._feature_size, 2048),
normalization(2048),
nn.SiLU(),
nn.Linear(2048, self.out_channels),
)
else:
raise NotImplementedError(f"Unexpected {pool} pooling")
def convert_to_fp16(self):
"""
Convert the torso of the model to float16.
"""
self.input_blocks.apply(convert_module_to_f16)
self.middle_block.apply(convert_module_to_f16)
def convert_to_fp32(self):
"""
Convert the torso of the model to float32.
"""
self.input_blocks.apply(convert_module_to_f32)
self.middle_block.apply(convert_module_to_f32)
def forward(self, x, timesteps):
"""
Apply the model to an input batch.
:param x: an [N x C x ...] Tensor of inputs.
:param timesteps: a 1-D batch of timesteps.
:return: an [N x K] Tensor of outputs.
"""
emb = self.time_embed(timestep_embedding(timesteps, self.model_channels))
results = []
h = x.type(self.dtype)
for module in self.input_blocks:
h = module(h, emb)
if self.pool.startswith("spatial"):
results.append(h.type(x.dtype).mean(dim=(2, 3)))
h = self.middle_block(h, emb)
if self.pool.startswith("spatial"):
results.append(h.type(x.dtype).mean(dim=(2, 3)))
h = th.cat(results, axis=-1)
return self.out(h)
else:
h = h.type(x.dtype)
return self.out(h)
@@ -0,0 +1,267 @@
# adopted from
# https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
# and
# https://github.com/lucidrains/denoising-diffusion-pytorch/blob/7706bdfc6f527f58d33f84b7b522e61e6e3164b3/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py
# and
# https://github.com/openai/guided-diffusion/blob/0ba878e517b276c45d1195eb29f6f5f72659a05b/guided_diffusion/nn.py
#
# thanks!
import os
import math
import torch
import torch.nn as nn
import numpy as np
from einops import repeat
from model.BrownianBridge.base.util import instantiate_from_config
def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3):
if schedule == "linear":
betas = (
torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=torch.float64) ** 2
)
elif schedule == "cosine":
timesteps = (
torch.arange(n_timestep + 1, dtype=torch.float64) / n_timestep + cosine_s
)
alphas = timesteps / (1 + cosine_s) * np.pi / 2
alphas = torch.cos(alphas).pow(2)
alphas = alphas / alphas[0]
betas = 1 - alphas[1:] / alphas[:-1]
betas = np.clip(betas, a_min=0, a_max=0.999)
elif schedule == "sqrt_linear":
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64)
elif schedule == "sqrt":
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64) ** 0.5
else:
raise ValueError(f"schedule '{schedule}' unknown.")
return betas.numpy()
def make_ddim_timesteps(ddim_discr_method, num_ddim_timesteps, num_ddpm_timesteps, verbose=True):
if ddim_discr_method == 'uniform':
c = num_ddpm_timesteps // num_ddim_timesteps
ddim_timesteps = np.asarray(list(range(0, num_ddpm_timesteps, c)))
elif ddim_discr_method == 'quad':
ddim_timesteps = ((np.linspace(0, np.sqrt(num_ddpm_timesteps * .8), num_ddim_timesteps)) ** 2).astype(int)
else:
raise NotImplementedError(f'There is no ddim discretization method called "{ddim_discr_method}"')
# assert ddim_timesteps.shape[0] == num_ddim_timesteps
# add one to get the final alpha values right (the ones from first scale to data during sampling)
steps_out = ddim_timesteps + 1
if verbose:
print(f'Selected timesteps for ddim sampler: {steps_out}')
return steps_out
def make_ddim_sampling_parameters(alphacums, ddim_timesteps, eta, verbose=True):
# select alphas for computing the variance schedule
alphas = alphacums[ddim_timesteps]
alphas_prev = np.asarray([alphacums[0]] + alphacums[ddim_timesteps[:-1]].tolist())
# according the the formula provided in https://arxiv.org/abs/2010.02502
sigmas = eta * np.sqrt((1 - alphas_prev) / (1 - alphas) * (1 - alphas / alphas_prev))
if verbose:
print(f'Selected alphas for ddim sampler: a_t: {alphas}; a_(t-1): {alphas_prev}')
print(f'For the chosen value of eta, which is {eta}, '
f'this results in the following sigma_t schedule for ddim sampler {sigmas}')
return sigmas, alphas, alphas_prev
def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
"""
Create a beta schedule that discretizes the given alpha_t_bar function,
which defines the cumulative product of (1-beta) over time from t = [0,1].
:param num_diffusion_timesteps: the number of betas to produce.
:param alpha_bar: a lambda that takes an argument t from 0 to 1 and
produces the cumulative product of (1-beta) up to that
part of the diffusion process.
:param max_beta: the maximum beta to use; use values lower than 1 to
prevent singularities.
"""
betas = []
for i in range(num_diffusion_timesteps):
t1 = i / num_diffusion_timesteps
t2 = (i + 1) / num_diffusion_timesteps
betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
return np.array(betas)
def extract_into_tensor(a, t, x_shape):
b, *_ = t.shape
out = a.gather(-1, t)
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
def checkpoint(func, inputs, params, flag):
"""
Evaluate a function without caching intermediate activations, allowing for
reduced memory at the expense of extra compute in the backward pass.
:param func: the function to evaluate.
:param inputs: the argument sequence to pass to `func`.
:param params: a sequence of parameters `func` depends on but does not
explicitly take as arguments.
:param flag: if False, disable gradient checkpointing.
"""
if flag:
args = tuple(inputs) + tuple(params)
return CheckpointFunction.apply(func, len(inputs), *args)
else:
return func(*inputs)
class CheckpointFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, run_function, length, *args):
ctx.run_function = run_function
ctx.input_tensors = list(args[:length])
ctx.input_params = list(args[length:])
with torch.no_grad():
output_tensors = ctx.run_function(*ctx.input_tensors)
return output_tensors
@staticmethod
def backward(ctx, *output_grads):
ctx.input_tensors = [x.detach().requires_grad_(True) for x in ctx.input_tensors]
with torch.enable_grad():
# Fixes a bug where the first op in run_function modifies the
# Tensor storage in place, which is not allowed for detach()'d
# Tensors.
shallow_copies = [x.view_as(x) for x in ctx.input_tensors]
output_tensors = ctx.run_function(*shallow_copies)
input_grads = torch.autograd.grad(
output_tensors,
ctx.input_tensors + ctx.input_params,
output_grads,
allow_unused=True,
)
del ctx.input_tensors
del ctx.input_params
del output_tensors
return (None, None) + input_grads
def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
"""
Create sinusoidal timestep embeddings.
:param timesteps: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:param dim: the dimension of the output.
:param max_period: controls the minimum frequency of the embeddings.
:return: an [N x dim] Tensor of positional embeddings.
"""
if not repeat_only:
half = dim // 2
freqs = torch.exp(
-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
).to(device=timesteps.device)
args = timesteps[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
else:
embedding = repeat(timesteps, 'b -> b d', d=dim)
return embedding
def zero_module(module):
"""
Zero out the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().zero_()
return module
def scale_module(module, scale):
"""
Scale the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().mul_(scale)
return module
def mean_flat(tensor):
"""
Take the mean over all non-batch dimensions.
"""
return tensor.mean(dim=list(range(1, len(tensor.shape))))
def normalization(channels):
"""
Make a standard normalization layer.
:param channels: number of input channels.
:return: an nn.Module for normalization.
"""
return GroupNorm32(32, channels)
# PyTorch 1.7 has SiLU, but we support PyTorch 1.5.
class SiLU(nn.Module):
def forward(self, x):
return x * torch.sigmoid(x)
class GroupNorm32(nn.GroupNorm):
def forward(self, x):
return super().forward(x.float()).type(x.dtype)
def conv_nd(dims, *args, **kwargs):
"""
Create a 1D, 2D, or 3D convolution module.
"""
if dims == 1:
return nn.Conv1d(*args, **kwargs)
elif dims == 2:
return nn.Conv2d(*args, **kwargs)
elif dims == 3:
return nn.Conv3d(*args, **kwargs)
raise ValueError(f"unsupported dimensions: {dims}")
def linear(*args, **kwargs):
"""
Create a linear module.
"""
return nn.Linear(*args, **kwargs)
def avg_pool_nd(dims, *args, **kwargs):
"""
Create a 1D, 2D, or 3D average pooling module.
"""
if dims == 1:
return nn.AvgPool1d(*args, **kwargs)
elif dims == 2:
return nn.AvgPool2d(*args, **kwargs)
elif dims == 3:
return nn.AvgPool3d(*args, **kwargs)
raise ValueError(f"unsupported dimensions: {dims}")
class HybridConditioner(nn.Module):
def __init__(self, c_concat_config, c_crossattn_config):
super().__init__()
self.concat_conditioner = instantiate_from_config(c_concat_config)
self.crossattn_conditioner = instantiate_from_config(c_crossattn_config)
def forward(self, c_concat, c_crossattn):
c_concat = self.concat_conditioner(c_concat)
c_crossattn = self.crossattn_conditioner(c_crossattn)
return {'c_concat': [c_concat], 'c_crossattn': [c_crossattn]}
def noise_like(shape, device, repeat=False):
repeat_noise = lambda: torch.randn((1, *shape[1:]), device=device).repeat(shape[0], *((1,) * (len(shape) - 1)))
noise = lambda: torch.randn(shape, device=device)
return repeat_noise() if repeat else noise()
@@ -0,0 +1,76 @@
import torch
from torch import nn
class LitEma(nn.Module):
def __init__(self, model, decay=0.9999, use_num_upates=True):
super().__init__()
if decay < 0.0 or decay > 1.0:
raise ValueError('Decay must be between 0 and 1')
self.m_name2s_name = {}
self.register_buffer('decay', torch.tensor(decay, dtype=torch.float32))
self.register_buffer('num_updates', torch.tensor(0,dtype=torch.int) if use_num_upates
else torch.tensor(-1,dtype=torch.int))
for name, p in model.named_parameters():
if p.requires_grad:
#remove as '.'-character is not allowed in buffers
s_name = name.replace('.','')
self.m_name2s_name.update({name:s_name})
self.register_buffer(s_name,p.clone().detach().data)
self.collected_params = []
def forward(self,model):
decay = self.decay
if self.num_updates >= 0:
self.num_updates += 1
decay = min(self.decay,(1 + self.num_updates) / (10 + self.num_updates))
one_minus_decay = 1.0 - decay
with torch.no_grad():
m_param = dict(model.named_parameters())
shadow_params = dict(self.named_buffers())
for key in m_param:
if m_param[key].requires_grad:
sname = self.m_name2s_name[key]
shadow_params[sname] = shadow_params[sname].type_as(m_param[key])
shadow_params[sname].sub_(one_minus_decay * (shadow_params[sname] - m_param[key]))
else:
assert not key in self.m_name2s_name
def copy_to(self, model):
m_param = dict(model.named_parameters())
shadow_params = dict(self.named_buffers())
for key in m_param:
if m_param[key].requires_grad:
m_param[key].data.copy_(shadow_params[self.m_name2s_name[key]].data)
else:
assert not key in self.m_name2s_name
def store(self, parameters):
"""
Save the current parameters for restoring later.
Args:
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
temporarily stored.
"""
self.collected_params = [param.clone() for param in parameters]
def restore(self, parameters):
"""
Restore the parameters stored with the `store` method.
Useful to validate the model with EMA parameters without affecting the
original optimization process. Store the parameters before the
`copy_to` method. After validation (or model saving), use this to
restore the former parameters.
Args:
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
updated with the stored parameters.
"""
for c_param, param in zip(self.collected_params, parameters):
param.data.copy_(c_param.data)
@@ -0,0 +1,134 @@
import pdb
import torch
import torch.nn as nn
from functools import partial
from einops import rearrange, repeat
from model.BrownianBridge.base.modules.x_transformer import Encoder, TransformerWrapper # TODO: can we directly rely on lucidrains code and simply add this as a reuirement? --> test
class AbstractEncoder(nn.Module):
def __init__(self):
super().__init__()
def encode(self, *args, **kwargs):
raise NotImplementedError
class ClassEmbedder(nn.Module):
def __init__(self, embed_dim, n_classes=1000, key='class'):
super().__init__()
self.key = key
self.embedding = nn.Embedding(n_classes, embed_dim)
def forward(self, batch, key=None):
if key is None:
key = self.key
# this is for use in crossattn
c = batch[key][:, None]
c = self.embedding(c)
return c
class TransformerEmbedder(AbstractEncoder):
"""Some transformer encoder layers"""
def __init__(self, n_embed, n_layer, vocab_size, max_seq_len=77, device="cuda"):
super().__init__()
self.device = device
self.transformer = TransformerWrapper(num_tokens=vocab_size, max_seq_len=max_seq_len,
attn_layers=Encoder(dim=n_embed, depth=n_layer))
def forward(self, tokens):
tokens = tokens.to(self.device) # meh
z = self.transformer(tokens, return_embeddings=True)
return z
def encode(self, x):
return self(x)
class BERTTokenizer(AbstractEncoder):
""" Uses a pretrained BERT tokenizer by huggingface. Vocab size: 30522 (?)"""
def __init__(self, device="cuda", vq_interface=True, max_length=77):
super().__init__()
from transformers import BertTokenizerFast # TODO: add to reuquirements
self.tokenizer = BertTokenizerFast.from_pretrained("bert-base-uncased")
self.device = device
self.vq_interface = vq_interface
self.max_length = max_length
def forward(self, text):
batch_encoding = self.tokenizer(text, truncation=True, max_length=self.max_length, return_length=True,
return_overflowing_tokens=False, padding="max_length", return_tensors="pt")
tokens = batch_encoding["input_ids"].to(self.device)
return tokens
@torch.no_grad()
def encode(self, text):
tokens = self(text)
if not self.vq_interface:
return tokens
return None, None, [None, None, tokens]
def decode(self, text):
return text
class BERTEmbedder(AbstractEncoder):
"""Uses the BERT tokenizr model and add some transformer encoder layers"""
def __init__(self, n_embed, n_layer, vocab_size=30522, max_seq_len=77,
device="cuda",use_tokenizer=True, embedding_dropout=0.0):
super().__init__()
self.use_tknz_fn = use_tokenizer
if self.use_tknz_fn:
self.tknz_fn = BERTTokenizer(vq_interface=False, max_length=max_seq_len)
self.device = device
self.transformer = TransformerWrapper(num_tokens=vocab_size, max_seq_len=max_seq_len,
attn_layers=Encoder(dim=n_embed, depth=n_layer),
emb_dropout=embedding_dropout)
def forward(self, text):
if self.use_tknz_fn:
tokens = self.tknz_fn(text)#.to(self.device)
else:
tokens = text
z = self.transformer(tokens, return_embeddings=True)
return z
def encode(self, text):
# output of length 77
return self(text)
class SpatialRescaler(nn.Module):
def __init__(self,
n_stages=1,
method='bilinear',
multiplier=0.5,
in_channels=3,
out_channels=None,
bias=False):
super().__init__()
self.n_stages = n_stages
assert self.n_stages >= 0
assert method in ['nearest','linear','bilinear','trilinear','bicubic','area']
self.multiplier = multiplier
self.interpolator = partial(torch.nn.functional.interpolate, mode=method)
self.remap_output = out_channels is not None
if self.remap_output:
print(f'Spatial Rescaler mapping from {in_channels} to {out_channels} channels after resizing.')
self.channel_mapper = nn.Conv2d(in_channels,out_channels,1,bias=bias)
def forward(self,x):
for stage in range(self.n_stages):
x = self.interpolator(x, scale_factor=self.multiplier)
if self.remap_output:
x = self.channel_mapper(x)
return x
def encode(self, x):
return self(x)
@@ -0,0 +1,352 @@
import torch
from torch import nn, einsum
import torch.nn.functional
from einops import rearrange, repeat
from einops.layers.torch import Rearrange, Reduce
from inspect import isfunction
# Code adapted from https://github.com/lucidrains/vit-pytorch/blob/main/vit_pytorch/max_vit.py
def exists(val):
return val is not None
def default(val, d):
if exists(val):
return val
return d() if isfunction(d) else d
class PreNormResidual(nn.Module):
def __init__(self, dim, fn):
super().__init__()
self.norm = nn.LayerNorm(dim)
self.fn = fn
def forward(self, x, c=None):
if exists(c):
return self.fn(self.norm(x), self.norm(c)) + x
return self.fn(self.norm(x)) + x
class SqueezeExcitation(nn.Module):
def __init__(self, dim, shrinkage_rate = 0.25):
super().__init__()
hidden_dim = int(dim * shrinkage_rate)
self.gate = nn.Sequential(
Reduce('b c h w -> b c', 'mean'),
nn.Linear(dim, hidden_dim, bias = False),
nn.SiLU(),
nn.Linear(hidden_dim, dim, bias = False),
nn.Sigmoid(),
Rearrange('b c -> b c 1 1')
)
def forward(self, x):
return x * self.gate(x)
class FeedForward(nn.Module):
def __init__(self, dim, mult = 4, dropout = 0.):
super().__init__()
inner_dim = int(dim * mult)
self.net = nn.Sequential(
nn.Linear(dim, inner_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(inner_dim, dim),
nn.Dropout(dropout)
)
def forward(self, x):
return self.net(x)
class Attention(nn.Module):
def __init__(
self,
dim,
dim_head = 32,
dropout = 0.,
window_size = 7
):
super().__init__()
assert (dim % dim_head) == 0, 'dimension should be divisible by dimension per head'
self.heads = dim // dim_head
self.scale = dim_head ** -0.5
self.to_q = nn.Linear(dim, dim, bias = False)
self.to_k = nn.Linear(dim, dim, bias = False)
self.to_v = nn.Linear(dim, dim, bias = False)
self.attend = nn.Sequential(
nn.Softmax(dim = -1),
nn.Dropout(dropout)
)
self.to_out = nn.Sequential(
nn.Linear(dim, dim, bias = False),
nn.Dropout(dropout)
)
# relative positional bias
self.rel_pos_bias = nn.Embedding((2 * window_size - 1) ** 2, self.heads)
pos = torch.arange(window_size)
grid = torch.stack(torch.meshgrid(pos, pos, indexing = 'ij'))
grid = rearrange(grid, 'c i j -> (i j) c')
rel_pos = rearrange(grid, 'i ... -> i 1 ...') - rearrange(grid, 'j ... -> 1 j ...')
rel_pos += window_size - 1
rel_pos_indices = (rel_pos * torch.tensor([2 * window_size - 1, 1])).sum(dim = -1)
self.register_buffer('rel_pos_indices', rel_pos_indices, persistent = False)
def forward(self, x, c=None):
c = default(c, x)
batch, height, width, window_height, window_width, _, device, h = *x.shape, x.device, self.heads
# flatten
x = rearrange(x, 'b x y w1 w2 d -> (b x y) (w1 w2) d')
c = rearrange(c, 'b x y w1 w2 d -> (b x y) (w1 w2) d')
# project for queries, keys, values
q = self.to_q(x)
k = self.to_k(c)
v = self.to_v(c)
# split heads
q, k, v = map(lambda t: rearrange(t, 'b n (h d ) -> b h n d', h = h), (q, k, v))
# scale
q = q * self.scale
# sim
sim = einsum('b h i d, b h j d -> b h i j', q, k)
# add positional bias
bias = self.rel_pos_bias(self.rel_pos_indices)
sim = sim + rearrange(bias, 'i j h -> h i j')
# attention
attn = self.attend(sim)
# aggregate
out = einsum('b h i j, b h j d -> b h i d', attn, v)
# merge heads
out = rearrange(out, 'b h (w1 w2) d -> b w1 w2 (h d)', w1 = window_height, w2 = window_width)
# combine heads out
out = self.to_out(out)
return rearrange(out, '(b x y) ... -> b x y ...', x = height, y = width)
class Dropsample(nn.Module):
def __init__(self, prob = 0):
super().__init__()
self.prob = prob
def forward(self, x):
device = x.device
if self.prob == 0. or (not self.training):
return x
keep_mask = torch.FloatTensor((x.shape[0], 1, 1, 1), device = device).uniform_() > self.prob
return x * keep_mask / (1 - self.prob)
class MBConvResidual(nn.Module):
def __init__(self, fn, dropout = 0.):
super().__init__()
self.fn = fn
self.dropsample = Dropsample(dropout)
def forward(self, x):
out = self.fn(x)
out = self.dropsample(out)
return out + x
def MBConv(
dim_in,
dim_out,
*,
downsample,
expansion_rate = 4,
shrinkage_rate = 0.25,
dropout = 0.
):
hidden_dim = int(expansion_rate * dim_out)
stride = 2 if downsample else 1
net = nn.Sequential(
nn.Conv2d(dim_in, hidden_dim, 1),
nn.BatchNorm2d(hidden_dim),
nn.GELU(),
nn.Conv2d(hidden_dim, hidden_dim, 3, stride = stride, padding = 1, groups = hidden_dim),
nn.BatchNorm2d(hidden_dim),
nn.GELU(),
SqueezeExcitation(hidden_dim, shrinkage_rate = shrinkage_rate),
nn.Conv2d(hidden_dim, dim_out, 1),
nn.BatchNorm2d(dim_out)
)
if dim_in == dim_out and not downsample:
net = MBConvResidual(net, dropout = dropout)
return net
class MaxAttentionBlock(nn.Module):
def __init__(self, in_channels, heads=8, dim_head=64, dropout=0., window_size=8):
super().__init__()
w = window_size
layer_dim = dim_head * heads
self.rearrange_block_in = Rearrange('b d (x w1) (y w2) -> b x y w1 w2 d', w1 = w, w2 = w) # block-like attention
self.attn_block = PreNormResidual(layer_dim, Attention(dim = layer_dim, dim_head = dim_head, dropout = dropout, window_size = w))
self.ff_block = PreNormResidual(layer_dim, FeedForward(dim = layer_dim, dropout = dropout))
self.rearrange_block_out = Rearrange('b x y w1 w2 d -> b d (x w1) (y w2)')
self.rearrange_grid_in = Rearrange('b d (w1 x) (w2 y) -> b x y w1 w2 d', w1 = w, w2 = w) # grid-like attention
self.attn_grid = PreNormResidual(layer_dim, Attention(dim = layer_dim, dim_head = dim_head, dropout = dropout, window_size = w))
self.ff_grid = PreNormResidual(layer_dim, FeedForward(dim = layer_dim, dropout = dropout))
self.rearrange_grid_out = Rearrange('b x y w1 w2 d -> b d (w1 x) (w2 y)')
def forward(self, x):
# block attention
x = self.rearrange_block_in(x)
x = self.attn_block(x)
x = self.ff_block(x)
x = self.rearrange_block_out(x)
# grid attention
x = self.rearrange_grid_in(x)
x = self.attn_grid(x)
x = self.ff_grid(x)
x = self.rearrange_grid_out(x)
## output stage
return x
class SpatialCrossAttentionWithMax(nn.Module):
def __init__(self, in_channels, heads=8, dim_head=64, ctx_dim=None, dropout=0., window_size=8):
super().__init__()
w = window_size
layer_dim = dim_head * heads
if ctx_dim == None:
self.proj_in = MBConv(layer_dim*2, layer_dim, downsample=False)
else:
self.proj_in = MBConv(ctx_dim, layer_dim, downsample=False)
self.rearrange_block_in = Rearrange('b d (x w1) (y w2) -> b x y w1 w2 d', w1 = w, w2 = w) # block-like attention
self.attn_block = PreNormResidual(layer_dim, Attention(dim = layer_dim, dim_head = dim_head, dropout = dropout, window_size = w))
self.ff_block = PreNormResidual(layer_dim, FeedForward(dim = layer_dim, dropout = dropout))
self.rearrange_block_out = Rearrange('b x y w1 w2 d -> b d (x w1) (y w2)')
self.rearrange_grid_in = Rearrange('b d (w1 x) (w2 y) -> b x y w1 w2 d', w1 = w, w2 = w) # grid-like attention
self.attn_grid = PreNormResidual(layer_dim, Attention(dim = layer_dim, dim_head = dim_head, dropout = dropout, window_size = w))
self.ff_grid = PreNormResidual(layer_dim, FeedForward(dim = layer_dim, dropout = dropout))
self.rearrange_grid_out = Rearrange('b x y w1 w2 d -> b d (w1 x) (w2 y)')
self.out_conv = nn.Sequential(
SqueezeExcitation(dim=layer_dim*2),
nn.Conv2d(layer_dim*2, layer_dim, kernel_size=3, padding=1)
)
def forward(self, x, context=None):
context = default(context, x)
# MBConv
c = self.proj_in(context)
# block attention
x = self.rearrange_block_in(x)
c = self.rearrange_block_in(c)
x = self.attn_block(x, c)
x = self.ff_block(x)
x = self.rearrange_block_out(x)
c = self.rearrange_block_out(c)
# grid attention
x = self.rearrange_grid_in(x)
c = self.rearrange_grid_in(c)
x = self.attn_grid(x, c)
x = self.ff_grid(x)
x = self.rearrange_grid_out(x)
return x
class SpatialTransformerWithMax(nn.Module):
"""
Transformer block for image-like data.
First, project the input (aka embedding) to inner_dim (d) using conv1x1
Then reshape to b, t, d.
Then apply standard transformer action (BasicTransformerBlock).
Finally, reshape to image and pass to output conv1x1 layer, to restore the channel size of input.
The dims of the input and output of the block are the same (arg in_channels).
"""
def __init__(self, in_channels, n_heads, d_head, dropout=0., context_dim=None, w=2):
super().__init__()
self.in_channels = in_channels
self.context_dim = context_dim
inner_dim = n_heads * d_head
self.proj_in = MBConv(context_dim, inner_dim, downsample=False)
self.rearrange_block_in = Rearrange('b d (x w1) (y w2) -> b x y w1 w2 d', w1 = w, w2 = w) # block-like attention
self.attn_block = PreNormResidual(inner_dim, Attention(dim = inner_dim, dim_head = d_head, dropout = dropout, window_size = w))
self.ff_block = PreNormResidual(inner_dim, FeedForward(dim = inner_dim, dropout = dropout))
self.rearrange_block_out = Rearrange('b x y w1 w2 d -> b d (x w1) (y w2)')
self.rearrange_grid_in = Rearrange('b d (w1 x) (w2 y) -> b x y w1 w2 d', w1 = w, w2 = w) # grid-like attention
self.attn_grid = PreNormResidual(inner_dim, Attention(dim = inner_dim, dim_head = d_head, dropout = dropout, window_size = w))
self.ff_grid = PreNormResidual(inner_dim, FeedForward(dim = inner_dim, dropout = dropout))
self.rearrange_grid_out = Rearrange('b x y w1 w2 d -> b d (w1 x) (w2 y)')
def forward(self, x, context=None):
context = default(context, x)
# down sample context if necessary
# this is due to the implementation of max crossattn here
if context.shape[2] != x.shape[2]:
stride = context.shape[2] // x.shape[2]
context = torch.nn.functional.avg_pool2d(context, kernel_size=stride, stride=stride)
# MBConv
c = self.proj_in(context)
# block attention
x = self.rearrange_block_in(x)
c = self.rearrange_block_in(c)
x = self.attn_block(x, c)
x = self.ff_block(x)
x = self.rearrange_block_out(x)
c = self.rearrange_block_out(c)
# grid attention
x = self.rearrange_grid_in(x)
c = self.rearrange_grid_in(c)
x = self.attn_grid(x, c)
x = self.ff_grid(x)
x = self.rearrange_grid_out(x)
return x
@@ -0,0 +1,641 @@
"""shout-out to https://github.com/lucidrains/x-transformers/tree/main/x_transformers"""
import torch
from torch import nn, einsum
import torch.nn.functional as F
from functools import partial
from inspect import isfunction
from collections import namedtuple
from einops import rearrange, repeat, reduce
# constants
DEFAULT_DIM_HEAD = 64
Intermediates = namedtuple('Intermediates', [
'pre_softmax_attn',
'post_softmax_attn'
])
LayerIntermediates = namedtuple('Intermediates', [
'hiddens',
'attn_intermediates'
])
class AbsolutePositionalEmbedding(nn.Module):
def __init__(self, dim, max_seq_len):
super().__init__()
self.emb = nn.Embedding(max_seq_len, dim)
self.init_()
def init_(self):
nn.init.normal_(self.emb.weight, std=0.02)
def forward(self, x):
n = torch.arange(x.shape[1], device=x.device)
return self.emb(n)[None, :, :]
class FixedPositionalEmbedding(nn.Module):
def __init__(self, dim):
super().__init__()
inv_freq = 1. / (10000 ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer('inv_freq', inv_freq)
def forward(self, x, seq_dim=1, offset=0):
t = torch.arange(x.shape[seq_dim], device=x.device).type_as(self.inv_freq) + offset
sinusoid_inp = torch.einsum('i , j -> i j', t, self.inv_freq)
emb = torch.cat((sinusoid_inp.sin(), sinusoid_inp.cos()), dim=-1)
return emb[None, :, :]
# helpers
def exists(val):
return val is not None
def default(val, d):
if exists(val):
return val
return d() if isfunction(d) else d
def always(val):
def inner(*args, **kwargs):
return val
return inner
def not_equals(val):
def inner(x):
return x != val
return inner
def equals(val):
def inner(x):
return x == val
return inner
def max_neg_value(tensor):
return -torch.finfo(tensor.dtype).max
# keyword argument helpers
def pick_and_pop(keys, d):
values = list(map(lambda key: d.pop(key), keys))
return dict(zip(keys, values))
def group_dict_by_key(cond, d):
return_val = [dict(), dict()]
for key in d.keys():
match = bool(cond(key))
ind = int(not match)
return_val[ind][key] = d[key]
return (*return_val,)
def string_begins_with(prefix, str):
return str.startswith(prefix)
def group_by_key_prefix(prefix, d):
return group_dict_by_key(partial(string_begins_with, prefix), d)
def groupby_prefix_and_trim(prefix, d):
kwargs_with_prefix, kwargs = group_dict_by_key(partial(string_begins_with, prefix), d)
kwargs_without_prefix = dict(map(lambda x: (x[0][len(prefix):], x[1]), tuple(kwargs_with_prefix.items())))
return kwargs_without_prefix, kwargs
# classes
class Scale(nn.Module):
def __init__(self, value, fn):
super().__init__()
self.value = value
self.fn = fn
def forward(self, x, **kwargs):
x, *rest = self.fn(x, **kwargs)
return (x * self.value, *rest)
class Rezero(nn.Module):
def __init__(self, fn):
super().__init__()
self.fn = fn
self.g = nn.Parameter(torch.zeros(1))
def forward(self, x, **kwargs):
x, *rest = self.fn(x, **kwargs)
return (x * self.g, *rest)
class ScaleNorm(nn.Module):
def __init__(self, dim, eps=1e-5):
super().__init__()
self.scale = dim ** -0.5
self.eps = eps
self.g = nn.Parameter(torch.ones(1))
def forward(self, x):
norm = torch.norm(x, dim=-1, keepdim=True) * self.scale
return x / norm.clamp(min=self.eps) * self.g
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-8):
super().__init__()
self.scale = dim ** -0.5
self.eps = eps
self.g = nn.Parameter(torch.ones(dim))
def forward(self, x):
norm = torch.norm(x, dim=-1, keepdim=True) * self.scale
return x / norm.clamp(min=self.eps) * self.g
class Residual(nn.Module):
def forward(self, x, residual):
return x + residual
class GRUGating(nn.Module):
def __init__(self, dim):
super().__init__()
self.gru = nn.GRUCell(dim, dim)
def forward(self, x, residual):
gated_output = self.gru(
rearrange(x, 'b n d -> (b n) d'),
rearrange(residual, 'b n d -> (b n) d')
)
return gated_output.reshape_as(x)
# feedforward
class GEGLU(nn.Module):
def __init__(self, dim_in, dim_out):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out * 2)
def forward(self, x):
x, gate = self.proj(x).chunk(2, dim=-1)
return x * F.gelu(gate)
class FeedForward(nn.Module):
def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0.):
super().__init__()
inner_dim = int(dim * mult)
dim_out = default(dim_out, dim)
project_in = nn.Sequential(
nn.Linear(dim, inner_dim),
nn.GELU()
) if not glu else GEGLU(dim, inner_dim)
self.net = nn.Sequential(
project_in,
nn.Dropout(dropout),
nn.Linear(inner_dim, dim_out)
)
def forward(self, x):
return self.net(x)
# attention.
class Attention(nn.Module):
def __init__(
self,
dim,
dim_head=DEFAULT_DIM_HEAD,
heads=8,
causal=False,
mask=None,
talking_heads=False,
sparse_topk=None,
use_entmax15=False,
num_mem_kv=0,
dropout=0.,
on_attn=False
):
super().__init__()
if use_entmax15:
raise NotImplementedError("Check out entmax activation instead of softmax activation!")
self.scale = dim_head ** -0.5
self.heads = heads
self.causal = causal
self.mask = mask
inner_dim = dim_head * heads
self.to_q = nn.Linear(dim, inner_dim, bias=False)
self.to_k = nn.Linear(dim, inner_dim, bias=False)
self.to_v = nn.Linear(dim, inner_dim, bias=False)
self.dropout = nn.Dropout(dropout)
# talking heads
self.talking_heads = talking_heads
if talking_heads:
self.pre_softmax_proj = nn.Parameter(torch.randn(heads, heads))
self.post_softmax_proj = nn.Parameter(torch.randn(heads, heads))
# explicit topk sparse attention
self.sparse_topk = sparse_topk
# entmax
#self.attn_fn = entmax15 if use_entmax15 else F.softmax
self.attn_fn = F.softmax
# add memory key / values
self.num_mem_kv = num_mem_kv
if num_mem_kv > 0:
self.mem_k = nn.Parameter(torch.randn(heads, num_mem_kv, dim_head))
self.mem_v = nn.Parameter(torch.randn(heads, num_mem_kv, dim_head))
# attention on attention
self.attn_on_attn = on_attn
self.to_out = nn.Sequential(nn.Linear(inner_dim, dim * 2), nn.GLU()) if on_attn else nn.Linear(inner_dim, dim)
def forward(
self,
x,
context=None,
mask=None,
context_mask=None,
rel_pos=None,
sinusoidal_emb=None,
prev_attn=None,
mem=None
):
b, n, _, h, talking_heads, device = *x.shape, self.heads, self.talking_heads, x.device
kv_input = default(context, x)
q_input = x
k_input = kv_input
v_input = kv_input
if exists(mem):
k_input = torch.cat((mem, k_input), dim=-2)
v_input = torch.cat((mem, v_input), dim=-2)
if exists(sinusoidal_emb):
# in shortformer, the query would start at a position offset depending on the past cached memory
offset = k_input.shape[-2] - q_input.shape[-2]
q_input = q_input + sinusoidal_emb(q_input, offset=offset)
k_input = k_input + sinusoidal_emb(k_input)
q = self.to_q(q_input)
k = self.to_k(k_input)
v = self.to_v(v_input)
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h=h), (q, k, v))
input_mask = None
if any(map(exists, (mask, context_mask))):
q_mask = default(mask, lambda: torch.ones((b, n), device=device).bool())
k_mask = q_mask if not exists(context) else context_mask
k_mask = default(k_mask, lambda: torch.ones((b, k.shape[-2]), device=device).bool())
q_mask = rearrange(q_mask, 'b i -> b () i ()')
k_mask = rearrange(k_mask, 'b j -> b () () j')
input_mask = q_mask * k_mask
if self.num_mem_kv > 0:
mem_k, mem_v = map(lambda t: repeat(t, 'h n d -> b h n d', b=b), (self.mem_k, self.mem_v))
k = torch.cat((mem_k, k), dim=-2)
v = torch.cat((mem_v, v), dim=-2)
if exists(input_mask):
input_mask = F.pad(input_mask, (self.num_mem_kv, 0), value=True)
dots = einsum('b h i d, b h j d -> b h i j', q, k) * self.scale
mask_value = max_neg_value(dots)
if exists(prev_attn):
dots = dots + prev_attn
pre_softmax_attn = dots
if talking_heads:
dots = einsum('b h i j, h k -> b k i j', dots, self.pre_softmax_proj).contiguous()
if exists(rel_pos):
dots = rel_pos(dots)
if exists(input_mask):
dots.masked_fill_(~input_mask, mask_value)
del input_mask
if self.causal:
i, j = dots.shape[-2:]
r = torch.arange(i, device=device)
mask = rearrange(r, 'i -> () () i ()') < rearrange(r, 'j -> () () () j')
mask = F.pad(mask, (j - i, 0), value=False)
dots.masked_fill_(mask, mask_value)
del mask
if exists(self.sparse_topk) and self.sparse_topk < dots.shape[-1]:
top, _ = dots.topk(self.sparse_topk, dim=-1)
vk = top[..., -1].unsqueeze(-1).expand_as(dots)
mask = dots < vk
dots.masked_fill_(mask, mask_value)
del mask
attn = self.attn_fn(dots, dim=-1)
post_softmax_attn = attn
attn = self.dropout(attn)
if talking_heads:
attn = einsum('b h i j, h k -> b k i j', attn, self.post_softmax_proj).contiguous()
out = einsum('b h i j, b h j d -> b h i d', attn, v)
out = rearrange(out, 'b h n d -> b n (h d)')
intermediates = Intermediates(
pre_softmax_attn=pre_softmax_attn,
post_softmax_attn=post_softmax_attn
)
return self.to_out(out), intermediates
class AttentionLayers(nn.Module):
def __init__(
self,
dim,
depth,
heads=8,
causal=False,
cross_attend=False,
only_cross=False,
use_scalenorm=False,
use_rmsnorm=False,
use_rezero=False,
rel_pos_num_buckets=32,
rel_pos_max_distance=128,
position_infused_attn=False,
custom_layers=None,
sandwich_coef=None,
par_ratio=None,
residual_attn=False,
cross_residual_attn=False,
macaron=False,
pre_norm=True,
gate_residual=False,
**kwargs
):
super().__init__()
ff_kwargs, kwargs = groupby_prefix_and_trim('ff_', kwargs)
attn_kwargs, _ = groupby_prefix_and_trim('attn_', kwargs)
dim_head = attn_kwargs.get('dim_head', DEFAULT_DIM_HEAD)
self.dim = dim
self.depth = depth
self.layers = nn.ModuleList([])
self.has_pos_emb = position_infused_attn
self.pia_pos_emb = FixedPositionalEmbedding(dim) if position_infused_attn else None
self.rotary_pos_emb = always(None)
assert rel_pos_num_buckets <= rel_pos_max_distance, 'number of relative position buckets must be less than the relative position max distance'
self.rel_pos = None
self.pre_norm = pre_norm
self.residual_attn = residual_attn
self.cross_residual_attn = cross_residual_attn
norm_class = ScaleNorm if use_scalenorm else nn.LayerNorm
norm_class = RMSNorm if use_rmsnorm else norm_class
norm_fn = partial(norm_class, dim)
norm_fn = nn.Identity if use_rezero else norm_fn
branch_fn = Rezero if use_rezero else None
if cross_attend and not only_cross:
default_block = ('a', 'c', 'f')
elif cross_attend and only_cross:
default_block = ('c', 'f')
else:
default_block = ('a', 'f')
if macaron:
default_block = ('f',) + default_block
if exists(custom_layers):
layer_types = custom_layers
elif exists(par_ratio):
par_depth = depth * len(default_block)
assert 1 < par_ratio <= par_depth, 'par ratio out of range'
default_block = tuple(filter(not_equals('f'), default_block))
par_attn = par_depth // par_ratio
depth_cut = par_depth * 2 // 3 # 2 / 3 attention layer cutoff suggested by PAR paper
par_width = (depth_cut + depth_cut // par_attn) // par_attn
assert len(default_block) <= par_width, 'default block is too large for par_ratio'
par_block = default_block + ('f',) * (par_width - len(default_block))
par_head = par_block * par_attn
layer_types = par_head + ('f',) * (par_depth - len(par_head))
elif exists(sandwich_coef):
assert sandwich_coef > 0 and sandwich_coef <= depth, 'sandwich coefficient should be less than the depth'
layer_types = ('a',) * sandwich_coef + default_block * (depth - sandwich_coef) + ('f',) * sandwich_coef
else:
layer_types = default_block * depth
self.layer_types = layer_types
self.num_attn_layers = len(list(filter(equals('a'), layer_types)))
for layer_type in self.layer_types:
if layer_type == 'a':
layer = Attention(dim, heads=heads, causal=causal, **attn_kwargs)
elif layer_type == 'c':
layer = Attention(dim, heads=heads, **attn_kwargs)
elif layer_type == 'f':
layer = FeedForward(dim, **ff_kwargs)
layer = layer if not macaron else Scale(0.5, layer)
else:
raise Exception(f'invalid layer type {layer_type}')
if isinstance(layer, Attention) and exists(branch_fn):
layer = branch_fn(layer)
if gate_residual:
residual_fn = GRUGating(dim)
else:
residual_fn = Residual()
self.layers.append(nn.ModuleList([
norm_fn(),
layer,
residual_fn
]))
def forward(
self,
x,
context=None,
mask=None,
context_mask=None,
mems=None,
return_hiddens=False
):
hiddens = []
intermediates = []
prev_attn = None
prev_cross_attn = None
mems = mems.copy() if exists(mems) else [None] * self.num_attn_layers
for ind, (layer_type, (norm, block, residual_fn)) in enumerate(zip(self.layer_types, self.layers)):
is_last = ind == (len(self.layers) - 1)
if layer_type == 'a':
hiddens.append(x)
layer_mem = mems.pop(0)
residual = x
if self.pre_norm:
x = norm(x)
if layer_type == 'a':
out, inter = block(x, mask=mask, sinusoidal_emb=self.pia_pos_emb, rel_pos=self.rel_pos,
prev_attn=prev_attn, mem=layer_mem)
elif layer_type == 'c':
out, inter = block(x, context=context, mask=mask, context_mask=context_mask, prev_attn=prev_cross_attn)
elif layer_type == 'f':
out = block(x)
x = residual_fn(out, residual)
if layer_type in ('a', 'c'):
intermediates.append(inter)
if layer_type == 'a' and self.residual_attn:
prev_attn = inter.pre_softmax_attn
elif layer_type == 'c' and self.cross_residual_attn:
prev_cross_attn = inter.pre_softmax_attn
if not self.pre_norm and not is_last:
x = norm(x)
if return_hiddens:
intermediates = LayerIntermediates(
hiddens=hiddens,
attn_intermediates=intermediates
)
return x, intermediates
return x
class Encoder(AttentionLayers):
def __init__(self, **kwargs):
assert 'causal' not in kwargs, 'cannot set causality on encoder'
super().__init__(causal=False, **kwargs)
class TransformerWrapper(nn.Module):
def __init__(
self,
*,
num_tokens,
max_seq_len,
attn_layers,
emb_dim=None,
max_mem_len=0.,
emb_dropout=0.,
num_memory_tokens=None,
tie_embedding=False,
use_pos_emb=True
):
super().__init__()
assert isinstance(attn_layers, AttentionLayers), 'attention layers must be one of Encoder or Decoder'
dim = attn_layers.dim
emb_dim = default(emb_dim, dim)
self.max_seq_len = max_seq_len
self.max_mem_len = max_mem_len
self.num_tokens = num_tokens
self.token_emb = nn.Embedding(num_tokens, emb_dim)
self.pos_emb = AbsolutePositionalEmbedding(emb_dim, max_seq_len) if (
use_pos_emb and not attn_layers.has_pos_emb) else always(0)
self.emb_dropout = nn.Dropout(emb_dropout)
self.project_emb = nn.Linear(emb_dim, dim) if emb_dim != dim else nn.Identity()
self.attn_layers = attn_layers
self.norm = nn.LayerNorm(dim)
self.init_()
self.to_logits = nn.Linear(dim, num_tokens) if not tie_embedding else lambda t: t @ self.token_emb.weight.t()
# memory tokens (like [cls]) from Memory Transformers paper
num_memory_tokens = default(num_memory_tokens, 0)
self.num_memory_tokens = num_memory_tokens
if num_memory_tokens > 0:
self.memory_tokens = nn.Parameter(torch.randn(num_memory_tokens, dim))
# let funnel encoder know number of memory tokens, if specified
if hasattr(attn_layers, 'num_memory_tokens'):
attn_layers.num_memory_tokens = num_memory_tokens
def init_(self):
nn.init.normal_(self.token_emb.weight, std=0.02)
def forward(
self,
x,
return_embeddings=False,
mask=None,
return_mems=False,
return_attn=False,
mems=None,
**kwargs
):
b, n, device, num_mem = *x.shape, x.device, self.num_memory_tokens
x = self.token_emb(x)
x += self.pos_emb(x)
x = self.emb_dropout(x)
x = self.project_emb(x)
if num_mem > 0:
mem = repeat(self.memory_tokens, 'n d -> b n d', b=b)
x = torch.cat((mem, x), dim=1)
# auto-handle masking after appending memory tokens
if exists(mask):
mask = F.pad(mask, (num_mem, 0), value=True)
x, intermediates = self.attn_layers(x, mask=mask, mems=mems, return_hiddens=True, **kwargs)
x = self.norm(x)
mem, x = x[:, :num_mem], x[:, num_mem:]
out = self.to_logits(x) if not return_embeddings else x
if return_mems:
hiddens = intermediates.hiddens
new_mems = list(map(lambda pair: torch.cat(pair, dim=-2), zip(mems, hiddens))) if exists(mems) else hiddens
new_mems = list(map(lambda t: t[..., -self.max_mem_len:, :].detach(), new_mems))
return out, new_mems
if return_attn:
attn_maps = list(map(lambda t: t.post_softmax_attn, intermediates.attn_intermediates))
return out, attn_maps
return out
+203
View File
@@ -0,0 +1,203 @@
import importlib
import torch
import numpy as np
from collections import abc
from einops import rearrange
from functools import partial
import multiprocessing as mp
from threading import Thread
from queue import Queue
from inspect import isfunction
from PIL import Image, ImageDraw, ImageFont
def log_txt_as_img(wh, xc, size=10):
# wh a tuple of (width, height)
# xc a list of captions to plot
b = len(xc)
txts = list()
for bi in range(b):
txt = Image.new("RGB", wh, color="white")
draw = ImageDraw.Draw(txt)
font = ImageFont.truetype('data/DejaVuSans.ttf', size=size)
nc = int(40 * (wh[0] / 256))
lines = "\n".join(xc[bi][start:start + nc] for start in range(0, len(xc[bi]), nc))
try:
draw.text((0, 0), lines, fill="black", font=font)
except UnicodeEncodeError:
print("Cant encode string for logging. Skipping.")
txt = np.array(txt).transpose(2, 0, 1) / 127.5 - 1.0
txts.append(txt)
txts = np.stack(txts)
txts = torch.tensor(txts)
return txts
def ismap(x):
if not isinstance(x, torch.Tensor):
return False
return (len(x.shape) == 4) and (x.shape[1] > 3)
def isimage(x):
if not isinstance(x, torch.Tensor):
return False
return (len(x.shape) == 4) and (x.shape[1] == 3 or x.shape[1] == 1)
def exists(x):
return x is not None
def default(val, d):
if exists(val):
return val
return d() if isfunction(d) else d
def mean_flat(tensor):
"""
https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/nn.py#L86
Take the mean over all non-batch dimensions.
"""
return tensor.mean(dim=list(range(1, len(tensor.shape))))
def count_params(model, verbose=False):
total_params = sum(p.numel() for p in model.parameters())
if verbose:
print(f"{model.__class__.__name__} has {total_params * 1.e-6:.2f} M params.")
return total_params
def instantiate_from_config(config):
if not "target" in config:
if config == '__is_first_stage__':
return None
elif config == "__is_unconditional__":
return None
raise KeyError("Expected key `target` to instantiate.")
return get_obj_from_str(config["target"])(**config.get("params", dict()))
def get_obj_from_str(string, reload=False):
module, cls = string.rsplit(".", 1)
if reload:
module_imp = importlib.import_module(module)
importlib.reload(module_imp)
return getattr(importlib.import_module(module, package=None), cls)
def _do_parallel_data_prefetch(func, Q, data, idx, idx_to_fn=False):
# create dummy dataset instance
# run prefetching
if idx_to_fn:
res = func(data, worker_id=idx)
else:
res = func(data)
Q.put([idx, res])
Q.put("Done")
def parallel_data_prefetch(
func: callable, data, n_proc, target_data_type="ndarray", cpu_intensive=True, use_worker_id=False
):
# if target_data_type not in ["ndarray", "list"]:
# raise ValueError(
# "Data, which is passed to parallel_data_prefetch has to be either of type list or ndarray."
# )
if isinstance(data, np.ndarray) and target_data_type == "list":
raise ValueError("list expected but function got ndarray.")
elif isinstance(data, abc.Iterable):
if isinstance(data, dict):
print(
f'WARNING:"data" argument passed to parallel_data_prefetch is a dict: Using only its values and disregarding keys.'
)
data = list(data.values())
if target_data_type == "ndarray":
data = np.asarray(data)
else:
data = list(data)
else:
raise TypeError(
f"The data, that shall be processed parallel has to be either an np.ndarray or an Iterable, but is actually {type(data)}."
)
if cpu_intensive:
Q = mp.Queue(1000)
proc = mp.Process
else:
Q = Queue(1000)
proc = Thread
# spawn processes
if target_data_type == "ndarray":
arguments = [
[func, Q, part, i, use_worker_id]
for i, part in enumerate(np.array_split(data, n_proc))
]
else:
step = (
int(len(data) / n_proc + 1)
if len(data) % n_proc != 0
else int(len(data) / n_proc)
)
arguments = [
[func, Q, part, i, use_worker_id]
for i, part in enumerate(
[data[i: i + step] for i in range(0, len(data), step)]
)
]
processes = []
for i in range(n_proc):
p = proc(target=_do_parallel_data_prefetch, args=arguments[i])
processes += [p]
# start processes
print(f"Start prefetching...")
import time
start = time.time()
gather_res = [[] for _ in range(n_proc)]
try:
for p in processes:
p.start()
k = 0
while k < n_proc:
# get result
res = Q.get()
if res == "Done":
k += 1
else:
gather_res[res[0]] = res[1]
except Exception as e:
print("Exception: ", e)
for p in processes:
p.terminate()
raise e
finally:
for p in processes:
p.join()
print(f"Prefetching complete. [{time.time() - start} sec.]")
if target_data_type == 'ndarray':
if not isinstance(gather_res[0], np.ndarray):
return np.concatenate([np.asarray(r) for r in gather_res], axis=0)
# order outputs
return np.concatenate(gather_res, axis=0)
elif target_data_type == 'list':
out = []
for r in gather_res:
out.extend(r)
return out
else:
return gather_res
+776
View File
@@ -0,0 +1,776 @@
# pytorch_diffusion + derived encoder decoder
import math
import torch
import torch.nn as nn
import numpy as np
def get_timestep_embedding(timesteps, embedding_dim):
"""
This matches the implementation in Denoising Diffusion Probabilistic Models:
From Fairseq.
Build sinusoidal embeddings.
This matches the implementation in tensor2tensor, but differs slightly
from the description in Section 3.5 of "Attention Is All You Need".
"""
assert len(timesteps.shape) == 1
half_dim = embedding_dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb)
emb = emb.to(device=timesteps.device)
emb = timesteps.float()[:, None] * emb[None, :]
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1)
if embedding_dim % 2 == 1: # zero pad
emb = torch.nn.functional.pad(emb, (0,1,0,0))
return emb
def nonlinearity(x):
# swish
return x*torch.sigmoid(x)
def Normalize(in_channels):
return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
class Upsample(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
self.conv = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x):
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
if self.with_conv:
x = self.conv(x)
return x
class Downsample(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
# no asymmetric padding in torch conv, must do it ourselves
self.conv = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=3,
stride=2,
padding=0)
def forward(self, x):
if self.with_conv:
pad = (0,1,0,1)
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
x = self.conv(x)
else:
x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2)
return x
class ResnetBlock(nn.Module):
def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False,
dropout, temb_channels=512):
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.use_conv_shortcut = conv_shortcut
self.norm1 = Normalize(in_channels)
self.conv1 = torch.nn.Conv2d(in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
if temb_channels > 0:
self.temb_proj = torch.nn.Linear(temb_channels,
out_channels)
self.norm2 = Normalize(out_channels)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = torch.nn.Conv2d(out_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
self.conv_shortcut = torch.nn.Conv2d(in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
else:
self.nin_shortcut = torch.nn.Conv2d(in_channels,
out_channels,
kernel_size=1,
stride=1,
padding=0)
def forward(self, x, temb):
h = x
h = self.norm1(h)
h = nonlinearity(h)
h = self.conv1(h)
if temb is not None:
h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None]
h = self.norm2(h)
h = nonlinearity(h)
h = self.dropout(h)
h = self.conv2(h)
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
x = self.conv_shortcut(x)
else:
x = self.nin_shortcut(x)
return x+h
class AttnBlock(nn.Module):
def __init__(self, in_channels):
super().__init__()
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.k = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.v = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.proj_out = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
def forward(self, x):
h_ = x
h_ = self.norm(h_)
q = self.q(h_)
k = self.k(h_)
v = self.v(h_)
# compute attention
b,c,h,w = q.shape
q = q.reshape(b,c,h*w)
q = q.permute(0,2,1) # b,hw,c
k = k.reshape(b,c,h*w) # b,c,hw
w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
w_ = w_ * (int(c)**(-0.5))
w_ = torch.nn.functional.softmax(w_, dim=2)
# attend to values
v = v.reshape(b,c,h*w)
w_ = w_.permute(0,2,1) # b,hw,hw (first hw of k, second of q)
h_ = torch.bmm(v,w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
h_ = h_.reshape(b,c,h,w)
h_ = self.proj_out(h_)
return x+h_
class Model(nn.Module):
def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
resolution, use_timestep=True):
super().__init__()
self.ch = ch
self.temb_ch = self.ch*4
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
self.use_timestep = use_timestep
if self.use_timestep:
# timestep embedding
self.temb = nn.Module()
self.temb.dense = nn.ModuleList([
torch.nn.Linear(self.ch,
self.temb_ch),
torch.nn.Linear(self.temb_ch,
self.temb_ch),
])
# downsampling
self.conv_in = torch.nn.Conv2d(in_channels,
self.ch,
kernel_size=3,
stride=1,
padding=1)
curr_res = resolution
in_ch_mult = (1,)+tuple(ch_mult)
self.down = nn.ModuleList()
for i_level in range(self.num_resolutions):
block = nn.ModuleList()
attn = nn.ModuleList()
block_in = ch*in_ch_mult[i_level]
block_out = ch*ch_mult[i_level]
for i_block in range(self.num_res_blocks):
block.append(ResnetBlock(in_channels=block_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout))
block_in = block_out
if curr_res in attn_resolutions:
attn.append(AttnBlock(block_in))
down = nn.Module()
down.block = block
down.attn = attn
if i_level != self.num_resolutions-1:
down.downsample = Downsample(block_in, resamp_with_conv)
curr_res = curr_res // 2
self.down.append(down)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
self.mid.attn_1 = AttnBlock(block_in)
self.mid.block_2 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
# upsampling
self.up = nn.ModuleList()
for i_level in reversed(range(self.num_resolutions)):
block = nn.ModuleList()
attn = nn.ModuleList()
block_out = ch*ch_mult[i_level]
skip_in = ch*ch_mult[i_level]
for i_block in range(self.num_res_blocks+1):
if i_block == self.num_res_blocks:
skip_in = ch*in_ch_mult[i_level]
block.append(ResnetBlock(in_channels=block_in+skip_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout))
block_in = block_out
if curr_res in attn_resolutions:
attn.append(AttnBlock(block_in))
up = nn.Module()
up.block = block
up.attn = attn
if i_level != 0:
up.upsample = Upsample(block_in, resamp_with_conv)
curr_res = curr_res * 2
self.up.insert(0, up) # prepend to get consistent order
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
out_ch,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x, t=None):
#assert x.shape[2] == x.shape[3] == self.resolution
if self.use_timestep:
# timestep embedding
assert t is not None
temb = get_timestep_embedding(t, self.ch)
temb = self.temb.dense[0](temb)
temb = nonlinearity(temb)
temb = self.temb.dense[1](temb)
else:
temb = None
# downsampling
hs = [self.conv_in(x)]
for i_level in range(self.num_resolutions):
for i_block in range(self.num_res_blocks):
h = self.down[i_level].block[i_block](hs[-1], temb)
if len(self.down[i_level].attn) > 0:
h = self.down[i_level].attn[i_block](h)
hs.append(h)
if i_level != self.num_resolutions-1:
hs.append(self.down[i_level].downsample(hs[-1]))
# middle
h = hs[-1]
h = self.mid.block_1(h, temb)
h = self.mid.attn_1(h)
h = self.mid.block_2(h, temb)
# upsampling
for i_level in reversed(range(self.num_resolutions)):
for i_block in range(self.num_res_blocks+1):
h = self.up[i_level].block[i_block](
torch.cat([h, hs.pop()], dim=1), temb)
if len(self.up[i_level].attn) > 0:
h = self.up[i_level].attn[i_block](h)
if i_level != 0:
h = self.up[i_level].upsample(h)
# end
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
return h
class Encoder(nn.Module):
def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
resolution, z_channels, double_z=True, **ignore_kwargs):
super().__init__()
self.ch = ch
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
# downsampling
self.conv_in = torch.nn.Conv2d(in_channels,
self.ch,
kernel_size=3,
stride=1,
padding=1)
curr_res = resolution
in_ch_mult = (1,)+tuple(ch_mult)
self.down = nn.ModuleList()
for i_level in range(self.num_resolutions):
block = nn.ModuleList()
attn = nn.ModuleList()
block_in = ch*in_ch_mult[i_level]
block_out = ch*ch_mult[i_level]
for i_block in range(self.num_res_blocks):
block.append(ResnetBlock(in_channels=block_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout))
block_in = block_out
if curr_res in attn_resolutions:
attn.append(AttnBlock(block_in))
down = nn.Module()
down.block = block
down.attn = attn
if i_level != self.num_resolutions-1:
down.downsample = Downsample(block_in, resamp_with_conv)
curr_res = curr_res // 2
self.down.append(down)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
self.mid.attn_1 = AttnBlock(block_in)
self.mid.block_2 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
2*z_channels if double_z else z_channels,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x):
#assert x.shape[2] == x.shape[3] == self.resolution, "{}, {}, {}".format(x.shape[2], x.shape[3], self.resolution)
# timestep embedding
temb = None
# downsampling
hs = [self.conv_in(x)]
for i_level in range(self.num_resolutions):
for i_block in range(self.num_res_blocks):
h = self.down[i_level].block[i_block](hs[-1], temb)
if len(self.down[i_level].attn) > 0:
h = self.down[i_level].attn[i_block](h)
hs.append(h)
if i_level != self.num_resolutions-1:
hs.append(self.down[i_level].downsample(hs[-1]))
# middle
h = hs[-1]
h = self.mid.block_1(h, temb)
h = self.mid.attn_1(h)
h = self.mid.block_2(h, temb)
# end
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
return h
class Decoder(nn.Module):
def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
resolution, z_channels, give_pre_end=False, **ignorekwargs):
super().__init__()
self.ch = ch
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
self.give_pre_end = give_pre_end
# compute in_ch_mult, block_in and curr_res at lowest res
in_ch_mult = (1,)+tuple(ch_mult)
block_in = ch*ch_mult[self.num_resolutions-1]
curr_res = resolution // 2**(self.num_resolutions-1)
self.z_shape = (1,z_channels,curr_res,curr_res)
print("Working with z of shape {} = {} dimensions.".format(
self.z_shape, np.prod(self.z_shape)))
# z to block_in
self.conv_in = torch.nn.Conv2d(z_channels,
block_in,
kernel_size=3,
stride=1,
padding=1)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
self.mid.attn_1 = AttnBlock(block_in)
self.mid.block_2 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
# upsampling
self.up = nn.ModuleList()
for i_level in reversed(range(self.num_resolutions)):
block = nn.ModuleList()
attn = nn.ModuleList()
block_out = ch*ch_mult[i_level]
for i_block in range(self.num_res_blocks+1):
block.append(ResnetBlock(in_channels=block_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout))
block_in = block_out
if curr_res in attn_resolutions:
attn.append(AttnBlock(block_in))
up = nn.Module()
up.block = block
up.attn = attn
if i_level != 0:
up.upsample = Upsample(block_in, resamp_with_conv)
curr_res = curr_res * 2
self.up.insert(0, up) # prepend to get consistent order
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
out_ch,
kernel_size=3,
stride=1,
padding=1)
def forward(self, z):
#assert z.shape[1:] == self.z_shape[1:]
self.last_z_shape = z.shape
# timestep embedding
temb = None
# z to block_in
h = self.conv_in(z)
# middle
h = self.mid.block_1(h, temb)
h = self.mid.attn_1(h)
h = self.mid.block_2(h, temb)
# upsampling
for i_level in reversed(range(self.num_resolutions)):
for i_block in range(self.num_res_blocks+1):
h = self.up[i_level].block[i_block](h, temb)
if len(self.up[i_level].attn) > 0:
h = self.up[i_level].attn[i_block](h)
if i_level != 0:
h = self.up[i_level].upsample(h)
# end
if self.give_pre_end:
return h
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
return h
class VUNet(nn.Module):
def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
attn_resolutions, dropout=0.0, resamp_with_conv=True,
in_channels, c_channels,
resolution, z_channels, use_timestep=False, **ignore_kwargs):
super().__init__()
self.ch = ch
self.temb_ch = self.ch*4
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.use_timestep = use_timestep
if self.use_timestep:
# timestep embedding
self.temb = nn.Module()
self.temb.dense = nn.ModuleList([
torch.nn.Linear(self.ch,
self.temb_ch),
torch.nn.Linear(self.temb_ch,
self.temb_ch),
])
# downsampling
self.conv_in = torch.nn.Conv2d(c_channels,
self.ch,
kernel_size=3,
stride=1,
padding=1)
curr_res = resolution
in_ch_mult = (1,)+tuple(ch_mult)
self.down = nn.ModuleList()
for i_level in range(self.num_resolutions):
block = nn.ModuleList()
attn = nn.ModuleList()
block_in = ch*in_ch_mult[i_level]
block_out = ch*ch_mult[i_level]
for i_block in range(self.num_res_blocks):
block.append(ResnetBlock(in_channels=block_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout))
block_in = block_out
if curr_res in attn_resolutions:
attn.append(AttnBlock(block_in))
down = nn.Module()
down.block = block
down.attn = attn
if i_level != self.num_resolutions-1:
down.downsample = Downsample(block_in, resamp_with_conv)
curr_res = curr_res // 2
self.down.append(down)
self.z_in = torch.nn.Conv2d(z_channels,
block_in,
kernel_size=1,
stride=1,
padding=0)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock(in_channels=2*block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
self.mid.attn_1 = AttnBlock(block_in)
self.mid.block_2 = ResnetBlock(in_channels=block_in,
out_channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout)
# upsampling
self.up = nn.ModuleList()
for i_level in reversed(range(self.num_resolutions)):
block = nn.ModuleList()
attn = nn.ModuleList()
block_out = ch*ch_mult[i_level]
skip_in = ch*ch_mult[i_level]
for i_block in range(self.num_res_blocks+1):
if i_block == self.num_res_blocks:
skip_in = ch*in_ch_mult[i_level]
block.append(ResnetBlock(in_channels=block_in+skip_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout))
block_in = block_out
if curr_res in attn_resolutions:
attn.append(AttnBlock(block_in))
up = nn.Module()
up.block = block
up.attn = attn
if i_level != 0:
up.upsample = Upsample(block_in, resamp_with_conv)
curr_res = curr_res * 2
self.up.insert(0, up) # prepend to get consistent order
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
out_ch,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x, z):
#assert x.shape[2] == x.shape[3] == self.resolution
if self.use_timestep:
# timestep embedding
assert t is not None
temb = get_timestep_embedding(t, self.ch)
temb = self.temb.dense[0](temb)
temb = nonlinearity(temb)
temb = self.temb.dense[1](temb)
else:
temb = None
# downsampling
hs = [self.conv_in(x)]
for i_level in range(self.num_resolutions):
for i_block in range(self.num_res_blocks):
h = self.down[i_level].block[i_block](hs[-1], temb)
if len(self.down[i_level].attn) > 0:
h = self.down[i_level].attn[i_block](h)
hs.append(h)
if i_level != self.num_resolutions-1:
hs.append(self.down[i_level].downsample(hs[-1]))
# middle
h = hs[-1]
z = self.z_in(z)
h = torch.cat((h,z),dim=1)
h = self.mid.block_1(h, temb)
h = self.mid.attn_1(h)
h = self.mid.block_2(h, temb)
# upsampling
for i_level in reversed(range(self.num_resolutions)):
for i_block in range(self.num_res_blocks+1):
h = self.up[i_level].block[i_block](
torch.cat([h, hs.pop()], dim=1), temb)
if len(self.up[i_level].attn) > 0:
h = self.up[i_level].attn[i_block](h)
if i_level != 0:
h = self.up[i_level].upsample(h)
# end
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
return h
class SimpleDecoder(nn.Module):
def __init__(self, in_channels, out_channels, *args, **kwargs):
super().__init__()
self.model = nn.ModuleList([nn.Conv2d(in_channels, in_channels, 1),
ResnetBlock(in_channels=in_channels,
out_channels=2 * in_channels,
temb_channels=0, dropout=0.0),
ResnetBlock(in_channels=2 * in_channels,
out_channels=4 * in_channels,
temb_channels=0, dropout=0.0),
ResnetBlock(in_channels=4 * in_channels,
out_channels=2 * in_channels,
temb_channels=0, dropout=0.0),
nn.Conv2d(2*in_channels, in_channels, 1),
Upsample(in_channels, with_conv=True)])
# end
self.norm_out = Normalize(in_channels)
self.conv_out = torch.nn.Conv2d(in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x):
for i, layer in enumerate(self.model):
if i in [1,2,3]:
x = layer(x, None)
else:
x = layer(x)
h = self.norm_out(x)
h = nonlinearity(h)
x = self.conv_out(h)
return x
class UpsampleDecoder(nn.Module):
def __init__(self, in_channels, out_channels, ch, num_res_blocks, resolution,
ch_mult=(2,2), dropout=0.0):
super().__init__()
# upsampling
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
block_in = in_channels
curr_res = resolution // 2 ** (self.num_resolutions - 1)
self.res_blocks = nn.ModuleList()
self.upsample_blocks = nn.ModuleList()
for i_level in range(self.num_resolutions):
res_block = []
block_out = ch * ch_mult[i_level]
for i_block in range(self.num_res_blocks + 1):
res_block.append(ResnetBlock(in_channels=block_in,
out_channels=block_out,
temb_channels=self.temb_ch,
dropout=dropout))
block_in = block_out
self.res_blocks.append(nn.ModuleList(res_block))
if i_level != self.num_resolutions - 1:
self.upsample_blocks.append(Upsample(block_in, True))
curr_res = curr_res * 2
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
out_channels,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x):
# upsampling
h = x
for k, i_level in enumerate(range(self.num_resolutions)):
for i_block in range(self.num_res_blocks + 1):
h = self.res_blocks[i_level][i_block](h, None)
if i_level != self.num_resolutions - 1:
h = self.upsample_blocks[k](h)
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
return h
+329
View File
@@ -0,0 +1,329 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from torch import einsum
from einops import rearrange
class VectorQuantizer(nn.Module):
"""
see https://github.com/MishaLaskin/vqvae/blob/d761a999e2267766400dc646d82d3ac3657771d4/models/quantizer.py
____________________________________________
Discretization bottleneck part of the VQ-VAE.
Inputs:
- n_e : number of embeddings
- e_dim : dimension of embedding
- beta : commitment cost used in loss term, beta * ||z_e(x)-sg[e]||^2
_____________________________________________
"""
# NOTE: this class contains a bug regarding beta; see VectorQuantizer2 for
# a fix and use legacy=False to apply that fix. VectorQuantizer2 can be
# used wherever VectorQuantizer has been used before and is additionally
# more efficient.
def __init__(self, n_e, e_dim, beta):
super(VectorQuantizer, self).__init__()
self.n_e = n_e
self.e_dim = e_dim
self.beta = beta
self.embedding = nn.Embedding(self.n_e, self.e_dim)
self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e)
def forward(self, z):
"""
Inputs the output of the encoder network z and maps it to a discrete
one-hot vector that is the index of the closest embedding vector e_j
z (continuous) -> z_q (discrete)
z.shape = (batch, channel, height, width)
quantization pipeline:
1. get encoder input (B,C,H,W)
2. flatten input to (B*H*W,C)
"""
# reshape z -> (batch, height, width, channel) and flatten
z = z.permute(0, 2, 3, 1).contiguous()
z_flattened = z.view(-1, self.e_dim)
# distances from z to embeddings e_j (z - e)^2 = z^2 + e^2 - 2 e * z
d = torch.sum(z_flattened ** 2, dim=1, keepdim=True) + \
torch.sum(self.embedding.weight**2, dim=1) - 2 * \
torch.matmul(z_flattened, self.embedding.weight.t())
## could possible replace this here
# #\start...
# find closest encodings
min_encoding_indices = torch.argmin(d, dim=1).unsqueeze(1)
min_encodings = torch.zeros(
min_encoding_indices.shape[0], self.n_e).to(z)
min_encodings.scatter_(1, min_encoding_indices, 1)
# dtype min encodings: torch.float32
# min_encodings shape: torch.Size([2048, 512])
# min_encoding_indices.shape: torch.Size([2048, 1])
# get quantized latent vectors
z_q = torch.matmul(min_encodings, self.embedding.weight).view(z.shape)
#.........\end
# with:
# .........\start
#min_encoding_indices = torch.argmin(d, dim=1)
#z_q = self.embedding(min_encoding_indices)
# ......\end......... (TODO)
# compute loss for embedding
loss = torch.mean((z_q.detach()-z)**2) + self.beta * \
torch.mean((z_q - z.detach()) ** 2)
# preserve gradients
z_q = z + (z_q - z).detach()
# perplexity
e_mean = torch.mean(min_encodings, dim=0)
perplexity = torch.exp(-torch.sum(e_mean * torch.log(e_mean + 1e-10)))
# reshape back to match original input shape
z_q = z_q.permute(0, 3, 1, 2).contiguous()
return z_q, loss, (perplexity, min_encodings, min_encoding_indices)
def get_codebook_entry(self, indices, shape):
# shape specifying (batch, height, width, channel)
# TODO: check for more easy handling with nn.Embedding
min_encodings = torch.zeros(indices.shape[0], self.n_e).to(indices)
min_encodings.scatter_(1, indices[:,None], 1)
# get quantized latent vectors
z_q = torch.matmul(min_encodings.float(), self.embedding.weight)
if shape is not None:
z_q = z_q.view(shape)
# reshape back to match original input shape
z_q = z_q.permute(0, 3, 1, 2).contiguous()
return z_q
class GumbelQuantize(nn.Module):
"""
credit to @karpathy: https://github.com/karpathy/deep-vector-quantization/blob/main/model.py (thanks!)
Gumbel Softmax trick quantizer
Categorical Reparameterization with Gumbel-Softmax, Jang et al. 2016
https://arxiv.org/abs/1611.01144
"""
def __init__(self, num_hiddens, embedding_dim, n_embed, straight_through=True,
kl_weight=5e-4, temp_init=1.0, use_vqinterface=True,
remap=None, unknown_index="random"):
super().__init__()
self.embedding_dim = embedding_dim
self.n_embed = n_embed
self.straight_through = straight_through
self.temperature = temp_init
self.kl_weight = kl_weight
self.proj = nn.Conv2d(num_hiddens, n_embed, 1)
self.embed = nn.Embedding(n_embed, embedding_dim)
self.use_vqinterface = use_vqinterface
self.remap = remap
if self.remap is not None:
self.register_buffer("used", torch.tensor(np.load(self.remap)))
self.re_embed = self.used.shape[0]
self.unknown_index = unknown_index # "random" or "extra" or integer
if self.unknown_index == "extra":
self.unknown_index = self.re_embed
self.re_embed = self.re_embed+1
print(f"Remapping {self.n_embed} indices to {self.re_embed} indices. "
f"Using {self.unknown_index} for unknown indices.")
else:
self.re_embed = n_embed
def remap_to_used(self, inds):
ishape = inds.shape
assert len(ishape)>1
inds = inds.reshape(ishape[0],-1)
used = self.used.to(inds)
match = (inds[:,:,None]==used[None,None,...]).long()
new = match.argmax(-1)
unknown = match.sum(2)<1
if self.unknown_index == "random":
new[unknown]=torch.randint(0,self.re_embed,size=new[unknown].shape).to(device=new.device)
else:
new[unknown] = self.unknown_index
return new.reshape(ishape)
def unmap_to_all(self, inds):
ishape = inds.shape
assert len(ishape)>1
inds = inds.reshape(ishape[0],-1)
used = self.used.to(inds)
if self.re_embed > self.used.shape[0]: # extra token
inds[inds>=self.used.shape[0]] = 0 # simply set to zero
back=torch.gather(used[None,:][inds.shape[0]*[0],:], 1, inds)
return back.reshape(ishape)
def forward(self, z, temp=None, return_logits=False):
# force hard = True when we are in eval mode, as we must quantize. actually, always true seems to work
hard = self.straight_through if self.training else True
temp = self.temperature if temp is None else temp
logits = self.proj(z)
if self.remap is not None:
# continue only with used logits
full_zeros = torch.zeros_like(logits)
logits = logits[:,self.used,...]
soft_one_hot = F.gumbel_softmax(logits, tau=temp, dim=1, hard=hard)
if self.remap is not None:
# go back to all entries but unused set to zero
full_zeros[:,self.used,...] = soft_one_hot
soft_one_hot = full_zeros
z_q = einsum('b n h w, n d -> b d h w', soft_one_hot, self.embed.weight)
# + kl divergence to the prior loss
qy = F.softmax(logits, dim=1)
diff = self.kl_weight * torch.sum(qy * torch.log(qy * self.n_embed + 1e-10), dim=1).mean()
ind = soft_one_hot.argmax(dim=1)
if self.remap is not None:
ind = self.remap_to_used(ind)
if self.use_vqinterface:
if return_logits:
return z_q, diff, (None, None, ind), logits
return z_q, diff, (None, None, ind)
return z_q, diff, ind
def get_codebook_entry(self, indices, shape):
b, h, w, c = shape
assert b*h*w == indices.shape[0]
indices = rearrange(indices, '(b h w) -> b h w', b=b, h=h, w=w)
if self.remap is not None:
indices = self.unmap_to_all(indices)
one_hot = F.one_hot(indices, num_classes=self.n_embed).permute(0, 3, 1, 2).float()
z_q = einsum('b n h w, n d -> b d h w', one_hot, self.embed.weight)
return z_q
class VectorQuantizer2(nn.Module):
"""
Improved version over VectorQuantizer, can be used as a drop-in replacement. Mostly
avoids costly matrix multiplications and allows for post-hoc remapping of indices.
"""
# NOTE: due to a bug the beta term was applied to the wrong term. for
# backwards compatibility we use the buggy version by default, but you can
# specify legacy=False to fix it.
def __init__(self, n_e, e_dim, beta, remap=None, unknown_index="random",
sane_index_shape=False, legacy=False):
super().__init__()
self.n_e = n_e
self.e_dim = e_dim
self.beta = beta
self.legacy = legacy
self.embedding = nn.Embedding(self.n_e, self.e_dim)
self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e)
self.remap = remap
if self.remap is not None:
self.register_buffer("used", torch.tensor(np.load(self.remap)))
self.re_embed = self.used.shape[0]
self.unknown_index = unknown_index # "random" or "extra" or integer
if self.unknown_index == "extra":
self.unknown_index = self.re_embed
self.re_embed = self.re_embed+1
print(f"Remapping {self.n_e} indices to {self.re_embed} indices. "
f"Using {self.unknown_index} for unknown indices.")
else:
self.re_embed = n_e
self.sane_index_shape = sane_index_shape
def remap_to_used(self, inds):
ishape = inds.shape
assert len(ishape)>1
inds = inds.reshape(ishape[0],-1)
used = self.used.to(inds)
match = (inds[:,:,None]==used[None,None,...]).long()
new = match.argmax(-1)
unknown = match.sum(2)<1
if self.unknown_index == "random":
new[unknown]=torch.randint(0,self.re_embed,size=new[unknown].shape).to(device=new.device)
else:
new[unknown] = self.unknown_index
return new.reshape(ishape)
def unmap_to_all(self, inds):
ishape = inds.shape
assert len(ishape)>1
inds = inds.reshape(ishape[0],-1)
used = self.used.to(inds)
if self.re_embed > self.used.shape[0]: # extra token
inds[inds>=self.used.shape[0]] = 0 # simply set to zero
back=torch.gather(used[None,:][inds.shape[0]*[0],:], 1, inds)
return back.reshape(ishape)
def forward(self, z, temp=None, rescale_logits=False, return_logits=False):
assert temp is None or temp==1.0, "Only for interface compatible with Gumbel"
assert rescale_logits==False, "Only for interface compatible with Gumbel"
assert return_logits==False, "Only for interface compatible with Gumbel"
# reshape z -> (batch, height, width, channel) and flatten
z = rearrange(z, 'b c h w -> b h w c').contiguous()
z_flattened = z.view(-1, self.e_dim)
# distances from z to embeddings e_j (z - e)^2 = z^2 + e^2 - 2 e * z
d = torch.sum(z_flattened ** 2, dim=1, keepdim=True) + \
torch.sum(self.embedding.weight**2, dim=1) - 2 * \
torch.einsum('bd,dn->bn', z_flattened, rearrange(self.embedding.weight, 'n d -> d n'))
min_encoding_indices = torch.argmin(d, dim=1)
z_q = self.embedding(min_encoding_indices).view(z.shape)
perplexity = None
min_encodings = None
# compute loss for embedding
if not self.legacy:
loss = self.beta * torch.mean((z_q.detach()-z)**2) + \
torch.mean((z_q - z.detach()) ** 2)
else:
loss = torch.mean((z_q.detach()-z)**2) + self.beta * \
torch.mean((z_q - z.detach()) ** 2)
# preserve gradients
z_q = z + (z_q - z).detach()
# reshape back to match original input shape
z_q = rearrange(z_q, 'b h w c -> b c h w').contiguous()
if self.remap is not None:
min_encoding_indices = min_encoding_indices.reshape(z.shape[0],-1) # add batch axis
min_encoding_indices = self.remap_to_used(min_encoding_indices)
min_encoding_indices = min_encoding_indices.reshape(-1,1) # flatten
if self.sane_index_shape:
min_encoding_indices = min_encoding_indices.reshape(
z_q.shape[0], z_q.shape[2], z_q.shape[3])
return z_q, loss, (perplexity, min_encodings, min_encoding_indices)
def get_codebook_entry(self, indices, shape):
# shape specifying (batch, height, width, channel)
if self.remap is not None:
indices = indices.reshape(shape[0],-1) # add batch axis
indices = self.unmap_to_all(indices)
indices = indices.reshape(-1) # flatten again
# get quantized latent vectors
z_q = self.embedding(indices)
if shape is not None:
z_q = z_q.view(shape)
# reshape back to match original input shape
z_q = z_q.permute(0, 3, 1, 2).contiguous()
return z_q
+485
View File
@@ -0,0 +1,485 @@
import pdb
import torch
import torch.nn.functional as F
import pytorch_lightning as pl
from torch.optim.lr_scheduler import LambdaLR,StepLR
import numpy as np
from packaging import version
from model.BrownianBridge.base.modules.ema import LitEma
from contextlib import contextmanager
import omegaconf
from model.VQGAN.model import Encoder, Decoder
from model.BrownianBridge.base.modules.diffusionmodules.model import *
from model.VQGAN.quantize import VectorQuantizer2 as VectorQuantizer
from model.VQGAN.quantize import GumbelQuantize
from model.BrownianBridge.base.util import instantiate_from_config
import argparse
#from utils import dict2namespace, namespace2dict
import importlib
def dict2namespace(config):
namespace = argparse.Namespace()
for key, value in config.items():
if isinstance(value, dict) or isinstance(value, omegaconf.dictconfig.DictConfig):
new_value = dict2namespace(value)
print("convertng to namespaces")
else:
new_value = value
setattr(namespace, key, new_value)
return namespace
def get_obj_from_str(string, reload=False):
module, cls = string.rsplit(".", 1)
if reload:
module_imp = importlib.import_module(module)
importlib.reload(module_imp)
return getattr(importlib.import_module(module, package=None), cls)
def instantiate_from_config(config):
# pdb.set_trace()
if not "target" in config:
raise KeyError("Expected key `target` to instantiate.")
if config.__contains__('params'):
return get_obj_from_str(config["target"])(**vars(config['params']))
else:
return get_obj_from_str(config["target"])()
class VQFlowNet(pl.LightningModule):
def __init__(self,
ddconfig,
lossconfig,
n_embed,
embed_dim,
ckpt_path=None,
ignore_keys=[],
image_key="image",
colorize_nlabels=None,
monitor=None,
batch_resize_range=None,
scheduler_config=None,
lr_g_factor=1.0,
remap=None,
sane_index_shape=False, # tell vector quantizer to return indices as bhw
use_ema=False
):
super().__init__()
try:
ddconfig = dict2namespace(ddconfig)
lossconfig = dict2namespace(lossconfig)
except:
pass
self.embed_dim = embed_dim # 3
self.n_embed = n_embed # 8192 * 2
self.image_key = image_key # 'image'
self.encoder = FlowEncoder(**vars(ddconfig))
self.decoder = FlowDecoderWithResidual(**vars(ddconfig))
self.loss = instantiate_from_config(vars(lossconfig))
self.quantize = VectorQuantizer(n_embed, embed_dim, beta=0.25,
remap=remap,
sane_index_shape=sane_index_shape)
self.quant_conv = torch.nn.Conv2d(ddconfig.z_channels, embed_dim, 1)
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig.z_channels, 1)
if colorize_nlabels is not None:
assert type(colorize_nlabels)==int
self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1))
if monitor is not None:
self.monitor = monitor
self.batch_resize_range = batch_resize_range
if self.batch_resize_range is not None:
print(f"{self.__class__.__name__}: Using per-batch resizing in range {batch_resize_range}.")
self.use_ema = use_ema
if self.use_ema:
self.model_ema = LitEma(self)
print(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.")
if ckpt_path is not None:
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
self.scheduler_config = scheduler_config
self.lr_g_factor = lr_g_factor
self.h0 = None
self.w0 = None
self.h_padded = None
self.w_padded = None
self.pad_h = 0
self.pad_w = 0
@contextmanager
def ema_scope(self, context=None):
if self.use_ema:
self.model_ema.store(self.parameters())
self.model_ema.copy_to(self)
if context is not None:
print(f"{context}: Switched to EMA weights")
try:
yield None
finally:
if self.use_ema:
self.model_ema.restore(self.parameters())
if context is not None:
print(f"{context}: Restored training weights")
def init_from_ckpt(self, path, ignore_keys=list()):
sd = torch.load(path, map_location="cpu")["state_dict"]
keys = list(sd.keys())
for k in keys:
for ik in ignore_keys:
if k.startswith(ik):
print("Deleting key {} from state_dict.".format(k))
del sd[k]
missing, unexpected = self.load_state_dict(sd, strict=False)
#print(f"Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys")
if len(missing) > 0:
print(f"Missing Keys: {missing}")
print(f"Unexpected Keys: {unexpected}")
def on_train_batch_end(self, *args, **kwargs):
if self.use_ema:
self.model_ema(self)
def encode(self, x, ret_feature=False):
'''
Set ret_feature = True when encoding conditions in ddpm
'''
# Pad the input first so its size is deividable by 8.
# this is to tolerate different f values, various size inputs,
# and some operations in the DDPM unet model.
self.h0, self.w0 = x.shape[2:]
# 8: window size for max vit
# 2**(nr-1): f
# 4: factor of downsampling in DDPM unet
min_side = 8 *2**(self.encoder.num_resolutions-1) * 4
if self.h0 % min_side != 0:
pad_h = min_side - (self.h0 % min_side)
if pad_h == self.h0: # this is to avoid padding 256 patches
pad_h = 0
x = F.pad(x, (0, 0, 0, pad_h), mode='reflect')
self.h_padded = True
self.pad_h = pad_h
if self.w0 % min_side != 0:
pad_w = min_side - (self.w0 % min_side)
if pad_w == self.w0:
pad_w = 0
x = F.pad(x, (0, pad_w, 0, 0), mode='reflect')
self.w_padded = True
self.pad_w = pad_w
h, phi_list = self.encoder(x, ret_feature)
h = self.quant_conv(h)
quant, emb_loss, info = self.quantize(h)
return quant, emb_loss, info, phi_list
def encode_to_prequant(self, x):
h = self.encoder(x)
h = self.quant_conv(h)
return h
def decode(self, quant, x_prev, x_next, phi_list=None):
cond_dict = dict(
phi_list = phi_list,
frame_prev = F.pad(x_prev, (0, self.pad_w, 0, self.pad_h), mode='reflect'),
frame_next = F.pad(x_next, (0, self.pad_w, 0, self.pad_h), mode='reflect')
)
"""
cond_dict = dict(
phi_prev_list = self.encode(x_prev, ret_feature=True)[-1],
phi_next_list = self.encode(x_next, ret_feature=True)[-1],
frame_prev = F.pad(x_prev, (0, self.pad_w, 0, self.pad_h), mode='reflect'),
frame_next = F.pad(x_next, (0, self.pad_w, 0, self.pad_h), mode='reflect')
)
"""
quant = self.post_quant_conv(quant)
dec = self.decoder(quant, cond_dict)
# check if image is padded and return the original part only
if self.h_padded:
dec = dec[:, :, 0:self.h0, :]
if self.w_padded:
dec = dec[:, :, :, 0:self.w0]
return dec
def decode_code(self, code_b):
quant_b = self.quantize.embed_code(code_b)
dec = self.decode(quant_b)
return dec
def forward(self, x, x_prev,x_next, return_pred_indices=False):
inputs = torch.cat([x_prev,x,x_next],0) ## B3 C H W
quant, diff, (_,_,ind), phi_list = self.encode(inputs)
#quant= self.encode(input)
dec = self.decode(quant, x_prev, x_next,phi_list)
#dec = self.decode(quant, x_prev, x_next)
if return_pred_indices:
return dec, diff, ind
return dec,diff
def get_input(self, batch, k):
x = batch[k]
if len(x.shape) == 3:
x = x[..., None]
x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format).float()
if self.batch_resize_range is not None:
lower_size = self.batch_resize_range[0]
upper_size = self.batch_resize_range[1]
if self.global_step <= 4:
# do the first few batches with max size to avoid later oom
new_resize = upper_size
else:
new_resize = np.random.choice(np.arange(lower_size, upper_size+16, 16))
if new_resize != x.shape[2]:
x = F.interpolate(x, size=new_resize, mode="bicubic")
x = x.detach()
return x
def training_step(self, batch, batch_idx, optimizer_idx):
self.train()
# https://github.com/pytorch/pytorch/issues/37142
# try not to fool the heuristics
x = self.get_input(batch, self.image_key)
x_prev = self.get_input(batch, 'prev_frame')
x_next = self.get_input(batch, 'next_frame')
xrec, qloss = self(x,x_prev,x_next)
if optimizer_idx == 0:
# autoencoder
aeloss, log_dict_ae = self.loss(qloss, x, xrec, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train",
flow_list = None,flow_gt=None,x_prev = x_prev,x_next = x_next)
self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=True)
return aeloss
if optimizer_idx == 1:
# discriminator
discloss, log_dict_disc = self.loss(qloss, x, xrec, optimizer_idx, self.global_step,
last_layer=self.get_last_layer(), split="train",x_prev = x_prev,x_next = x_next)
self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=True)
return discloss
def validation_step(self, batch, batch_idx):
log_dict = self._validation_step(batch, batch_idx)
with self.ema_scope():
log_dict_ema = self._validation_step(batch, batch_idx, suffix="_ema")
return log_dict
def _validation_step(self, batch, batch_idx, suffix=""):
self.eval()
x = self.get_input(batch, self.image_key)
x_prev = self.get_input(batch, 'prev_frame')
x_next = self.get_input(batch, 'next_frame')
xrec, qloss = self(x, x_prev, x_next)
aeloss, log_dict_ae = self.loss(qloss, x, xrec, 0,
self.global_step,
last_layer=self.get_last_layer(),
split="val"+suffix,
)
discloss, log_dict_disc = self.loss(qloss, x, xrec, 1,
self.global_step,
last_layer=self.get_last_layer(),
split="val"+suffix,
)
rec_loss = log_dict_ae[f"val{suffix}/rec_loss"]
self.log(f"val{suffix}/rec_loss", rec_loss,
prog_bar=True, logger=True, on_step=False, on_epoch=True, sync_dist=True)
self.log(f"val{suffix}/aeloss", aeloss,
prog_bar=True, logger=True, on_step=False, on_epoch=True, sync_dist=True)
if version.parse(pl.__version__) >= version.parse('1.4.0'):
del log_dict_ae[f"val{suffix}/rec_loss"]
self.log_dict(log_dict_ae)
self.log_dict(log_dict_disc)
return self.log_dict
def configure_optimizers(self):
lr_d = self.learning_rate
lr_g = self.lr_g_factor*self.learning_rate
print("lr_d", lr_d)
print("lr_g", lr_g)
opt_ae = torch.optim.Adam(list(self.encoder.parameters())+
list(self.decoder.parameters())+
list(self.quantize.parameters())+
list(self.quant_conv.parameters())+
list(self.post_quant_conv.parameters()),
lr=lr_g, betas=(0.5, 0.9))
opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(),
lr=lr_d, betas=(0.5, 0.9))
if self.scheduler_config is not None:
scheduler = instantiate_from_config(self.scheduler_config)
print("Setting up LambdaLR scheduler...")
scheduler = [
{
'scheduler': LambdaLR(opt_ae, lr_lambda=scheduler.schedule),
'interval': 'step',
'frequency': 1
},
{
'scheduler': LambdaLR(opt_disc, lr_lambda=scheduler.schedule),
'interval': 'step',
'frequency': 1
},
]
return [opt_ae,opt_disc ], scheduler
return [opt_ae, opt_disc], []
def get_last_layer(self):
return self.decoder.conv_out.weight
def log_images(self, batch, only_inputs=False, plot_ema=False, **kwargs):
log = dict()
x = self.get_input(batch, self.image_key)
x_prev = self.get_input(batch, 'prev_frame')
x_next = self.get_input(batch, 'next_frame')
x = x.to(self.device)
if only_inputs:
log["inputs"] = x
return log
xrec, _ = self(x, x_prev, x_next)
if x.shape[1] > 3:
# colorize with random projection
assert xrec.shape[1] > 3
x = self.to_rgb(x)
xrec = self.to_rgb(xrec)
log["inputs"] = x
log["reconstructions"] = xrec
if plot_ema:
with self.ema_scope():
xrec_ema, _ = self(x, x_prev, x_next)
if x.shape[1] > 3: xrec_ema = self.to_rgb(xrec_ema)
log["reconstructions_ema"] = xrec_ema
return log
def to_rgb(self, x):
assert self.image_key == "segmentation"
if not hasattr(self, "colorize"):
self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))
x = F.conv2d(x, weight=self.colorize)
x = 2.*(x-x.min())/(x.max()-x.min()) - 1.
return x
def get_flow(self, img0, img1,feats):
return self.decoder.get_flow(img0,img1,feats)
class VQFlowNetInterface(VQFlowNet):
def __init__(self, **kwargs):
super().__init__(**kwargs)
def encode(self, x, ret_feature=False):
'''
Set ret_feature = True when encoding conditions in ddpm
'''
# Pad the input first so its size is deividable by 8.
# this is to tolerate different f values, various size inputs,
# and some operations in the DDPM unet model.
self.h0, self.w0 = x.shape[2:]
# 8: window size for max vit
# 2**(nr-1): f
# 4: factor of downsampling in DDPM unet
min_side = 8 * 2**(self.encoder.num_resolutions-1) * 4
#min_side = 256
if self.h0 % min_side != 0:
pad_h = min_side - (self.h0 % min_side)
if pad_h == self.h0: # this is to avoid padding 256 patches
pad_h = 0
x = F.pad(x, (0, 0, 0, pad_h), mode='reflect')
self.h_padded = True
self.pad_h = pad_h
else:
self.h_padded = False
self.pad_h = 0
if self.w0 % min_side != 0:
pad_w = min_side - (self.w0 % min_side)
if pad_w == self.w0:
pad_w = 0
x = F.pad(x, (0, pad_w, 0, 0), mode='reflect')
self.w_padded = True
self.pad_w = pad_w
else:
self.w_padded = False
self.pad_w = 0
h, phi_list = self.encoder(x, ret_feature)
h = self.quant_conv(h) ## before quantization
return h, phi_list
def decode(self, h, x_prev, x_next, phi_list, force_not_quantize=False,scale = 0.5):
# also go through quantization layer
if not force_not_quantize:
quant, emb_loss, info = self.quantize(h)
else:
quant = h
cond_dict = dict(
phi_list = phi_list,
frame_prev = F.pad(x_prev, (0, self.pad_w, 0, self.pad_h), mode='reflect'),
frame_next = F.pad(x_next, (0, self.pad_w, 0, self.pad_h), mode='reflect')
)
quant = self.post_quant_conv(quant)
tmp_list = []
b,c,h,w = x_prev.shape
with torch.no_grad():
if scale < 1:
b,c,h,w = F.interpolate(F.pad(x_prev, (0, self.pad_w, 0, self.pad_h), mode='reflect'), scale_factor=scale, mode="bilinear", align_corners=False).shape
img0_down_ = F.interpolate(x_prev, scale_factor=scale, mode="bilinear", align_corners=False)
img1_down_ = F.interpolate(x_next, scale_factor=scale, mode="bilinear", align_corners=False)
_,_,h_,w_ = img0_down_.shape
img0_down = torch.zeros(b,c,h,w).to(img0_down_.device)
img1_down = torch.zeros(b,c,h,w).to(img1_down_.device)
img0_down[:,:,:h_,:w_] = img0_down_
img1_down[:,:,:h_,:w_] = img1_down_
_,tmp_list = self.encoder(torch.cat([img0_down,torch.zeros_like(img0_down),img1_down]))
flow_down = self.get_flow(img0_down, img1_down,tmp_list[:-2])
flow = F.interpolate(flow_down, scale_factor=1/scale, mode="bilinear", align_corners=False) * 1/scale
else:
flow = None
dec = self.decoder(quant, cond_dict,flow)
# check if image is padded and return the original part only
if self.h_padded:
dec = dec[:, :, 0:self.h0, :]
if self.w_padded:
dec = dec[:, :, :, 0:self.w0]
return dec
+17
View File
@@ -0,0 +1,17 @@
from inspect import isfunction
def extract(a, t, x_shape):
b, *_ = t.shape
out = a.gather(-1, t)
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
def exists(x):
return x is not None
def default(val, d):
if exists(val):
return val
return d() if isfunction(d) else d
+11
View File
@@ -0,0 +1,11 @@
from .tlbvfi_node import TLBVFI_VFI
NODE_CLASS_MAPPINGS = {
"TLBVFI_VFI": TLBVFI_VFI
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TLBVFI_VFI": "TLBVFI Frame Interpolation"
}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+2
View File
@@ -0,0 +1,2 @@
pytorch_lightning
cupy-cuda11x
+159
View File
@@ -0,0 +1,159 @@
import torch
import os
import sys
from pathlib import Path
import yaml
import argparse
import folder_paths
from comfy.utils import ProgressBar
import numpy as np
from tqdm import tqdm
# --- Robust Path Handling ---
if 'interpolation' not in folder_paths.folder_names_and_paths:
new_path = os.path.join(folder_paths.models_dir, 'interpolation')
os.makedirs(new_path, exist_ok=True)
folder_paths.folder_names_and_paths['interpolation'] = ([new_path], {'.pth', '.ckpt'})
# --- Helper Functions ---
def dict2namespace(config):
"""Converts a dictionary to a namespace for easier access."""
namespace = argparse.Namespace()
for key, value in config.items():
if isinstance(value, dict):
new_value = dict2namespace(value)
else:
new_value = value
setattr(namespace, key, new_value)
return namespace
def find_models(folder_type: str, extensions: list) -> list:
"""Recursively finds all model files with given extensions in the specified folder type."""
model_list = []
base_paths = folder_paths.get_folder_paths(folder_type)
for base_path in base_paths:
for root, _, files in os.walk(base_path, followlinks=True):
for file in files:
if any(file.lower().endswith(ext) for ext in extensions):
relative_path = os.path.relpath(os.path.join(root, file), base_path)
model_list.append(relative_path.replace("\\", "/"))
return sorted(list(set(model_list)))
# --- TLBVFI Setup ---
try:
current_path = Path(__file__).parent
tlbvfi_path = current_path / "TLBVFI"
if tlbvfi_path.is_dir():
sys.path.insert(0, str(tlbvfi_path))
from model.BrownianBridge.LatentBrownianBridgeModel import LatentBrownianBridgeModel
else:
raise ImportError("TLBVFI directory not found.")
except ImportError as e:
print("-------------------------------------------------------------------")
print(f"Error: {e}")
print("Could not import TLBVFI model from ComfyUI-TLBVFI node.")
print("Please follow the setup instructions in the README.md file.")
print("-------------------------------------------------------------------")
raise
# --- Main Node Class ---
class TLBVFI_VFI:
@classmethod
def INPUT_TYPES(s):
# We only need the main model file now.
unet_models = find_models("interpolation", [".pth"])
if not unet_models:
raise Exception("No TLBVFI UNet models (.pth) found in 'ComfyUI/models/interpolation/'. Please download 'vimeo_unet.pth'.")
return {
"required": {
"images": ("IMAGE", ),
"model_name": (unet_models, ),
"times_to_interpolate": ("INT", {"default": 1, "min": 1, "max": 4, "step": 1}),
"gpu_id": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "interpolate"
CATEGORY = "frame_interpolation/TLBVFI" # Updated Category for better organization
def interpolate(self, images, model_name, times_to_interpolate, gpu_id):
# --- Setup ---
device = torch.device(f"cuda:{gpu_id}" if torch.cuda.is_available() else "cpu")
# --- Load Config ---
tlbvfi_repo_path = Path(__file__).parent / "TLBVFI"
config_path = tlbvfi_repo_path / "configs" / "Template-LBBDM-video.yaml"
if not config_path.exists():
raise FileNotFoundError(f"Config file not found at {config_path}. Make sure the TLBVFI repo is cloned correctly.")
with open(config_path, 'r') as f:
config = yaml.load(f, Loader=yaml.FullLoader)
nconfig = dict2namespace(config)
# --- Simplified Model Loading ---
model_path = folder_paths.get_full_path("interpolation", model_name)
if not model_path or not os.path.exists(model_path):
raise FileNotFoundError(f"Model file {model_name} not found.")
# Prevent the model from trying to load a VQGAN checkpoint during initialization
nconfig.model.VQGAN.params.ckpt_path = None
# 1. Initialize the model structure with random weights.
model = LatentBrownianBridgeModel(nconfig.model).to(device)
# 2. Load the entire state dict from the single .pth file.
# This will populate both the VQGAN and the UNet with the correct weights.
checkpoint = torch.load(model_path, map_location=device)
# The state dict might be nested under a 'model' key.
state_dict_to_load = checkpoint.get('model', checkpoint)
model.load_state_dict(state_dict_to_load)
model.eval()
# --- Prepare Images ---
image_tensors = images.permute(0, 3, 1, 2).float()
image_tensors = (image_tensors * 2.0) - 1.0
if len(image_tensors) < 2:
print("TLBVFI Warning: Not enough images to interpolate. Returning original images.")
return (images, )
gui_pbar = ProgressBar(len(image_tensors) - 1)
output_frames = [image_tensors[0:1]]
# --- Main Interpolation Loop ---
for i in tqdm(range(len(image_tensors) - 1), desc="TLBVFI Interpolating"):
frame1 = image_tensors[i].unsqueeze(0).to(device)
frame2 = image_tensors[i+1].unsqueeze(0).to(device)
current_frames = [frame1, frame2]
for _ in range(times_to_interpolate):
temp_frames = [current_frames[0]]
for j in range(len(current_frames) - 1):
with torch.no_grad():
mid_frame = model.sample(current_frames[j], current_frames[j+1], disable_progress=True)
temp_frames.extend([mid_frame, current_frames[j+1]])
current_frames = temp_frames
for frame in current_frames[1:]:
output_frames.append(frame.cpu())
gui_pbar.update(1)
final_tensors = torch.cat(output_frames, dim=0)
# --- Convert back to ComfyUI's expected format ---
final_tensors = (final_tensors + 1.0) / 2.0
final_tensors = final_tensors.clamp(0, 1)
final_tensors = final_tensors.permute(0, 2, 3, 1)
return (final_tensors, )