Initial commit
This commit is contained in:
@@ -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},
|
||||
}
|
||||
|
||||
```
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -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"}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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']
|
||||
@@ -0,0 +1,2 @@
|
||||
pytorch_lightning
|
||||
cupy-cuda11x
|
||||
+159
@@ -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, )
|
||||
Reference in New Issue
Block a user