v3 schema and code refinments

This commit is contained in:
BobRandomNumber
2026-02-17 21:21:56 -05:00
committed by GitHub
parent aaedb24bcc
commit ba9f941b12
30 changed files with 688 additions and 5932 deletions
+24 -70
View File
@@ -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},
}
```
+77 -475
View File
@@ -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
+6 -7
View File
@@ -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):
View File
-722
View File
@@ -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
-203
View File
@@ -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
View File
+43 -585
View File
@@ -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
+13 -7
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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"
-2
View File
@@ -1,2 +0,0 @@
pytorch-lightning
cupy-cuda12x
+92 -118
View File
@@ -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()