v3 schema and code refinments
This commit is contained in:
@@ -1,103 +1,58 @@
|
||||
# ComfyUI-TLBVFI
|
||||
|
||||
A LLM coded node pack for ComfyUI that provides video frame interpolation using the **TLB-VFI** model.
|
||||
A 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.
|
||||
This is a wrapper for the [TLB-VFI: Temporal-Aware Latent Brownian Bridge Diffusion for Video Frame Interpolation](https://github.com/ZonglinL/TLBVFI) project.
|
||||
|
||||
## 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.
|
||||
|
||||
---
|
||||
- **Zero-Dependency**: All non-standard requirements (CuPy, PyTorch-Lightning, etc.) have been removed or replaced with native implementations.
|
||||
- **Efficient Batching**: Supports processing multiple frame pairs simultaneously.
|
||||
|
||||
## ⚙️ 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.
|
||||
Clone this repository into your `ComfyUI/custom_nodes/` directory:
|
||||
|
||||
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
|
||||
### Step 2: Download the Pre-trained Model
|
||||
Download `vimeo_unet.pth` from the official repository:
|
||||
- **[ucfzl/TLBVFI on Hugging Face](https://huggingface.co/ucfzl/TLBVFI/tree/main)**
|
||||
|
||||
```bash
|
||||
# Navigate into the newly created custom node directory
|
||||
cd ComfyUI/custom_nodes/ComfyUI-TLBVFI/
|
||||
### Step 3: Place Model in the `interpolation` Folder
|
||||
Place the downloaded `.pth` file into `ComfyUI/models/interpolation/`.
|
||||
|
||||
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
|
||||
└── 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.
|
||||
|
||||
---
|
||||
1. Add the **TLBVFI Frame Interpolation** node from the `frame_interpolation/TLBVFI` category.
|
||||
2. Select the correct model from the **`model_name`** dropdown.
|
||||
3. **`times_to_interpolate`**: Sets how many new frames are generated between pairs (1 = double FPS, 2 = quadruple, etc.).
|
||||
4. **`diffusion_steps`**: Controls the refinement quality. Higher values (e.g., 20-50) improve quality at the cost of speed.
|
||||
5. **`batch_size`**: Number of pairs to process at once. Increase if you have sufficient VRAM for a speed boost.
|
||||
6. **`flow_scale`**: Resolution for motion analysis. Use `0.5` for most videos; lower values handle fast motion better.
|
||||
|
||||
## 🧠 How It Works
|
||||
|
||||
This node uses a two-stage **latent diffusion** process:
|
||||
1. **VQGAN**: Compresses input frames into a latent space.
|
||||
2. **UNet (Brownian Bridge)**: Operates in latent space to diffuse and generate the in-between representation using a reverse diffusion process.
|
||||
3. **VQGAN Decoder**: Reconstructs the generated latent back into a full-resolution image.
|
||||
|
||||
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.
|
||||
## 🙏 Acknowledgements
|
||||
|
||||
This approach is highly efficient and allows for the generation of high-quality, temporally consistent frames.
|
||||
All credit for the architecture and research goes to the original authors of TLB-VFI.
|
||||
|
||||
---
|
||||
|
||||
## 🙏 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/)
|
||||
- **Original GitHub**: [ZonglinL/TLBVFI](https://github.com/ZonglinL/TLBVFI)
|
||||
- **Project Page**: [https://zonglinl.github.io/tlbvfi_page/](https://zonglinl.github.io/tlbvfi_page/)
|
||||
|
||||
```bibtex
|
||||
@article{lyu2025tlbvfitemporalawarelatentbrownian,
|
||||
@@ -108,5 +63,4 @@ This node is a wrapper implementation for ComfyUI. All credit for the model arch
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
}
|
||||
|
||||
```
|
||||
|
||||
@@ -4,59 +4,11 @@ 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.layers import DropPath, to_2tuple, trunc_normal_
|
||||
from model.utils import 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,
|
||||
@@ -64,7 +16,6 @@ def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
|
||||
nn.PReLU(out_planes)
|
||||
)
|
||||
|
||||
|
||||
class Conv2(nn.Module):
|
||||
def __init__(self, in_planes, out_planes, stride=2):
|
||||
super().__init__()
|
||||
@@ -76,7 +27,6 @@ class Conv2(nn.Module):
|
||||
x = self.conv2(x)
|
||||
return x
|
||||
|
||||
|
||||
class IFBlock(nn.Module):
|
||||
def __init__(self, in_planes, scale=1, c=64):
|
||||
super().__init__()
|
||||
@@ -86,14 +36,8 @@ class IFBlock(nn.Module):
|
||||
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),
|
||||
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)
|
||||
|
||||
@@ -108,7 +52,6 @@ class IFBlock(nn.Module):
|
||||
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__()
|
||||
@@ -129,182 +72,99 @@ class IFNet(nn.Module):
|
||||
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
|
||||
corr_dim, flow_dim, motion_dim, hidden_dim = c, c, c, c
|
||||
self.n_iters, self.r = n_iters, 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),
|
||||
self.occl_convs = nn.Sequential(nn.Conv2d(2 * context_dim, hidden_dim, 1), nn.PReLU(hidden_dim),
|
||||
nn.Conv2d(hidden_dim, hidden_dim, 1), nn.PReLU(hidden_dim),
|
||||
nn.Conv2d(hidden_dim, 1, 1), nn.Sigmoid())
|
||||
self.corr_convs = nn.Sequential(nn.Conv2d(self.n_pts, hidden_dim, 1), nn.PReLU(hidden_dim),
|
||||
nn.Conv2d(hidden_dim, corr_dim, 1), 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)
|
||||
eps = 1e-5
|
||||
x_f = x.float()
|
||||
norm = x_f.pow(2).sum(dim=dim, keepdim=True).add(eps).sqrt()
|
||||
return (x_f / norm).type(x.dtype)
|
||||
|
||||
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
|
||||
x0_unfold = F.unfold(x0, kernel_size=(self.r * 2 + 1), padding=1).view(B, C * self.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)
|
||||
contents0, contents1 = warp(x0_unfold, flow0).view(B, C, self.n_pts, H, W), warp(x1_unfold, flow1).view(B, C, self.n_pts, H, W)
|
||||
fea0, fea1 = contents0[:, :, self.n_pts // 2, :, :], contents1[:, :, self.n_pts // 2, :, :]
|
||||
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
|
||||
|
||||
fea_view, contents0, contents1 = self.L2normalize(fea_view, -1), self.L2normalize(contents0, -1), self.L2normalize(contents1, -1)
|
||||
corr0, corr1 = torch.einsum('bic,bjc->bij', fea_view, contents0), torch.einsum('bic,bjc->bij', fea_view, contents1)
|
||||
corr0 = self.corr_convs(corr0.view(B, H, W, self.n_pts).permute(0, 3, 1, 2).contiguous())
|
||||
corr1 = self.corr_convs(corr1.view(B, H, W, self.n_pts).permute(0, 3, 1, 2).contiguous())
|
||||
flow0_fea, flow1_fea = self.flow_convs(flow0), self.flow_convs(flow1)
|
||||
motion0, motion1 = self.motion_convs(torch.cat([corr0, flow0_fea], 1)), self.motion_convs(torch.cat([corr1, flow1_fea], 1))
|
||||
flow0 = flow0 + self.flow_head(self.gru(torch.cat([fea, fea0, motion0, flow0], 1)))
|
||||
flow1 = flow1 + self.flow_head(self.gru(torch.cat([fea, fea1, motion1, flow1], 1)))
|
||||
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=0.25, mode="bilinear", align_corners=False) * 0.25
|
||||
flow = self.rf_block4(s_4[:bs], s_4[bs:], flow[:, :2], flow[:, 2:4])
|
||||
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 = self.rf_block3(s_3[:bs], s_3[bs:], flow[:, :2], flow[:, 2: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 = self.rf_block2(s_2[:bs], s_2[bs:], flow[:, :2], flow[:, 2:4])
|
||||
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
|
||||
flow = self.rf_block1(s_1[:bs], s_1[bs:], flow[:, :2], flow[:, 2:4])
|
||||
c0, c1 = [s_1[:bs], s_2[:bs], s_3[:bs], s_4[:bs]], [s_1[bs:], s_2[bs:], s_3[bs:], s_4[bs:]]
|
||||
return flow, self.warp_fea(c0, flow[:, :2]), self.warp_fea(c1, flow[:, 2:4])
|
||||
|
||||
def warp_fea(self, feas, flow):
|
||||
outs = []
|
||||
for i, fea in enumerate(feas):
|
||||
out = warp(fea, flow)
|
||||
outs.append(out)
|
||||
for fea in feas:
|
||||
outs.append(warp(fea, flow))
|
||||
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
|
||||
|
||||
|
||||
c0, c1 = [s_1[:bs], s_2[:bs], s_3[:bs], s_4[:bs]], [s_1[bs:], s_2[bs:], s_3[bs:], s_4[bs:]]
|
||||
return flow, self.warp_fea(c0, flow[:, :2]), self.warp_fea(c1, flow[:, 2:4])
|
||||
|
||||
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.conv1, self.conv2 = Conv2(3, c, 1), Conv2(c, 2 * c)
|
||||
self.conv3, self.conv4 = Conv2(2 * c, 4 * c), 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)
|
||||
@@ -312,337 +172,79 @@ class FlowRefineNet_Multis(nn.Module):
|
||||
|
||||
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
|
||||
inp = torch.cat([x0, x1], 0)
|
||||
s_1, s_2, s_3, s_4 = self.conv1(inp), self.conv2(self.conv1(inp)), self.conv3(self.conv2(self.conv1(inp))), self.conv4(self.conv3(self.conv2(self.conv1(inp))))
|
||||
c0, c1 = [s_1[:bs], s_2[:bs], s_3[:bs], s_4[:bs]], [s_1[bs:], s_2[bs:], s_3[bs:], s_4[bs:]]
|
||||
return flow, self.warp_fea(c0, flow[:, :2]), self.warp_fea(c1, flow[:, 2:4])
|
||||
|
||||
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
|
||||
inp = torch.cat([x0, x1], 0)
|
||||
s_1 = self.conv1(inp)
|
||||
s_2 = self.conv2(s_1)
|
||||
s_3 = self.conv3(s_2)
|
||||
s_4 = self.conv4(s_3)
|
||||
flow = F.interpolate(flow, scale_factor=0.25, mode="bilinear", align_corners=False) * 0.25
|
||||
flow = self.rf_block4(s_4[:bs], s_4[bs:], flow[:, :2], flow[:, 2:4])
|
||||
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 = self.rf_block3(s_3[:bs], s_3[bs:], flow[:, :2], flow[:, 2: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 = self.rf_block2(s_2[:bs], s_2[bs:], flow[:, :2], flow[:, 2:4])
|
||||
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
|
||||
flow = self.rf_block1(s_1[:bs], s_1[bs:], flow[:, :2], flow[:, 2:4])
|
||||
c0, c1 = [s_1[:bs], s_2[:bs], s_3[:bs], s_4[:bs]], [s_1[bs:], s_2[bs:], s_3[bs:], s_4[bs:]]
|
||||
return flow, self.warp_fea(c0, flow[:, :2]), self.warp_fea(c1, flow[:, 2:4])
|
||||
|
||||
def warp_fea(self, feas, flow):
|
||||
outs = []
|
||||
for i, fea in enumerate(feas):
|
||||
out = warp(fea, flow)
|
||||
outs.append(out)
|
||||
for fea in feas:
|
||||
outs.append(warp(fea, flow))
|
||||
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
|
||||
|
||||
c, height, width, window_size, embed_dim = 24, 192, 192, 8, 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.fuse_block = nn.Sequential(nn.Conv2d(12, 2*c, 3, 1, 1), nn.LeakyReLU(0.2, True),
|
||||
nn.Conv2d(2*c, 2*c, 3, 1, 1), nn.LeakyReLU(0.2, 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]], \
|
||||
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)
|
||||
if 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)
|
||||
|
||||
flow, _ = self.flownet(imgs)
|
||||
flow, _, _ = 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, _ = 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))
|
||||
|
||||
warped_img0, warped_img1 = warp(img0, flow[:, :2]), warp(img1, flow[:, 2:])
|
||||
x = self.fuse_block(torch.cat([img0, img1, warped_img0, warped_img1], 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())
|
||||
|
||||
res, mask = torch.sigmoid(refine_output[:, :3]) * 2 - 1, torch.sigmoid(refine_output[:, 3:4])
|
||||
return torch.clamp(warped_img0 * mask + warped_img1 * (1 - mask) + res, 0, 1), flow
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,25 +1,24 @@
|
||||
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()))
|
||||
k = (str(tenFlow.device), str(tenFlow.size()), str(tenInput.dtype))
|
||||
if k not in backwarp_tenGrid:
|
||||
tenHorizontal = torch.linspace(-1.0, 1.0, tenFlow.shape[3], device=device).view(
|
||||
tenHorizontal = torch.linspace(-1.0, 1.0, tenFlow.shape[3], device=tenFlow.device, dtype=tenInput.dtype).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(
|
||||
tenVertical = torch.linspace(-1.0, 1.0, tenFlow.shape[2], device=tenFlow.device, dtype=tenInput.dtype).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)
|
||||
[tenHorizontal, tenVertical], 1)
|
||||
|
||||
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)
|
||||
g = (backwarp_tenGrid[k].float() + tenFlow.float()).permute(0, 2, 3, 1)
|
||||
return torch.nn.functional.grid_sample(input=tenInput.float(), grid=g, mode='bilinear', padding_mode='border', align_corners=True).type(tenInput.dtype)
|
||||
|
||||
|
||||
def flow_reversal(flow):
|
||||
|
||||
@@ -1,722 +0,0 @@
|
||||
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(r'(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(r'(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(r'(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));
|
||||
@@ -1,5 +1,3 @@
|
||||
import pdb
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
@@ -7,16 +5,18 @@ 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
|
||||
|
||||
# Helper for extracting schedule values
|
||||
def extract(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||
|
||||
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
|
||||
@@ -27,13 +27,8 @@ class BrownianBridgeModel(nn.Module):
|
||||
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
|
||||
@@ -52,12 +47,13 @@ class BrownianBridgeModel(nn.Module):
|
||||
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
|
||||
|
||||
m_tminus = np.append(0, m_t[:-1])
|
||||
variance_t = 2. * (m_t - m_t ** 2) * self.max_var
|
||||
variance_tminus = np.append(0., variance_t[:-1])
|
||||
variance_t_tminus = variance_t - variance_tminus * ((1. - m_t) / (1. - m_tminus)) ** 2
|
||||
posterior_variance_t = variance_t_tminus * variance_tminus / variance_t
|
||||
|
||||
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))
|
||||
@@ -78,196 +74,22 @@ class BrownianBridgeModel(nn.Module):
|
||||
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)
|
||||
sigma_t = torch.sqrt(var_t.clamp(min=0))
|
||||
x0_recon = (x_t - m_t * y - sigma_t * objective_recon) / (1. - m_t).clamp(min=1e-12)
|
||||
elif self.objective == 'ysubx':
|
||||
x0_recon = y - objective_recon
|
||||
|
||||
elif self.objective == 'BB':
|
||||
x0_recon = -objective_recon + x_t ## if predicting xt - x0
|
||||
x0_recon = -objective_recon + x_t
|
||||
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
|
||||
@@ -287,46 +109,37 @@ class BrownianBridgeModel(nn.Module):
|
||||
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
|
||||
m_t = extract(self.m_t, t, x_t.shape).float()
|
||||
m_nt = extract(self.m_t, n_t, x_t.shape).float()
|
||||
var_t = extract(self.variance_t, t, x_t.shape).float()
|
||||
var_nt = extract(self.variance_t, n_t, x_t.shape).float()
|
||||
|
||||
eps = 1e-12
|
||||
sigma2_t = (var_t - var_nt * (1. - m_t) ** 2 / (1. - m_nt).clamp(min=eps) ** 2) * var_nt / var_t.clamp(min=eps)
|
||||
sigma_t = torch.sqrt(sigma2_t.clamp(min=0)) * 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
|
||||
noise = torch.randn_like(x_t).float()
|
||||
x_t_f = x_t.float()
|
||||
x0_recon_f = x0_recon.float()
|
||||
y_f = y.float()
|
||||
|
||||
x_tminus_mean = (1. - m_nt) * x0_recon_f + m_nt * y_f + torch.sqrt(((var_nt - sigma2_t) / var_t.clamp(min=eps)).clamp(min=0)) * \
|
||||
(x_t_f - (1. - m_t) * x0_recon_f - m_t * y_f)
|
||||
|
||||
return (x_tminus_mean + sigma_t * noise).type(x_t.dtype), x0_recon
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample_loop(self, y, context=None, clip_denoised=True, sample_mid_step=False):
|
||||
def p_sample_loop(self, y, context=None, clip_denoised=True):
|
||||
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
|
||||
|
||||
|
||||
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)
|
||||
def sample(self, y, z, context_y=None, context_z=None, clip_denoised=True):
|
||||
return self.p_sample_loop(y, z, context_y, context_z, clip_denoised)
|
||||
|
||||
@@ -7,7 +7,6 @@ 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
|
||||
|
||||
|
||||
@@ -31,29 +30,14 @@ class LatentBrownianBridgeModel(BrownianBridgeModel):
|
||||
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()
|
||||
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())
|
||||
@@ -103,22 +87,28 @@ class LatentBrownianBridgeModel(BrownianBridgeModel):
|
||||
|
||||
@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))
|
||||
# VQGAN encoder works in float32
|
||||
y_f32 = y.float()
|
||||
z_f32 = z.float()
|
||||
x = torch.zeros_like(y_f32)
|
||||
latent,phi_list = self.encode(torch.cat([y_f32,x,z_f32], dim=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,
|
||||
|
||||
# Determine the computation dtype from the UNet
|
||||
compute_dtype = self.denoise_fn.time_embed[0].weight.dtype
|
||||
|
||||
# Pass disable_progress to the sampling loop, ensuring inputs match compute_dtype
|
||||
imgs,one_step_imgs = self.latent_p_sample_loop(latent = latent.to(compute_dtype),
|
||||
y = latent.to(compute_dtype),
|
||||
context = latent.to(compute_dtype),
|
||||
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)
|
||||
# VQGAN decoder works in float32
|
||||
out = self.decode(imgs[-1].detach().float(), y_f32, z_f32, phi_list, scale=scale)
|
||||
return out
|
||||
|
||||
@torch.no_grad()
|
||||
|
||||
@@ -74,8 +74,10 @@ def zero_module(module):
|
||||
return module
|
||||
|
||||
|
||||
from model.BrownianBridge.base.modules.diffusionmodules.util import GroupNorm32
|
||||
|
||||
def Normalize(in_channels):
|
||||
return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
return GroupNorm32(32, in_channels, eps=1e-6, affine=True)
|
||||
|
||||
|
||||
class LinearAttention(nn.Module):
|
||||
@@ -90,7 +92,7 @@ class LinearAttention(nn.Module):
|
||||
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)
|
||||
k = k.float().softmax(dim=-1).type(k.dtype)
|
||||
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)
|
||||
@@ -138,7 +140,7 @@ class SpatialSelfAttention(nn.Module):
|
||||
w_ = torch.einsum('bij,bjk->bik', q, k)
|
||||
|
||||
w_ = w_ * (int(c)**(-0.5))
|
||||
w_ = torch.nn.functional.softmax(w_, dim=2)
|
||||
w_ = torch.nn.functional.softmax(w_.float(), dim=2).type(w_.dtype)
|
||||
|
||||
# attend to values
|
||||
v = rearrange(v, 'b c h w -> b c (h w)')
|
||||
@@ -189,7 +191,7 @@ class CrossAttention(nn.Module):
|
||||
sim.masked_fill_(~mask, max_neg_value)
|
||||
|
||||
# attention, what we cannot get enough of
|
||||
attn = sim.softmax(dim=-1)
|
||||
attn = sim.float().softmax(dim=-1).type(sim.dtype)
|
||||
|
||||
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)
|
||||
@@ -337,7 +339,7 @@ class SpatialCrossAttentionWithPosEmb(nn.Module):
|
||||
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)
|
||||
attn = sim.float().softmax(dim=-1).type(sim.dtype)
|
||||
|
||||
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)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,10 +1,5 @@
|
||||
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
|
||||
@@ -21,45 +16,6 @@ from model.BrownianBridge.base.modules.diffusionmodules.util import (
|
||||
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.
|
||||
@@ -79,7 +35,6 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
|
||||
"""
|
||||
|
||||
def forward(self, x, emb, context=None):
|
||||
# pdb.set_trace()
|
||||
for layer in self:
|
||||
if isinstance(layer, TimestepBlock):
|
||||
x = layer(x, emb)
|
||||
@@ -93,10 +48,6 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
|
||||
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):
|
||||
@@ -121,26 +72,9 @@ class Upsample(nn.Module):
|
||||
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):
|
||||
@@ -166,17 +100,6 @@ class Downsample(nn.Module):
|
||||
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__(
|
||||
@@ -244,12 +167,6 @@ class ResBlock(TimestepBlock):
|
||||
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
|
||||
)
|
||||
@@ -281,8 +198,6 @@ class ResBlock(TimestepBlock):
|
||||
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__(
|
||||
@@ -306,17 +221,14 @@ class AttentionBlock(nn.Module):
|
||||
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
|
||||
return checkpoint(self._forward, (x,), self.parameters(), self.use_checkpoint)
|
||||
|
||||
def _forward(self, x):
|
||||
b, c, *spatial = x.shape
|
||||
@@ -327,41 +239,12 @@ class AttentionBlock(nn.Module):
|
||||
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)
|
||||
@@ -369,31 +252,18 @@ class QKVAttentionLegacy(nn.Module):
|
||||
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)
|
||||
@@ -403,44 +273,15 @@ class QKVAttention(nn.Module):
|
||||
"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__(
|
||||
@@ -466,10 +307,10 @@ class UNetModel(nn.Module):
|
||||
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
|
||||
use_spatial_transformer=False,
|
||||
transformer_depth=1,
|
||||
context_dim=None,
|
||||
n_embed=None,
|
||||
legacy=True,
|
||||
condition_key="concat",
|
||||
):
|
||||
@@ -479,9 +320,6 @@ class UNetModel(nn.Module):
|
||||
|
||||
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
|
||||
@@ -553,7 +391,6 @@ class UNetModel(nn.Module):
|
||||
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(
|
||||
@@ -604,7 +441,6 @@ class UNetModel(nn.Module):
|
||||
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(
|
||||
@@ -662,7 +498,6 @@ class UNetModel(nn.Module):
|
||||
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(
|
||||
@@ -708,46 +543,15 @@ class UNetModel(nn.Module):
|
||||
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)
|
||||
t_emb = t_emb.type(self.time_embed[0].weight.dtype)
|
||||
emb = self.time_embed(t_emb)
|
||||
|
||||
if self.num_classes is not None:
|
||||
@@ -755,9 +559,10 @@ class UNetModel(nn.Module):
|
||||
emb = emb + self.label_emb(y)
|
||||
|
||||
if self.condition_key != 'nocond':
|
||||
if context is not None:
|
||||
context = context.type(self.time_embed[0].weight.dtype)
|
||||
x = th.cat([x, context], dim=1)
|
||||
#x = th.cat([x,cond],dim = 1) ## cat with previous path
|
||||
h = x.type(self.dtype)
|
||||
h = x.type(self.time_embed[0].weight.dtype)
|
||||
for module in self.input_blocks:
|
||||
h = module(h, emb, context)
|
||||
hs.append(h)
|
||||
@@ -773,222 +578,3 @@ class UNetModel(nn.Module):
|
||||
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)
|
||||
|
||||
|
||||
@@ -4,95 +4,15 @@
|
||||
# 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 torch.nn.functional as F
|
||||
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)
|
||||
@@ -103,11 +23,6 @@ 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)
|
||||
@@ -205,15 +120,11 @@ def normalization(channels):
|
||||
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)
|
||||
weight = self.weight.float() if self.weight is not None else None
|
||||
bias = self.bias.float() if self.bias is not None else None
|
||||
return F.group_norm(x.float(), self.num_groups, weight, bias, self.eps).type(x.dtype)
|
||||
|
||||
def conv_nd(dims, *args, **kwargs):
|
||||
"""
|
||||
@@ -248,20 +159,7 @@ def avg_pool_nd(dims, *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()
|
||||
return repeat_noise() if repeat else noise()
|
||||
|
||||
@@ -1,76 +0,0 @@
|
||||
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)
|
||||
@@ -1,134 +0,0 @@
|
||||
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)
|
||||
@@ -82,11 +82,7 @@ class Attention(nn.Module):
|
||||
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.dropout = dropout
|
||||
self.to_out = nn.Sequential(
|
||||
nn.Linear(dim, dim, bias = False),
|
||||
nn.Dropout(dropout)
|
||||
@@ -139,7 +135,9 @@ class Attention(nn.Module):
|
||||
|
||||
# attention
|
||||
|
||||
attn = self.attend(sim)
|
||||
attn = torch.nn.functional.softmax(sim.float(), dim=-1).type(sim.dtype)
|
||||
if self.dropout > 0:
|
||||
attn = torch.nn.functional.dropout(attn, p=self.dropout, training=self.training)
|
||||
|
||||
# aggregate
|
||||
|
||||
|
||||
@@ -1,641 +0,0 @@
|
||||
"""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
|
||||
|
||||
@@ -1,203 +0,0 @@
|
||||
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
|
||||
+43
-585
@@ -1,50 +1,21 @@
|
||||
# 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
|
||||
|
||||
import torch.nn.functional as F
|
||||
from model.BrownianBridge.base.modules.diffusionmodules.util import GroupNorm32
|
||||
|
||||
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)
|
||||
|
||||
return GroupNorm32(32, 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)
|
||||
self.conv = torch.nn.Conv2d(in_channels, in_channels, 3, 1, 1)
|
||||
|
||||
def forward(self, x):
|
||||
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
|
||||
@@ -52,29 +23,21 @@ class Upsample(nn.Module):
|
||||
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)
|
||||
self.conv = torch.nn.Conv2d(in_channels, in_channels, 3, 2, 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 = torch.nn.functional.pad(x, (0,1,0,1), 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):
|
||||
@@ -85,43 +48,26 @@ class ResnetBlock(nn.Module):
|
||||
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.conv1 = torch.nn.Conv2d(in_channels, out_channels, 3, 1, 1)
|
||||
|
||||
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)
|
||||
self.conv2 = torch.nn.Conv2d(out_channels, out_channels, 3, 1, 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)
|
||||
self.conv_shortcut = torch.nn.Conv2d(in_channels, out_channels, 3, 1, 1)
|
||||
else:
|
||||
self.nin_shortcut = torch.nn.Conv2d(in_channels,
|
||||
out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
self.nin_shortcut = torch.nn.Conv2d(in_channels, out_channels, 1, 1, 0)
|
||||
|
||||
def forward(self, x, temb):
|
||||
h = x
|
||||
h = self.norm1(h)
|
||||
h = self.norm1(x)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv1(h)
|
||||
|
||||
if temb is not None:
|
||||
h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None]
|
||||
# VQGAN usually doesn't use temb in this node's inference path
|
||||
pass
|
||||
|
||||
h = self.norm2(h)
|
||||
h = nonlinearity(h)
|
||||
@@ -136,241 +82,53 @@ class ResnetBlock(nn.Module):
|
||||
|
||||
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)
|
||||
|
||||
self.q = torch.nn.Conv2d(in_channels, in_channels, 1)
|
||||
self.k = torch.nn.Conv2d(in_channels, in_channels, 1)
|
||||
self.v = torch.nn.Conv2d(in_channels, in_channels, 1)
|
||||
self.proj_out = torch.nn.Conv2d(in_channels, in_channels, 1)
|
||||
|
||||
def forward(self, x):
|
||||
h_ = x
|
||||
h_ = self.norm(h_)
|
||||
h_ = self.norm(x)
|
||||
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)
|
||||
q = q.reshape(b,c,h*w).permute(0,2,1)
|
||||
k = k.reshape(b,c,h*w)
|
||||
w_ = torch.bmm(q.float(), k.float()) * (int(c)**(-0.5))
|
||||
w_ = torch.nn.functional.softmax(w_, dim=2).type(x.dtype)
|
||||
|
||||
# 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_ = torch.bmm(v, w_.permute(0,2,1)).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
|
||||
self.conv_in = torch.nn.Conv2d(in_channels, self.ch, 3, 1, 1)
|
||||
|
||||
# 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()
|
||||
curr_res = resolution
|
||||
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.append(ResnetBlock(in_channels=block_in, out_channels=block_out, temb_channels=0, dropout=dropout))
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(AttnBlock(block_in))
|
||||
@@ -382,108 +140,63 @@ class Encoder(nn.Module):
|
||||
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.block_1 = ResnetBlock(in_channels=block_in, out_channels=block_in, temb_channels=0, 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)
|
||||
self.mid.block_2 = ResnetBlock(in_channels=block_in, out_channels=block_in, temb_channels=0, 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)
|
||||
|
||||
self.conv_out = torch.nn.Conv2d(block_in, 2*z_channels if double_z else z_channels, 3, 1, 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)
|
||||
h = self.down[i_level].block[i_block](hs[-1], None)
|
||||
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.block_1(h, None)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h, temb)
|
||||
h = self.mid.block_2(h, None)
|
||||
|
||||
# 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)
|
||||
self.conv_in = torch.nn.Conv2d(z_channels, block_in, 3, 1, 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.block_1 = ResnetBlock(in_channels=block_in, out_channels=block_in, temb_channels=0, 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)
|
||||
self.mid.block_2 = ResnetBlock(in_channels=block_in, out_channels=block_in, temb_channels=0, 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.append(ResnetBlock(in_channels=block_in, out_channels=block_out, temb_channels=0, dropout=dropout))
|
||||
block_in = block_out
|
||||
if curr_res in attn_resolutions:
|
||||
attn.append(AttnBlock(block_in))
|
||||
@@ -493,41 +206,25 @@ class Decoder(nn.Module):
|
||||
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
|
||||
self.up.insert(0, up)
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = torch.nn.Conv2d(block_in,
|
||||
out_ch,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
self.conv_out = torch.nn.Conv2d(block_in, out_ch, 3, 1, 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.block_1(h, None)
|
||||
h = self.mid.attn_1(h)
|
||||
h = self.mid.block_2(h, temb)
|
||||
h = self.mid.block_2(h, None)
|
||||
|
||||
# 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)
|
||||
h = self.up[i_level].block[i_block](h, None)
|
||||
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
|
||||
|
||||
@@ -535,242 +232,3 @@ class Decoder(nn.Module):
|
||||
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
|
||||
|
||||
|
||||
@@ -45,10 +45,13 @@ class VectorQuantizer(nn.Module):
|
||||
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
|
||||
z_f = z_flattened.float()
|
||||
emb_f = self.embedding.weight.float()
|
||||
|
||||
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())
|
||||
d = torch.sum(z_f ** 2, dim=1, keepdim=True) + \
|
||||
torch.sum(emb_f**2, dim=1) - 2 * \
|
||||
torch.matmul(z_f, emb_f.t())
|
||||
d = d.type(z_flattened.dtype)
|
||||
|
||||
## could possible replace this here
|
||||
# #\start...
|
||||
@@ -187,7 +190,7 @@ class GumbelQuantize(nn.Module):
|
||||
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)
|
||||
qy = F.softmax(logits.float(), dim=1).type(logits.dtype)
|
||||
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)
|
||||
@@ -276,10 +279,13 @@ class VectorQuantizer2(nn.Module):
|
||||
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
|
||||
z_f = z_flattened.float()
|
||||
emb_f = self.embedding.weight.float()
|
||||
|
||||
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'))
|
||||
d = torch.sum(z_f ** 2, dim=1, keepdim=True) + \
|
||||
torch.sum(emb_f**2, dim=1) - 2 * \
|
||||
torch.einsum('bd,dn->bn', z_f, rearrange(emb_f, 'n d -> d n'))
|
||||
d = d.type(z_flattened.dtype)
|
||||
|
||||
min_encoding_indices = torch.argmin(d, dim=1)
|
||||
z_q = self.embedding(min_encoding_indices).view(z.shape)
|
||||
|
||||
+6
-215
@@ -1,55 +1,17 @@
|
||||
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
|
||||
from model.utils import instantiate_from_config, dict2namespace
|
||||
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):
|
||||
class VQFlowNet(torch.nn.Module):
|
||||
def __init__(self,
|
||||
ddconfig,
|
||||
lossconfig,
|
||||
@@ -65,7 +27,6 @@ class VQFlowNet(pl.LightningModule):
|
||||
lr_g_factor=1.0,
|
||||
remap=None,
|
||||
sane_index_shape=False, # tell vector quantizer to return indices as bhw
|
||||
use_ema=False
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -95,11 +56,6 @@ class VQFlowNet(pl.LightningModule):
|
||||
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
|
||||
@@ -111,21 +67,6 @@ class VQFlowNet(pl.LightningModule):
|
||||
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())
|
||||
@@ -140,10 +81,6 @@ class VQFlowNet(pl.LightningModule):
|
||||
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
|
||||
@@ -231,154 +168,9 @@ class VQFlowNet(pl.LightningModule):
|
||||
|
||||
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"):
|
||||
@@ -462,14 +254,13 @@ class VQFlowNetInterface(VQFlowNet):
|
||||
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 = torch.zeros(b,c,h,w, device=img0_down_.device, dtype=img0_down_.dtype)
|
||||
img1_down = torch.zeros(b,c,h,w, device=img1_down_.device, dtype=img1_down_.dtype)
|
||||
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]))
|
||||
# Ensure zeros_like matches dtype/device
|
||||
_,tmp_list = self.encoder(torch.cat([img0_down,torch.zeros_like(img0_down),img1_down], dim=0))
|
||||
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:
|
||||
|
||||
+66
-11
@@ -1,17 +1,72 @@
|
||||
from inspect import isfunction
|
||||
import importlib
|
||||
import argparse
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import collections.abc
|
||||
from itertools import repeat
|
||||
|
||||
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 extract(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||
def instantiate_from_config(config):
|
||||
if not isinstance(config, dict) or 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 dict2namespace(config):
|
||||
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 exists(x):
|
||||
return x is not None
|
||||
# --- Native Replacements for timm utilities ---
|
||||
|
||||
def to_2tuple(x):
|
||||
if isinstance(x, collections.abc.Iterable):
|
||||
return x
|
||||
return tuple(repeat(x, 2))
|
||||
|
||||
def default(val, d):
|
||||
if exists(val):
|
||||
return val
|
||||
return d() if isfunction(d) else d
|
||||
def trunc_normal_(tensor, mean=0., std=1., a=-2., b=2.):
|
||||
# Simple implementation of truncated normal
|
||||
with torch.no_grad():
|
||||
size = tensor.shape
|
||||
tmp = tensor.new_empty(size + (4,)).normal_()
|
||||
valid = (tmp < b) & (tmp > a)
|
||||
ind = valid.max(-1, keepdim=True)[1]
|
||||
tensor.data.copy_(tmp.gather(-1, ind).squeeze(-1))
|
||||
tensor.data.mul_(std).add_(mean)
|
||||
return tensor
|
||||
|
||||
def drop_path(x, drop_prob: float = 0., training: bool = False, scale_by_keep: bool = True):
|
||||
if drop_prob == 0. or not training:
|
||||
return x
|
||||
keep_prob = 1 - drop_prob
|
||||
shape = (x.shape[0],) + (1,) * (x.ndim - 1)
|
||||
random_tensor = x.new_empty(shape).bernoulli_(keep_prob)
|
||||
if keep_prob > 0.0 and scale_by_keep:
|
||||
random_tensor.div_(keep_prob)
|
||||
return x * random_tensor
|
||||
|
||||
class DropPath(nn.Module):
|
||||
def __init__(self, drop_prob: float = 0., scale_by_keep: bool = True):
|
||||
super(DropPath, self).__init__()
|
||||
self.drop_prob = drop_prob
|
||||
self.scale_by_keep = scale_by_keep
|
||||
|
||||
def forward(self, x):
|
||||
return drop_path(x, self.drop_prob, self.training, self.scale_by_keep)
|
||||
|
||||
def extra_repr(self):
|
||||
return f'drop_prob={round(self.drop_prob,3):0.3f}'
|
||||
|
||||
+2
-10
@@ -1,11 +1,3 @@
|
||||
from .tlbvfi_node import TLBVFI_VFI
|
||||
from .tlbvfi_node import comfy_entrypoint
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"TLBVFI_VFI": TLBVFI_VFI
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"TLBVFI_VFI": "TLBVFI Frame Interpolation"
|
||||
}
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
__all__ = ['comfy_entrypoint']
|
||||
|
||||
+2
-3
@@ -1,13 +1,12 @@
|
||||
[project]
|
||||
name = "TLBVFI"
|
||||
description = "wrapper for the TLB-VFI: Temporal-Aware Latent Brownian Bridge Diffusion for Video Frame Interpolation project"
|
||||
version = "0.1.4"
|
||||
version = "0.2.0"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = ["pytorch_lightning", "cupy-cuda11x"]
|
||||
dependencies = []
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/BobRandomNumber/ComfyUI-TLBVFI"
|
||||
# Used by Comfy Registry https://registry.comfy.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "bobrandomnumber"
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
pytorch-lightning
|
||||
cupy-cuda12x
|
||||
+92
-118
@@ -3,37 +3,21 @@ 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
|
||||
from comfy_api.latest import io, ComfyExtension
|
||||
|
||||
# --- Robust Path Handling ---
|
||||
# Setup models directory for frame interpolation
|
||||
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
|
||||
_CURRENT_MODEL = None
|
||||
_CURRENT_MODEL_KEY = None
|
||||
|
||||
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:
|
||||
@@ -42,118 +26,108 @@ def find_models(folder_type: str, extensions: list) -> list:
|
||||
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:
|
||||
class TLBVFI_VFI(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
# We only need the main model file now.
|
||||
def define_schema(cls) -> io.Schema:
|
||||
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 io.Schema(
|
||||
node_id="TLBVFI_VFI",
|
||||
display_name="TLBVFI Frame Interpolation",
|
||||
category="frame_interpolation/TLBVFI",
|
||||
description="Temporal-Aware Latent Brownian Bridge for Video Frame Interpolation",
|
||||
inputs=[
|
||||
io.Image.Input("images"),
|
||||
io.Combo.Input("model_name", options=unet_models if unet_models else ["No models found"]),
|
||||
io.Int.Input("times_to_interpolate", default=1, min=1, max=4, step=1),
|
||||
io.Int.Input("diffusion_steps", default=10, min=1, max=100, step=1),
|
||||
io.Int.Input("batch_size", default=2, min=1, max=64),
|
||||
io.Float.Input("flow_scale", default=0.5, min=0.1, max=1.0, step=0.1),
|
||||
],
|
||||
outputs=[io.Image.Output()]
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "interpolate"
|
||||
CATEGORY = "frame_interpolation/TLBVFI" # Updated Category for better organization
|
||||
@classmethod
|
||||
def execute(cls, images, model_name, times_to_interpolate, diffusion_steps, batch_size, flow_scale) -> io.NodeOutput:
|
||||
from comfy.utils import ProgressBar
|
||||
from tqdm import tqdm
|
||||
import gc
|
||||
|
||||
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)
|
||||
if model_name == "No models found":
|
||||
raise Exception("No TLBVFI UNet models found. Please download 'vimeo_unet.pth' to models/interpolation.")
|
||||
|
||||
# --- 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.")
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
current_path = Path(__file__).parent
|
||||
tlbvfi_path = current_path / "TLBVFI"
|
||||
if str(tlbvfi_path) not in sys.path:
|
||||
sys.path.insert(0, str(tlbvfi_path))
|
||||
|
||||
from model.BrownianBridge.LatentBrownianBridgeModel import LatentBrownianBridgeModel
|
||||
from model.utils import dict2namespace
|
||||
|
||||
global _CURRENT_MODEL, _CURRENT_MODEL_KEY
|
||||
cache_key = (model_name, diffusion_steps)
|
||||
|
||||
if _CURRENT_MODEL_KEY == cache_key and _CURRENT_MODEL is not None:
|
||||
model = _CURRENT_MODEL
|
||||
else:
|
||||
if _CURRENT_MODEL is not None:
|
||||
_CURRENT_MODEL = None
|
||||
gc.collect()
|
||||
if torch.cuda.is_available(): torch.cuda.empty_cache()
|
||||
|
||||
model_path = folder_paths.get_full_path("interpolation", model_name)
|
||||
if not model_path: raise FileNotFoundError(f"Model file {model_name} not found.")
|
||||
|
||||
config_path = tlbvfi_path / "configs" / "Template-LBBDM-video.yaml"
|
||||
with open(config_path, 'r') as f:
|
||||
config = yaml.load(f, Loader=yaml.FullLoader)
|
||||
nconfig = dict2namespace(config)
|
||||
nconfig.model.VQGAN.params.ckpt_path = None
|
||||
nconfig.model.BB.params.sample_step = diffusion_steps
|
||||
|
||||
# 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)
|
||||
model = LatentBrownianBridgeModel(nconfig.model).to(device)
|
||||
checkpoint = torch.load(model_path, map_location=device)
|
||||
model.load_state_dict(checkpoint.get('model', checkpoint))
|
||||
model.float().eval()
|
||||
_CURRENT_MODEL, _CURRENT_MODEL_KEY = model, cache_key
|
||||
|
||||
# 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, )
|
||||
return io.NodeOutput(images)
|
||||
|
||||
gui_pbar = ProgressBar(len(image_tensors) - 1)
|
||||
num_pairs = len(image_tensors) - 1
|
||||
gui_pbar = ProgressBar(num_pairs)
|
||||
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)
|
||||
with torch.no_grad():
|
||||
for i in tqdm(range(0, num_pairs, batch_size), desc="TLBVFI Interpolating"):
|
||||
current_batch_size = min(batch_size, num_pairs - i)
|
||||
f1_batch = image_tensors[i : i + current_batch_size].to(device)
|
||||
f2_batch = image_tensors[i + 1 : i + 1 + current_batch_size].to(device)
|
||||
|
||||
current_frames = [f1_batch, f2_batch]
|
||||
for _ in range(times_to_interpolate):
|
||||
temp_frames = [current_frames[0]]
|
||||
for j in range(len(current_frames) - 1):
|
||||
mid_frame = model.sample(current_frames[j], current_frames[j+1], scale=flow_scale, disable_progress=True)
|
||||
mid_frame = torch.nan_to_num(mid_frame, nan=0.0, posinf=1.0, neginf=-1.0).cpu()
|
||||
temp_frames.extend([mid_frame, current_frames[j+1].cpu()])
|
||||
current_frames = temp_frames
|
||||
|
||||
for b in range(current_batch_size):
|
||||
for k in range(1, len(current_frames)):
|
||||
output_frames.append(current_frames[k][b:b+1])
|
||||
gui_pbar.update(current_batch_size)
|
||||
|
||||
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, )
|
||||
return io.NodeOutput(final_tensors.clamp(0, 1).permute(0, 2, 3, 1))
|
||||
|
||||
class TLBVFIExtension(ComfyExtension):
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [TLBVFI_VFI]
|
||||
|
||||
async def comfy_entrypoint() -> TLBVFIExtension:
|
||||
return TLBVFIExtension()
|
||||
|
||||
Reference in New Issue
Block a user