Commit ·
c8c00f0
0
Parent(s):
Clean commit without binary (image) files
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +35 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/autoencoder.py +584 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/__init__.py +1 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/__pycache__/__init__.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/__pycache__/clip.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/__pycache__/model.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/__pycache__/simple_tokenizer.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/bpe_simple_vocab_16e6.txt.gz +3 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/clip.py +237 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/model.py +434 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/simple_tokenizer.py +132 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/__init__.py +46 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/__pycache__/__init__.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/__pycache__/diffusion_utils.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/__pycache__/gaussian_diffusion.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/__pycache__/respace.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/diffusion_utils.py +88 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/gaussian_diffusion.py +903 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/respace.py +129 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/timestep_sampler.py +150 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/models_add_cross_concate.py +498 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/test.py +198 -0
- ArtiAgent - DefectDiffu/engine/DefectDiffu/train.py +231 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/LICENSE +201 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/README.md +163 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/__init__.py +0 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/__pycache__/__init__.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/config/GroundingDINO_SwinB.py +43 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/config/GroundingDINO_SwinT_OGC.py +43 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/datasets/__init__.py +0 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/datasets/__pycache__/__init__.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/datasets/__pycache__/transforms.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/datasets/transforms.py +311 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/__init__.py +15 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/__pycache__/__init__.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/__pycache__/groundingdino.cpython-310.pyc +0 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/backbone/__init__.py +1 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/backbone/backbone.py +221 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/backbone/position_encoding.py +186 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/backbone/swin_transformer.py +802 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/bertwarper.py +273 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_attn.h +64 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_attn_cpu.cpp +43 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_attn_cpu.h +35 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_attn_cuda.cu +156 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_attn_cuda.h +33 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_im2col_cuda.cuh +1327 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/cuda_version.cu +7 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/vision.cpp +58 -0
- ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/fuse_modules.py +297 -0
.gitattributes
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
*.7z filter=lfs diff=lfs merge=lfs -text
|
| 2 |
+
*.arrow filter=lfs diff=lfs merge=lfs -text
|
| 3 |
+
*.bin filter=lfs diff=lfs merge=lfs -text
|
| 4 |
+
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
| 5 |
+
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
| 6 |
+
*.ftz filter=lfs diff=lfs merge=lfs -text
|
| 7 |
+
*.gz filter=lfs diff=lfs merge=lfs -text
|
| 8 |
+
*.h5 filter=lfs diff=lfs merge=lfs -text
|
| 9 |
+
*.joblib filter=lfs diff=lfs merge=lfs -text
|
| 10 |
+
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
| 11 |
+
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
| 12 |
+
*.model filter=lfs diff=lfs merge=lfs -text
|
| 13 |
+
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
| 14 |
+
*.npy filter=lfs diff=lfs merge=lfs -text
|
| 15 |
+
*.npz filter=lfs diff=lfs merge=lfs -text
|
| 16 |
+
*.onnx filter=lfs diff=lfs merge=lfs -text
|
| 17 |
+
*.ot filter=lfs diff=lfs merge=lfs -text
|
| 18 |
+
*.parquet filter=lfs diff=lfs merge=lfs -text
|
| 19 |
+
*.pb filter=lfs diff=lfs merge=lfs -text
|
| 20 |
+
*.pickle filter=lfs diff=lfs merge=lfs -text
|
| 21 |
+
*.pkl filter=lfs diff=lfs merge=lfs -text
|
| 22 |
+
*.pt filter=lfs diff=lfs merge=lfs -text
|
| 23 |
+
*.pth filter=lfs diff=lfs merge=lfs -text
|
| 24 |
+
*.rar filter=lfs diff=lfs merge=lfs -text
|
| 25 |
+
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
| 26 |
+
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
| 27 |
+
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
| 28 |
+
*.tar filter=lfs diff=lfs merge=lfs -text
|
| 29 |
+
*.tflite filter=lfs diff=lfs merge=lfs -text
|
| 30 |
+
*.tgz filter=lfs diff=lfs merge=lfs -text
|
| 31 |
+
*.wasm filter=lfs diff=lfs merge=lfs -text
|
| 32 |
+
*.xz filter=lfs diff=lfs merge=lfs -text
|
| 33 |
+
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
+
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
+
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/autoencoder.py
ADDED
|
@@ -0,0 +1,584 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import numpy as np
|
| 4 |
+
from einops import rearrange
|
| 5 |
+
import os
|
| 6 |
+
import torchvision.transforms as transforms
|
| 7 |
+
from torchvision.utils import save_image
|
| 8 |
+
import os
|
| 9 |
+
from PIL import Image
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
os.environ["CUDA_VISIBLE_DEVICES"] = "1"
|
| 13 |
+
class LinearAttention(nn.Module):
|
| 14 |
+
def __init__(self, dim, heads=4, dim_head=32):
|
| 15 |
+
super().__init__()
|
| 16 |
+
self.heads = heads
|
| 17 |
+
hidden_dim = dim_head * heads
|
| 18 |
+
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
|
| 19 |
+
self.to_out = nn.Conv2d(hidden_dim, dim, 1)
|
| 20 |
+
|
| 21 |
+
def forward(self, x):
|
| 22 |
+
b, c, h, w = x.shape
|
| 23 |
+
qkv = self.to_qkv(x)
|
| 24 |
+
q, k, v = rearrange(qkv, 'b (qkv heads c) h w -> qkv b heads c (h w)', heads = self.heads, qkv=3)
|
| 25 |
+
k = k.softmax(dim=-1)
|
| 26 |
+
context = torch.einsum('bhdn,bhen->bhde', k, v)
|
| 27 |
+
out = torch.einsum('bhde,bhdn->bhen', context, q)
|
| 28 |
+
out = rearrange(out, 'b heads c (h w) -> b (heads c) h w', heads=self.heads, h=h, w=w)
|
| 29 |
+
return self.to_out(out)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def nonlinearity(x):
|
| 33 |
+
# swish
|
| 34 |
+
return x*torch.sigmoid(x)
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def Normalize(in_channels, num_groups=32):
|
| 38 |
+
return torch.nn.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
class Upsample(nn.Module):
|
| 42 |
+
def __init__(self, in_channels, with_conv):
|
| 43 |
+
super().__init__()
|
| 44 |
+
self.with_conv = with_conv
|
| 45 |
+
if self.with_conv:
|
| 46 |
+
self.conv = torch.nn.Conv2d(in_channels,
|
| 47 |
+
in_channels,
|
| 48 |
+
kernel_size=3,
|
| 49 |
+
stride=1,
|
| 50 |
+
padding=1)
|
| 51 |
+
|
| 52 |
+
def forward(self, x):
|
| 53 |
+
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
|
| 54 |
+
if self.with_conv:
|
| 55 |
+
x = self.conv(x)
|
| 56 |
+
return x
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
class Downsample(nn.Module):
|
| 60 |
+
def __init__(self, in_channels, with_conv):
|
| 61 |
+
super().__init__()
|
| 62 |
+
self.with_conv = with_conv
|
| 63 |
+
if self.with_conv:
|
| 64 |
+
# no asymmetric padding in torch conv, must do it ourselves
|
| 65 |
+
self.conv = torch.nn.Conv2d(in_channels,
|
| 66 |
+
in_channels,
|
| 67 |
+
kernel_size=3,
|
| 68 |
+
stride=2,
|
| 69 |
+
padding=0)
|
| 70 |
+
|
| 71 |
+
def forward(self, x):
|
| 72 |
+
if self.with_conv:
|
| 73 |
+
pad = (0,1,0,1)
|
| 74 |
+
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
|
| 75 |
+
x = self.conv(x)
|
| 76 |
+
else:
|
| 77 |
+
x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2)
|
| 78 |
+
return x
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
class ResnetBlock(nn.Module):
|
| 82 |
+
def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False,
|
| 83 |
+
dropout, temb_channels=512):
|
| 84 |
+
super().__init__()
|
| 85 |
+
self.in_channels = in_channels
|
| 86 |
+
out_channels = in_channels if out_channels is None else out_channels
|
| 87 |
+
self.out_channels = out_channels
|
| 88 |
+
self.use_conv_shortcut = conv_shortcut
|
| 89 |
+
|
| 90 |
+
self.norm1 = Normalize(in_channels)
|
| 91 |
+
self.conv1 = torch.nn.Conv2d(in_channels,
|
| 92 |
+
out_channels,
|
| 93 |
+
kernel_size=3,
|
| 94 |
+
stride=1,
|
| 95 |
+
padding=1)
|
| 96 |
+
if temb_channels > 0:
|
| 97 |
+
self.temb_proj = torch.nn.Linear(temb_channels,
|
| 98 |
+
out_channels)
|
| 99 |
+
self.norm2 = Normalize(out_channels)
|
| 100 |
+
self.dropout = torch.nn.Dropout(dropout)
|
| 101 |
+
self.conv2 = torch.nn.Conv2d(out_channels,
|
| 102 |
+
out_channels,
|
| 103 |
+
kernel_size=3,
|
| 104 |
+
stride=1,
|
| 105 |
+
padding=1)
|
| 106 |
+
if self.in_channels != self.out_channels:
|
| 107 |
+
if self.use_conv_shortcut:
|
| 108 |
+
self.conv_shortcut = torch.nn.Conv2d(in_channels,
|
| 109 |
+
out_channels,
|
| 110 |
+
kernel_size=3,
|
| 111 |
+
stride=1,
|
| 112 |
+
padding=1)
|
| 113 |
+
else:
|
| 114 |
+
self.nin_shortcut = torch.nn.Conv2d(in_channels,
|
| 115 |
+
out_channels,
|
| 116 |
+
kernel_size=1,
|
| 117 |
+
stride=1,
|
| 118 |
+
padding=0)
|
| 119 |
+
|
| 120 |
+
def forward(self, x, temb):
|
| 121 |
+
h = x
|
| 122 |
+
h = self.norm1(h)
|
| 123 |
+
h = nonlinearity(h)
|
| 124 |
+
h = self.conv1(h)
|
| 125 |
+
|
| 126 |
+
if temb is not None:
|
| 127 |
+
h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None]
|
| 128 |
+
|
| 129 |
+
h = self.norm2(h)
|
| 130 |
+
h = nonlinearity(h)
|
| 131 |
+
h = self.dropout(h)
|
| 132 |
+
h = self.conv2(h)
|
| 133 |
+
|
| 134 |
+
if self.in_channels != self.out_channels:
|
| 135 |
+
if self.use_conv_shortcut:
|
| 136 |
+
x = self.conv_shortcut(x)
|
| 137 |
+
else:
|
| 138 |
+
x = self.nin_shortcut(x)
|
| 139 |
+
|
| 140 |
+
return x+h
|
| 141 |
+
|
| 142 |
+
|
| 143 |
+
class LinAttnBlock(LinearAttention):
|
| 144 |
+
"""to match AttnBlock usage"""
|
| 145 |
+
def __init__(self, in_channels):
|
| 146 |
+
super().__init__(dim=in_channels, heads=1, dim_head=in_channels)
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
class AttnBlock(nn.Module):
|
| 150 |
+
def __init__(self, in_channels):
|
| 151 |
+
super().__init__()
|
| 152 |
+
self.in_channels = in_channels
|
| 153 |
+
|
| 154 |
+
self.norm = Normalize(in_channels)
|
| 155 |
+
self.q = torch.nn.Conv2d(in_channels,
|
| 156 |
+
in_channels,
|
| 157 |
+
kernel_size=1,
|
| 158 |
+
stride=1,
|
| 159 |
+
padding=0)
|
| 160 |
+
self.k = torch.nn.Conv2d(in_channels,
|
| 161 |
+
in_channels,
|
| 162 |
+
kernel_size=1,
|
| 163 |
+
stride=1,
|
| 164 |
+
padding=0)
|
| 165 |
+
self.v = torch.nn.Conv2d(in_channels,
|
| 166 |
+
in_channels,
|
| 167 |
+
kernel_size=1,
|
| 168 |
+
stride=1,
|
| 169 |
+
padding=0)
|
| 170 |
+
self.proj_out = torch.nn.Conv2d(in_channels,
|
| 171 |
+
in_channels,
|
| 172 |
+
kernel_size=1,
|
| 173 |
+
stride=1,
|
| 174 |
+
padding=0)
|
| 175 |
+
|
| 176 |
+
|
| 177 |
+
def forward(self, x):
|
| 178 |
+
h_ = x
|
| 179 |
+
h_ = self.norm(h_)
|
| 180 |
+
q = self.q(h_)
|
| 181 |
+
k = self.k(h_)
|
| 182 |
+
v = self.v(h_)
|
| 183 |
+
|
| 184 |
+
# compute attention
|
| 185 |
+
b,c,h,w = q.shape
|
| 186 |
+
q = q.reshape(b,c,h*w)
|
| 187 |
+
q = q.permute(0,2,1) # b,hw,c
|
| 188 |
+
k = k.reshape(b,c,h*w) # b,c,hw
|
| 189 |
+
w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
|
| 190 |
+
w_ = w_ * (int(c)**(-0.5))
|
| 191 |
+
w_ = torch.nn.functional.softmax(w_, dim=2)
|
| 192 |
+
|
| 193 |
+
# attend to values
|
| 194 |
+
v = v.reshape(b,c,h*w)
|
| 195 |
+
w_ = w_.permute(0,2,1) # b,hw,hw (first hw of k, second of q)
|
| 196 |
+
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]
|
| 197 |
+
h_ = h_.reshape(b,c,h,w)
|
| 198 |
+
|
| 199 |
+
h_ = self.proj_out(h_)
|
| 200 |
+
|
| 201 |
+
return x+h_
|
| 202 |
+
|
| 203 |
+
|
| 204 |
+
def make_attn(in_channels, attn_type="vanilla"):
|
| 205 |
+
assert attn_type in ["vanilla", "linear", "none"], f'attn_type {attn_type} unknown'
|
| 206 |
+
print(f"making attention of type '{attn_type}' with {in_channels} in_channels")
|
| 207 |
+
if attn_type == "vanilla":
|
| 208 |
+
return AttnBlock(in_channels)
|
| 209 |
+
elif attn_type == "none":
|
| 210 |
+
return nn.Identity(in_channels)
|
| 211 |
+
else:
|
| 212 |
+
return LinAttnBlock(in_channels)
|
| 213 |
+
|
| 214 |
+
|
| 215 |
+
class Encoder(nn.Module):
|
| 216 |
+
def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
|
| 217 |
+
attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
|
| 218 |
+
resolution, z_channels, double_z=True, use_linear_attn=False, attn_type="vanilla",
|
| 219 |
+
**ignore_kwargs):
|
| 220 |
+
super().__init__()
|
| 221 |
+
if use_linear_attn: attn_type = "linear"
|
| 222 |
+
self.ch = ch
|
| 223 |
+
self.temb_ch = 0
|
| 224 |
+
self.num_resolutions = len(ch_mult)
|
| 225 |
+
self.num_res_blocks = num_res_blocks
|
| 226 |
+
self.resolution = resolution
|
| 227 |
+
self.in_channels = in_channels
|
| 228 |
+
|
| 229 |
+
# downsampling
|
| 230 |
+
self.conv_in = torch.nn.Conv2d(in_channels,
|
| 231 |
+
self.ch,
|
| 232 |
+
kernel_size=3,
|
| 233 |
+
stride=1,
|
| 234 |
+
padding=1)
|
| 235 |
+
|
| 236 |
+
curr_res = resolution
|
| 237 |
+
in_ch_mult = (1,)+tuple(ch_mult)
|
| 238 |
+
self.in_ch_mult = in_ch_mult
|
| 239 |
+
self.down = nn.ModuleList()
|
| 240 |
+
for i_level in range(self.num_resolutions):
|
| 241 |
+
block = nn.ModuleList()
|
| 242 |
+
attn = nn.ModuleList()
|
| 243 |
+
block_in = ch*in_ch_mult[i_level]
|
| 244 |
+
block_out = ch*ch_mult[i_level]
|
| 245 |
+
for i_block in range(self.num_res_blocks):
|
| 246 |
+
block.append(ResnetBlock(in_channels=block_in,
|
| 247 |
+
out_channels=block_out,
|
| 248 |
+
temb_channels=self.temb_ch,
|
| 249 |
+
dropout=dropout))
|
| 250 |
+
block_in = block_out
|
| 251 |
+
if curr_res in attn_resolutions:
|
| 252 |
+
attn.append(make_attn(block_in, attn_type=attn_type))
|
| 253 |
+
down = nn.Module()
|
| 254 |
+
down.block = block
|
| 255 |
+
down.attn = attn
|
| 256 |
+
if i_level != self.num_resolutions-1:
|
| 257 |
+
down.downsample = Downsample(block_in, resamp_with_conv)
|
| 258 |
+
curr_res = curr_res // 2
|
| 259 |
+
self.down.append(down)
|
| 260 |
+
|
| 261 |
+
# middle
|
| 262 |
+
self.mid = nn.Module()
|
| 263 |
+
self.mid.block_1 = ResnetBlock(in_channels=block_in,
|
| 264 |
+
out_channels=block_in,
|
| 265 |
+
temb_channels=self.temb_ch,
|
| 266 |
+
dropout=dropout)
|
| 267 |
+
self.mid.attn_1 = make_attn(block_in, attn_type=attn_type)
|
| 268 |
+
self.mid.block_2 = ResnetBlock(in_channels=block_in,
|
| 269 |
+
out_channels=block_in,
|
| 270 |
+
temb_channels=self.temb_ch,
|
| 271 |
+
dropout=dropout)
|
| 272 |
+
|
| 273 |
+
# end
|
| 274 |
+
self.norm_out = Normalize(block_in)
|
| 275 |
+
self.conv_out = torch.nn.Conv2d(block_in,
|
| 276 |
+
2*z_channels if double_z else z_channels,
|
| 277 |
+
kernel_size=3,
|
| 278 |
+
stride=1,
|
| 279 |
+
padding=1)
|
| 280 |
+
|
| 281 |
+
def forward(self, x):
|
| 282 |
+
# timestep embedding
|
| 283 |
+
temb = None
|
| 284 |
+
|
| 285 |
+
# downsampling
|
| 286 |
+
hs = [self.conv_in(x)]
|
| 287 |
+
for i_level in range(self.num_resolutions):
|
| 288 |
+
for i_block in range(self.num_res_blocks):
|
| 289 |
+
h = self.down[i_level].block[i_block](hs[-1], temb)
|
| 290 |
+
if len(self.down[i_level].attn) > 0:
|
| 291 |
+
h = self.down[i_level].attn[i_block](h)
|
| 292 |
+
hs.append(h)
|
| 293 |
+
if i_level != self.num_resolutions-1:
|
| 294 |
+
hs.append(self.down[i_level].downsample(hs[-1]))
|
| 295 |
+
|
| 296 |
+
# middle
|
| 297 |
+
h = hs[-1]
|
| 298 |
+
h = self.mid.block_1(h, temb)
|
| 299 |
+
h = self.mid.attn_1(h)
|
| 300 |
+
h = self.mid.block_2(h, temb)
|
| 301 |
+
|
| 302 |
+
# end
|
| 303 |
+
h = self.norm_out(h)
|
| 304 |
+
h = nonlinearity(h)
|
| 305 |
+
h = self.conv_out(h)
|
| 306 |
+
return h
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
class Decoder(nn.Module):
|
| 310 |
+
def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
|
| 311 |
+
attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
|
| 312 |
+
resolution, z_channels, give_pre_end=False, tanh_out=False, use_linear_attn=False,
|
| 313 |
+
attn_type="vanilla", **ignorekwargs):
|
| 314 |
+
super().__init__()
|
| 315 |
+
if use_linear_attn: attn_type = "linear"
|
| 316 |
+
self.ch = ch
|
| 317 |
+
self.temb_ch = 0
|
| 318 |
+
self.num_resolutions = len(ch_mult)
|
| 319 |
+
self.num_res_blocks = num_res_blocks
|
| 320 |
+
self.resolution = resolution
|
| 321 |
+
self.in_channels = in_channels
|
| 322 |
+
self.give_pre_end = give_pre_end
|
| 323 |
+
self.tanh_out = tanh_out
|
| 324 |
+
|
| 325 |
+
# compute in_ch_mult, block_in and curr_res at lowest res
|
| 326 |
+
in_ch_mult = (1,)+tuple(ch_mult)
|
| 327 |
+
block_in = ch*ch_mult[self.num_resolutions-1]
|
| 328 |
+
curr_res = resolution // 2**(self.num_resolutions-1)
|
| 329 |
+
self.z_shape = (1,z_channels,curr_res,curr_res)
|
| 330 |
+
print("Working with z of shape {} = {} dimensions.".format(
|
| 331 |
+
self.z_shape, np.prod(self.z_shape)))
|
| 332 |
+
|
| 333 |
+
# z to block_in
|
| 334 |
+
self.conv_in = torch.nn.Conv2d(z_channels,
|
| 335 |
+
block_in,
|
| 336 |
+
kernel_size=3,
|
| 337 |
+
stride=1,
|
| 338 |
+
padding=1)
|
| 339 |
+
|
| 340 |
+
# middle
|
| 341 |
+
self.mid = nn.Module()
|
| 342 |
+
self.mid.block_1 = ResnetBlock(in_channels=block_in,
|
| 343 |
+
out_channels=block_in,
|
| 344 |
+
temb_channels=self.temb_ch,
|
| 345 |
+
dropout=dropout)
|
| 346 |
+
self.mid.attn_1 = make_attn(block_in, attn_type=attn_type)
|
| 347 |
+
self.mid.block_2 = ResnetBlock(in_channels=block_in,
|
| 348 |
+
out_channels=block_in,
|
| 349 |
+
temb_channels=self.temb_ch,
|
| 350 |
+
dropout=dropout)
|
| 351 |
+
|
| 352 |
+
# upsampling
|
| 353 |
+
self.up = nn.ModuleList()
|
| 354 |
+
for i_level in reversed(range(self.num_resolutions)):
|
| 355 |
+
block = nn.ModuleList()
|
| 356 |
+
attn = nn.ModuleList()
|
| 357 |
+
block_out = ch*ch_mult[i_level]
|
| 358 |
+
for i_block in range(self.num_res_blocks+1):
|
| 359 |
+
block.append(ResnetBlock(in_channels=block_in,
|
| 360 |
+
out_channels=block_out,
|
| 361 |
+
temb_channels=self.temb_ch,
|
| 362 |
+
dropout=dropout))
|
| 363 |
+
block_in = block_out
|
| 364 |
+
if curr_res in attn_resolutions:
|
| 365 |
+
attn.append(make_attn(block_in, attn_type=attn_type))
|
| 366 |
+
up = nn.Module()
|
| 367 |
+
up.block = block
|
| 368 |
+
up.attn = attn
|
| 369 |
+
if i_level != 0:
|
| 370 |
+
up.upsample = Upsample(block_in, resamp_with_conv)
|
| 371 |
+
curr_res = curr_res * 2
|
| 372 |
+
self.up.insert(0, up) # prepend to get consistent order
|
| 373 |
+
|
| 374 |
+
# end
|
| 375 |
+
self.norm_out = Normalize(block_in)
|
| 376 |
+
self.conv_out = torch.nn.Conv2d(block_in,
|
| 377 |
+
out_ch,
|
| 378 |
+
kernel_size=3,
|
| 379 |
+
stride=1,
|
| 380 |
+
padding=1)
|
| 381 |
+
|
| 382 |
+
def forward(self, z):
|
| 383 |
+
#assert z.shape[1:] == self.z_shape[1:]
|
| 384 |
+
self.last_z_shape = z.shape
|
| 385 |
+
|
| 386 |
+
# timestep embedding
|
| 387 |
+
temb = None
|
| 388 |
+
|
| 389 |
+
# z to block_in
|
| 390 |
+
h = self.conv_in(z)
|
| 391 |
+
|
| 392 |
+
# middle
|
| 393 |
+
h = self.mid.block_1(h, temb)
|
| 394 |
+
h = self.mid.attn_1(h)
|
| 395 |
+
h = self.mid.block_2(h, temb)
|
| 396 |
+
|
| 397 |
+
# upsampling
|
| 398 |
+
for i_level in reversed(range(self.num_resolutions)):
|
| 399 |
+
for i_block in range(self.num_res_blocks+1):
|
| 400 |
+
h = self.up[i_level].block[i_block](h, temb)
|
| 401 |
+
if len(self.up[i_level].attn) > 0:
|
| 402 |
+
h = self.up[i_level].attn[i_block](h)
|
| 403 |
+
if i_level != 0:
|
| 404 |
+
h = self.up[i_level].upsample(h)
|
| 405 |
+
|
| 406 |
+
# end
|
| 407 |
+
if self.give_pre_end:
|
| 408 |
+
return h
|
| 409 |
+
|
| 410 |
+
h = self.norm_out(h)
|
| 411 |
+
h = nonlinearity(h)
|
| 412 |
+
h = self.conv_out(h)
|
| 413 |
+
if self.tanh_out:
|
| 414 |
+
h = torch.tanh(h)
|
| 415 |
+
return h
|
| 416 |
+
|
| 417 |
+
|
| 418 |
+
class FrozenAutoencoderKL(nn.Module):
|
| 419 |
+
def __init__(self, ddconfig, embed_dim, pretrained_path, scale_factor=0.18215):
|
| 420 |
+
super().__init__()
|
| 421 |
+
print(f'Create autoencoder with scale_factor={scale_factor}')
|
| 422 |
+
self.encoder = Encoder(**ddconfig)
|
| 423 |
+
self.decoder = Decoder(**ddconfig)
|
| 424 |
+
assert ddconfig["double_z"]
|
| 425 |
+
self.quant_conv = torch.nn.Conv2d(2 * ddconfig["z_channels"], 2 * embed_dim, 1)
|
| 426 |
+
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
|
| 427 |
+
self.embed_dim = embed_dim
|
| 428 |
+
self.scale_factor = scale_factor
|
| 429 |
+
m, u = self.load_state_dict(torch.load(pretrained_path, map_location='cpu'))
|
| 430 |
+
assert len(m) == 0 and len(u) == 0
|
| 431 |
+
self.eval()
|
| 432 |
+
self.requires_grad_(False)
|
| 433 |
+
|
| 434 |
+
def encode_moments(self, x):
|
| 435 |
+
h = self.encoder(x)
|
| 436 |
+
moments = self.quant_conv(h)
|
| 437 |
+
return moments
|
| 438 |
+
|
| 439 |
+
def sample(self, moments):
|
| 440 |
+
mean, logvar = torch.chunk(moments, 2, dim=1)
|
| 441 |
+
logvar = torch.clamp(logvar, -30.0, 20.0)
|
| 442 |
+
std = torch.exp(0.5 * logvar)
|
| 443 |
+
z = mean + std * torch.randn_like(mean)
|
| 444 |
+
z = self.scale_factor * z
|
| 445 |
+
return z
|
| 446 |
+
|
| 447 |
+
def encode(self, x):
|
| 448 |
+
moments = self.encode_moments(x)
|
| 449 |
+
z = self.sample(moments)
|
| 450 |
+
return z
|
| 451 |
+
|
| 452 |
+
def decode(self, z):
|
| 453 |
+
z = (1. / self.scale_factor) * z
|
| 454 |
+
z = self.post_quant_conv(z)
|
| 455 |
+
dec = self.decoder(z)
|
| 456 |
+
return dec
|
| 457 |
+
|
| 458 |
+
def forward(self, inputs, fn):
|
| 459 |
+
if fn == 'encode_moments':
|
| 460 |
+
return self.encode_moments(inputs)
|
| 461 |
+
elif fn == 'encode':
|
| 462 |
+
return self.encode(inputs)
|
| 463 |
+
elif fn == 'decode':
|
| 464 |
+
return self.decode(inputs)
|
| 465 |
+
else:
|
| 466 |
+
raise NotImplementedError
|
| 467 |
+
|
| 468 |
+
|
| 469 |
+
def get_model(pretrained_path, scale_factor=0.18215):
|
| 470 |
+
ddconfig = dict(
|
| 471 |
+
double_z=True,
|
| 472 |
+
z_channels=4,
|
| 473 |
+
resolution=256,
|
| 474 |
+
in_channels=3,
|
| 475 |
+
out_ch=3,
|
| 476 |
+
ch=128,
|
| 477 |
+
ch_mult=[1, 2, 4, 4],
|
| 478 |
+
num_res_blocks=2,
|
| 479 |
+
attn_resolutions=[],
|
| 480 |
+
dropout=0.0
|
| 481 |
+
)
|
| 482 |
+
return FrozenAutoencoderKL(ddconfig, 4, pretrained_path, scale_factor)
|
| 483 |
+
|
| 484 |
+
|
| 485 |
+
def main():
|
| 486 |
+
import torchvision.transforms as transforms
|
| 487 |
+
from torchvision.utils import save_image
|
| 488 |
+
import os
|
| 489 |
+
from PIL import Image
|
| 490 |
+
|
| 491 |
+
model = get_model('/data/CASIA/Projects/sqf/Diffsion_Model/self_operate/save_auto/checkpoint/autoencoder_kl.pth')
|
| 492 |
+
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
| 493 |
+
model = model.to(device)
|
| 494 |
+
|
| 495 |
+
scale_factor = 0.18215
|
| 496 |
+
T = transforms.Compose([transforms.Resize(256), transforms.CenterCrop(256), transforms.ToTensor()])
|
| 497 |
+
path = '/data/CASIA/Projects/sqf/Diffsion_Model/self_operate/auto_data/test'
|
| 498 |
+
fnames = os.listdir(path)
|
| 499 |
+
for fname in fnames:
|
| 500 |
+
p = os.path.join(path, fname)
|
| 501 |
+
img = Image.open(p)
|
| 502 |
+
img = T(img)
|
| 503 |
+
img = img * 2. - 1
|
| 504 |
+
img = img[None, ...]
|
| 505 |
+
img = img.to(device)
|
| 506 |
+
|
| 507 |
+
# with torch.cuda.amp.autocast():
|
| 508 |
+
# moments = model.encode_moments(img)
|
| 509 |
+
# mean, logvar = torch.chunk(moments, 2, dim=1)
|
| 510 |
+
# logvar = torch.clamp(logvar, -30.0, 20.0)
|
| 511 |
+
# std = torch.exp(0.5 * logvar)
|
| 512 |
+
# zs = [(mean + std * torch.randn_like(mean)) * scale_factor for _ in range(4)]
|
| 513 |
+
# recons = [model.decode(z) for z in zs]
|
| 514 |
+
|
| 515 |
+
with torch.cuda.amp.autocast():
|
| 516 |
+
print('test encode & decode')
|
| 517 |
+
recons = [model.decode(model.encode(img)) for _ in range(4)]
|
| 518 |
+
|
| 519 |
+
out = torch.cat([img, *recons], dim=0)
|
| 520 |
+
out = (out + 1) * 0.5
|
| 521 |
+
save_image(out, '/data/CASIA/Projects/sqf/Diffsion_Model/self_operate/save_auto/img/' + f'recons_{fname}')
|
| 522 |
+
|
| 523 |
+
|
| 524 |
+
from torch.utils.data import Dataset, DataLoader
|
| 525 |
+
from torchvision.transforms import ToTensor, Compose, CenterCrop, Resize, RandomCrop, Lambda
|
| 526 |
+
from torch.optim import Adam
|
| 527 |
+
class Dataset_self(Dataset):
|
| 528 |
+
def __init__(self,img_root, preprocess):
|
| 529 |
+
self.img_root = img_root
|
| 530 |
+
self.img_process = preprocess
|
| 531 |
+
self.img = []
|
| 532 |
+
for name_img in os.listdir(self.img_root):
|
| 533 |
+
self.img.append(self.img_root + '/' + name_img)
|
| 534 |
+
|
| 535 |
+
def __len__(self):
|
| 536 |
+
return len(self.img)
|
| 537 |
+
|
| 538 |
+
def __getitem__(self, idx):
|
| 539 |
+
img_path = self.img[idx]
|
| 540 |
+
image = Image.open(img_path).convert('RGB')
|
| 541 |
+
image = self.img_process(image)
|
| 542 |
+
return image
|
| 543 |
+
|
| 544 |
+
def train():
|
| 545 |
+
model = get_model('assets/stable-diffusion/autoencoder_kl.pth')
|
| 546 |
+
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
| 547 |
+
model = model.to(device)
|
| 548 |
+
scale_factor = 0.18215
|
| 549 |
+
path_test = ['/data/CASIA/Projects/sqf/Diffsion_Model/self_operate/auto_data/test/389.png', '/data/CASIA/Projects/sqf/Diffsion_Model/self_operate/auto_data/test/390.png']
|
| 550 |
+
T = transforms.Compose([transforms.Resize(256), transforms.CenterCrop(256), transforms.ToTensor(), Lambda(lambda t: (t * 2) - 1)])
|
| 551 |
+
dataset = Dataset_self(img_root= '/data/CASIA/Projects/sqf/Diffsion_Model/self_operate/auto_data/train/good', preprocess=T)
|
| 552 |
+
dataloader = DataLoader(dataset, batch_size=4, shuffle=True)
|
| 553 |
+
criterion = nn.MSELoss()
|
| 554 |
+
optimizer = Adam(model.parameters(), lr=1e-5)
|
| 555 |
+
epochs = 100000
|
| 556 |
+
for epoch in range(epochs):
|
| 557 |
+
step = 0
|
| 558 |
+
for batch in dataloader:
|
| 559 |
+
step += 1
|
| 560 |
+
optimizer.zero_grad()
|
| 561 |
+
batch = batch.to(device)
|
| 562 |
+
output_model = model.decode(model.encode(batch))
|
| 563 |
+
loss = criterion(output_model, batch)
|
| 564 |
+
loss.backward()
|
| 565 |
+
optimizer.step()
|
| 566 |
+
print(epoch, loss.item())
|
| 567 |
+
torch.save({
|
| 568 |
+
'model_state_dict': model.state_dict(),
|
| 569 |
+
'optimizer_state_dict': optimizer.state_dict(),
|
| 570 |
+
}, '/data/CASIA/Projects/sqf/Diffsion_Model/self_operate/save_auto/checkpoint/model_hazelnut_last.pth')
|
| 571 |
+
for i in range(2):
|
| 572 |
+
img_path = path_test[i]
|
| 573 |
+
img = Image.open(img_path)
|
| 574 |
+
img = T(img)
|
| 575 |
+
img = img[None, ...]
|
| 576 |
+
img = img.to(device)
|
| 577 |
+
out = model.decode(model.encode(img))
|
| 578 |
+
out = (out + 1) * 0.5
|
| 579 |
+
save_image(out, '/data/CASIA/Projects/sqf/Diffsion_Model/self_operate/save_auto/img/' + f'recons_{epoch}_{i}.png')
|
| 580 |
+
|
| 581 |
+
|
| 582 |
+
if __name__ == "__main__":
|
| 583 |
+
# train()
|
| 584 |
+
main()
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
from .clip import *
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/__pycache__/__init__.cpython-310.pyc
ADDED
|
Binary file (214 Bytes). View file
|
|
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/__pycache__/clip.cpython-310.pyc
ADDED
|
Binary file (8.86 kB). View file
|
|
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/__pycache__/model.cpython-310.pyc
ADDED
|
Binary file (15.2 kB). View file
|
|
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/__pycache__/simple_tokenizer.cpython-310.pyc
ADDED
|
Binary file (5.74 kB). View file
|
|
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/bpe_simple_vocab_16e6.txt.gz
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:924691ac288e54409236115652ad4aa250f48203de50a9e4722a6ecd48d6804a
|
| 3 |
+
size 1356917
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/clip.py
ADDED
|
@@ -0,0 +1,237 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import hashlib
|
| 2 |
+
import os
|
| 3 |
+
import urllib
|
| 4 |
+
import warnings
|
| 5 |
+
from typing import Any, Union, List
|
| 6 |
+
from pkg_resources import packaging
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
from PIL import Image
|
| 10 |
+
from torchvision.transforms import Compose, Resize, CenterCrop, ToTensor, Normalize
|
| 11 |
+
from tqdm import tqdm
|
| 12 |
+
|
| 13 |
+
from .model import build_model
|
| 14 |
+
from .simple_tokenizer import SimpleTokenizer as _Tokenizer
|
| 15 |
+
|
| 16 |
+
try:
|
| 17 |
+
from torchvision.transforms import InterpolationMode
|
| 18 |
+
BICUBIC = InterpolationMode.BICUBIC
|
| 19 |
+
except ImportError:
|
| 20 |
+
BICUBIC = Image.BICUBIC
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
if packaging.version.parse(torch.__version__) < packaging.version.parse("1.7.1"):
|
| 24 |
+
warnings.warn("PyTorch version 1.7.1 or higher is recommended")
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
__all__ = ["available_models", "load", "tokenize"]
|
| 28 |
+
_tokenizer = _Tokenizer()
|
| 29 |
+
|
| 30 |
+
_MODELS = {
|
| 31 |
+
"RN50": "https://openaipublic.azureedge.net/clip/models/afeb0e10f9e5a86da6080e35cf09123aca3b358a0c3e3b6c78a7b63bc04b6762/RN50.pt",
|
| 32 |
+
"RN101": "https://openaipublic.azureedge.net/clip/models/8fa8567bab74a42d41c5915025a8e4538c3bdbe8804a470a72f30b0d94fab599/RN101.pt",
|
| 33 |
+
"RN50x4": "https://openaipublic.azureedge.net/clip/models/7e526bd135e493cef0776de27d5f42653e6b4c8bf9e0f653bb11773263205fdd/RN50x4.pt",
|
| 34 |
+
"RN50x16": "https://openaipublic.azureedge.net/clip/models/52378b407f34354e150460fe41077663dd5b39c54cd0bfd2b27167a4a06ec9aa/RN50x16.pt",
|
| 35 |
+
"RN50x64": "https://openaipublic.azureedge.net/clip/models/be1cfb55d75a9666199fb2206c106743da0f6468c9d327f3e0d0a543a9919d9c/RN50x64.pt",
|
| 36 |
+
"ViT-B/32": "https://openaipublic.azureedge.net/clip/models/40d365715913c9da98579312b702a82c18be219cc2a73407c4526f58eba950af/ViT-B-32.pt",
|
| 37 |
+
"ViT-B/16": "https://openaipublic.azureedge.net/clip/models/5806e77cd80f8b59890b7e101eabd078d9fb84e6937f9e85e4ecb61988df416f/ViT-B-16.pt",
|
| 38 |
+
"ViT-L/14": "https://openaipublic.azureedge.net/clip/models/b8cca3fd41ae0c99ba7e8951adf17d267cdb84cd88be6f7c2e0eca1737a03836/ViT-L-14.pt",
|
| 39 |
+
"ViT-L/14@336px": "https://openaipublic.azureedge.net/clip/models/3035c92b350959924f9f00213499208652fc7ea050643e8b385c2dac08641f02/ViT-L-14-336px.pt",
|
| 40 |
+
}
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def _download(url: str, root: str):
|
| 44 |
+
os.makedirs(root, exist_ok=True)
|
| 45 |
+
filename = os.path.basename(url)
|
| 46 |
+
|
| 47 |
+
expected_sha256 = url.split("/")[-2]
|
| 48 |
+
download_target = os.path.join(root, filename)
|
| 49 |
+
|
| 50 |
+
if os.path.exists(download_target) and not os.path.isfile(download_target):
|
| 51 |
+
raise RuntimeError(f"{download_target} exists and is not a regular file")
|
| 52 |
+
|
| 53 |
+
if os.path.isfile(download_target):
|
| 54 |
+
if hashlib.sha256(open(download_target, "rb").read()).hexdigest() == expected_sha256:
|
| 55 |
+
return download_target
|
| 56 |
+
else:
|
| 57 |
+
warnings.warn(f"{download_target} exists, but the SHA256 checksum does not match; re-downloading the file")
|
| 58 |
+
|
| 59 |
+
with urllib.request.urlopen(url) as source, open(download_target, "wb") as output:
|
| 60 |
+
with tqdm(total=int(source.info().get("Content-Length")), ncols=80, unit='iB', unit_scale=True, unit_divisor=1024) as loop:
|
| 61 |
+
while True:
|
| 62 |
+
buffer = source.read(8192)
|
| 63 |
+
if not buffer:
|
| 64 |
+
break
|
| 65 |
+
|
| 66 |
+
output.write(buffer)
|
| 67 |
+
loop.update(len(buffer))
|
| 68 |
+
|
| 69 |
+
if hashlib.sha256(open(download_target, "rb").read()).hexdigest() != expected_sha256:
|
| 70 |
+
raise RuntimeError("Model has been downloaded but the SHA256 checksum does not not match")
|
| 71 |
+
|
| 72 |
+
return download_target
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
def _convert_image_to_rgb(image):
|
| 76 |
+
return image.convert("RGB")
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
def _transform(n_px):
|
| 80 |
+
return Compose([
|
| 81 |
+
Resize(n_px, interpolation=BICUBIC),
|
| 82 |
+
CenterCrop(n_px),
|
| 83 |
+
_convert_image_to_rgb,
|
| 84 |
+
ToTensor(),
|
| 85 |
+
Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711)),
|
| 86 |
+
])
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def available_models() -> List[str]:
|
| 90 |
+
"""Returns the names of available CLIP models"""
|
| 91 |
+
return list(_MODELS.keys())
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
def load(name: str, device: Union[str, torch.device] = "cuda" if torch.cuda.is_available() else "cpu", jit: bool = False, download_root: str = None):
|
| 95 |
+
"""Load a CLIP model
|
| 96 |
+
|
| 97 |
+
Parameters
|
| 98 |
+
----------
|
| 99 |
+
name : str
|
| 100 |
+
A model name listed by `clip.available_models()`, or the path to a model checkpoint containing the state_dict
|
| 101 |
+
|
| 102 |
+
device : Union[str, torch.device]
|
| 103 |
+
The device to put the loaded model
|
| 104 |
+
|
| 105 |
+
jit : bool
|
| 106 |
+
Whether to load the optimized JIT model or more hackable non-JIT model (default).
|
| 107 |
+
|
| 108 |
+
download_root: str
|
| 109 |
+
path to download the model files; by default, it uses "~/.cache/clip"
|
| 110 |
+
|
| 111 |
+
Returns
|
| 112 |
+
-------
|
| 113 |
+
model : torch.nn.Module
|
| 114 |
+
The CLIP model
|
| 115 |
+
|
| 116 |
+
preprocess : Callable[[PIL.Image], torch.Tensor]
|
| 117 |
+
A torchvision transform that converts a PIL image into a tensor that the returned model can take as its input
|
| 118 |
+
"""
|
| 119 |
+
if name in _MODELS:
|
| 120 |
+
model_path = _download(_MODELS[name], download_root or os.path.expanduser("~/.cache/clip"))
|
| 121 |
+
elif os.path.isfile(name):
|
| 122 |
+
model_path = name
|
| 123 |
+
else:
|
| 124 |
+
raise RuntimeError(f"Model {name} not found; available models = {available_models()}")
|
| 125 |
+
|
| 126 |
+
with open(model_path, 'rb') as opened_file:
|
| 127 |
+
try:
|
| 128 |
+
# loading JIT archive
|
| 129 |
+
model = torch.jit.load(opened_file, map_location=device if jit else "cpu").eval()
|
| 130 |
+
state_dict = None
|
| 131 |
+
except RuntimeError:
|
| 132 |
+
# loading saved state dict
|
| 133 |
+
if jit:
|
| 134 |
+
warnings.warn(f"File {model_path} is not a JIT archive. Loading as a state dict instead")
|
| 135 |
+
jit = False
|
| 136 |
+
state_dict = torch.load(opened_file, map_location="cpu")
|
| 137 |
+
|
| 138 |
+
if not jit:
|
| 139 |
+
model = build_model(state_dict or model.state_dict()).to(device)
|
| 140 |
+
if str(device) == "cpu":
|
| 141 |
+
model.float()
|
| 142 |
+
return model, _transform(model.visual.input_resolution)
|
| 143 |
+
|
| 144 |
+
# patch the device names
|
| 145 |
+
device_holder = torch.jit.trace(lambda: torch.ones([]).to(torch.device(device)), example_inputs=[])
|
| 146 |
+
device_node = [n for n in device_holder.graph.findAllNodes("prim::Constant") if "Device" in repr(n)][-1]
|
| 147 |
+
|
| 148 |
+
def patch_device(module):
|
| 149 |
+
try:
|
| 150 |
+
graphs = [module.graph] if hasattr(module, "graph") else []
|
| 151 |
+
except RuntimeError:
|
| 152 |
+
graphs = []
|
| 153 |
+
|
| 154 |
+
if hasattr(module, "forward1"):
|
| 155 |
+
graphs.append(module.forward1.graph)
|
| 156 |
+
|
| 157 |
+
for graph in graphs:
|
| 158 |
+
for node in graph.findAllNodes("prim::Constant"):
|
| 159 |
+
if "value" in node.attributeNames() and str(node["value"]).startswith("cuda"):
|
| 160 |
+
node.copyAttributes(device_node)
|
| 161 |
+
|
| 162 |
+
model.apply(patch_device)
|
| 163 |
+
patch_device(model.encode_image)
|
| 164 |
+
patch_device(model.encode_text)
|
| 165 |
+
|
| 166 |
+
# patch dtype to float32 on CPU
|
| 167 |
+
if str(device) == "cpu":
|
| 168 |
+
float_holder = torch.jit.trace(lambda: torch.ones([]).float(), example_inputs=[])
|
| 169 |
+
float_input = list(float_holder.graph.findNode("aten::to").inputs())[1]
|
| 170 |
+
float_node = float_input.node()
|
| 171 |
+
|
| 172 |
+
def patch_float(module):
|
| 173 |
+
try:
|
| 174 |
+
graphs = [module.graph] if hasattr(module, "graph") else []
|
| 175 |
+
except RuntimeError:
|
| 176 |
+
graphs = []
|
| 177 |
+
|
| 178 |
+
if hasattr(module, "forward1"):
|
| 179 |
+
graphs.append(module.forward1.graph)
|
| 180 |
+
|
| 181 |
+
for graph in graphs:
|
| 182 |
+
for node in graph.findAllNodes("aten::to"):
|
| 183 |
+
inputs = list(node.inputs())
|
| 184 |
+
for i in [1, 2]: # dtype can be the second or third argument to aten::to()
|
| 185 |
+
if inputs[i].node()["value"] == 5:
|
| 186 |
+
inputs[i].node().copyAttributes(float_node)
|
| 187 |
+
|
| 188 |
+
model.apply(patch_float)
|
| 189 |
+
patch_float(model.encode_image)
|
| 190 |
+
patch_float(model.encode_text)
|
| 191 |
+
|
| 192 |
+
model.float()
|
| 193 |
+
|
| 194 |
+
return model, _transform(model.input_resolution.item())
|
| 195 |
+
|
| 196 |
+
|
| 197 |
+
def tokenize(texts: Union[str, List[str]], context_length: int = 77, truncate: bool = False) -> Union[torch.IntTensor, torch.LongTensor]:
|
| 198 |
+
"""
|
| 199 |
+
Returns the tokenized representation of given input string(s)
|
| 200 |
+
|
| 201 |
+
Parameters
|
| 202 |
+
----------
|
| 203 |
+
texts : Union[str, List[str]]
|
| 204 |
+
An input string or a list of input strings to tokenize
|
| 205 |
+
|
| 206 |
+
context_length : int
|
| 207 |
+
The context length to use; all CLIP models use 77 as the context length
|
| 208 |
+
|
| 209 |
+
truncate: bool
|
| 210 |
+
Whether to truncate the text in case its encoding is longer than the context length
|
| 211 |
+
|
| 212 |
+
Returns
|
| 213 |
+
-------
|
| 214 |
+
A two-dimensional tensor containing the resulting tokens, shape = [number of input strings, context_length].
|
| 215 |
+
We return LongTensor when torch version is <1.8.0, since older index_select requires indices to be long.
|
| 216 |
+
"""
|
| 217 |
+
if isinstance(texts, str):
|
| 218 |
+
texts = [texts]
|
| 219 |
+
|
| 220 |
+
sot_token = _tokenizer.encoder["<|startoftext|>"]
|
| 221 |
+
eot_token = _tokenizer.encoder["<|endoftext|>"]
|
| 222 |
+
all_tokens = [[sot_token] + _tokenizer.encode(text) + [eot_token] for text in texts]
|
| 223 |
+
if packaging.version.parse(torch.__version__) < packaging.version.parse("1.8.0"):
|
| 224 |
+
result = torch.zeros(len(all_tokens), context_length, dtype=torch.long)
|
| 225 |
+
else:
|
| 226 |
+
result = torch.zeros(len(all_tokens), context_length, dtype=torch.int)
|
| 227 |
+
|
| 228 |
+
for i, tokens in enumerate(all_tokens):
|
| 229 |
+
if len(tokens) > context_length:
|
| 230 |
+
if truncate:
|
| 231 |
+
tokens = tokens[:context_length]
|
| 232 |
+
tokens[-1] = eot_token
|
| 233 |
+
else:
|
| 234 |
+
raise RuntimeError(f"Input {texts[i]} is too long for context length {context_length}")
|
| 235 |
+
result[i, :len(tokens)] = torch.tensor(tokens)
|
| 236 |
+
|
| 237 |
+
return result
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/model.py
ADDED
|
@@ -0,0 +1,434 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from collections import OrderedDict
|
| 2 |
+
from typing import Tuple, Union
|
| 3 |
+
|
| 4 |
+
import numpy as np
|
| 5 |
+
import torch
|
| 6 |
+
import torch.nn.functional as F
|
| 7 |
+
from torch import nn
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class Bottleneck(nn.Module):
|
| 11 |
+
expansion = 4
|
| 12 |
+
|
| 13 |
+
def __init__(self, inplanes, planes, stride=1):
|
| 14 |
+
super().__init__()
|
| 15 |
+
|
| 16 |
+
# all conv layers have stride 1. an avgpool is performed after the second convolution when stride > 1
|
| 17 |
+
self.conv1 = nn.Conv2d(inplanes, planes, 1, bias=False)
|
| 18 |
+
self.bn1 = nn.BatchNorm2d(planes)
|
| 19 |
+
self.relu1 = nn.ReLU(inplace=True)
|
| 20 |
+
|
| 21 |
+
self.conv2 = nn.Conv2d(planes, planes, 3, padding=1, bias=False)
|
| 22 |
+
self.bn2 = nn.BatchNorm2d(planes)
|
| 23 |
+
self.relu2 = nn.ReLU(inplace=True)
|
| 24 |
+
|
| 25 |
+
self.avgpool = nn.AvgPool2d(stride) if stride > 1 else nn.Identity()
|
| 26 |
+
|
| 27 |
+
self.conv3 = nn.Conv2d(planes, planes * self.expansion, 1, bias=False)
|
| 28 |
+
self.bn3 = nn.BatchNorm2d(planes * self.expansion)
|
| 29 |
+
self.relu3 = nn.ReLU(inplace=True)
|
| 30 |
+
|
| 31 |
+
self.downsample = None
|
| 32 |
+
self.stride = stride
|
| 33 |
+
|
| 34 |
+
if stride > 1 or inplanes != planes * Bottleneck.expansion:
|
| 35 |
+
# downsampling layer is prepended with an avgpool, and the subsequent convolution has stride 1
|
| 36 |
+
self.downsample = nn.Sequential(OrderedDict([
|
| 37 |
+
("-1", nn.AvgPool2d(stride)),
|
| 38 |
+
("0", nn.Conv2d(inplanes, planes * self.expansion, 1, stride=1, bias=False)),
|
| 39 |
+
("1", nn.BatchNorm2d(planes * self.expansion))
|
| 40 |
+
]))
|
| 41 |
+
|
| 42 |
+
def forward(self, x: torch.Tensor):
|
| 43 |
+
identity = x
|
| 44 |
+
|
| 45 |
+
out = self.relu1(self.bn1(self.conv1(x)))
|
| 46 |
+
out = self.relu2(self.bn2(self.conv2(out)))
|
| 47 |
+
out = self.avgpool(out)
|
| 48 |
+
out = self.bn3(self.conv3(out))
|
| 49 |
+
|
| 50 |
+
if self.downsample is not None:
|
| 51 |
+
identity = self.downsample(x)
|
| 52 |
+
|
| 53 |
+
out += identity
|
| 54 |
+
out = self.relu3(out)
|
| 55 |
+
return out
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
class AttentionPool2d(nn.Module):
|
| 59 |
+
def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None):
|
| 60 |
+
super().__init__()
|
| 61 |
+
self.positional_embedding = nn.Parameter(torch.randn(spacial_dim ** 2 + 1, embed_dim) / embed_dim ** 0.5)
|
| 62 |
+
self.k_proj = nn.Linear(embed_dim, embed_dim)
|
| 63 |
+
self.q_proj = nn.Linear(embed_dim, embed_dim)
|
| 64 |
+
self.v_proj = nn.Linear(embed_dim, embed_dim)
|
| 65 |
+
self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim)
|
| 66 |
+
self.num_heads = num_heads
|
| 67 |
+
|
| 68 |
+
def forward(self, x):
|
| 69 |
+
x = x.flatten(start_dim=2).permute(2, 0, 1) # NCHW -> (HW)NC
|
| 70 |
+
x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (HW+1)NC
|
| 71 |
+
x = x + self.positional_embedding[:, None, :].to(x.dtype) # (HW+1)NC
|
| 72 |
+
x, _ = F.multi_head_attention_forward(
|
| 73 |
+
query=x[:1], key=x, value=x,
|
| 74 |
+
embed_dim_to_check=x.shape[-1],
|
| 75 |
+
num_heads=self.num_heads,
|
| 76 |
+
q_proj_weight=self.q_proj.weight,
|
| 77 |
+
k_proj_weight=self.k_proj.weight,
|
| 78 |
+
v_proj_weight=self.v_proj.weight,
|
| 79 |
+
in_proj_weight=None,
|
| 80 |
+
in_proj_bias=torch.cat([self.q_proj.bias, self.k_proj.bias, self.v_proj.bias]),
|
| 81 |
+
bias_k=None,
|
| 82 |
+
bias_v=None,
|
| 83 |
+
add_zero_attn=False,
|
| 84 |
+
dropout_p=0,
|
| 85 |
+
out_proj_weight=self.c_proj.weight,
|
| 86 |
+
out_proj_bias=self.c_proj.bias,
|
| 87 |
+
use_separate_proj_weight=True,
|
| 88 |
+
training=self.training,
|
| 89 |
+
need_weights=False
|
| 90 |
+
)
|
| 91 |
+
return x.squeeze(0)
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
class ModifiedResNet(nn.Module):
|
| 95 |
+
"""
|
| 96 |
+
A ResNet class that is similar to torchvision's but contains the following changes:
|
| 97 |
+
- There are now 3 "stem" convolutions as opposed to 1, with an average pool instead of a max pool.
|
| 98 |
+
- Performs anti-aliasing strided convolutions, where an avgpool is prepended to convolutions with stride > 1
|
| 99 |
+
- The final pooling layer is a QKV attention instead of an average pool
|
| 100 |
+
"""
|
| 101 |
+
|
| 102 |
+
def __init__(self, layers, output_dim, heads, input_resolution=224, width=64):
|
| 103 |
+
super().__init__()
|
| 104 |
+
self.output_dim = output_dim
|
| 105 |
+
self.input_resolution = input_resolution
|
| 106 |
+
|
| 107 |
+
# the 3-layer stem
|
| 108 |
+
self.conv1 = nn.Conv2d(3, width // 2, kernel_size=3, stride=2, padding=1, bias=False)
|
| 109 |
+
self.bn1 = nn.BatchNorm2d(width // 2)
|
| 110 |
+
self.relu1 = nn.ReLU(inplace=True)
|
| 111 |
+
self.conv2 = nn.Conv2d(width // 2, width // 2, kernel_size=3, padding=1, bias=False)
|
| 112 |
+
self.bn2 = nn.BatchNorm2d(width // 2)
|
| 113 |
+
self.relu2 = nn.ReLU(inplace=True)
|
| 114 |
+
self.conv3 = nn.Conv2d(width // 2, width, kernel_size=3, padding=1, bias=False)
|
| 115 |
+
self.bn3 = nn.BatchNorm2d(width)
|
| 116 |
+
self.relu3 = nn.ReLU(inplace=True)
|
| 117 |
+
self.avgpool = nn.AvgPool2d(2)
|
| 118 |
+
|
| 119 |
+
# residual layers
|
| 120 |
+
self._inplanes = width # this is a *mutable* variable used during construction
|
| 121 |
+
self.layer1 = self._make_layer(width, layers[0])
|
| 122 |
+
self.layer2 = self._make_layer(width * 2, layers[1], stride=2)
|
| 123 |
+
self.layer3 = self._make_layer(width * 4, layers[2], stride=2)
|
| 124 |
+
self.layer4 = self._make_layer(width * 8, layers[3], stride=2)
|
| 125 |
+
|
| 126 |
+
embed_dim = width * 32 # the ResNet feature dimension
|
| 127 |
+
self.attnpool = AttentionPool2d(input_resolution // 32, embed_dim, heads, output_dim)
|
| 128 |
+
|
| 129 |
+
def _make_layer(self, planes, blocks, stride=1):
|
| 130 |
+
layers = [Bottleneck(self._inplanes, planes, stride)]
|
| 131 |
+
|
| 132 |
+
self._inplanes = planes * Bottleneck.expansion
|
| 133 |
+
for _ in range(1, blocks):
|
| 134 |
+
layers.append(Bottleneck(self._inplanes, planes))
|
| 135 |
+
|
| 136 |
+
return nn.Sequential(*layers)
|
| 137 |
+
|
| 138 |
+
def forward(self, x):
|
| 139 |
+
def stem(x):
|
| 140 |
+
x = self.relu1(self.bn1(self.conv1(x)))
|
| 141 |
+
x = self.relu2(self.bn2(self.conv2(x)))
|
| 142 |
+
x = self.relu3(self.bn3(self.conv3(x)))
|
| 143 |
+
x = self.avgpool(x)
|
| 144 |
+
return x
|
| 145 |
+
|
| 146 |
+
x = x.type(self.conv1.weight.dtype)
|
| 147 |
+
x = stem(x)
|
| 148 |
+
x = self.layer1(x)
|
| 149 |
+
x = self.layer2(x)
|
| 150 |
+
x = self.layer3(x)
|
| 151 |
+
x = self.layer4(x)
|
| 152 |
+
x = self.attnpool(x)
|
| 153 |
+
|
| 154 |
+
return x
|
| 155 |
+
|
| 156 |
+
|
| 157 |
+
class LayerNorm(nn.LayerNorm):
|
| 158 |
+
"""Subclass torch's LayerNorm to handle fp16."""
|
| 159 |
+
|
| 160 |
+
def forward(self, x: torch.Tensor):
|
| 161 |
+
orig_type = x.dtype
|
| 162 |
+
ret = super().forward(x.type(torch.float32))
|
| 163 |
+
return ret.type(orig_type)
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
class QuickGELU(nn.Module):
|
| 167 |
+
def forward(self, x: torch.Tensor):
|
| 168 |
+
return x * torch.sigmoid(1.702 * x)
|
| 169 |
+
|
| 170 |
+
|
| 171 |
+
class ResidualAttentionBlock(nn.Module):
|
| 172 |
+
def __init__(self, d_model: int, n_head: int, attn_mask: torch.Tensor = None):
|
| 173 |
+
super().__init__()
|
| 174 |
+
|
| 175 |
+
self.attn = nn.MultiheadAttention(d_model, n_head)
|
| 176 |
+
self.ln_1 = LayerNorm(d_model)
|
| 177 |
+
self.mlp = nn.Sequential(OrderedDict([
|
| 178 |
+
("c_fc", nn.Linear(d_model, d_model * 4)),
|
| 179 |
+
("gelu", QuickGELU()),
|
| 180 |
+
("c_proj", nn.Linear(d_model * 4, d_model))
|
| 181 |
+
]))
|
| 182 |
+
self.ln_2 = LayerNorm(d_model)
|
| 183 |
+
self.attn_mask = attn_mask
|
| 184 |
+
|
| 185 |
+
def attention(self, x: torch.Tensor):
|
| 186 |
+
self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None
|
| 187 |
+
return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0]
|
| 188 |
+
|
| 189 |
+
def forward(self, x: torch.Tensor):
|
| 190 |
+
x = x + self.attention(self.ln_1(x))
|
| 191 |
+
x = x + self.mlp(self.ln_2(x))
|
| 192 |
+
return x
|
| 193 |
+
|
| 194 |
+
|
| 195 |
+
class Transformer(nn.Module):
|
| 196 |
+
def __init__(self, width: int, layers: int, heads: int, attn_mask: torch.Tensor = None):
|
| 197 |
+
super().__init__()
|
| 198 |
+
self.width = width
|
| 199 |
+
self.layers = layers
|
| 200 |
+
self.resblocks = nn.Sequential(*[ResidualAttentionBlock(width, heads, attn_mask) for _ in range(layers)])
|
| 201 |
+
|
| 202 |
+
def forward(self, x: torch.Tensor):
|
| 203 |
+
return self.resblocks(x)
|
| 204 |
+
|
| 205 |
+
|
| 206 |
+
class VisionTransformer(nn.Module):
|
| 207 |
+
def __init__(self, input_resolution: int, patch_size: int, width: int, layers: int, heads: int, output_dim: int):
|
| 208 |
+
super().__init__()
|
| 209 |
+
self.input_resolution = input_resolution
|
| 210 |
+
self.output_dim = output_dim
|
| 211 |
+
self.conv1 = nn.Conv2d(in_channels=3, out_channels=width, kernel_size=patch_size, stride=patch_size, bias=False)
|
| 212 |
+
|
| 213 |
+
scale = width ** -0.5
|
| 214 |
+
self.class_embedding = nn.Parameter(scale * torch.randn(width))
|
| 215 |
+
self.positional_embedding = nn.Parameter(scale * torch.randn((input_resolution // patch_size) ** 2 + 1, width))
|
| 216 |
+
self.ln_pre = LayerNorm(width)
|
| 217 |
+
|
| 218 |
+
self.transformer = Transformer(width, layers, heads)
|
| 219 |
+
|
| 220 |
+
self.ln_post = LayerNorm(width)
|
| 221 |
+
self.proj = nn.Parameter(scale * torch.randn(width, output_dim))
|
| 222 |
+
|
| 223 |
+
def forward(self, x: torch.Tensor):
|
| 224 |
+
x = self.conv1(x) # shape = [*, width, grid, grid]
|
| 225 |
+
x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2]
|
| 226 |
+
x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width]
|
| 227 |
+
x = torch.cat([self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), x], dim=1) # shape = [*, grid ** 2 + 1, width]
|
| 228 |
+
x = x + self.positional_embedding.to(x.dtype)
|
| 229 |
+
x = self.ln_pre(x)
|
| 230 |
+
|
| 231 |
+
x = x.permute(1, 0, 2) # NLD -> LND
|
| 232 |
+
x = self.transformer(x)
|
| 233 |
+
x = x.permute(1, 0, 2) # LND -> NLD
|
| 234 |
+
|
| 235 |
+
x = self.ln_post(x[:, 0, :])
|
| 236 |
+
|
| 237 |
+
if self.proj is not None:
|
| 238 |
+
x = x @ self.proj
|
| 239 |
+
|
| 240 |
+
return x
|
| 241 |
+
|
| 242 |
+
|
| 243 |
+
class CLIP(nn.Module):
|
| 244 |
+
def __init__(self,
|
| 245 |
+
embed_dim: int,
|
| 246 |
+
# vision
|
| 247 |
+
image_resolution: int,
|
| 248 |
+
vision_layers: Union[Tuple[int, int, int, int], int],
|
| 249 |
+
vision_width: int,
|
| 250 |
+
vision_patch_size: int,
|
| 251 |
+
# text
|
| 252 |
+
context_length: int,
|
| 253 |
+
vocab_size: int,
|
| 254 |
+
transformer_width: int,
|
| 255 |
+
transformer_heads: int,
|
| 256 |
+
transformer_layers: int
|
| 257 |
+
):
|
| 258 |
+
super().__init__()
|
| 259 |
+
|
| 260 |
+
self.context_length = context_length
|
| 261 |
+
|
| 262 |
+
if isinstance(vision_layers, (tuple, list)):
|
| 263 |
+
vision_heads = vision_width * 32 // 64
|
| 264 |
+
self.visual = ModifiedResNet(
|
| 265 |
+
layers=vision_layers,
|
| 266 |
+
output_dim=embed_dim,
|
| 267 |
+
heads=vision_heads,
|
| 268 |
+
input_resolution=image_resolution,
|
| 269 |
+
width=vision_width
|
| 270 |
+
)
|
| 271 |
+
else:
|
| 272 |
+
vision_heads = vision_width // 64
|
| 273 |
+
self.visual = VisionTransformer(
|
| 274 |
+
input_resolution=image_resolution,
|
| 275 |
+
patch_size=vision_patch_size,
|
| 276 |
+
width=vision_width,
|
| 277 |
+
layers=vision_layers,
|
| 278 |
+
heads=vision_heads,
|
| 279 |
+
output_dim=embed_dim
|
| 280 |
+
)
|
| 281 |
+
|
| 282 |
+
self.transformer = Transformer(
|
| 283 |
+
width=transformer_width,
|
| 284 |
+
layers=transformer_layers,
|
| 285 |
+
heads=transformer_heads,
|
| 286 |
+
attn_mask=self.build_attention_mask()
|
| 287 |
+
)
|
| 288 |
+
|
| 289 |
+
self.vocab_size = vocab_size
|
| 290 |
+
self.token_embedding = nn.Embedding(vocab_size, transformer_width)
|
| 291 |
+
self.positional_embedding = nn.Parameter(torch.empty(self.context_length, transformer_width))
|
| 292 |
+
self.ln_final = LayerNorm(transformer_width)
|
| 293 |
+
|
| 294 |
+
self.text_projection = nn.Parameter(torch.empty(transformer_width, embed_dim))
|
| 295 |
+
self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))
|
| 296 |
+
|
| 297 |
+
self.initialize_parameters()
|
| 298 |
+
|
| 299 |
+
def initialize_parameters(self):
|
| 300 |
+
nn.init.normal_(self.token_embedding.weight, std=0.02)
|
| 301 |
+
nn.init.normal_(self.positional_embedding, std=0.01)
|
| 302 |
+
|
| 303 |
+
if isinstance(self.visual, ModifiedResNet):
|
| 304 |
+
if self.visual.attnpool is not None:
|
| 305 |
+
std = self.visual.attnpool.c_proj.in_features ** -0.5
|
| 306 |
+
nn.init.normal_(self.visual.attnpool.q_proj.weight, std=std)
|
| 307 |
+
nn.init.normal_(self.visual.attnpool.k_proj.weight, std=std)
|
| 308 |
+
nn.init.normal_(self.visual.attnpool.v_proj.weight, std=std)
|
| 309 |
+
nn.init.normal_(self.visual.attnpool.c_proj.weight, std=std)
|
| 310 |
+
|
| 311 |
+
for resnet_block in [self.visual.layer1, self.visual.layer2, self.visual.layer3, self.visual.layer4]:
|
| 312 |
+
for name, param in resnet_block.named_parameters():
|
| 313 |
+
if name.endswith("bn3.weight"):
|
| 314 |
+
nn.init.zeros_(param)
|
| 315 |
+
|
| 316 |
+
proj_std = (self.transformer.width ** -0.5) * ((2 * self.transformer.layers) ** -0.5)
|
| 317 |
+
attn_std = self.transformer.width ** -0.5
|
| 318 |
+
fc_std = (2 * self.transformer.width) ** -0.5
|
| 319 |
+
for block in self.transformer.resblocks:
|
| 320 |
+
nn.init.normal_(block.attn.in_proj_weight, std=attn_std)
|
| 321 |
+
nn.init.normal_(block.attn.out_proj.weight, std=proj_std)
|
| 322 |
+
nn.init.normal_(block.mlp.c_fc.weight, std=fc_std)
|
| 323 |
+
nn.init.normal_(block.mlp.c_proj.weight, std=proj_std)
|
| 324 |
+
|
| 325 |
+
if self.text_projection is not None:
|
| 326 |
+
nn.init.normal_(self.text_projection, std=self.transformer.width ** -0.5)
|
| 327 |
+
|
| 328 |
+
def build_attention_mask(self):
|
| 329 |
+
# lazily create causal attention mask, with full attention between the vision tokens
|
| 330 |
+
# pytorch uses additive attention mask; fill with -inf
|
| 331 |
+
mask = torch.empty(self.context_length, self.context_length)
|
| 332 |
+
mask.fill_(float("-inf"))
|
| 333 |
+
mask.triu_(1) # zero out the lower diagonal
|
| 334 |
+
return mask
|
| 335 |
+
|
| 336 |
+
@property
|
| 337 |
+
def dtype(self):
|
| 338 |
+
return self.visual.conv1.weight.dtype
|
| 339 |
+
|
| 340 |
+
def encode_image(self, image):
|
| 341 |
+
return self.visual(image.type(self.dtype))
|
| 342 |
+
|
| 343 |
+
def encode_text(self, text):
|
| 344 |
+
x = self.token_embedding(text).type(self.dtype) # [batch_size, n_ctx, d_model]
|
| 345 |
+
|
| 346 |
+
x = x + self.positional_embedding.type(self.dtype)
|
| 347 |
+
x = x.permute(1, 0, 2) # NLD -> LND
|
| 348 |
+
x = self.transformer(x)
|
| 349 |
+
x = x.permute(1, 0, 2) # LND -> NLD
|
| 350 |
+
x = self.ln_final(x).type(self.dtype)
|
| 351 |
+
|
| 352 |
+
x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection
|
| 353 |
+
|
| 354 |
+
return x
|
| 355 |
+
|
| 356 |
+
def forward(self, image, text):
|
| 357 |
+
image_features = self.encode_image(image)
|
| 358 |
+
text_features = self.encode_text(text)
|
| 359 |
+
|
| 360 |
+
# normalized features
|
| 361 |
+
image_features = image_features / image_features.norm(dim=1, keepdim=True)
|
| 362 |
+
text_features = text_features / text_features.norm(dim=1, keepdim=True)
|
| 363 |
+
|
| 364 |
+
# cosine similarity as logits
|
| 365 |
+
logit_scale = self.logit_scale.exp()
|
| 366 |
+
logits_per_image = logit_scale * image_features @ text_features.t()
|
| 367 |
+
logits_per_text = logits_per_image.t()
|
| 368 |
+
|
| 369 |
+
# shape = [global_batch_size, global_batch_size]
|
| 370 |
+
return logits_per_image, logits_per_text
|
| 371 |
+
|
| 372 |
+
|
| 373 |
+
def convert_weights(model: nn.Module):
|
| 374 |
+
"""Convert applicable model parameters to fp16"""
|
| 375 |
+
|
| 376 |
+
def _convert_weights_to_fp16(l):
|
| 377 |
+
if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Linear)):
|
| 378 |
+
l.weight.data = l.weight.data.half()
|
| 379 |
+
if l.bias is not None:
|
| 380 |
+
l.bias.data = l.bias.data.half()
|
| 381 |
+
|
| 382 |
+
if isinstance(l, nn.MultiheadAttention):
|
| 383 |
+
for attr in [*[f"{s}_proj_weight" for s in ["in", "q", "k", "v"]], "in_proj_bias", "bias_k", "bias_v"]:
|
| 384 |
+
tensor = getattr(l, attr)
|
| 385 |
+
if tensor is not None:
|
| 386 |
+
tensor.data = tensor.data.half()
|
| 387 |
+
|
| 388 |
+
for name in ["text_projection", "proj"]:
|
| 389 |
+
if hasattr(l, name):
|
| 390 |
+
attr = getattr(l, name)
|
| 391 |
+
if attr is not None:
|
| 392 |
+
attr.data = attr.data.half()
|
| 393 |
+
|
| 394 |
+
model.apply(_convert_weights_to_fp16)
|
| 395 |
+
|
| 396 |
+
|
| 397 |
+
def build_model(state_dict: dict):
|
| 398 |
+
vit = "visual.proj" in state_dict
|
| 399 |
+
|
| 400 |
+
if vit:
|
| 401 |
+
vision_width = state_dict["visual.conv1.weight"].shape[0]
|
| 402 |
+
vision_layers = len([k for k in state_dict.keys() if k.startswith("visual.") and k.endswith(".attn.in_proj_weight")])
|
| 403 |
+
vision_patch_size = state_dict["visual.conv1.weight"].shape[-1]
|
| 404 |
+
grid_size = round((state_dict["visual.positional_embedding"].shape[0] - 1) ** 0.5)
|
| 405 |
+
image_resolution = vision_patch_size * grid_size
|
| 406 |
+
else:
|
| 407 |
+
counts: list = [len(set(k.split(".")[2] for k in state_dict if k.startswith(f"visual.layer{b}"))) for b in [1, 2, 3, 4]]
|
| 408 |
+
vision_layers = tuple(counts)
|
| 409 |
+
vision_width = state_dict["visual.layer1.0.conv1.weight"].shape[0]
|
| 410 |
+
output_width = round((state_dict["visual.attnpool.positional_embedding"].shape[0] - 1) ** 0.5)
|
| 411 |
+
vision_patch_size = None
|
| 412 |
+
assert output_width ** 2 + 1 == state_dict["visual.attnpool.positional_embedding"].shape[0]
|
| 413 |
+
image_resolution = output_width * 32
|
| 414 |
+
|
| 415 |
+
embed_dim = state_dict["text_projection"].shape[1]
|
| 416 |
+
context_length = state_dict["positional_embedding"].shape[0]
|
| 417 |
+
vocab_size = state_dict["token_embedding.weight"].shape[0]
|
| 418 |
+
transformer_width = state_dict["ln_final.weight"].shape[0]
|
| 419 |
+
transformer_heads = transformer_width // 64
|
| 420 |
+
transformer_layers = len(set(k.split(".")[2] for k in state_dict if k.startswith("transformer.resblocks")))
|
| 421 |
+
|
| 422 |
+
model = CLIP(
|
| 423 |
+
embed_dim,
|
| 424 |
+
image_resolution, vision_layers, vision_width, vision_patch_size,
|
| 425 |
+
context_length, vocab_size, transformer_width, transformer_heads, transformer_layers
|
| 426 |
+
)
|
| 427 |
+
|
| 428 |
+
for key in ["input_resolution", "context_length", "vocab_size"]:
|
| 429 |
+
if key in state_dict:
|
| 430 |
+
del state_dict[key]
|
| 431 |
+
|
| 432 |
+
convert_weights(model)
|
| 433 |
+
model.load_state_dict(state_dict)
|
| 434 |
+
return model.eval()
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/clip/simple_tokenizer.py
ADDED
|
@@ -0,0 +1,132 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import gzip
|
| 2 |
+
import html
|
| 3 |
+
import os
|
| 4 |
+
from functools import lru_cache
|
| 5 |
+
|
| 6 |
+
import ftfy
|
| 7 |
+
import regex as re
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
@lru_cache()
|
| 11 |
+
def default_bpe():
|
| 12 |
+
return os.path.join(os.path.dirname(os.path.abspath(__file__)), "bpe_simple_vocab_16e6.txt.gz")
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
@lru_cache()
|
| 16 |
+
def bytes_to_unicode():
|
| 17 |
+
"""
|
| 18 |
+
Returns list of utf-8 byte and a corresponding list of unicode strings.
|
| 19 |
+
The reversible bpe codes work on unicode strings.
|
| 20 |
+
This means you need a large # of unicode characters in your vocab if you want to avoid UNKs.
|
| 21 |
+
When you're at something like a 10B token dataset you end up needing around 5K for decent coverage.
|
| 22 |
+
This is a signficant percentage of your normal, say, 32K bpe vocab.
|
| 23 |
+
To avoid that, we want lookup tables between utf-8 bytes and unicode strings.
|
| 24 |
+
And avoids mapping to whitespace/control characters the bpe code barfs on.
|
| 25 |
+
"""
|
| 26 |
+
bs = list(range(ord("!"), ord("~")+1))+list(range(ord("¡"), ord("¬")+1))+list(range(ord("®"), ord("ÿ")+1))
|
| 27 |
+
cs = bs[:]
|
| 28 |
+
n = 0
|
| 29 |
+
for b in range(2**8):
|
| 30 |
+
if b not in bs:
|
| 31 |
+
bs.append(b)
|
| 32 |
+
cs.append(2**8+n)
|
| 33 |
+
n += 1
|
| 34 |
+
cs = [chr(n) for n in cs]
|
| 35 |
+
return dict(zip(bs, cs))
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def get_pairs(word):
|
| 39 |
+
"""Return set of symbol pairs in a word.
|
| 40 |
+
Word is represented as tuple of symbols (symbols being variable-length strings).
|
| 41 |
+
"""
|
| 42 |
+
pairs = set()
|
| 43 |
+
prev_char = word[0]
|
| 44 |
+
for char in word[1:]:
|
| 45 |
+
pairs.add((prev_char, char))
|
| 46 |
+
prev_char = char
|
| 47 |
+
return pairs
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def basic_clean(text):
|
| 51 |
+
text = ftfy.fix_text(text)
|
| 52 |
+
text = html.unescape(html.unescape(text))
|
| 53 |
+
return text.strip()
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def whitespace_clean(text):
|
| 57 |
+
text = re.sub(r'\s+', ' ', text)
|
| 58 |
+
text = text.strip()
|
| 59 |
+
return text
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
class SimpleTokenizer(object):
|
| 63 |
+
def __init__(self, bpe_path: str = default_bpe()):
|
| 64 |
+
self.byte_encoder = bytes_to_unicode()
|
| 65 |
+
self.byte_decoder = {v: k for k, v in self.byte_encoder.items()}
|
| 66 |
+
merges = gzip.open(bpe_path).read().decode("utf-8").split('\n')
|
| 67 |
+
merges = merges[1:49152-256-2+1]
|
| 68 |
+
merges = [tuple(merge.split()) for merge in merges]
|
| 69 |
+
vocab = list(bytes_to_unicode().values())
|
| 70 |
+
vocab = vocab + [v+'</w>' for v in vocab]
|
| 71 |
+
for merge in merges:
|
| 72 |
+
vocab.append(''.join(merge))
|
| 73 |
+
vocab.extend(['<|startoftext|>', '<|endoftext|>'])
|
| 74 |
+
self.encoder = dict(zip(vocab, range(len(vocab))))
|
| 75 |
+
self.decoder = {v: k for k, v in self.encoder.items()}
|
| 76 |
+
self.bpe_ranks = dict(zip(merges, range(len(merges))))
|
| 77 |
+
self.cache = {'<|startoftext|>': '<|startoftext|>', '<|endoftext|>': '<|endoftext|>'}
|
| 78 |
+
self.pat = re.compile(r"""<\|startoftext\|>|<\|endoftext\|>|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""", re.IGNORECASE)
|
| 79 |
+
|
| 80 |
+
def bpe(self, token):
|
| 81 |
+
if token in self.cache:
|
| 82 |
+
return self.cache[token]
|
| 83 |
+
word = tuple(token[:-1]) + ( token[-1] + '</w>',)
|
| 84 |
+
pairs = get_pairs(word)
|
| 85 |
+
|
| 86 |
+
if not pairs:
|
| 87 |
+
return token+'</w>'
|
| 88 |
+
|
| 89 |
+
while True:
|
| 90 |
+
bigram = min(pairs, key = lambda pair: self.bpe_ranks.get(pair, float('inf')))
|
| 91 |
+
if bigram not in self.bpe_ranks:
|
| 92 |
+
break
|
| 93 |
+
first, second = bigram
|
| 94 |
+
new_word = []
|
| 95 |
+
i = 0
|
| 96 |
+
while i < len(word):
|
| 97 |
+
try:
|
| 98 |
+
j = word.index(first, i)
|
| 99 |
+
new_word.extend(word[i:j])
|
| 100 |
+
i = j
|
| 101 |
+
except:
|
| 102 |
+
new_word.extend(word[i:])
|
| 103 |
+
break
|
| 104 |
+
|
| 105 |
+
if word[i] == first and i < len(word)-1 and word[i+1] == second:
|
| 106 |
+
new_word.append(first+second)
|
| 107 |
+
i += 2
|
| 108 |
+
else:
|
| 109 |
+
new_word.append(word[i])
|
| 110 |
+
i += 1
|
| 111 |
+
new_word = tuple(new_word)
|
| 112 |
+
word = new_word
|
| 113 |
+
if len(word) == 1:
|
| 114 |
+
break
|
| 115 |
+
else:
|
| 116 |
+
pairs = get_pairs(word)
|
| 117 |
+
word = ' '.join(word)
|
| 118 |
+
self.cache[token] = word
|
| 119 |
+
return word
|
| 120 |
+
|
| 121 |
+
def encode(self, text):
|
| 122 |
+
bpe_tokens = []
|
| 123 |
+
text = whitespace_clean(basic_clean(text)).lower()
|
| 124 |
+
for token in re.findall(self.pat, text):
|
| 125 |
+
token = ''.join(self.byte_encoder[b] for b in token.encode('utf-8'))
|
| 126 |
+
bpe_tokens.extend(self.encoder[bpe_token] for bpe_token in self.bpe(token).split(' '))
|
| 127 |
+
return bpe_tokens
|
| 128 |
+
|
| 129 |
+
def decode(self, tokens):
|
| 130 |
+
text = ''.join([self.decoder[token] for token in tokens])
|
| 131 |
+
text = bytearray([self.byte_decoder[c] for c in text]).decode('utf-8', errors="replace").replace('</w>', ' ')
|
| 132 |
+
return text
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/__init__.py
ADDED
|
@@ -0,0 +1,46 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Modified from OpenAI's diffusion repos
|
| 2 |
+
# GLIDE: https://github.com/openai/glide-text2im/blob/main/glide_text2im/gaussian_diffusion.py
|
| 3 |
+
# ADM: https://github.com/openai/guided-diffusion/blob/main/guided_diffusion
|
| 4 |
+
# IDDPM: https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
|
| 5 |
+
|
| 6 |
+
from . import gaussian_diffusion as gd
|
| 7 |
+
from .respace import SpacedDiffusion, space_timesteps
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def create_diffusion(
|
| 11 |
+
timestep_respacing,
|
| 12 |
+
noise_schedule="linear",
|
| 13 |
+
use_kl=False,
|
| 14 |
+
sigma_small=False,
|
| 15 |
+
predict_xstart=False,
|
| 16 |
+
learn_sigma=True,
|
| 17 |
+
rescale_learned_sigmas=False,
|
| 18 |
+
diffusion_steps=1000
|
| 19 |
+
):
|
| 20 |
+
betas = gd.get_named_beta_schedule(noise_schedule, diffusion_steps)
|
| 21 |
+
if use_kl:
|
| 22 |
+
loss_type = gd.LossType.RESCALED_KL
|
| 23 |
+
elif rescale_learned_sigmas:
|
| 24 |
+
loss_type = gd.LossType.RESCALED_MSE
|
| 25 |
+
else:
|
| 26 |
+
loss_type = gd.LossType.MSE
|
| 27 |
+
if timestep_respacing is None or timestep_respacing == "":
|
| 28 |
+
timestep_respacing = [diffusion_steps]
|
| 29 |
+
return SpacedDiffusion(
|
| 30 |
+
use_timesteps=space_timesteps(diffusion_steps, timestep_respacing),
|
| 31 |
+
betas=betas,
|
| 32 |
+
model_mean_type=(
|
| 33 |
+
gd.ModelMeanType.EPSILON if not predict_xstart else gd.ModelMeanType.START_X
|
| 34 |
+
),
|
| 35 |
+
model_var_type=(
|
| 36 |
+
(
|
| 37 |
+
gd.ModelVarType.FIXED_LARGE
|
| 38 |
+
if not sigma_small
|
| 39 |
+
else gd.ModelVarType.FIXED_SMALL
|
| 40 |
+
)
|
| 41 |
+
if not learn_sigma
|
| 42 |
+
else gd.ModelVarType.LEARNED_RANGE
|
| 43 |
+
),
|
| 44 |
+
loss_type=loss_type
|
| 45 |
+
# rescale_timesteps=rescale_timesteps,
|
| 46 |
+
)
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/__pycache__/__init__.cpython-310.pyc
ADDED
|
Binary file (1.06 kB). View file
|
|
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/__pycache__/diffusion_utils.cpython-310.pyc
ADDED
|
Binary file (2.88 kB). View file
|
|
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/__pycache__/gaussian_diffusion.cpython-310.pyc
ADDED
|
Binary file (25 kB). View file
|
|
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/__pycache__/respace.cpython-310.pyc
ADDED
|
Binary file (5.02 kB). View file
|
|
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/diffusion_utils.py
ADDED
|
@@ -0,0 +1,88 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Modified from OpenAI's diffusion repos
|
| 2 |
+
# GLIDE: https://github.com/openai/glide-text2im/blob/main/glide_text2im/gaussian_diffusion.py
|
| 3 |
+
# ADM: https://github.com/openai/guided-diffusion/blob/main/guided_diffusion
|
| 4 |
+
# IDDPM: https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
|
| 5 |
+
|
| 6 |
+
import torch as th
|
| 7 |
+
import numpy as np
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def normal_kl(mean1, logvar1, mean2, logvar2):
|
| 11 |
+
"""
|
| 12 |
+
Compute the KL divergence between two gaussians.
|
| 13 |
+
Shapes are automatically broadcasted, so batches can be compared to
|
| 14 |
+
scalars, among other use cases.
|
| 15 |
+
"""
|
| 16 |
+
tensor = None
|
| 17 |
+
for obj in (mean1, logvar1, mean2, logvar2):
|
| 18 |
+
if isinstance(obj, th.Tensor):
|
| 19 |
+
tensor = obj
|
| 20 |
+
break
|
| 21 |
+
assert tensor is not None, "at least one argument must be a Tensor"
|
| 22 |
+
|
| 23 |
+
# Force variances to be Tensors. Broadcasting helps convert scalars to
|
| 24 |
+
# Tensors, but it does not work for th.exp().
|
| 25 |
+
logvar1, logvar2 = [
|
| 26 |
+
x if isinstance(x, th.Tensor) else th.tensor(x).to(tensor)
|
| 27 |
+
for x in (logvar1, logvar2)
|
| 28 |
+
]
|
| 29 |
+
|
| 30 |
+
return 0.5 * (
|
| 31 |
+
-1.0
|
| 32 |
+
+ logvar2
|
| 33 |
+
- logvar1
|
| 34 |
+
+ th.exp(logvar1 - logvar2)
|
| 35 |
+
+ ((mean1 - mean2) ** 2) * th.exp(-logvar2)
|
| 36 |
+
)
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def approx_standard_normal_cdf(x):
|
| 40 |
+
"""
|
| 41 |
+
A fast approximation of the cumulative distribution function of the
|
| 42 |
+
standard normal.
|
| 43 |
+
"""
|
| 44 |
+
return 0.5 * (1.0 + th.tanh(np.sqrt(2.0 / np.pi) * (x + 0.044715 * th.pow(x, 3))))
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def continuous_gaussian_log_likelihood(x, *, means, log_scales):
|
| 48 |
+
"""
|
| 49 |
+
Compute the log-likelihood of a continuous Gaussian distribution.
|
| 50 |
+
:param x: the targets
|
| 51 |
+
:param means: the Gaussian mean Tensor.
|
| 52 |
+
:param log_scales: the Gaussian log stddev Tensor.
|
| 53 |
+
:return: a tensor like x of log probabilities (in nats).
|
| 54 |
+
"""
|
| 55 |
+
centered_x = x - means
|
| 56 |
+
inv_stdv = th.exp(-log_scales)
|
| 57 |
+
normalized_x = centered_x * inv_stdv
|
| 58 |
+
log_probs = th.distributions.Normal(th.zeros_like(x), th.ones_like(x)).log_prob(normalized_x)
|
| 59 |
+
return log_probs
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def discretized_gaussian_log_likelihood(x, *, means, log_scales):
|
| 63 |
+
"""
|
| 64 |
+
Compute the log-likelihood of a Gaussian distribution discretizing to a
|
| 65 |
+
given image.
|
| 66 |
+
:param x: the target images. It is assumed that this was uint8 values,
|
| 67 |
+
rescaled to the range [-1, 1].
|
| 68 |
+
:param means: the Gaussian mean Tensor.
|
| 69 |
+
:param log_scales: the Gaussian log stddev Tensor.
|
| 70 |
+
:return: a tensor like x of log probabilities (in nats).
|
| 71 |
+
"""
|
| 72 |
+
assert x.shape == means.shape == log_scales.shape
|
| 73 |
+
centered_x = x - means
|
| 74 |
+
inv_stdv = th.exp(-log_scales)
|
| 75 |
+
plus_in = inv_stdv * (centered_x + 1.0 / 255.0)
|
| 76 |
+
cdf_plus = approx_standard_normal_cdf(plus_in)
|
| 77 |
+
min_in = inv_stdv * (centered_x - 1.0 / 255.0)
|
| 78 |
+
cdf_min = approx_standard_normal_cdf(min_in)
|
| 79 |
+
log_cdf_plus = th.log(cdf_plus.clamp(min=1e-12))
|
| 80 |
+
log_one_minus_cdf_min = th.log((1.0 - cdf_min).clamp(min=1e-12))
|
| 81 |
+
cdf_delta = cdf_plus - cdf_min
|
| 82 |
+
log_probs = th.where(
|
| 83 |
+
x < -0.999,
|
| 84 |
+
log_cdf_plus,
|
| 85 |
+
th.where(x > 0.999, log_one_minus_cdf_min, th.log(cdf_delta.clamp(min=1e-12))),
|
| 86 |
+
)
|
| 87 |
+
assert log_probs.shape == x.shape
|
| 88 |
+
return log_probs
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/gaussian_diffusion.py
ADDED
|
@@ -0,0 +1,903 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Modified from OpenAI's diffusion repos
|
| 2 |
+
# GLIDE: https://github.com/openai/glide-text2im/blob/main/glide_text2im/gaussian_diffusion.py
|
| 3 |
+
# ADM: https://github.com/openai/guided-diffusion/blob/main/guided_diffusion
|
| 4 |
+
# IDDPM: https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
import math
|
| 8 |
+
|
| 9 |
+
import numpy as np
|
| 10 |
+
import torch as th
|
| 11 |
+
import enum
|
| 12 |
+
|
| 13 |
+
from .diffusion_utils import discretized_gaussian_log_likelihood, normal_kl
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def mean_flat(tensor):
|
| 17 |
+
"""
|
| 18 |
+
Take the mean over all non-batch dimensions.
|
| 19 |
+
"""
|
| 20 |
+
return tensor.mean(dim=list(range(1, len(tensor.shape))))
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
class ModelMeanType(enum.Enum):
|
| 24 |
+
"""
|
| 25 |
+
Which type of output the model predicts.
|
| 26 |
+
"""
|
| 27 |
+
|
| 28 |
+
PREVIOUS_X = enum.auto() # the model predicts x_{t-1}
|
| 29 |
+
START_X = enum.auto() # the model predicts x_0
|
| 30 |
+
EPSILON = enum.auto() # the model predicts epsilon
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
class ModelVarType(enum.Enum):
|
| 34 |
+
"""
|
| 35 |
+
What is used as the model's output variance.
|
| 36 |
+
The LEARNED_RANGE option has been added to allow the model to predict
|
| 37 |
+
values between FIXED_SMALL and FIXED_LARGE, making its job easier.
|
| 38 |
+
"""
|
| 39 |
+
|
| 40 |
+
LEARNED = enum.auto()
|
| 41 |
+
FIXED_SMALL = enum.auto()
|
| 42 |
+
FIXED_LARGE = enum.auto()
|
| 43 |
+
LEARNED_RANGE = enum.auto()
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
class LossType(enum.Enum):
|
| 47 |
+
MSE = enum.auto() # use raw MSE loss (and KL when learning variances)
|
| 48 |
+
RESCALED_MSE = (
|
| 49 |
+
enum.auto()
|
| 50 |
+
) # use raw MSE loss (with RESCALED_KL when learning variances)
|
| 51 |
+
KL = enum.auto() # use the variational lower-bound
|
| 52 |
+
RESCALED_KL = enum.auto() # like KL, but rescale to estimate the full VLB
|
| 53 |
+
|
| 54 |
+
def is_vb(self):
|
| 55 |
+
return self == LossType.KL or self == LossType.RESCALED_KL
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def _warmup_beta(beta_start, beta_end, num_diffusion_timesteps, warmup_frac):
|
| 59 |
+
betas = beta_end * np.ones(num_diffusion_timesteps, dtype=np.float64)
|
| 60 |
+
warmup_time = int(num_diffusion_timesteps * warmup_frac)
|
| 61 |
+
betas[:warmup_time] = np.linspace(beta_start, beta_end, warmup_time, dtype=np.float64)
|
| 62 |
+
return betas
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
def get_beta_schedule(beta_schedule, *, beta_start, beta_end, num_diffusion_timesteps):
|
| 66 |
+
"""
|
| 67 |
+
This is the deprecated API for creating beta schedules.
|
| 68 |
+
See get_named_beta_schedule() for the new library of schedules.
|
| 69 |
+
"""
|
| 70 |
+
if beta_schedule == "quad":
|
| 71 |
+
betas = (
|
| 72 |
+
np.linspace(
|
| 73 |
+
beta_start ** 0.5,
|
| 74 |
+
beta_end ** 0.5,
|
| 75 |
+
num_diffusion_timesteps,
|
| 76 |
+
dtype=np.float64,
|
| 77 |
+
)
|
| 78 |
+
** 2
|
| 79 |
+
)
|
| 80 |
+
elif beta_schedule == "linear":
|
| 81 |
+
betas = np.linspace(beta_start, beta_end, num_diffusion_timesteps, dtype=np.float64)
|
| 82 |
+
elif beta_schedule == "warmup10":
|
| 83 |
+
betas = _warmup_beta(beta_start, beta_end, num_diffusion_timesteps, 0.1)
|
| 84 |
+
elif beta_schedule == "warmup50":
|
| 85 |
+
betas = _warmup_beta(beta_start, beta_end, num_diffusion_timesteps, 0.5)
|
| 86 |
+
elif beta_schedule == "const":
|
| 87 |
+
betas = beta_end * np.ones(num_diffusion_timesteps, dtype=np.float64)
|
| 88 |
+
elif beta_schedule == "jsd": # 1/T, 1/(T-1), 1/(T-2), ..., 1
|
| 89 |
+
betas = 1.0 / np.linspace(
|
| 90 |
+
num_diffusion_timesteps, 1, num_diffusion_timesteps, dtype=np.float64
|
| 91 |
+
)
|
| 92 |
+
else:
|
| 93 |
+
raise NotImplementedError(beta_schedule)
|
| 94 |
+
assert betas.shape == (num_diffusion_timesteps,)
|
| 95 |
+
return betas
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
def get_named_beta_schedule(schedule_name, num_diffusion_timesteps):
|
| 99 |
+
"""
|
| 100 |
+
Get a pre-defined beta schedule for the given name.
|
| 101 |
+
The beta schedule library consists of beta schedules which remain similar
|
| 102 |
+
in the limit of num_diffusion_timesteps.
|
| 103 |
+
Beta schedules may be added, but should not be removed or changed once
|
| 104 |
+
they are committed to maintain backwards compatibility.
|
| 105 |
+
"""
|
| 106 |
+
if schedule_name == "linear":
|
| 107 |
+
# Linear schedule from Ho et al, extended to work for any number of
|
| 108 |
+
# diffusion steps.
|
| 109 |
+
scale = 1000 / num_diffusion_timesteps
|
| 110 |
+
return get_beta_schedule(
|
| 111 |
+
"linear",
|
| 112 |
+
beta_start=scale * 0.0001,
|
| 113 |
+
beta_end=scale * 0.02,
|
| 114 |
+
num_diffusion_timesteps=num_diffusion_timesteps,
|
| 115 |
+
)
|
| 116 |
+
elif schedule_name == "squaredcos_cap_v2":
|
| 117 |
+
return betas_for_alpha_bar(
|
| 118 |
+
num_diffusion_timesteps,
|
| 119 |
+
lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2,
|
| 120 |
+
)
|
| 121 |
+
else:
|
| 122 |
+
raise NotImplementedError(f"unknown beta schedule: {schedule_name}")
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
|
| 126 |
+
"""
|
| 127 |
+
Create a beta schedule that discretizes the given alpha_t_bar function,
|
| 128 |
+
which defines the cumulative product of (1-beta) over time from t = [0,1].
|
| 129 |
+
:param num_diffusion_timesteps: the number of betas to produce.
|
| 130 |
+
:param alpha_bar: a lambda that takes an argument t from 0 to 1 and
|
| 131 |
+
produces the cumulative product of (1-beta) up to that
|
| 132 |
+
part of the diffusion process.
|
| 133 |
+
:param max_beta: the maximum beta to use; use values lower than 1 to
|
| 134 |
+
prevent singularities.
|
| 135 |
+
"""
|
| 136 |
+
betas = []
|
| 137 |
+
for i in range(num_diffusion_timesteps):
|
| 138 |
+
t1 = i / num_diffusion_timesteps
|
| 139 |
+
t2 = (i + 1) / num_diffusion_timesteps
|
| 140 |
+
betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
|
| 141 |
+
return np.array(betas)
|
| 142 |
+
|
| 143 |
+
|
| 144 |
+
class GaussianDiffusion:
|
| 145 |
+
"""
|
| 146 |
+
Utilities for training and sampling diffusion models.
|
| 147 |
+
Original ported from this codebase:
|
| 148 |
+
https://github.com/hojonathanho/diffusion/blob/1e0dceb3b3495bbe19116a5e1b3596cd0706c543/diffusion_tf/diffusion_utils_2.py#L42
|
| 149 |
+
:param betas: a 1-D numpy array of betas for each diffusion timestep,
|
| 150 |
+
starting at T and going to 1.
|
| 151 |
+
"""
|
| 152 |
+
|
| 153 |
+
def __init__(
|
| 154 |
+
self,
|
| 155 |
+
*,
|
| 156 |
+
betas,
|
| 157 |
+
model_mean_type,
|
| 158 |
+
model_var_type,
|
| 159 |
+
loss_type
|
| 160 |
+
):
|
| 161 |
+
|
| 162 |
+
self.model_mean_type = model_mean_type
|
| 163 |
+
self.model_var_type = model_var_type
|
| 164 |
+
self.loss_type = loss_type
|
| 165 |
+
|
| 166 |
+
# Use float64 for accuracy.
|
| 167 |
+
betas = np.array(betas, dtype=np.float64)
|
| 168 |
+
self.betas = betas
|
| 169 |
+
assert len(betas.shape) == 1, "betas must be 1-D"
|
| 170 |
+
assert (betas > 0).all() and (betas <= 1).all()
|
| 171 |
+
|
| 172 |
+
self.num_timesteps = int(betas.shape[0])
|
| 173 |
+
|
| 174 |
+
alphas = 1.0 - betas
|
| 175 |
+
self.alphas_cumprod = np.cumprod(alphas, axis=0)
|
| 176 |
+
self.alphas_cumprod_prev = np.append(1.0, self.alphas_cumprod[:-1])
|
| 177 |
+
self.alphas_cumprod_next = np.append(self.alphas_cumprod[1:], 0.0)
|
| 178 |
+
assert self.alphas_cumprod_prev.shape == (self.num_timesteps,)
|
| 179 |
+
|
| 180 |
+
# calculations for diffusion q(x_t | x_{t-1}) and others
|
| 181 |
+
self.sqrt_alphas_cumprod = np.sqrt(self.alphas_cumprod)
|
| 182 |
+
self.sqrt_one_minus_alphas_cumprod = np.sqrt(1.0 - self.alphas_cumprod)
|
| 183 |
+
self.log_one_minus_alphas_cumprod = np.log(1.0 - self.alphas_cumprod)
|
| 184 |
+
self.sqrt_recip_alphas_cumprod = np.sqrt(1.0 / self.alphas_cumprod)
|
| 185 |
+
self.sqrt_recipm1_alphas_cumprod = np.sqrt(1.0 / self.alphas_cumprod - 1)
|
| 186 |
+
|
| 187 |
+
# calculations for posterior q(x_{t-1} | x_t, x_0)
|
| 188 |
+
self.posterior_variance = (
|
| 189 |
+
betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
|
| 190 |
+
)
|
| 191 |
+
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
| 192 |
+
self.posterior_log_variance_clipped = np.log(
|
| 193 |
+
np.append(self.posterior_variance[1], self.posterior_variance[1:])
|
| 194 |
+
) if len(self.posterior_variance) > 1 else np.array([])
|
| 195 |
+
|
| 196 |
+
self.posterior_mean_coef1 = (
|
| 197 |
+
betas * np.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
|
| 198 |
+
)
|
| 199 |
+
self.posterior_mean_coef2 = (
|
| 200 |
+
(1.0 - self.alphas_cumprod_prev) * np.sqrt(alphas) / (1.0 - self.alphas_cumprod)
|
| 201 |
+
)
|
| 202 |
+
|
| 203 |
+
# self.defect_w = th.nn.Parameter(torch.FloatTensor(1), requires_grad=True)
|
| 204 |
+
|
| 205 |
+
|
| 206 |
+
def q_mean_variance(self, x_start, t):
|
| 207 |
+
"""
|
| 208 |
+
Get the distribution q(x_t | x_0).
|
| 209 |
+
:param x_start: the [N x C x ...] tensor of noiseless inputs.
|
| 210 |
+
:param t: the number of diffusion steps (minus 1). Here, 0 means one step.
|
| 211 |
+
:return: A tuple (mean, variance, log_variance), all of x_start's shape.
|
| 212 |
+
"""
|
| 213 |
+
mean = _extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
|
| 214 |
+
variance = _extract_into_tensor(1.0 - self.alphas_cumprod, t, x_start.shape)
|
| 215 |
+
log_variance = _extract_into_tensor(self.log_one_minus_alphas_cumprod, t, x_start.shape)
|
| 216 |
+
return mean, variance, log_variance
|
| 217 |
+
|
| 218 |
+
def q_sample(self, x_start, t, noise=None):
|
| 219 |
+
"""
|
| 220 |
+
Diffuse the data for a given number of diffusion steps.
|
| 221 |
+
In other words, sample from q(x_t | x_0).
|
| 222 |
+
:param x_start: the initial data batch.
|
| 223 |
+
:param t: the number of diffusion steps (minus 1). Here, 0 means one step.
|
| 224 |
+
:param noise: if specified, the split-out normal noise.
|
| 225 |
+
:return: A noisy version of x_start.
|
| 226 |
+
"""
|
| 227 |
+
if noise is None:
|
| 228 |
+
noise = th.randn_like(x_start)
|
| 229 |
+
assert noise.shape == x_start.shape
|
| 230 |
+
return (
|
| 231 |
+
_extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
|
| 232 |
+
+ _extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
|
| 233 |
+
)
|
| 234 |
+
|
| 235 |
+
def q_posterior_mean_variance(self, x_start, x_t, t):
|
| 236 |
+
"""
|
| 237 |
+
Compute the mean and variance of the diffusion posterior:
|
| 238 |
+
q(x_{t-1} | x_t, x_0)
|
| 239 |
+
"""
|
| 240 |
+
assert x_start.shape == x_t.shape
|
| 241 |
+
posterior_mean = (
|
| 242 |
+
_extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start
|
| 243 |
+
+ _extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t
|
| 244 |
+
)
|
| 245 |
+
posterior_variance = _extract_into_tensor(self.posterior_variance, t, x_t.shape)
|
| 246 |
+
posterior_log_variance_clipped = _extract_into_tensor(
|
| 247 |
+
self.posterior_log_variance_clipped, t, x_t.shape
|
| 248 |
+
)
|
| 249 |
+
assert (
|
| 250 |
+
posterior_mean.shape[0]
|
| 251 |
+
== posterior_variance.shape[0]
|
| 252 |
+
== posterior_log_variance_clipped.shape[0]
|
| 253 |
+
== x_start.shape[0]
|
| 254 |
+
)
|
| 255 |
+
return posterior_mean, posterior_variance, posterior_log_variance_clipped
|
| 256 |
+
|
| 257 |
+
def p_mean_variance(self, model, x, t, clip_denoised=True, denoised_fn=None, model_kwargs=None):
|
| 258 |
+
"""
|
| 259 |
+
Apply the model to get p(x_{t-1} | x_t), as well as a prediction of
|
| 260 |
+
the initial x, x_0.
|
| 261 |
+
:param model: the model, which takes a signal and a batch of timesteps
|
| 262 |
+
as input.
|
| 263 |
+
:param x: the [N x C x ...] tensor at time t.
|
| 264 |
+
:param t: a 1-D Tensor of timesteps.
|
| 265 |
+
:param clip_denoised: if True, clip the denoised signal into [-1, 1].
|
| 266 |
+
:param denoised_fn: if not None, a function which applies to the
|
| 267 |
+
x_start prediction before it is used to sample. Applies before
|
| 268 |
+
clip_denoised.
|
| 269 |
+
:param model_kwargs: if not None, a dict of extra keyword arguments to
|
| 270 |
+
pass to the model. This can be used for conditioning.
|
| 271 |
+
:return: a dict with the following keys:
|
| 272 |
+
- 'mean': the model mean output.
|
| 273 |
+
- 'variance': the model variance output.
|
| 274 |
+
- 'log_variance': the log of 'variance'.
|
| 275 |
+
- 'pred_xstart': the prediction for x_0.
|
| 276 |
+
"""
|
| 277 |
+
if model_kwargs is None:
|
| 278 |
+
model_kwargs = {}
|
| 279 |
+
new_mask=None
|
| 280 |
+
B, C = x.shape[:2]
|
| 281 |
+
assert t.shape == (B,)
|
| 282 |
+
if model_kwargs == {}:
|
| 283 |
+
model_output = model(x, t, **model_kwargs)
|
| 284 |
+
else:
|
| 285 |
+
model_output, new_mask, _ = model(x, t, **model_kwargs)
|
| 286 |
+
if isinstance(model_output, tuple):
|
| 287 |
+
model_output, extra = model_output
|
| 288 |
+
else:
|
| 289 |
+
extra = None
|
| 290 |
+
|
| 291 |
+
if self.model_var_type in [ModelVarType.LEARNED, ModelVarType.LEARNED_RANGE]:
|
| 292 |
+
assert model_output.shape == (B, C * 2, *x.shape[2:])
|
| 293 |
+
model_output, model_var_values = th.split(model_output, C, dim=1)
|
| 294 |
+
min_log = _extract_into_tensor(self.posterior_log_variance_clipped, t, x.shape)
|
| 295 |
+
max_log = _extract_into_tensor(np.log(self.betas), t, x.shape)
|
| 296 |
+
# The model_var_values is [-1, 1] for [min_var, max_var].
|
| 297 |
+
frac = (model_var_values + 1) / 2
|
| 298 |
+
model_log_variance = frac * max_log + (1 - frac) * min_log
|
| 299 |
+
model_variance = th.exp(model_log_variance)
|
| 300 |
+
else:
|
| 301 |
+
model_variance, model_log_variance = {
|
| 302 |
+
# for fixedlarge, we set the initial (log-)variance like so
|
| 303 |
+
# to get a better decoder log likelihood.
|
| 304 |
+
ModelVarType.FIXED_LARGE: (
|
| 305 |
+
np.append(self.posterior_variance[1], self.betas[1:]),
|
| 306 |
+
np.log(np.append(self.posterior_variance[1], self.betas[1:])),
|
| 307 |
+
),
|
| 308 |
+
ModelVarType.FIXED_SMALL: (
|
| 309 |
+
self.posterior_variance,
|
| 310 |
+
self.posterior_log_variance_clipped,
|
| 311 |
+
),
|
| 312 |
+
}[self.model_var_type]
|
| 313 |
+
model_variance = _extract_into_tensor(model_variance, t, x.shape)
|
| 314 |
+
model_log_variance = _extract_into_tensor(model_log_variance, t, x.shape)
|
| 315 |
+
|
| 316 |
+
def process_xstart(x):
|
| 317 |
+
if denoised_fn is not None:
|
| 318 |
+
x = denoised_fn(x)
|
| 319 |
+
if clip_denoised:
|
| 320 |
+
return x.clamp(-1, 1)
|
| 321 |
+
return x
|
| 322 |
+
|
| 323 |
+
if self.model_mean_type == ModelMeanType.START_X:
|
| 324 |
+
pred_xstart = process_xstart(model_output)
|
| 325 |
+
else:
|
| 326 |
+
pred_xstart = process_xstart(
|
| 327 |
+
self._predict_xstart_from_eps(x_t=x, t=t, eps=model_output)
|
| 328 |
+
)
|
| 329 |
+
model_mean, _, _ = self.q_posterior_mean_variance(x_start=pred_xstart, x_t=x, t=t)
|
| 330 |
+
|
| 331 |
+
assert model_mean.shape == model_log_variance.shape == pred_xstart.shape == x.shape
|
| 332 |
+
if new_mask is None:
|
| 333 |
+
new_mask = pred_xstart
|
| 334 |
+
return {
|
| 335 |
+
"mean": model_mean,
|
| 336 |
+
"variance": model_variance,
|
| 337 |
+
"log_variance": model_log_variance,
|
| 338 |
+
"pred_xstart": pred_xstart,
|
| 339 |
+
"extra": extra,
|
| 340 |
+
"mask":new_mask,
|
| 341 |
+
}
|
| 342 |
+
|
| 343 |
+
def _predict_xstart_from_eps(self, x_t, t, eps):
|
| 344 |
+
assert x_t.shape == eps.shape
|
| 345 |
+
return (
|
| 346 |
+
_extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t
|
| 347 |
+
- _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * eps
|
| 348 |
+
)
|
| 349 |
+
|
| 350 |
+
def _predict_eps_from_xstart(self, x_t, t, pred_xstart):
|
| 351 |
+
return (
|
| 352 |
+
_extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - pred_xstart
|
| 353 |
+
) / _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
|
| 354 |
+
|
| 355 |
+
def condition_mean(self, cond_fn, p_mean_var, x, t, model_kwargs=None):
|
| 356 |
+
"""
|
| 357 |
+
Compute the mean for the previous step, given a function cond_fn that
|
| 358 |
+
computes the gradient of a conditional log probability with respect to
|
| 359 |
+
x. In particular, cond_fn computes grad(log(p(y|x))), and we want to
|
| 360 |
+
condition on y.
|
| 361 |
+
This uses the conditioning strategy from Sohl-Dickstein et al. (2015).
|
| 362 |
+
"""
|
| 363 |
+
gradient = cond_fn(x, t, **model_kwargs)
|
| 364 |
+
new_mean = p_mean_var["mean"].float() + p_mean_var["variance"] * gradient.float()
|
| 365 |
+
return new_mean
|
| 366 |
+
|
| 367 |
+
|
| 368 |
+
def condition_score(self, cond_fn, p_mean_var, x, t, model_kwargs=None):
|
| 369 |
+
"""
|
| 370 |
+
Compute what the p_mean_variance output would have been, should the
|
| 371 |
+
model's score function be conditioned by cond_fn.
|
| 372 |
+
See condition_mean() for details on cond_fn.
|
| 373 |
+
Unlike condition_mean(), this instead uses the conditioning strategy
|
| 374 |
+
from Song et al (2020).
|
| 375 |
+
"""
|
| 376 |
+
alpha_bar = _extract_into_tensor(self.alphas_cumprod, t, x.shape)
|
| 377 |
+
|
| 378 |
+
eps = self._predict_eps_from_xstart(x, t, p_mean_var["pred_xstart"])
|
| 379 |
+
eps = eps - (1 - alpha_bar).sqrt() * cond_fn(x, t, **model_kwargs)
|
| 380 |
+
|
| 381 |
+
out = p_mean_var.copy()
|
| 382 |
+
out["pred_xstart"] = self._predict_xstart_from_eps(x, t, eps)
|
| 383 |
+
out["mean"], _, _ = self.q_posterior_mean_variance(x_start=out["pred_xstart"], x_t=x, t=t)
|
| 384 |
+
return out
|
| 385 |
+
|
| 386 |
+
def p_sample(
|
| 387 |
+
self,
|
| 388 |
+
model,
|
| 389 |
+
x,
|
| 390 |
+
t,
|
| 391 |
+
clip_denoised=True,
|
| 392 |
+
denoised_fn=None,
|
| 393 |
+
cond_fn=None,
|
| 394 |
+
model_kwargs=None,
|
| 395 |
+
):
|
| 396 |
+
"""
|
| 397 |
+
Sample x_{t-1} from the model at the given timestep.
|
| 398 |
+
:param model: the model to sample from.
|
| 399 |
+
:param x: the current tensor at x_{t-1}.
|
| 400 |
+
:param t: the value of t, starting at 0 for the first diffusion step.
|
| 401 |
+
:param clip_denoised: if True, clip the x_start prediction to [-1, 1].
|
| 402 |
+
:param denoised_fn: if not None, a function which applies to the
|
| 403 |
+
x_start prediction before it is used to sample.
|
| 404 |
+
:param cond_fn: if not None, this is a gradient function that acts
|
| 405 |
+
similarly to the model.
|
| 406 |
+
:param model_kwargs: if not None, a dict of extra keyword arguments to
|
| 407 |
+
pass to the model. This can be used for conditioning.
|
| 408 |
+
:return: a dict containing the following keys:
|
| 409 |
+
- 'sample': a random sample from the model.
|
| 410 |
+
- 'pred_xstart': a prediction of x_0.
|
| 411 |
+
"""
|
| 412 |
+
out = self.p_mean_variance(
|
| 413 |
+
model,
|
| 414 |
+
x,
|
| 415 |
+
t,
|
| 416 |
+
clip_denoised=clip_denoised,
|
| 417 |
+
denoised_fn=denoised_fn,
|
| 418 |
+
model_kwargs=model_kwargs,
|
| 419 |
+
)
|
| 420 |
+
noise = th.randn_like(x)
|
| 421 |
+
nonzero_mask = (
|
| 422 |
+
(t != 0).float().view(-1, *([1] * (len(x.shape) - 1)))
|
| 423 |
+
) # no noise when t == 0
|
| 424 |
+
if cond_fn is not None:
|
| 425 |
+
out["mean"] = self.condition_mean(cond_fn, out, x, t, model_kwargs=model_kwargs)
|
| 426 |
+
sample = out["mean"] + nonzero_mask * th.exp(0.5 * out["log_variance"]) * noise
|
| 427 |
+
return {"sample": sample, "pred_xstart": out["pred_xstart"], "mask":out["mask"]}
|
| 428 |
+
|
| 429 |
+
def p_sample_loop(
|
| 430 |
+
self,
|
| 431 |
+
model,
|
| 432 |
+
shape,
|
| 433 |
+
noise=None,
|
| 434 |
+
clip_denoised=True,
|
| 435 |
+
denoised_fn=None,
|
| 436 |
+
cond_fn=None,
|
| 437 |
+
model_kwargs=None,
|
| 438 |
+
device=None,
|
| 439 |
+
progress=False,
|
| 440 |
+
):
|
| 441 |
+
"""
|
| 442 |
+
Generate samples from the model.
|
| 443 |
+
:param model: the model module.
|
| 444 |
+
:param shape: the shape of the samples, (N, C, H, W).
|
| 445 |
+
:param noise: if specified, the noise from the encoder to sample.
|
| 446 |
+
Should be of the same shape as `shape`.
|
| 447 |
+
:param clip_denoised: if True, clip x_start predictions to [-1, 1].
|
| 448 |
+
:param denoised_fn: if not None, a function which applies to the
|
| 449 |
+
x_start prediction before it is used to sample.
|
| 450 |
+
:param cond_fn: if not None, this is a gradient function that acts
|
| 451 |
+
similarly to the model.
|
| 452 |
+
:param model_kwargs: if not None, a dict of extra keyword arguments to
|
| 453 |
+
pass to the model. This can be used for conditioning.
|
| 454 |
+
:param device: if specified, the device to create the samples on.
|
| 455 |
+
If not specified, use a model parameter's device.
|
| 456 |
+
:param progress: if True, show a tqdm progress bar.
|
| 457 |
+
:return: a non-differentiable batch of samples.
|
| 458 |
+
"""
|
| 459 |
+
final = None
|
| 460 |
+
mask = 0
|
| 461 |
+
num = 0
|
| 462 |
+
for sample in self.p_sample_loop_progressive(
|
| 463 |
+
model,
|
| 464 |
+
shape,
|
| 465 |
+
noise=noise,
|
| 466 |
+
clip_denoised=clip_denoised,
|
| 467 |
+
denoised_fn=denoised_fn,
|
| 468 |
+
cond_fn=cond_fn,
|
| 469 |
+
model_kwargs=model_kwargs,
|
| 470 |
+
device=device,
|
| 471 |
+
progress=progress,
|
| 472 |
+
):
|
| 473 |
+
|
| 474 |
+
num += 1
|
| 475 |
+
final = sample
|
| 476 |
+
if num > 45:
|
| 477 |
+
mask += sample["mask"]
|
| 478 |
+
return final["sample"], mask / 5
|
| 479 |
+
|
| 480 |
+
def p_sample_loop_progressive(
|
| 481 |
+
self,
|
| 482 |
+
model,
|
| 483 |
+
shape,
|
| 484 |
+
noise=None,
|
| 485 |
+
clip_denoised=True,
|
| 486 |
+
denoised_fn=None,
|
| 487 |
+
cond_fn=None,
|
| 488 |
+
model_kwargs=None,
|
| 489 |
+
device=None,
|
| 490 |
+
progress=False,
|
| 491 |
+
):
|
| 492 |
+
"""
|
| 493 |
+
Generate samples from the model and yield intermediate samples from
|
| 494 |
+
each timestep of diffusion.
|
| 495 |
+
Arguments are the same as p_sample_loop().
|
| 496 |
+
Returns a generator over dicts, where each dict is the return value of
|
| 497 |
+
p_sample().
|
| 498 |
+
"""
|
| 499 |
+
if device is None:
|
| 500 |
+
device = next(model.parameters()).device
|
| 501 |
+
assert isinstance(shape, (tuple, list))
|
| 502 |
+
if noise is not None:
|
| 503 |
+
img = noise
|
| 504 |
+
else:
|
| 505 |
+
img = th.randn(*shape, device=device)
|
| 506 |
+
indices = list(range(self.num_timesteps))[::-1]
|
| 507 |
+
|
| 508 |
+
if progress:
|
| 509 |
+
# Lazy import so that we don't depend on tqdm.
|
| 510 |
+
from tqdm.auto import tqdm
|
| 511 |
+
|
| 512 |
+
indices = tqdm(indices)
|
| 513 |
+
|
| 514 |
+
for i in indices:
|
| 515 |
+
t = th.tensor([i] * shape[0], device=device)
|
| 516 |
+
with th.no_grad():
|
| 517 |
+
out = self.p_sample(
|
| 518 |
+
model,
|
| 519 |
+
img,
|
| 520 |
+
t,
|
| 521 |
+
clip_denoised=clip_denoised,
|
| 522 |
+
denoised_fn=denoised_fn,
|
| 523 |
+
cond_fn=cond_fn,
|
| 524 |
+
model_kwargs=model_kwargs,
|
| 525 |
+
)
|
| 526 |
+
yield out
|
| 527 |
+
img = out["sample"]
|
| 528 |
+
|
| 529 |
+
def ddim_sample(
|
| 530 |
+
self,
|
| 531 |
+
model,
|
| 532 |
+
x,
|
| 533 |
+
t,
|
| 534 |
+
clip_denoised=True,
|
| 535 |
+
denoised_fn=None,
|
| 536 |
+
cond_fn=None,
|
| 537 |
+
model_kwargs=None,
|
| 538 |
+
eta=0.0,
|
| 539 |
+
):
|
| 540 |
+
"""
|
| 541 |
+
Sample x_{t-1} from the model using DDIM.
|
| 542 |
+
Same usage as p_sample().
|
| 543 |
+
"""
|
| 544 |
+
out = self.p_mean_variance(
|
| 545 |
+
model,
|
| 546 |
+
x,
|
| 547 |
+
t,
|
| 548 |
+
clip_denoised=clip_denoised,
|
| 549 |
+
denoised_fn=denoised_fn,
|
| 550 |
+
model_kwargs=model_kwargs,
|
| 551 |
+
)
|
| 552 |
+
if cond_fn is not None:
|
| 553 |
+
out = self.condition_score(cond_fn, out, x, t, model_kwargs=model_kwargs)
|
| 554 |
+
|
| 555 |
+
# Usually our model outputs epsilon, but we re-derive it
|
| 556 |
+
# in case we used x_start or x_prev prediction.
|
| 557 |
+
eps = self._predict_eps_from_xstart(x, t, out["pred_xstart"])
|
| 558 |
+
|
| 559 |
+
alpha_bar = _extract_into_tensor(self.alphas_cumprod, t, x.shape)
|
| 560 |
+
alpha_bar_prev = _extract_into_tensor(self.alphas_cumprod_prev, t, x.shape)
|
| 561 |
+
sigma = (
|
| 562 |
+
eta
|
| 563 |
+
* th.sqrt((1 - alpha_bar_prev) / (1 - alpha_bar))
|
| 564 |
+
* th.sqrt(1 - alpha_bar / alpha_bar_prev)
|
| 565 |
+
)
|
| 566 |
+
# Equation 12.
|
| 567 |
+
noise = th.randn_like(x)
|
| 568 |
+
mean_pred = (
|
| 569 |
+
out["pred_xstart"] * th.sqrt(alpha_bar_prev)
|
| 570 |
+
+ th.sqrt(1 - alpha_bar_prev - sigma ** 2) * eps
|
| 571 |
+
)
|
| 572 |
+
nonzero_mask = (
|
| 573 |
+
(t != 0).float().view(-1, *([1] * (len(x.shape) - 1)))
|
| 574 |
+
) # no noise when t == 0
|
| 575 |
+
sample = mean_pred + nonzero_mask * sigma * noise
|
| 576 |
+
return {"sample": sample, "pred_xstart": out["pred_xstart"]}
|
| 577 |
+
|
| 578 |
+
def ddim_reverse_sample(
|
| 579 |
+
self,
|
| 580 |
+
model,
|
| 581 |
+
x,
|
| 582 |
+
t,
|
| 583 |
+
clip_denoised=True,
|
| 584 |
+
denoised_fn=None,
|
| 585 |
+
cond_fn=None,
|
| 586 |
+
model_kwargs=None,
|
| 587 |
+
eta=0.0,
|
| 588 |
+
):
|
| 589 |
+
"""
|
| 590 |
+
Sample x_{t+1} from the model using DDIM reverse ODE.
|
| 591 |
+
"""
|
| 592 |
+
assert eta == 0.0, "Reverse ODE only for deterministic path"
|
| 593 |
+
out = self.p_mean_variance(
|
| 594 |
+
model,
|
| 595 |
+
x,
|
| 596 |
+
t,
|
| 597 |
+
clip_denoised=clip_denoised,
|
| 598 |
+
denoised_fn=denoised_fn,
|
| 599 |
+
model_kwargs=model_kwargs,
|
| 600 |
+
)
|
| 601 |
+
if cond_fn is not None:
|
| 602 |
+
out = self.condition_score(cond_fn, out, x, t, model_kwargs=model_kwargs)
|
| 603 |
+
# Usually our model outputs epsilon, but we re-derive it
|
| 604 |
+
# in case we used x_start or x_prev prediction.
|
| 605 |
+
eps = (
|
| 606 |
+
_extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x.shape) * x
|
| 607 |
+
- out["pred_xstart"]
|
| 608 |
+
) / _extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x.shape)
|
| 609 |
+
alpha_bar_next = _extract_into_tensor(self.alphas_cumprod_next, t, x.shape)
|
| 610 |
+
|
| 611 |
+
# Equation 12. reversed
|
| 612 |
+
mean_pred = out["pred_xstart"] * th.sqrt(alpha_bar_next) + th.sqrt(1 - alpha_bar_next) * eps
|
| 613 |
+
|
| 614 |
+
return {"sample": mean_pred, "pred_xstart": out["pred_xstart"]}
|
| 615 |
+
|
| 616 |
+
def ddim_sample_loop(
|
| 617 |
+
self,
|
| 618 |
+
model,
|
| 619 |
+
shape,
|
| 620 |
+
noise=None,
|
| 621 |
+
clip_denoised=True,
|
| 622 |
+
denoised_fn=None,
|
| 623 |
+
cond_fn=None,
|
| 624 |
+
model_kwargs=None,
|
| 625 |
+
device=None,
|
| 626 |
+
progress=False,
|
| 627 |
+
eta=0.0,
|
| 628 |
+
):
|
| 629 |
+
"""
|
| 630 |
+
Generate samples from the model using DDIM.
|
| 631 |
+
Same usage as p_sample_loop().
|
| 632 |
+
"""
|
| 633 |
+
final = None
|
| 634 |
+
for sample in self.ddim_sample_loop_progressive(
|
| 635 |
+
model,
|
| 636 |
+
shape,
|
| 637 |
+
noise=noise,
|
| 638 |
+
clip_denoised=clip_denoised,
|
| 639 |
+
denoised_fn=denoised_fn,
|
| 640 |
+
cond_fn=cond_fn,
|
| 641 |
+
model_kwargs=model_kwargs,
|
| 642 |
+
device=device,
|
| 643 |
+
progress=progress,
|
| 644 |
+
eta=eta,
|
| 645 |
+
):
|
| 646 |
+
final = sample
|
| 647 |
+
return final["sample"]
|
| 648 |
+
|
| 649 |
+
def ddim_sample_loop_progressive(
|
| 650 |
+
self,
|
| 651 |
+
model,
|
| 652 |
+
shape,
|
| 653 |
+
noise=None,
|
| 654 |
+
clip_denoised=True,
|
| 655 |
+
denoised_fn=None,
|
| 656 |
+
cond_fn=None,
|
| 657 |
+
model_kwargs=None,
|
| 658 |
+
device=None,
|
| 659 |
+
progress=False,
|
| 660 |
+
eta=0.0,
|
| 661 |
+
):
|
| 662 |
+
"""
|
| 663 |
+
Use DDIM to sample from the model and yield intermediate samples from
|
| 664 |
+
each timestep of DDIM.
|
| 665 |
+
Same usage as p_sample_loop_progressive().
|
| 666 |
+
"""
|
| 667 |
+
if device is None:
|
| 668 |
+
device = next(model.parameters()).device
|
| 669 |
+
assert isinstance(shape, (tuple, list))
|
| 670 |
+
if noise is not None:
|
| 671 |
+
img = noise
|
| 672 |
+
else:
|
| 673 |
+
img = th.randn(*shape, device=device)
|
| 674 |
+
indices = list(range(self.num_timesteps))[::-1]
|
| 675 |
+
|
| 676 |
+
if progress:
|
| 677 |
+
# Lazy import so that we don't depend on tqdm.
|
| 678 |
+
from tqdm.auto import tqdm
|
| 679 |
+
|
| 680 |
+
indices = tqdm(indices)
|
| 681 |
+
|
| 682 |
+
for i in indices:
|
| 683 |
+
t = th.tensor([i] * shape[0], device=device)
|
| 684 |
+
with th.no_grad():
|
| 685 |
+
out = self.ddim_sample(
|
| 686 |
+
model,
|
| 687 |
+
img,
|
| 688 |
+
t,
|
| 689 |
+
clip_denoised=clip_denoised,
|
| 690 |
+
denoised_fn=denoised_fn,
|
| 691 |
+
cond_fn=cond_fn,
|
| 692 |
+
model_kwargs=model_kwargs,
|
| 693 |
+
eta=eta,
|
| 694 |
+
)
|
| 695 |
+
yield out
|
| 696 |
+
img = out["sample"]
|
| 697 |
+
|
| 698 |
+
def _vb_terms_bpd(
|
| 699 |
+
self, model, x_start, x_t, t, clip_denoised=True, model_kwargs=None
|
| 700 |
+
):
|
| 701 |
+
"""
|
| 702 |
+
Get a term for the variational lower-bound.
|
| 703 |
+
The resulting units are bits (rather than nats, as one might expect).
|
| 704 |
+
This allows for comparison to other papers.
|
| 705 |
+
:return: a dict with the following keys:
|
| 706 |
+
- 'output': a shape [N] tensor of NLLs or KLs.
|
| 707 |
+
- 'pred_xstart': the x_0 predictions.
|
| 708 |
+
"""
|
| 709 |
+
true_mean, _, true_log_variance_clipped = self.q_posterior_mean_variance(
|
| 710 |
+
x_start=x_start, x_t=x_t, t=t
|
| 711 |
+
)
|
| 712 |
+
out = self.p_mean_variance(
|
| 713 |
+
model, x_t, t, clip_denoised=clip_denoised, model_kwargs=model_kwargs
|
| 714 |
+
)
|
| 715 |
+
kl = normal_kl(
|
| 716 |
+
true_mean, true_log_variance_clipped, out["mean"], out["log_variance"]
|
| 717 |
+
)
|
| 718 |
+
kl = mean_flat(kl) / np.log(2.0)
|
| 719 |
+
|
| 720 |
+
decoder_nll = -discretized_gaussian_log_likelihood(
|
| 721 |
+
x_start, means=out["mean"], log_scales=0.5 * out["log_variance"]
|
| 722 |
+
)
|
| 723 |
+
assert decoder_nll.shape == x_start.shape
|
| 724 |
+
decoder_nll = mean_flat(decoder_nll) / np.log(2.0)
|
| 725 |
+
|
| 726 |
+
# At the first timestep return the decoder NLL,
|
| 727 |
+
# otherwise return KL(q(x_{t-1}|x_t,x_0) || p(x_{t-1}|x_t))
|
| 728 |
+
output = th.where((t == 0), decoder_nll, kl)
|
| 729 |
+
return {"output": output, "pred_xstart": out["pred_xstart"]}
|
| 730 |
+
|
| 731 |
+
def training_losses(self, model, x_start, t, model_kwargs=None, noise=None, label_mask=None, mask_resize=None, mask_att=None):
|
| 732 |
+
"""
|
| 733 |
+
Compute training losses for a single timestep.
|
| 734 |
+
:param model: the model to evaluate loss on.
|
| 735 |
+
:param x_start: the [N x C x ...] tensor of inputs.
|
| 736 |
+
:param t: a batch of timestep indices.
|
| 737 |
+
:param model_kwargs: if not None, a dict of extra keyword arguments to
|
| 738 |
+
pass to the model. This can be used for conditioning.
|
| 739 |
+
:param noise: if specified, the specific Gaussian noise to try to remove.
|
| 740 |
+
:param label_mask: tiao jie sun shi
|
| 741 |
+
:return: a dict with the key "loss" containing a tensor of shape [N].
|
| 742 |
+
Some mean or variance settings may also have other keys.
|
| 743 |
+
"""
|
| 744 |
+
|
| 745 |
+
if model_kwargs is None:
|
| 746 |
+
model_kwargs = {}
|
| 747 |
+
if noise is None:
|
| 748 |
+
noise = th.randn_like(x_start)
|
| 749 |
+
x_t = self.q_sample(x_start, t, noise=noise)
|
| 750 |
+
|
| 751 |
+
terms = {}
|
| 752 |
+
|
| 753 |
+
if self.loss_type == LossType.KL or self.loss_type == LossType.RESCALED_KL:
|
| 754 |
+
terms["loss"] = self._vb_terms_bpd(
|
| 755 |
+
model=model,
|
| 756 |
+
x_start=x_start,
|
| 757 |
+
x_t=x_t,
|
| 758 |
+
t=t,
|
| 759 |
+
clip_denoised=False,
|
| 760 |
+
model_kwargs=model_kwargs,
|
| 761 |
+
)["output"]
|
| 762 |
+
if self.loss_type == LossType.RESCALED_KL:
|
| 763 |
+
terms["loss"] *= self.num_timesteps
|
| 764 |
+
elif self.loss_type == LossType.MSE or self.loss_type == LossType.RESCALED_MSE:
|
| 765 |
+
model_output, att_mask, att_loss = model(x_t, t, **model_kwargs)
|
| 766 |
+
|
| 767 |
+
if self.model_var_type in [
|
| 768 |
+
ModelVarType.LEARNED,
|
| 769 |
+
ModelVarType.LEARNED_RANGE,
|
| 770 |
+
]:
|
| 771 |
+
B, C = x_t.shape[:2]
|
| 772 |
+
assert model_output.shape == (B, C * 2, *x_t.shape[2:])
|
| 773 |
+
model_output, model_var_values = th.split(model_output, C, dim=1)
|
| 774 |
+
# Learn the variance using the variational bound, but don't let
|
| 775 |
+
# it affect our mean prediction.
|
| 776 |
+
frozen_out = th.cat([model_output.detach(), model_var_values], dim=1)
|
| 777 |
+
terms["vb"] = self._vb_terms_bpd(
|
| 778 |
+
model=lambda *args, r=frozen_out: r,
|
| 779 |
+
x_start=x_start,
|
| 780 |
+
x_t=x_t,
|
| 781 |
+
t=t,
|
| 782 |
+
clip_denoised=False,
|
| 783 |
+
)["output"]
|
| 784 |
+
if self.loss_type == LossType.RESCALED_MSE:
|
| 785 |
+
# Divide by 1000 for equivalence with initial implementation.
|
| 786 |
+
# Without a factor of 1/1000, the VB term hurts the MSE term.
|
| 787 |
+
terms["vb"] *= self.num_timesteps / 1000.0
|
| 788 |
+
|
| 789 |
+
target = {
|
| 790 |
+
ModelMeanType.PREVIOUS_X: self.q_posterior_mean_variance(
|
| 791 |
+
x_start=x_start, x_t=x_t, t=t
|
| 792 |
+
)[0],
|
| 793 |
+
ModelMeanType.START_X: x_start,
|
| 794 |
+
ModelMeanType.EPSILON: noise,
|
| 795 |
+
}[self.model_mean_type]
|
| 796 |
+
assert model_output.shape == target.shape == x_start.shape
|
| 797 |
+
|
| 798 |
+
# print(att_loss.shape, mask_att.shape)
|
| 799 |
+
loss_defect = mean_flat(((th.mul(target, mask_resize) - th.mul(model_output, mask_resize)) ** 2))
|
| 800 |
+
# loss_back = mean_flat(((th.mul(target, 1 - mask_resize) - th.mul(model_output, 1 - mask_resize)) ** 2))
|
| 801 |
+
rat_loss = (th.sum(th.sum(th.mul(att_loss, 1 - mask_att), dim=-1), dim=-1)) / (th.sum(th.sum(th.mul(att_loss, mask_att), dim=-1), dim=-1) + 0.0001)
|
| 802 |
+
|
| 803 |
+
rat_loss[rat_loss > 8] = 8
|
| 804 |
+
rat_loss[rat_loss < 2] = 2
|
| 805 |
+
|
| 806 |
+
loss_att = rat_loss * loss_defect
|
| 807 |
+
|
| 808 |
+
terms["mse"] = mean_flat((target - model_output) ** 2)
|
| 809 |
+
terms["mask"] = mean_flat((att_mask - label_mask) ** 2)
|
| 810 |
+
if "vb" in terms:
|
| 811 |
+
terms["loss"] = terms["mse"] + terms["vb"] + 0.2 * terms["mask"] + loss_att
|
| 812 |
+
else:
|
| 813 |
+
terms["loss"] = terms["mse"] + 0.2 * terms["mask"] + loss_att
|
| 814 |
+
else:
|
| 815 |
+
raise NotImplementedError(self.loss_type)
|
| 816 |
+
|
| 817 |
+
return terms
|
| 818 |
+
|
| 819 |
+
def _prior_bpd(self, x_start):
|
| 820 |
+
"""
|
| 821 |
+
Get the prior KL term for the variational lower-bound, measured in
|
| 822 |
+
bits-per-dim.
|
| 823 |
+
This term can't be optimized, as it only depends on the encoder.
|
| 824 |
+
:param x_start: the [N x C x ...] tensor of inputs.
|
| 825 |
+
:return: a batch of [N] KL values (in bits), one per batch element.
|
| 826 |
+
"""
|
| 827 |
+
batch_size = x_start.shape[0]
|
| 828 |
+
t = th.tensor([self.num_timesteps - 1] * batch_size, device=x_start.device)
|
| 829 |
+
qt_mean, _, qt_log_variance = self.q_mean_variance(x_start, t)
|
| 830 |
+
kl_prior = normal_kl(
|
| 831 |
+
mean1=qt_mean, logvar1=qt_log_variance, mean2=0.0, logvar2=0.0
|
| 832 |
+
)
|
| 833 |
+
return mean_flat(kl_prior) / np.log(2.0)
|
| 834 |
+
|
| 835 |
+
def calc_bpd_loop(self, model, x_start, clip_denoised=True, model_kwargs=None):
|
| 836 |
+
"""
|
| 837 |
+
Compute the entire variational lower-bound, measured in bits-per-dim,
|
| 838 |
+
as well as other related quantities.
|
| 839 |
+
:param model: the model to evaluate loss on.
|
| 840 |
+
:param x_start: the [N x C x ...] tensor of inputs.
|
| 841 |
+
:param clip_denoised: if True, clip denoised samples.
|
| 842 |
+
:param model_kwargs: if not None, a dict of extra keyword arguments to
|
| 843 |
+
pass to the model. This can be used for conditioning.
|
| 844 |
+
:return: a dict containing the following keys:
|
| 845 |
+
- total_bpd: the total variational lower-bound, per batch element.
|
| 846 |
+
- prior_bpd: the prior term in the lower-bound.
|
| 847 |
+
- vb: an [N x T] tensor of terms in the lower-bound.
|
| 848 |
+
- xstart_mse: an [N x T] tensor of x_0 MSEs for each timestep.
|
| 849 |
+
- mse: an [N x T] tensor of epsilon MSEs for each timestep.
|
| 850 |
+
"""
|
| 851 |
+
device = x_start.device
|
| 852 |
+
batch_size = x_start.shape[0]
|
| 853 |
+
|
| 854 |
+
vb = []
|
| 855 |
+
xstart_mse = []
|
| 856 |
+
mse = []
|
| 857 |
+
for t in list(range(self.num_timesteps))[::-1]:
|
| 858 |
+
t_batch = th.tensor([t] * batch_size, device=device)
|
| 859 |
+
noise = th.randn_like(x_start)
|
| 860 |
+
x_t = self.q_sample(x_start=x_start, t=t_batch, noise=noise)
|
| 861 |
+
# Calculate VLB term at the current timestep
|
| 862 |
+
with th.no_grad():
|
| 863 |
+
out = self._vb_terms_bpd(
|
| 864 |
+
model,
|
| 865 |
+
x_start=x_start,
|
| 866 |
+
x_t=x_t,
|
| 867 |
+
t=t_batch,
|
| 868 |
+
clip_denoised=clip_denoised,
|
| 869 |
+
model_kwargs=model_kwargs,
|
| 870 |
+
)
|
| 871 |
+
vb.append(out["output"])
|
| 872 |
+
xstart_mse.append(mean_flat((out["pred_xstart"] - x_start) ** 2))
|
| 873 |
+
eps = self._predict_eps_from_xstart(x_t, t_batch, out["pred_xstart"])
|
| 874 |
+
mse.append(mean_flat((eps - noise) ** 2))
|
| 875 |
+
|
| 876 |
+
vb = th.stack(vb, dim=1)
|
| 877 |
+
xstart_mse = th.stack(xstart_mse, dim=1)
|
| 878 |
+
mse = th.stack(mse, dim=1)
|
| 879 |
+
|
| 880 |
+
prior_bpd = self._prior_bpd(x_start)
|
| 881 |
+
total_bpd = vb.sum(dim=1) + prior_bpd
|
| 882 |
+
return {
|
| 883 |
+
"total_bpd": total_bpd,
|
| 884 |
+
"prior_bpd": prior_bpd,
|
| 885 |
+
"vb": vb,
|
| 886 |
+
"xstart_mse": xstart_mse,
|
| 887 |
+
"mse": mse,
|
| 888 |
+
}
|
| 889 |
+
|
| 890 |
+
|
| 891 |
+
def _extract_into_tensor(arr, timesteps, broadcast_shape):
|
| 892 |
+
"""
|
| 893 |
+
Extract values from a 1-D numpy array for a batch of indices.
|
| 894 |
+
:param arr: the 1-D numpy array.
|
| 895 |
+
:param timesteps: a tensor of indices into the array to extract.
|
| 896 |
+
:param broadcast_shape: a larger shape of K dimensions with the batch
|
| 897 |
+
dimension equal to the length of timesteps.
|
| 898 |
+
:return: a tensor of shape [batch_size, 1, ...] where the shape has K dims.
|
| 899 |
+
"""
|
| 900 |
+
res = th.from_numpy(arr).to(device=timesteps.device)[timesteps].float()
|
| 901 |
+
while len(res.shape) < len(broadcast_shape):
|
| 902 |
+
res = res[..., None]
|
| 903 |
+
return res + th.zeros(broadcast_shape, device=timesteps.device)
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/respace.py
ADDED
|
@@ -0,0 +1,129 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Modified from OpenAI's diffusion repos
|
| 2 |
+
# GLIDE: https://github.com/openai/glide-text2im/blob/main/glide_text2im/gaussian_diffusion.py
|
| 3 |
+
# ADM: https://github.com/openai/guided-diffusion/blob/main/guided_diffusion
|
| 4 |
+
# IDDPM: https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
|
| 5 |
+
|
| 6 |
+
import numpy as np
|
| 7 |
+
import torch as th
|
| 8 |
+
|
| 9 |
+
from .gaussian_diffusion import GaussianDiffusion
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def space_timesteps(num_timesteps, section_counts):
|
| 13 |
+
"""
|
| 14 |
+
Create a list of timesteps to use from an original diffusion process,
|
| 15 |
+
given the number of timesteps we want to take from equally-sized portions
|
| 16 |
+
of the original process.
|
| 17 |
+
For example, if there's 300 timesteps and the section counts are [10,15,20]
|
| 18 |
+
then the first 100 timesteps are strided to be 10 timesteps, the second 100
|
| 19 |
+
are strided to be 15 timesteps, and the final 100 are strided to be 20.
|
| 20 |
+
If the stride is a string starting with "ddim", then the fixed striding
|
| 21 |
+
from the DDIM paper is used, and only one section is allowed.
|
| 22 |
+
:param num_timesteps: the number of diffusion steps in the original
|
| 23 |
+
process to divide up.
|
| 24 |
+
:param section_counts: either a list of numbers, or a string containing
|
| 25 |
+
comma-separated numbers, indicating the step count
|
| 26 |
+
per section. As a special case, use "ddimN" where N
|
| 27 |
+
is a number of steps to use the striding from the
|
| 28 |
+
DDIM paper.
|
| 29 |
+
:return: a set of diffusion steps from the original process to use.
|
| 30 |
+
"""
|
| 31 |
+
if isinstance(section_counts, str):
|
| 32 |
+
if section_counts.startswith("ddim"):
|
| 33 |
+
desired_count = int(section_counts[len("ddim") :])
|
| 34 |
+
for i in range(1, num_timesteps):
|
| 35 |
+
if len(range(0, num_timesteps, i)) == desired_count:
|
| 36 |
+
return set(range(0, num_timesteps, i))
|
| 37 |
+
raise ValueError(
|
| 38 |
+
f"cannot create exactly {num_timesteps} steps with an integer stride"
|
| 39 |
+
)
|
| 40 |
+
section_counts = [int(x) for x in section_counts.split(",")]
|
| 41 |
+
size_per = num_timesteps // len(section_counts)
|
| 42 |
+
extra = num_timesteps % len(section_counts)
|
| 43 |
+
start_idx = 0
|
| 44 |
+
all_steps = []
|
| 45 |
+
for i, section_count in enumerate(section_counts):
|
| 46 |
+
size = size_per + (1 if i < extra else 0)
|
| 47 |
+
if size < section_count:
|
| 48 |
+
raise ValueError(
|
| 49 |
+
f"cannot divide section of {size} steps into {section_count}"
|
| 50 |
+
)
|
| 51 |
+
if section_count <= 1:
|
| 52 |
+
frac_stride = 1
|
| 53 |
+
else:
|
| 54 |
+
frac_stride = (size - 1) / (section_count - 1)
|
| 55 |
+
cur_idx = 0.0
|
| 56 |
+
taken_steps = []
|
| 57 |
+
for _ in range(section_count):
|
| 58 |
+
taken_steps.append(start_idx + round(cur_idx))
|
| 59 |
+
cur_idx += frac_stride
|
| 60 |
+
all_steps += taken_steps
|
| 61 |
+
start_idx += size
|
| 62 |
+
return set(all_steps)
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
class SpacedDiffusion(GaussianDiffusion):
|
| 66 |
+
"""
|
| 67 |
+
A diffusion process which can skip steps in a base diffusion process.
|
| 68 |
+
:param use_timesteps: a collection (sequence or set) of timesteps from the
|
| 69 |
+
original diffusion process to retain.
|
| 70 |
+
:param kwargs: the kwargs to create the base diffusion process.
|
| 71 |
+
"""
|
| 72 |
+
|
| 73 |
+
def __init__(self, use_timesteps, **kwargs):
|
| 74 |
+
self.use_timesteps = set(use_timesteps)
|
| 75 |
+
self.timestep_map = []
|
| 76 |
+
self.original_num_steps = len(kwargs["betas"])
|
| 77 |
+
|
| 78 |
+
base_diffusion = GaussianDiffusion(**kwargs) # pylint: disable=missing-kwoa
|
| 79 |
+
last_alpha_cumprod = 1.0
|
| 80 |
+
new_betas = []
|
| 81 |
+
for i, alpha_cumprod in enumerate(base_diffusion.alphas_cumprod):
|
| 82 |
+
if i in self.use_timesteps:
|
| 83 |
+
new_betas.append(1 - alpha_cumprod / last_alpha_cumprod)
|
| 84 |
+
last_alpha_cumprod = alpha_cumprod
|
| 85 |
+
self.timestep_map.append(i)
|
| 86 |
+
kwargs["betas"] = np.array(new_betas)
|
| 87 |
+
super().__init__(**kwargs)
|
| 88 |
+
|
| 89 |
+
def p_mean_variance(
|
| 90 |
+
self, model, *args, **kwargs
|
| 91 |
+
): # pylint: disable=signature-differs
|
| 92 |
+
return super().p_mean_variance(self._wrap_model(model), *args, **kwargs)
|
| 93 |
+
|
| 94 |
+
def training_losses(
|
| 95 |
+
self, model, *args, **kwargs
|
| 96 |
+
): # pylint: disable=signature-differs
|
| 97 |
+
return super().training_losses(self._wrap_model(model), *args, **kwargs)
|
| 98 |
+
|
| 99 |
+
def condition_mean(self, cond_fn, *args, **kwargs):
|
| 100 |
+
return super().condition_mean(self._wrap_model(cond_fn), *args, **kwargs)
|
| 101 |
+
|
| 102 |
+
def condition_score(self, cond_fn, *args, **kwargs):
|
| 103 |
+
return super().condition_score(self._wrap_model(cond_fn), *args, **kwargs)
|
| 104 |
+
|
| 105 |
+
def _wrap_model(self, model):
|
| 106 |
+
if isinstance(model, _WrappedModel):
|
| 107 |
+
return model
|
| 108 |
+
return _WrappedModel(
|
| 109 |
+
model, self.timestep_map, self.original_num_steps
|
| 110 |
+
)
|
| 111 |
+
|
| 112 |
+
def _scale_timesteps(self, t):
|
| 113 |
+
# Scaling is done by the wrapped model.
|
| 114 |
+
return t
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
class _WrappedModel:
|
| 118 |
+
def __init__(self, model, timestep_map, original_num_steps):
|
| 119 |
+
self.model = model
|
| 120 |
+
self.timestep_map = timestep_map
|
| 121 |
+
# self.rescale_timesteps = rescale_timesteps
|
| 122 |
+
self.original_num_steps = original_num_steps
|
| 123 |
+
|
| 124 |
+
def __call__(self, x, ts, **kwargs):
|
| 125 |
+
map_tensor = th.tensor(self.timestep_map, device=ts.device, dtype=ts.dtype)
|
| 126 |
+
new_ts = map_tensor[ts]
|
| 127 |
+
# if self.rescale_timesteps:
|
| 128 |
+
# new_ts = new_ts.float() * (1000.0 / self.original_num_steps)
|
| 129 |
+
return self.model(x, new_ts, **kwargs)
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/diffusion/timestep_sampler.py
ADDED
|
@@ -0,0 +1,150 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Modified from OpenAI's diffusion repos
|
| 2 |
+
# GLIDE: https://github.com/openai/glide-text2im/blob/main/glide_text2im/gaussian_diffusion.py
|
| 3 |
+
# ADM: https://github.com/openai/guided-diffusion/blob/main/guided_diffusion
|
| 4 |
+
# IDDPM: https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
|
| 5 |
+
|
| 6 |
+
from abc import ABC, abstractmethod
|
| 7 |
+
|
| 8 |
+
import numpy as np
|
| 9 |
+
import torch as th
|
| 10 |
+
import torch.distributed as dist
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def create_named_schedule_sampler(name, diffusion):
|
| 14 |
+
"""
|
| 15 |
+
Create a ScheduleSampler from a library of pre-defined samplers.
|
| 16 |
+
:param name: the name of the sampler.
|
| 17 |
+
:param diffusion: the diffusion object to sample for.
|
| 18 |
+
"""
|
| 19 |
+
if name == "uniform":
|
| 20 |
+
return UniformSampler(diffusion)
|
| 21 |
+
elif name == "loss-second-moment":
|
| 22 |
+
return LossSecondMomentResampler(diffusion)
|
| 23 |
+
else:
|
| 24 |
+
raise NotImplementedError(f"unknown schedule sampler: {name}")
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
class ScheduleSampler(ABC):
|
| 28 |
+
"""
|
| 29 |
+
A distribution over timesteps in the diffusion process, intended to reduce
|
| 30 |
+
variance of the objective.
|
| 31 |
+
By default, samplers perform unbiased importance sampling, in which the
|
| 32 |
+
objective's mean is unchanged.
|
| 33 |
+
However, subclasses may override sample() to change how the resampled
|
| 34 |
+
terms are reweighted, allowing for actual changes in the objective.
|
| 35 |
+
"""
|
| 36 |
+
|
| 37 |
+
@abstractmethod
|
| 38 |
+
def weights(self):
|
| 39 |
+
"""
|
| 40 |
+
Get a numpy array of weights, one per diffusion step.
|
| 41 |
+
The weights needn't be normalized, but must be positive.
|
| 42 |
+
"""
|
| 43 |
+
|
| 44 |
+
def sample(self, batch_size, device):
|
| 45 |
+
"""
|
| 46 |
+
Importance-sample timesteps for a batch.
|
| 47 |
+
:param batch_size: the number of timesteps.
|
| 48 |
+
:param device: the torch device to save to.
|
| 49 |
+
:return: a tuple (timesteps, weights):
|
| 50 |
+
- timesteps: a tensor of timestep indices.
|
| 51 |
+
- weights: a tensor of weights to scale the resulting losses.
|
| 52 |
+
"""
|
| 53 |
+
w = self.weights()
|
| 54 |
+
p = w / np.sum(w)
|
| 55 |
+
indices_np = np.random.choice(len(p), size=(batch_size,), p=p)
|
| 56 |
+
indices = th.from_numpy(indices_np).long().to(device)
|
| 57 |
+
weights_np = 1 / (len(p) * p[indices_np])
|
| 58 |
+
weights = th.from_numpy(weights_np).float().to(device)
|
| 59 |
+
return indices, weights
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
class UniformSampler(ScheduleSampler):
|
| 63 |
+
def __init__(self, diffusion):
|
| 64 |
+
self.diffusion = diffusion
|
| 65 |
+
self._weights = np.ones([diffusion.num_timesteps])
|
| 66 |
+
|
| 67 |
+
def weights(self):
|
| 68 |
+
return self._weights
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
class LossAwareSampler(ScheduleSampler):
|
| 72 |
+
def update_with_local_losses(self, local_ts, local_losses):
|
| 73 |
+
"""
|
| 74 |
+
Update the reweighting using losses from a model.
|
| 75 |
+
Call this method from each rank with a batch of timesteps and the
|
| 76 |
+
corresponding losses for each of those timesteps.
|
| 77 |
+
This method will perform synchronization to make sure all of the ranks
|
| 78 |
+
maintain the exact same reweighting.
|
| 79 |
+
:param local_ts: an integer Tensor of timesteps.
|
| 80 |
+
:param local_losses: a 1D Tensor of losses.
|
| 81 |
+
"""
|
| 82 |
+
batch_sizes = [
|
| 83 |
+
th.tensor([0], dtype=th.int32, device=local_ts.device)
|
| 84 |
+
for _ in range(dist.get_world_size())
|
| 85 |
+
]
|
| 86 |
+
dist.all_gather(
|
| 87 |
+
batch_sizes,
|
| 88 |
+
th.tensor([len(local_ts)], dtype=th.int32, device=local_ts.device),
|
| 89 |
+
)
|
| 90 |
+
|
| 91 |
+
# Pad all_gather batches to be the maximum batch size.
|
| 92 |
+
batch_sizes = [x.item() for x in batch_sizes]
|
| 93 |
+
max_bs = max(batch_sizes)
|
| 94 |
+
|
| 95 |
+
timestep_batches = [th.zeros(max_bs).to(local_ts) for bs in batch_sizes]
|
| 96 |
+
loss_batches = [th.zeros(max_bs).to(local_losses) for bs in batch_sizes]
|
| 97 |
+
dist.all_gather(timestep_batches, local_ts)
|
| 98 |
+
dist.all_gather(loss_batches, local_losses)
|
| 99 |
+
timesteps = [
|
| 100 |
+
x.item() for y, bs in zip(timestep_batches, batch_sizes) for x in y[:bs]
|
| 101 |
+
]
|
| 102 |
+
losses = [x.item() for y, bs in zip(loss_batches, batch_sizes) for x in y[:bs]]
|
| 103 |
+
self.update_with_all_losses(timesteps, losses)
|
| 104 |
+
|
| 105 |
+
@abstractmethod
|
| 106 |
+
def update_with_all_losses(self, ts, losses):
|
| 107 |
+
"""
|
| 108 |
+
Update the reweighting using losses from a model.
|
| 109 |
+
Sub-classes should override this method to update the reweighting
|
| 110 |
+
using losses from the model.
|
| 111 |
+
This method directly updates the reweighting without synchronizing
|
| 112 |
+
between workers. It is called by update_with_local_losses from all
|
| 113 |
+
ranks with identical arguments. Thus, it should have deterministic
|
| 114 |
+
behavior to maintain state across workers.
|
| 115 |
+
:param ts: a list of int timesteps.
|
| 116 |
+
:param losses: a list of float losses, one per timestep.
|
| 117 |
+
"""
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
class LossSecondMomentResampler(LossAwareSampler):
|
| 121 |
+
def __init__(self, diffusion, history_per_term=10, uniform_prob=0.001):
|
| 122 |
+
self.diffusion = diffusion
|
| 123 |
+
self.history_per_term = history_per_term
|
| 124 |
+
self.uniform_prob = uniform_prob
|
| 125 |
+
self._loss_history = np.zeros(
|
| 126 |
+
[diffusion.num_timesteps, history_per_term], dtype=np.float64
|
| 127 |
+
)
|
| 128 |
+
self._loss_counts = np.zeros([diffusion.num_timesteps], dtype=np.int)
|
| 129 |
+
|
| 130 |
+
def weights(self):
|
| 131 |
+
if not self._warmed_up():
|
| 132 |
+
return np.ones([self.diffusion.num_timesteps], dtype=np.float64)
|
| 133 |
+
weights = np.sqrt(np.mean(self._loss_history ** 2, axis=-1))
|
| 134 |
+
weights /= np.sum(weights)
|
| 135 |
+
weights *= 1 - self.uniform_prob
|
| 136 |
+
weights += self.uniform_prob / len(weights)
|
| 137 |
+
return weights
|
| 138 |
+
|
| 139 |
+
def update_with_all_losses(self, ts, losses):
|
| 140 |
+
for t, loss in zip(ts, losses):
|
| 141 |
+
if self._loss_counts[t] == self.history_per_term:
|
| 142 |
+
# Shift out the oldest loss term.
|
| 143 |
+
self._loss_history[t, :-1] = self._loss_history[t, 1:]
|
| 144 |
+
self._loss_history[t, -1] = loss
|
| 145 |
+
else:
|
| 146 |
+
self._loss_history[t, self._loss_counts[t]] = loss
|
| 147 |
+
self._loss_counts[t] += 1
|
| 148 |
+
|
| 149 |
+
def _warmed_up(self):
|
| 150 |
+
return (self._loss_counts == self.history_per_term).all()
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/models_add_cross_concate.py
ADDED
|
@@ -0,0 +1,498 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
| 2 |
+
# All rights reserved.
|
| 3 |
+
|
| 4 |
+
# This source code is licensed under the license found in the
|
| 5 |
+
# LICENSE file in the root directory of this source tree.
|
| 6 |
+
# --------------------------------------------------------
|
| 7 |
+
# References:
|
| 8 |
+
# GLIDE: https://github.com/openai/glide-text2im
|
| 9 |
+
# MAE: https://github.com/facebookresearch/mae/blob/main/models_mae.py
|
| 10 |
+
# --------------------------------------------------------
|
| 11 |
+
|
| 12 |
+
import torch
|
| 13 |
+
import torch.nn as nn
|
| 14 |
+
import numpy as np
|
| 15 |
+
import math
|
| 16 |
+
from timm.models.vision_transformer import PatchEmbed, Attention, Mlp
|
| 17 |
+
from torch import einsum
|
| 18 |
+
from einops import rearrange, repeat
|
| 19 |
+
from autoencoder import *
|
| 20 |
+
|
| 21 |
+
def modulate(x, shift, scale):
|
| 22 |
+
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
#################################################################################
|
| 26 |
+
# Embedding Layers for Timesteps and Class Labels #
|
| 27 |
+
#################################################################################
|
| 28 |
+
|
| 29 |
+
class TimestepEmbedder(nn.Module):
|
| 30 |
+
"""
|
| 31 |
+
Embeds scalar timesteps into vector representations.
|
| 32 |
+
"""
|
| 33 |
+
def __init__(self, hidden_size, frequency_embedding_size=256):
|
| 34 |
+
super().__init__()
|
| 35 |
+
self.mlp = nn.Sequential(
|
| 36 |
+
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
| 37 |
+
nn.SiLU(),
|
| 38 |
+
nn.Linear(hidden_size, hidden_size, bias=True),
|
| 39 |
+
)
|
| 40 |
+
self.frequency_embedding_size = frequency_embedding_size
|
| 41 |
+
|
| 42 |
+
@staticmethod
|
| 43 |
+
def timestep_embedding(t, dim, max_period=10000):
|
| 44 |
+
"""
|
| 45 |
+
Create sinusoidal timestep embeddings.
|
| 46 |
+
:param t: a 1-D Tensor of N indices, one per batch element.
|
| 47 |
+
These may be fractional.
|
| 48 |
+
:param dim: the dimension of the output.
|
| 49 |
+
:param max_period: controls the minimum frequency of the embeddings.
|
| 50 |
+
:return: an (N, D) Tensor of positional embeddings.
|
| 51 |
+
"""
|
| 52 |
+
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
| 53 |
+
half = dim // 2
|
| 54 |
+
freqs = torch.exp(
|
| 55 |
+
-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
|
| 56 |
+
).to(device=t.device)
|
| 57 |
+
args = t[:, None].float() * freqs[None]
|
| 58 |
+
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
| 59 |
+
if dim % 2:
|
| 60 |
+
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
| 61 |
+
return embedding
|
| 62 |
+
|
| 63 |
+
def forward(self, t):
|
| 64 |
+
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
| 65 |
+
t_emb = self.mlp(t_freq)
|
| 66 |
+
return t_emb
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
#################################################################################
|
| 70 |
+
# Core DiT Model #
|
| 71 |
+
#################################################################################
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
class CrossAttention(nn.Module):
|
| 75 |
+
def __init__(self, query_dim, heads=8, dropout=0.):
|
| 76 |
+
super().__init__()
|
| 77 |
+
dim_head = query_dim / heads
|
| 78 |
+
|
| 79 |
+
self.scale = dim_head ** -0.5
|
| 80 |
+
self.heads = heads
|
| 81 |
+
self.to_q = nn.Linear(query_dim, query_dim, bias=True)
|
| 82 |
+
self.to_k = nn.Linear(query_dim, query_dim, bias=True)
|
| 83 |
+
self.to_v = nn.Linear(query_dim, query_dim, bias=True)
|
| 84 |
+
|
| 85 |
+
def forward(self, x, context=None):
|
| 86 |
+
h = self.heads
|
| 87 |
+
q = self.to_q(x)
|
| 88 |
+
k = self.to_k(context).unsqueeze(1)
|
| 89 |
+
v = self.to_v(context).unsqueeze(1)
|
| 90 |
+
|
| 91 |
+
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v))
|
| 92 |
+
sim = einsum('b i d, b j d -> b i j', q, k) * self.scale
|
| 93 |
+
|
| 94 |
+
# attention, what we cannot get enough of
|
| 95 |
+
attn = sim.softmax(dim=-2)
|
| 96 |
+
|
| 97 |
+
out = einsum('b i j, b j d -> b i d', attn, v)
|
| 98 |
+
out = rearrange(out, '(b h) n d -> b n (h d)', h=h)
|
| 99 |
+
attn_out = rearrange(attn, '(b h) n d -> b n (h d)', h=h)
|
| 100 |
+
return out, attn_out
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
class Cross_Norm(nn.Module):
|
| 104 |
+
def __init__(self, hidden_size, num_heads):
|
| 105 |
+
super().__init__()
|
| 106 |
+
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 107 |
+
self.cross_attention = CrossAttention(hidden_size, heads=num_heads)
|
| 108 |
+
|
| 109 |
+
def forward(self, x, c):
|
| 110 |
+
x = self.norm(x)
|
| 111 |
+
x = self.cross_attention(x, c)
|
| 112 |
+
|
| 113 |
+
return x
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
class DiTBlock(nn.Module):
|
| 118 |
+
"""
|
| 119 |
+
A DiT block with adaptive layer norm zero (adaLN-Zero) conditioning.
|
| 120 |
+
"""
|
| 121 |
+
def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, **block_kwargs):
|
| 122 |
+
super().__init__()
|
| 123 |
+
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 124 |
+
self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True, **block_kwargs)
|
| 125 |
+
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 126 |
+
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
| 127 |
+
approx_gelu = lambda: nn.GELU()
|
| 128 |
+
self.mlp = Mlp(in_features=hidden_size, hidden_features=mlp_hidden_dim, act_layer=approx_gelu, drop=0)
|
| 129 |
+
self.adaLN_modulation = nn.Sequential(
|
| 130 |
+
nn.SiLU(),
|
| 131 |
+
nn.Linear(hidden_size, 6 * hidden_size, bias=True)
|
| 132 |
+
)
|
| 133 |
+
|
| 134 |
+
def forward(self, x, c):
|
| 135 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1)
|
| 136 |
+
x = x + gate_msa.unsqueeze(1) * self.attn(modulate(self.norm1(x), shift_msa, scale_msa))
|
| 137 |
+
x = x + gate_mlp.unsqueeze(1) * self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))
|
| 138 |
+
return x
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
class FinalLayer(nn.Module):
|
| 142 |
+
"""
|
| 143 |
+
The final layer of DiT.
|
| 144 |
+
"""
|
| 145 |
+
def __init__(self, hidden_size, patch_size, out_channels):
|
| 146 |
+
super().__init__()
|
| 147 |
+
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 148 |
+
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
|
| 149 |
+
self.adaLN_modulation = nn.Sequential(
|
| 150 |
+
nn.SiLU(),
|
| 151 |
+
nn.Linear(hidden_size, 2 * hidden_size, bias=True)
|
| 152 |
+
)
|
| 153 |
+
|
| 154 |
+
def forward(self, x, c):
|
| 155 |
+
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
|
| 156 |
+
x = modulate(self.norm_final(x), shift, scale)
|
| 157 |
+
x = self.linear(x)
|
| 158 |
+
return x
|
| 159 |
+
|
| 160 |
+
|
| 161 |
+
class temp_Adaptive_Mask(nn.Module):
|
| 162 |
+
def __init__(self, hidden_size, patch_size, out_channels):
|
| 163 |
+
super().__init__()
|
| 164 |
+
self.out_channels = out_channels
|
| 165 |
+
self.patch_size = patch_size
|
| 166 |
+
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 167 |
+
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
|
| 168 |
+
self.mlp = nn.Sequential(
|
| 169 |
+
nn.SiLU(),
|
| 170 |
+
nn.Linear(hidden_size, 2 * hidden_size, bias=True),
|
| 171 |
+
nn.Linear(2 * hidden_size, hidden_size, bias=True)
|
| 172 |
+
)
|
| 173 |
+
|
| 174 |
+
|
| 175 |
+
def forward(self, x):
|
| 176 |
+
x = self.norm(x)
|
| 177 |
+
x = self.mlp(x)
|
| 178 |
+
x = self.linear(x)
|
| 179 |
+
h = w = int(x.shape[1] ** 0.5)
|
| 180 |
+
assert h * w == x.shape[1]
|
| 181 |
+
c = self.out_channels
|
| 182 |
+
p = self.patch_size
|
| 183 |
+
|
| 184 |
+
x = x.reshape(shape=(x.shape[0], h, w, p, p, c))
|
| 185 |
+
x = torch.einsum('nhwpqc->nchpwq', x)
|
| 186 |
+
x = x.reshape(shape=(x.shape[0], c, h * p, h * p))
|
| 187 |
+
return x
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
class DiT(nn.Module):
|
| 191 |
+
"""
|
| 192 |
+
Diffusion model with a Transformer backbone.
|
| 193 |
+
"""
|
| 194 |
+
def __init__(
|
| 195 |
+
self,
|
| 196 |
+
input_size=32,
|
| 197 |
+
patch_size=2,
|
| 198 |
+
in_channels=4,
|
| 199 |
+
hidden_size=1152,
|
| 200 |
+
depth=28,
|
| 201 |
+
num_heads=16,
|
| 202 |
+
mlp_ratio=4.0,
|
| 203 |
+
class_dropout_prob=0.1,
|
| 204 |
+
num_classes=1000,
|
| 205 |
+
learn_sigma=True,
|
| 206 |
+
):
|
| 207 |
+
super().__init__()
|
| 208 |
+
self.learn_sigma = learn_sigma
|
| 209 |
+
self.in_channels = in_channels
|
| 210 |
+
self.out_channels = in_channels * 2 if learn_sigma else in_channels
|
| 211 |
+
self.patch_size = patch_size
|
| 212 |
+
self.num_heads = num_heads
|
| 213 |
+
|
| 214 |
+
self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size, bias=True)
|
| 215 |
+
self.t_embedder = TimestepEmbedder(hidden_size)
|
| 216 |
+
self.y_embedders = nn.Linear(1024, 1152)
|
| 217 |
+
num_patches = self.x_embedder.num_patches
|
| 218 |
+
|
| 219 |
+
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, hidden_size), requires_grad=False)
|
| 220 |
+
|
| 221 |
+
self.blocks = nn.ModuleList([
|
| 222 |
+
DiTBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio) for _ in range(depth)
|
| 223 |
+
])
|
| 224 |
+
|
| 225 |
+
self.cross_defect = nn.ModuleList([
|
| 226 |
+
Cross_Norm(hidden_size, num_heads) for _ in range(10)
|
| 227 |
+
])
|
| 228 |
+
self.adapt_mask = temp_Adaptive_Mask(num_heads*10, patch_size, in_channels)
|
| 229 |
+
|
| 230 |
+
self.final_layer = FinalLayer(hidden_size, patch_size, self.out_channels)
|
| 231 |
+
self.initialize_weights()
|
| 232 |
+
|
| 233 |
+
def initialize_weights(self):
|
| 234 |
+
# Initialize transformer layers:
|
| 235 |
+
def _basic_init(module):
|
| 236 |
+
if isinstance(module, nn.Linear):
|
| 237 |
+
torch.nn.init.xavier_uniform_(module.weight)
|
| 238 |
+
if module.bias is not None:
|
| 239 |
+
nn.init.constant_(module.bias, 0)
|
| 240 |
+
self.apply(_basic_init)
|
| 241 |
+
|
| 242 |
+
# Initialize (and freeze) pos_embed by sin-cos embedding:
|
| 243 |
+
pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], int(self.x_embedder.num_patches ** 0.5))
|
| 244 |
+
self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0))
|
| 245 |
+
|
| 246 |
+
# Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
|
| 247 |
+
w = self.x_embedder.proj.weight.data
|
| 248 |
+
nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
|
| 249 |
+
nn.init.constant_(self.x_embedder.proj.bias, 0)
|
| 250 |
+
|
| 251 |
+
# Initialize timestep embedding MLP:
|
| 252 |
+
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
|
| 253 |
+
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
|
| 254 |
+
|
| 255 |
+
# Zero-out adaLN modulation layers in DiT blocks:
|
| 256 |
+
for block in self.blocks:
|
| 257 |
+
nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
|
| 258 |
+
nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
|
| 259 |
+
|
| 260 |
+
# Zero-out output layers:
|
| 261 |
+
nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)
|
| 262 |
+
nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
|
| 263 |
+
nn.init.constant_(self.final_layer.linear.weight, 0)
|
| 264 |
+
nn.init.constant_(self.final_layer.linear.bias, 0)
|
| 265 |
+
|
| 266 |
+
def unpatchify(self, x):
|
| 267 |
+
"""
|
| 268 |
+
x: (N, T, patch_size**2 * C)
|
| 269 |
+
imgs: (N, H, W, C)
|
| 270 |
+
"""
|
| 271 |
+
c = self.out_channels
|
| 272 |
+
p = self.x_embedder.patch_size[0]
|
| 273 |
+
h = w = int(x.shape[1] ** 0.5)
|
| 274 |
+
assert h * w == x.shape[1]
|
| 275 |
+
|
| 276 |
+
x = x.reshape(shape=(x.shape[0], h, w, p, p, c))
|
| 277 |
+
x = torch.einsum('nhwpqc->nchpwq', x)
|
| 278 |
+
imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p))
|
| 279 |
+
return imgs
|
| 280 |
+
|
| 281 |
+
def forward(self, x, t, y):
|
| 282 |
+
"""
|
| 283 |
+
Forward pass of DiT.
|
| 284 |
+
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
| 285 |
+
t: (N,) tensor of diffusion timesteps
|
| 286 |
+
y: (N,) tensor of class labels
|
| 287 |
+
"""
|
| 288 |
+
x = self.x_embedder(x) + self.pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
| 289 |
+
t = self.t_embedder(t) # (N, D)
|
| 290 |
+
y_defect = self.y_embedders(y[0])
|
| 291 |
+
y_class = self.y_embedders(y[1]) # (N, D)
|
| 292 |
+
y_all = self.y_embedders(y[2])
|
| 293 |
+
att_map = []
|
| 294 |
+
loss_att = 0
|
| 295 |
+
for i in range(28):
|
| 296 |
+
block = self.blocks[i]
|
| 297 |
+
if i < 10:
|
| 298 |
+
c = t + y_class
|
| 299 |
+
x = block(x, c) # (N, T, D)
|
| 300 |
+
elif i < 20:
|
| 301 |
+
cross = self.cross_defect[i - 10]
|
| 302 |
+
c = t + y_defect
|
| 303 |
+
|
| 304 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = block.adaLN_modulation(c).chunk(6,
|
| 305 |
+
dim=1)
|
| 306 |
+
x = x + gate_msa.unsqueeze(1) * block.attn(modulate(block.norm1(x), shift_msa, scale_msa))
|
| 307 |
+
|
| 308 |
+
cross_att, att_weight = cross(x, c)
|
| 309 |
+
loss_att += att_weight
|
| 310 |
+
att_map.append(att_weight)
|
| 311 |
+
x = x + cross_att
|
| 312 |
+
|
| 313 |
+
x = x + gate_mlp.unsqueeze(1) * block.mlp(modulate(block.norm2(x), shift_mlp, scale_mlp))
|
| 314 |
+
|
| 315 |
+
elif i < 28:
|
| 316 |
+
c = t + y_all
|
| 317 |
+
x = block(x, c)
|
| 318 |
+
|
| 319 |
+
x = self.final_layer(x, c) # (N, T, patch_size ** 2 * out_channels)
|
| 320 |
+
x = self.unpatchify(x) # (N, out_channels, H, W)
|
| 321 |
+
att_map = torch.cat(att_map, dim=-1)
|
| 322 |
+
att_mask = self.adapt_mask(att_map)
|
| 323 |
+
|
| 324 |
+
return x, att_mask, loss_att.resize(x.shape[0], x.shape[2]//2, x.shape[3]//2, 16).mean(dim=-1)
|
| 325 |
+
|
| 326 |
+
def forward_free_2(self, x, t, y, mask_temp=None):
|
| 327 |
+
"""
|
| 328 |
+
Forward pass of DiT.
|
| 329 |
+
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
| 330 |
+
t: (N,) tensor of diffusion timesteps
|
| 331 |
+
y: (N,) tensor of class labels
|
| 332 |
+
"""
|
| 333 |
+
x = self.x_embedder(x) + self.pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
| 334 |
+
t = self.t_embedder(t) # (N, D)
|
| 335 |
+
y_defect = self.y_embedders(torch.cat([y[0][0], y[1][0]], dim=0))
|
| 336 |
+
y_class = self.y_embedders(torch.cat([y[0][1], y[1][1]], dim=0)) # (N, D)
|
| 337 |
+
y_all = self.y_embedders(torch.cat([y[0][2], y[1][2]], dim=0))
|
| 338 |
+
att_map = []
|
| 339 |
+
loss_att = 0
|
| 340 |
+
for i in range(28):
|
| 341 |
+
block = self.blocks[i]
|
| 342 |
+
if i < 10:
|
| 343 |
+
c = t + y_class
|
| 344 |
+
x = block(x, c) # (N, T, D)
|
| 345 |
+
elif i < 20:
|
| 346 |
+
cross = self.cross_defect[i - 10]
|
| 347 |
+
c = t + y_defect
|
| 348 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = block.adaLN_modulation(c).chunk(6,
|
| 349 |
+
dim=1)
|
| 350 |
+
x = x + gate_msa.unsqueeze(1) * block.attn(modulate(block.norm1(x), shift_msa, scale_msa))
|
| 351 |
+
cross_att, att_weight = cross(x, c)
|
| 352 |
+
att_map.append(att_weight)
|
| 353 |
+
loss_att+=att_weight
|
| 354 |
+
x = x + cross_att
|
| 355 |
+
x = x + gate_mlp.unsqueeze(1) * block.mlp(modulate(block.norm2(x), shift_mlp, scale_mlp))
|
| 356 |
+
|
| 357 |
+
elif i < 28:
|
| 358 |
+
c = t + y_all
|
| 359 |
+
x = block(x, c)
|
| 360 |
+
|
| 361 |
+
x = self.final_layer(x, c) # (N, T, patch_size ** 2 * out_channels)
|
| 362 |
+
x = self.unpatchify(x) # (N, out_channels, H, W)
|
| 363 |
+
att_map = torch.cat(att_map, dim=-1)
|
| 364 |
+
att_mask = self.adapt_mask(att_map)
|
| 365 |
+
|
| 366 |
+
return x, att_mask, loss_att
|
| 367 |
+
|
| 368 |
+
def forward_with_cfg_2(self, x, t, y, cfg_scale):
|
| 369 |
+
"""
|
| 370 |
+
Forward pass of DiT, but also batches the unconditional forward pass for classifier-free guidance.
|
| 371 |
+
"""
|
| 372 |
+
# https://github.com/openai/glide-text2im/blob/main/notebooks/text2im.ipynb
|
| 373 |
+
half = x[: len(x) // 2]
|
| 374 |
+
combined = torch.cat([half, half], dim=0)
|
| 375 |
+
model_out, mask, _ = self.forward_free_2(combined, t, y)
|
| 376 |
+
# For exact reproducibility reasons, we apply classifier-free guidance on only
|
| 377 |
+
# three channels by default. The standard approach to cfg applies it to all channels.
|
| 378 |
+
# This can be done by uncommenting the following line and commenting-out the line following that.
|
| 379 |
+
# eps, rest = model_out[:, :self.in_channels], model_out[:, self.in_channels:]
|
| 380 |
+
eps, rest = model_out[:, :3], model_out[:, 3:]
|
| 381 |
+
cond_eps, uncond_eps = torch.split(eps, len(eps) // 2, dim=0)
|
| 382 |
+
half_eps = uncond_eps + cfg_scale * (cond_eps - uncond_eps)
|
| 383 |
+
eps = torch.cat([half_eps, half_eps], dim=0)
|
| 384 |
+
return torch.cat([eps, rest], dim=1), mask, _
|
| 385 |
+
|
| 386 |
+
def forward_free_3(self, x, t, y, mask_temp=None):
|
| 387 |
+
"""
|
| 388 |
+
Forward pass of DiT.
|
| 389 |
+
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
| 390 |
+
t: (N,) tensor of diffusion timesteps
|
| 391 |
+
y: (N,) tensor of class labels
|
| 392 |
+
"""
|
| 393 |
+
x = self.x_embedder(x) + self.pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
| 394 |
+
t = self.t_embedder(t) # (N, D)
|
| 395 |
+
y_defect = self.y_embedders(torch.cat([y[0][0], y[1][0], y[2][0]], dim=0))
|
| 396 |
+
y_class = self.y_embedders(torch.cat([y[0][1], y[1][1], y[2][1]], dim=0)) # (N, D)
|
| 397 |
+
y_all = self.y_embedders(torch.cat([y[0][2], y[1][2], y[2][2]], dim=0))
|
| 398 |
+
att_map = []
|
| 399 |
+
att_loss = 0
|
| 400 |
+
for i in range(28):
|
| 401 |
+
block = self.blocks[i]
|
| 402 |
+
if i < 10:
|
| 403 |
+
c = t + y_class
|
| 404 |
+
x = block(x, c) # (N, T, D)
|
| 405 |
+
elif i < 20:
|
| 406 |
+
cross = self.cross_defect[i - 10]
|
| 407 |
+
c = t + y_defect
|
| 408 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = block.adaLN_modulation(c).chunk(6,
|
| 409 |
+
dim=1)
|
| 410 |
+
x = x + gate_msa.unsqueeze(1) * block.attn(modulate(block.norm1(x), shift_msa, scale_msa))
|
| 411 |
+
cross_att, att_weight = cross(x, c)
|
| 412 |
+
att_map.append(att_weight)
|
| 413 |
+
att_loss += att_weight
|
| 414 |
+
x = x + cross_att
|
| 415 |
+
x = x + gate_mlp.unsqueeze(1) * block.mlp(modulate(block.norm2(x), shift_mlp, scale_mlp))
|
| 416 |
+
|
| 417 |
+
elif i < 28:
|
| 418 |
+
c = t + y_all
|
| 419 |
+
x = block(x, c)
|
| 420 |
+
|
| 421 |
+
x = self.final_layer(x, c) # (N, T, patch_size ** 2 * out_channels)
|
| 422 |
+
x = self.unpatchify(x) # (N, out_channels, H, W)
|
| 423 |
+
att_map = torch.cat(att_map, dim=-1)
|
| 424 |
+
att_mask = self.adapt_mask(att_map)
|
| 425 |
+
|
| 426 |
+
return x, att_mask, att_loss
|
| 427 |
+
|
| 428 |
+
def forward_with_cfg_3(self, x, t, y, cfg_scale):
|
| 429 |
+
"""
|
| 430 |
+
Forward pass of DiT, but also batches the unconditional forward pass for classifier-free guidance.
|
| 431 |
+
"""
|
| 432 |
+
# https://github.com/openai/glide-text2im/blob/main/notebooks/text2im.ipynb
|
| 433 |
+
half = x[: len(x) // 3]
|
| 434 |
+
combined = torch.cat([half, half, half], dim=0)
|
| 435 |
+
model_out, mask, _ = self.forward_free_3(combined, t, y)
|
| 436 |
+
# For exact reproducibility reasons, we apply classifier-free guidance on only
|
| 437 |
+
# three channels by default. The standard approach to cfg applies it to all channels.
|
| 438 |
+
# This can be done by uncommenting the following line and commenting-out the line following that.
|
| 439 |
+
# eps, rest = model_out[:, :self.in_channels], model_out[:, self.in_channels:]
|
| 440 |
+
eps, rest = model_out[:, :3], model_out[:, 3:]
|
| 441 |
+
cond_eps, uncond_eps_defect, uncond_eps = torch.split(eps, len(eps) // 3, dim=0)
|
| 442 |
+
half_eps = uncond_eps + cfg_scale * (cond_eps - uncond_eps_defect) + cfg_scale * (uncond_eps_defect - uncond_eps)
|
| 443 |
+
eps = torch.cat([half_eps, half_eps, half_eps], dim=0)
|
| 444 |
+
return torch.cat([eps, rest], dim=1), mask, _
|
| 445 |
+
|
| 446 |
+
#################################################################################
|
| 447 |
+
# Sine/Cosine Positional Embedding Functions #
|
| 448 |
+
#################################################################################
|
| 449 |
+
# https://github.com/facebookresearch/mae/blob/main/util/pos_embed.py
|
| 450 |
+
|
| 451 |
+
def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0):
|
| 452 |
+
"""
|
| 453 |
+
grid_size: int of the grid height and width
|
| 454 |
+
return:
|
| 455 |
+
pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)
|
| 456 |
+
"""
|
| 457 |
+
grid_h = np.arange(grid_size, dtype=np.float32)
|
| 458 |
+
grid_w = np.arange(grid_size, dtype=np.float32)
|
| 459 |
+
grid = np.meshgrid(grid_w, grid_h) # here w goes first
|
| 460 |
+
grid = np.stack(grid, axis=0)
|
| 461 |
+
|
| 462 |
+
grid = grid.reshape([2, 1, grid_size, grid_size])
|
| 463 |
+
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
|
| 464 |
+
if cls_token and extra_tokens > 0:
|
| 465 |
+
pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0)
|
| 466 |
+
return pos_embed
|
| 467 |
+
|
| 468 |
+
|
| 469 |
+
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
|
| 470 |
+
assert embed_dim % 2 == 0
|
| 471 |
+
|
| 472 |
+
# use half of dimensions to encode grid_h
|
| 473 |
+
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
|
| 474 |
+
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
|
| 475 |
+
|
| 476 |
+
emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
|
| 477 |
+
return emb
|
| 478 |
+
|
| 479 |
+
|
| 480 |
+
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
|
| 481 |
+
"""
|
| 482 |
+
embed_dim: output dimension for each position
|
| 483 |
+
pos: a list of positions to be encoded: size (M,)
|
| 484 |
+
out: (M, D)
|
| 485 |
+
"""
|
| 486 |
+
assert embed_dim % 2 == 0
|
| 487 |
+
omega = np.arange(embed_dim // 2, dtype=np.float64)
|
| 488 |
+
omega /= embed_dim / 2.
|
| 489 |
+
omega = 1. / 10000**omega # (D/2,)
|
| 490 |
+
|
| 491 |
+
pos = pos.reshape(-1) # (M,)
|
| 492 |
+
out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
|
| 493 |
+
|
| 494 |
+
emb_sin = np.sin(out) # (M, D/2)
|
| 495 |
+
emb_cos = np.cos(out) # (M, D/2)
|
| 496 |
+
|
| 497 |
+
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
|
| 498 |
+
return emb
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/test.py
ADDED
|
@@ -0,0 +1,198 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import argparse
|
| 3 |
+
import torch
|
| 4 |
+
import numpy as np
|
| 5 |
+
from torchvision.utils import save_image
|
| 6 |
+
from diffusers.models import AutoencoderKL
|
| 7 |
+
import clip.clip as clip
|
| 8 |
+
|
| 9 |
+
from models_add_cross_concate import DiT
|
| 10 |
+
from diffusion import create_diffusion
|
| 11 |
+
from autoencoder import *
|
| 12 |
+
|
| 13 |
+
# Enable TF32 for fast execution on modern NVIDIA GPUs
|
| 14 |
+
torch.backends.cuda.matmul.allow_tf32 = True
|
| 15 |
+
torch.backends.cudnn.allow_tf32 = True
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def rgb_to_gray(tensor):
|
| 19 |
+
r, g, b = tensor[:, 0], tensor[:, 1], tensor[:, 2]
|
| 20 |
+
gray = 0.299 * r + 0.587 * g + 0.114 * b
|
| 21 |
+
return gray
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def iterative_thresholding_batch(gray_tensor):
|
| 25 |
+
gray_np = gray_tensor.detach().cpu().numpy()
|
| 26 |
+
binarized = np.zeros_like(gray_np, dtype=np.uint8)
|
| 27 |
+
|
| 28 |
+
for i in range(gray_np.shape[0]):
|
| 29 |
+
img = gray_np[i]
|
| 30 |
+
T = img.mean()
|
| 31 |
+
prev_T = -1
|
| 32 |
+
|
| 33 |
+
while abs(T - prev_T) > 1e-4:
|
| 34 |
+
prev_T = T
|
| 35 |
+
G1 = img[img >= T]
|
| 36 |
+
G2 = img[img < T]
|
| 37 |
+
m1 = G1.mean() if G1.size > 0 else 0
|
| 38 |
+
m2 = G2.mean() if G2.size > 0 else 0
|
| 39 |
+
T = (m1 + m2) / 2
|
| 40 |
+
|
| 41 |
+
binarized[i] = (img >= T).astype(np.uint8)
|
| 42 |
+
|
| 43 |
+
return torch.from_numpy(binarized).to(gray_tensor.device)
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def binarize_tensor_iterative(x):
|
| 47 |
+
gray = rgb_to_gray(x)
|
| 48 |
+
binary = iterative_thresholding_batch(gray)
|
| 49 |
+
return binary.unsqueeze(1)
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def get_label(data_path):
|
| 53 |
+
"""Safely extracts defect labels from 4-level dataset hierarchy."""
|
| 54 |
+
label_list1 = []
|
| 55 |
+
if not os.path.exists(data_path):
|
| 56 |
+
return label_list1
|
| 57 |
+
|
| 58 |
+
for name_class in os.listdir(data_path):
|
| 59 |
+
img_dir = os.path.join(data_path, name_class, 'img')
|
| 60 |
+
if os.path.exists(img_dir) and os.path.isdir(img_dir):
|
| 61 |
+
for class_object in os.listdir(img_dir):
|
| 62 |
+
defect_dir = os.path.join(img_dir, class_object)
|
| 63 |
+
if os.path.isdir(defect_dir) and class_object != 'good':
|
| 64 |
+
label_list1.append(f"{class_object} {name_class}")
|
| 65 |
+
return label_list1
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def gen(args):
|
| 69 |
+
data_path = args.data
|
| 70 |
+
label_list = get_label(data_path)
|
| 71 |
+
|
| 72 |
+
if not label_list:
|
| 73 |
+
print(f"❌ No valid defect subfolders found in {data_path}. Please check directory structure.")
|
| 74 |
+
return
|
| 75 |
+
|
| 76 |
+
print(f"📋 Found defect categories to generate: {label_list}")
|
| 77 |
+
|
| 78 |
+
image_size = args.imagesize
|
| 79 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 80 |
+
latent_size = image_size // 8
|
| 81 |
+
|
| 82 |
+
# 1. Load CLIP model
|
| 83 |
+
model_clip, _ = clip.load('RN50', device)
|
| 84 |
+
|
| 85 |
+
# 2. Setup DiT architecture and weights
|
| 86 |
+
model = DiT(
|
| 87 |
+
depth=28, hidden_size=1152, patch_size=2,
|
| 88 |
+
num_heads=16, input_size=latent_size, num_classes=1000
|
| 89 |
+
).to(device)
|
| 90 |
+
|
| 91 |
+
print(f"📦 Loading checkpoint from: {args.ckpt}")
|
| 92 |
+
checkpoint = torch.load(args.ckpt, map_location=device)
|
| 93 |
+
if isinstance(checkpoint, dict) and 'model_state_dict' in checkpoint:
|
| 94 |
+
model.load_state_dict(checkpoint['model_state_dict'])
|
| 95 |
+
else:
|
| 96 |
+
model.load_state_dict(checkpoint)
|
| 97 |
+
|
| 98 |
+
model.eval()
|
| 99 |
+
|
| 100 |
+
# 3. Setup VAE and Diffusion pipeline
|
| 101 |
+
diffusion = create_diffusion(timestep_respacing="50")
|
| 102 |
+
vae = AutoencoderKL.from_pretrained(args.vae).to(device)
|
| 103 |
+
|
| 104 |
+
os.makedirs(args.out_dir, exist_ok=True)
|
| 105 |
+
num_img = args.batchsize
|
| 106 |
+
|
| 107 |
+
# 4. Generate specified number of output batches (replaces infinite while loop)
|
| 108 |
+
for sample_round in range(args.num_samples):
|
| 109 |
+
print(f"\n🔄 --- Generating Batch {sample_round + 1}/{args.num_samples} ---")
|
| 110 |
+
|
| 111 |
+
for c in label_list:
|
| 112 |
+
defect_name, class_name = c.split()[0], c.split()[1]
|
| 113 |
+
print(f"🎨 Generating defect: '{defect_name}' on object: '{class_name}'...")
|
| 114 |
+
|
| 115 |
+
# Prepare text embeddings for dual-branch CFG
|
| 116 |
+
y_null_product = torch.cat([clip.tokenize("a photo of good industry")] * num_img).to(device)
|
| 117 |
+
y_null_good = torch.cat([clip.tokenize(f"a photo of good {class_name}")] * num_img).to(device)
|
| 118 |
+
|
| 119 |
+
with torch.no_grad():
|
| 120 |
+
y_null_product = model_clip.encode_text(y_null_product)
|
| 121 |
+
y_null_good = model_clip.encode_text(y_null_good)
|
| 122 |
+
|
| 123 |
+
y_null_product = (y_null_product / y_null_product.norm(dim=-1, keepdim=True)).float()
|
| 124 |
+
y_null_good = (y_null_good / y_null_good.norm(dim=-1, keepdim=True)).float()
|
| 125 |
+
|
| 126 |
+
only_good = torch.cat([clip.tokenize("a photo of good")] * num_img).to(device)
|
| 127 |
+
defect = torch.cat([clip.tokenize(f"a photo of {defect_name}")] * num_img).to(device)
|
| 128 |
+
classes = torch.cat([clip.tokenize(f"a photo of {class_name}")] * num_img).to(device)
|
| 129 |
+
classes_industry = torch.cat([clip.tokenize("a photo of industry")] * num_img).to(device)
|
| 130 |
+
y_all = torch.cat([clip.tokenize(f"a photo of {c}")] * num_img).to(device)
|
| 131 |
+
|
| 132 |
+
with torch.no_grad():
|
| 133 |
+
only_good = model_clip.encode_text(only_good)
|
| 134 |
+
defect = model_clip.encode_text(defect)
|
| 135 |
+
classes = model_clip.encode_text(classes)
|
| 136 |
+
classes_industry = model_clip.encode_text(classes_industry)
|
| 137 |
+
y_all = model_clip.encode_text(y_all)
|
| 138 |
+
|
| 139 |
+
only_good = (only_good / only_good.norm(dim=-1, keepdim=True)).float()
|
| 140 |
+
defect = (defect / defect.norm(dim=-1, keepdim=True)).float()
|
| 141 |
+
classes_industry = (classes_industry / classes_industry.norm(dim=-1, keepdim=True)).float()
|
| 142 |
+
classes = (classes / classes.norm(dim=-1, keepdim=True)).float()
|
| 143 |
+
y_all = (y_all / y_all.norm(dim=-1, keepdim=True)).float()
|
| 144 |
+
|
| 145 |
+
y_defect_class = [defect, classes, y_all]
|
| 146 |
+
y_good_class = [only_good, classes, y_null_good]
|
| 147 |
+
|
| 148 |
+
z = torch.randn(num_img, 4, latent_size, latent_size, device=device)
|
| 149 |
+
z = torch.cat([z, z], 0)
|
| 150 |
+
|
| 151 |
+
y = [y_defect_class, y_good_class]
|
| 152 |
+
|
| 153 |
+
for num in np.arange(0.5, 3.0, 0.5):
|
| 154 |
+
model_kwargs = dict(y=y, cfg_scale=float(num))
|
| 155 |
+
|
| 156 |
+
with torch.no_grad():
|
| 157 |
+
samples, cross = diffusion.p_sample_loop(
|
| 158 |
+
model.forward_with_cfg_2,
|
| 159 |
+
z.shape,
|
| 160 |
+
z,
|
| 161 |
+
clip_denoised=False,
|
| 162 |
+
model_kwargs=model_kwargs,
|
| 163 |
+
progress=False,
|
| 164 |
+
device=device
|
| 165 |
+
)
|
| 166 |
+
|
| 167 |
+
img_gen, _ = samples.chunk(2, dim=0)
|
| 168 |
+
mask_gen, _ = cross.chunk(2, dim=0)
|
| 169 |
+
|
| 170 |
+
with torch.no_grad():
|
| 171 |
+
img_gen = vae.decode(img_gen / 0.18215).sample
|
| 172 |
+
mask_gen = vae.decode(mask_gen / 0.18215).sample
|
| 173 |
+
|
| 174 |
+
# Save generated images and binarized masks
|
| 175 |
+
img_path = os.path.join(args.out_dir, f"{class_name}_{defect_name}_cfg{num:.1f}_b{sample_round}.png")
|
| 176 |
+
mask_path = os.path.join(args.out_dir, f"{class_name}_{defect_name}_cfg{num:.1f}_b{sample_round}_mask.png")
|
| 177 |
+
|
| 178 |
+
save_image(img_gen, img_path, nrow=2, normalize=True)
|
| 179 |
+
|
| 180 |
+
mask_gen = binarize_tensor_iterative(mask_gen)
|
| 181 |
+
mask_gen = (mask_gen * 255).to(torch.uint8).float() / 255.0
|
| 182 |
+
save_image(mask_gen, mask_path, nrow=2, normalize=True)
|
| 183 |
+
|
| 184 |
+
print(f"\n✨ Generation complete! Synthetic pairs saved to: {os.path.abspath(args.out_dir)}")
|
| 185 |
+
|
| 186 |
+
|
| 187 |
+
if __name__ == "__main__":
|
| 188 |
+
parser = argparse.ArgumentParser()
|
| 189 |
+
parser.add_argument("--batchsize", type=int, default=2)
|
| 190 |
+
parser.add_argument("--num_samples", type=int, default=1, help="Number of sampling passes to run.")
|
| 191 |
+
parser.add_argument("--data", type=str, required=True)
|
| 192 |
+
parser.add_argument("--imagesize", type=int, choices=[256, 512], default=512)
|
| 193 |
+
parser.add_argument("--ckpt", type=str, required=True, help="Path to fine-tuned checkpoint.")
|
| 194 |
+
parser.add_argument("--vae", type=str, required=True, help="Path to VAE checkpoint.")
|
| 195 |
+
parser.add_argument("--out_dir", type=str, default="./generated_results", help="Directory to save generated samples.")
|
| 196 |
+
|
| 197 |
+
args = parser.parse_args()
|
| 198 |
+
gen(args)
|
ArtiAgent - DefectDiffu/engine/DefectDiffu/train.py
ADDED
|
@@ -0,0 +1,231 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import torch
|
| 3 |
+
from PIL import Image
|
| 4 |
+
from torch.utils.data import Dataset, DataLoader
|
| 5 |
+
from torchvision import transforms
|
| 6 |
+
from torchvision.transforms import Lambda
|
| 7 |
+
from diffusers.models import AutoencoderKL
|
| 8 |
+
import argparse
|
| 9 |
+
|
| 10 |
+
import clip.clip as clip
|
| 11 |
+
from models_add_cross_concate import DiT
|
| 12 |
+
from diffusion import create_diffusion
|
| 13 |
+
|
| 14 |
+
torch.backends.cuda.matmul.allow_tf32 = True
|
| 15 |
+
torch.backends.cudnn.allow_tf32 = True
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
# =========================================================================
|
| 19 |
+
# 1. TOP-LEVEL HELPER FUNCTION (Prevents Pickle Errors on Windows)
|
| 20 |
+
# =========================================================================
|
| 21 |
+
def scale_to_neg_one_to_one(t):
|
| 22 |
+
return (t * 2) - 1
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
# =========================================================================
|
| 26 |
+
# 2. DATASET CLASS FOR 4-LEVEL STRUCTURE
|
| 27 |
+
# =========================================================================
|
| 28 |
+
class Dataset_self(Dataset):
|
| 29 |
+
def __init__(self, img_root, preprocess):
|
| 30 |
+
self.img_root = img_root
|
| 31 |
+
self.img_process = preprocess
|
| 32 |
+
self.img = []
|
| 33 |
+
self.label_word = []
|
| 34 |
+
self.label_mask = []
|
| 35 |
+
|
| 36 |
+
# Parse 4-level structure: root / object_class / img / defect_type / image.png
|
| 37 |
+
for name_class in os.listdir(self.img_root): # e.g., 'vcsel'
|
| 38 |
+
class_path = os.path.join(self.img_root, name_class)
|
| 39 |
+
img_base = os.path.join(class_path, 'img')
|
| 40 |
+
gt_base = os.path.join(class_path, 'ground_truth')
|
| 41 |
+
|
| 42 |
+
if os.path.exists(img_base):
|
| 43 |
+
for defect in os.listdir(img_base): # e.g., 'good', 'scratch', 'bubble', 'crack'
|
| 44 |
+
defect_img_dir = os.path.join(img_base, defect)
|
| 45 |
+
defect_gt_dir = os.path.join(gt_base, defect)
|
| 46 |
+
|
| 47 |
+
if os.path.isdir(defect_img_dir):
|
| 48 |
+
for name_img in os.listdir(defect_img_dir):
|
| 49 |
+
if name_img.lower().endswith(('.png', '.jpg', '.jpeg')):
|
| 50 |
+
img_path = os.path.join(defect_img_dir, name_img)
|
| 51 |
+
|
| 52 |
+
# Flexible check: supports '000_mask.png' or '000.png'
|
| 53 |
+
base_name = os.path.splitext(name_img)[0]
|
| 54 |
+
mask_candidate_1 = os.path.join(defect_gt_dir, f"{base_name}_mask.png")
|
| 55 |
+
mask_candidate_2 = os.path.join(defect_gt_dir, name_img)
|
| 56 |
+
|
| 57 |
+
if os.path.exists(mask_candidate_1):
|
| 58 |
+
mask_path = mask_candidate_1
|
| 59 |
+
elif os.path.exists(mask_candidate_2):
|
| 60 |
+
mask_path = mask_candidate_2
|
| 61 |
+
else:
|
| 62 |
+
print(f"⚠️ Warning: Mask missing for image {img_path}")
|
| 63 |
+
continue
|
| 64 |
+
|
| 65 |
+
self.img.append(img_path)
|
| 66 |
+
self.label_word.append(f"{defect} {name_class}")
|
| 67 |
+
self.label_mask.append(mask_path)
|
| 68 |
+
|
| 69 |
+
print(f" Successfully loaded {len(self.img)} samples across all defect categories.")
|
| 70 |
+
|
| 71 |
+
def __len__(self):
|
| 72 |
+
return len(self.img)
|
| 73 |
+
|
| 74 |
+
def __getitem__(self, idx):
|
| 75 |
+
img_path = self.img[idx]
|
| 76 |
+
label_mask_path = self.label_mask[idx]
|
| 77 |
+
|
| 78 |
+
image = Image.open(img_path).convert('RGB')
|
| 79 |
+
label_mask_img = Image.open(label_mask_path).convert('RGB')
|
| 80 |
+
|
| 81 |
+
label_mask = self.img_process[1](label_mask_img)
|
| 82 |
+
mask_resize = self.img_process[2](label_mask_img)
|
| 83 |
+
mask_loss = self.img_process[3](label_mask_img)
|
| 84 |
+
|
| 85 |
+
mask_loss = mask_loss[0, :, :]
|
| 86 |
+
mask_loss[mask_loss != 0] = 1
|
| 87 |
+
mask_resize_res = torch.cat([mask_resize, mask_resize[0, :, :].unsqueeze(0)], dim=0)
|
| 88 |
+
|
| 89 |
+
label = self.label_word[idx]
|
| 90 |
+
image = self.img_process[0](image)
|
| 91 |
+
|
| 92 |
+
return image, label, label_mask, mask_resize_res, mask_loss
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
# =========================================================================
|
| 96 |
+
# 3. MAIN TRAINING LOGIC
|
| 97 |
+
# =========================================================================
|
| 98 |
+
def main(args):
|
| 99 |
+
device = "cuda"
|
| 100 |
+
model_clip, _ = clip.load('RN50', device)
|
| 101 |
+
|
| 102 |
+
data_path = args.data
|
| 103 |
+
image_size = args.imagesize
|
| 104 |
+
batch_size = args.batchsize
|
| 105 |
+
latent_size = image_size // 8
|
| 106 |
+
|
| 107 |
+
model = DiT(depth=28, hidden_size=1152, patch_size=2, num_heads=16, input_size=latent_size, num_classes=1000).to(device)
|
| 108 |
+
state_dict = torch.load(args.ckpt)
|
| 109 |
+
model.load_state_dict(state_dict, strict=False)
|
| 110 |
+
|
| 111 |
+
diffusion = create_diffusion(timestep_respacing="")
|
| 112 |
+
vae = AutoencoderKL.from_pretrained(args.vae).to(device)
|
| 113 |
+
opt = torch.optim.AdamW(model.parameters(), lr=1e-5, weight_decay=1e-8)
|
| 114 |
+
|
| 115 |
+
transform = transforms.Compose([
|
| 116 |
+
transforms.Resize(image_size),
|
| 117 |
+
transforms.CenterCrop(image_size),
|
| 118 |
+
transforms.ToTensor(),
|
| 119 |
+
Lambda(scale_to_neg_one_to_one),
|
| 120 |
+
])
|
| 121 |
+
|
| 122 |
+
transform_mask = transforms.Compose([
|
| 123 |
+
transforms.Resize(image_size),
|
| 124 |
+
transforms.CenterCrop(image_size),
|
| 125 |
+
transforms.ToTensor(),
|
| 126 |
+
Lambda(scale_to_neg_one_to_one),
|
| 127 |
+
])
|
| 128 |
+
|
| 129 |
+
transform_resize_mask = transforms.Compose([
|
| 130 |
+
transforms.ToTensor(),
|
| 131 |
+
transforms.Resize(latent_size),
|
| 132 |
+
transforms.CenterCrop(latent_size),
|
| 133 |
+
])
|
| 134 |
+
|
| 135 |
+
transform_mask_loss = transforms.Compose([
|
| 136 |
+
transforms.ToTensor(),
|
| 137 |
+
transforms.Resize(latent_size // 2),
|
| 138 |
+
transforms.CenterCrop(latent_size // 2),
|
| 139 |
+
])
|
| 140 |
+
|
| 141 |
+
dataset = Dataset_self(img_root=data_path, preprocess=[transform, transform_mask, transform_resize_mask, transform_mask_loss])
|
| 142 |
+
|
| 143 |
+
loader = DataLoader(
|
| 144 |
+
dataset,
|
| 145 |
+
batch_size=batch_size,
|
| 146 |
+
shuffle=True,
|
| 147 |
+
num_workers=0, # 0 workers avoids Windows multiprocessing crashes
|
| 148 |
+
pin_memory=True,
|
| 149 |
+
drop_last=True
|
| 150 |
+
)
|
| 151 |
+
|
| 152 |
+
model.train()
|
| 153 |
+
EPOCH = args.epochs
|
| 154 |
+
|
| 155 |
+
for epoch in range(EPOCH):
|
| 156 |
+
for x, y, mask, mask_resize, mask_loss in loader:
|
| 157 |
+
x = x.to(device)
|
| 158 |
+
mask = mask.to(device)
|
| 159 |
+
mask_resize = mask_resize.to(device)
|
| 160 |
+
mask_loss = mask_loss.to(device)
|
| 161 |
+
|
| 162 |
+
drop_rat = 0.2
|
| 163 |
+
if args.free == 2:
|
| 164 |
+
for i in range(len(y)):
|
| 165 |
+
c = y[i]
|
| 166 |
+
if c.split()[0] == 'good':
|
| 167 |
+
rat_1 = torch.rand(1)
|
| 168 |
+
if rat_1 < drop_rat:
|
| 169 |
+
y[i] = 'good industry'
|
| 170 |
+
else:
|
| 171 |
+
rat = torch.rand(1)
|
| 172 |
+
if rat < drop_rat:
|
| 173 |
+
y[i] = ('good ' + c.split()[1])
|
| 174 |
+
else:
|
| 175 |
+
for i in range(len(y)):
|
| 176 |
+
c = y[i]
|
| 177 |
+
if c.split()[0] != 'good':
|
| 178 |
+
rat_1 = torch.rand(1)
|
| 179 |
+
if rat_1 < drop_rat:
|
| 180 |
+
y[i] = ('good ' + c.split()[1])
|
| 181 |
+
|
| 182 |
+
defect = torch.cat([clip.tokenize(f"a photo of {c.split()[0]}") for c in y]).to(device)
|
| 183 |
+
classes = torch.cat([clip.tokenize(f"a photo of {c.split()[1]}") for c in y]).to(device)
|
| 184 |
+
y_all = torch.cat([clip.tokenize(f"a photo of {c}") for c in y]).to(device)
|
| 185 |
+
|
| 186 |
+
with torch.no_grad():
|
| 187 |
+
defect = model_clip.encode_text(defect)
|
| 188 |
+
classes = model_clip.encode_text(classes)
|
| 189 |
+
y_all = model_clip.encode_text(y_all)
|
| 190 |
+
|
| 191 |
+
defect /= defect.norm(dim=-1, keepdim=True)
|
| 192 |
+
defect = defect.float().to(device)
|
| 193 |
+
|
| 194 |
+
classes /= classes.norm(dim=-1, keepdim=True)
|
| 195 |
+
classes = classes.float().to(device)
|
| 196 |
+
|
| 197 |
+
y_all /= y_all.norm(dim=-1, keepdim=True)
|
| 198 |
+
y_all = y_all.float().to(device)
|
| 199 |
+
|
| 200 |
+
with torch.no_grad():
|
| 201 |
+
x = vae.encode(x).latent_dist.sample().mul_(0.18215)
|
| 202 |
+
mask_gt = vae.encode(mask).latent_dist.sample().mul_(0.18215)
|
| 203 |
+
|
| 204 |
+
t = torch.randint(0, diffusion.num_timesteps, (x.shape[0],), device=device)
|
| 205 |
+
model_kwargs = dict(y=[defect, classes, y_all])
|
| 206 |
+
loss_dict = diffusion.training_losses(model, x, t, model_kwargs, mask_resize=mask_resize, mask_att=mask_loss, label_mask=mask_gt)
|
| 207 |
+
loss = loss_dict["loss"].mean()
|
| 208 |
+
|
| 209 |
+
opt.zero_grad()
|
| 210 |
+
loss.backward()
|
| 211 |
+
opt.step()
|
| 212 |
+
print(f"Epoch {epoch} | Loss: {loss.item():.4f}")
|
| 213 |
+
|
| 214 |
+
if epoch % 100 == 0 and 2000 >= epoch >= 100:
|
| 215 |
+
os.makedirs('checkpoint', exist_ok=True)
|
| 216 |
+
torch.save({
|
| 217 |
+
'model_state_dict': model.state_dict(),
|
| 218 |
+
}, f'checkpoint/model_{epoch}.pth')
|
| 219 |
+
|
| 220 |
+
|
| 221 |
+
if __name__ == "__main__":
|
| 222 |
+
parser = argparse.ArgumentParser()
|
| 223 |
+
parser.add_argument("--batchsize", type=int, default=2)
|
| 224 |
+
parser.add_argument("--free", type=int, default=1)
|
| 225 |
+
parser.add_argument("--data", type=str, required=True)
|
| 226 |
+
parser.add_argument("--imagesize", type=int, choices=[256, 512], default=256)
|
| 227 |
+
parser.add_argument("--ckpt", type=str, required=True, help="Optional path to a DiT checkpoint.")
|
| 228 |
+
parser.add_argument("--vae", type=str, required=True, help="Optional path to a vae checkpoint.")
|
| 229 |
+
parser.add_argument("--epochs", type=int, default=501, help="Number of training epochs.")
|
| 230 |
+
args = parser.parse_args()
|
| 231 |
+
main(args)
|
ArtiAgent - DefectDiffu/src/GroundingDINO/LICENSE
ADDED
|
@@ -0,0 +1,201 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Apache License
|
| 2 |
+
Version 2.0, January 2004
|
| 3 |
+
http://www.apache.org/licenses/
|
| 4 |
+
|
| 5 |
+
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
| 6 |
+
|
| 7 |
+
1. Definitions.
|
| 8 |
+
|
| 9 |
+
"License" shall mean the terms and conditions for use, reproduction,
|
| 10 |
+
and distribution as defined by Sections 1 through 9 of this document.
|
| 11 |
+
|
| 12 |
+
"Licensor" shall mean the copyright owner or entity authorized by
|
| 13 |
+
the copyright owner that is granting the License.
|
| 14 |
+
|
| 15 |
+
"Legal Entity" shall mean the union of the acting entity and all
|
| 16 |
+
other entities that control, are controlled by, or are under common
|
| 17 |
+
control with that entity. For the purposes of this definition,
|
| 18 |
+
"control" means (i) the power, direct or indirect, to cause the
|
| 19 |
+
direction or management of such entity, whether by contract or
|
| 20 |
+
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
| 21 |
+
outstanding shares, or (iii) beneficial ownership of such entity.
|
| 22 |
+
|
| 23 |
+
"You" (or "Your") shall mean an individual or Legal Entity
|
| 24 |
+
exercising permissions granted by this License.
|
| 25 |
+
|
| 26 |
+
"Source" form shall mean the preferred form for making modifications,
|
| 27 |
+
including but not limited to software source code, documentation
|
| 28 |
+
source, and configuration files.
|
| 29 |
+
|
| 30 |
+
"Object" form shall mean any form resulting from mechanical
|
| 31 |
+
transformation or translation of a Source form, including but
|
| 32 |
+
not limited to compiled object code, generated documentation,
|
| 33 |
+
and conversions to other media types.
|
| 34 |
+
|
| 35 |
+
"Work" shall mean the work of authorship, whether in Source or
|
| 36 |
+
Object form, made available under the License, as indicated by a
|
| 37 |
+
copyright notice that is included in or attached to the work
|
| 38 |
+
(an example is provided in the Appendix below).
|
| 39 |
+
|
| 40 |
+
"Derivative Works" shall mean any work, whether in Source or Object
|
| 41 |
+
form, that is based on (or derived from) the Work and for which the
|
| 42 |
+
editorial revisions, annotations, elaborations, or other modifications
|
| 43 |
+
represent, as a whole, an original work of authorship. For the purposes
|
| 44 |
+
of this License, Derivative Works shall not include works that remain
|
| 45 |
+
separable from, or merely link (or bind by name) to the interfaces of,
|
| 46 |
+
the Work and Derivative Works thereof.
|
| 47 |
+
|
| 48 |
+
"Contribution" shall mean any work of authorship, including
|
| 49 |
+
the original version of the Work and any modifications or additions
|
| 50 |
+
to that Work or Derivative Works thereof, that is intentionally
|
| 51 |
+
submitted to Licensor for inclusion in the Work by the copyright owner
|
| 52 |
+
or by an individual or Legal Entity authorized to submit on behalf of
|
| 53 |
+
the copyright owner. For the purposes of this definition, "submitted"
|
| 54 |
+
means any form of electronic, verbal, or written communication sent
|
| 55 |
+
to the Licensor or its representatives, including but not limited to
|
| 56 |
+
communication on electronic mailing lists, source code control systems,
|
| 57 |
+
and issue tracking systems that are managed by, or on behalf of, the
|
| 58 |
+
Licensor for the purpose of discussing and improving the Work, but
|
| 59 |
+
excluding communication that is conspicuously marked or otherwise
|
| 60 |
+
designated in writing by the copyright owner as "Not a Contribution."
|
| 61 |
+
|
| 62 |
+
"Contributor" shall mean Licensor and any individual or Legal Entity
|
| 63 |
+
on behalf of whom a Contribution has been received by Licensor and
|
| 64 |
+
subsequently incorporated within the Work.
|
| 65 |
+
|
| 66 |
+
2. Grant of Copyright License. Subject to the terms and conditions of
|
| 67 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 68 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 69 |
+
copyright license to reproduce, prepare Derivative Works of,
|
| 70 |
+
publicly display, publicly perform, sublicense, and distribute the
|
| 71 |
+
Work and such Derivative Works in Source or Object form.
|
| 72 |
+
|
| 73 |
+
3. Grant of Patent License. Subject to the terms and conditions of
|
| 74 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 75 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 76 |
+
(except as stated in this section) patent license to make, have made,
|
| 77 |
+
use, offer to sell, sell, import, and otherwise transfer the Work,
|
| 78 |
+
where such license applies only to those patent claims licensable
|
| 79 |
+
by such Contributor that are necessarily infringed by their
|
| 80 |
+
Contribution(s) alone or by combination of their Contribution(s)
|
| 81 |
+
with the Work to which such Contribution(s) was submitted. If You
|
| 82 |
+
institute patent litigation against any entity (including a
|
| 83 |
+
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
| 84 |
+
or a Contribution incorporated within the Work constitutes direct
|
| 85 |
+
or contributory patent infringement, then any patent licenses
|
| 86 |
+
granted to You under this License for that Work shall terminate
|
| 87 |
+
as of the date such litigation is filed.
|
| 88 |
+
|
| 89 |
+
4. Redistribution. You may reproduce and distribute copies of the
|
| 90 |
+
Work or Derivative Works thereof in any medium, with or without
|
| 91 |
+
modifications, and in Source or Object form, provided that You
|
| 92 |
+
meet the following conditions:
|
| 93 |
+
|
| 94 |
+
(a) You must give any other recipients of the Work or
|
| 95 |
+
Derivative Works a copy of this License; and
|
| 96 |
+
|
| 97 |
+
(b) You must cause any modified files to carry prominent notices
|
| 98 |
+
stating that You changed the files; and
|
| 99 |
+
|
| 100 |
+
(c) You must retain, in the Source form of any Derivative Works
|
| 101 |
+
that You distribute, all copyright, patent, trademark, and
|
| 102 |
+
attribution notices from the Source form of the Work,
|
| 103 |
+
excluding those notices that do not pertain to any part of
|
| 104 |
+
the Derivative Works; and
|
| 105 |
+
|
| 106 |
+
(d) If the Work includes a "NOTICE" text file as part of its
|
| 107 |
+
distribution, then any Derivative Works that You distribute must
|
| 108 |
+
include a readable copy of the attribution notices contained
|
| 109 |
+
within such NOTICE file, excluding those notices that do not
|
| 110 |
+
pertain to any part of the Derivative Works, in at least one
|
| 111 |
+
of the following places: within a NOTICE text file distributed
|
| 112 |
+
as part of the Derivative Works; within the Source form or
|
| 113 |
+
documentation, if provided along with the Derivative Works; or,
|
| 114 |
+
within a display generated by the Derivative Works, if and
|
| 115 |
+
wherever such third-party notices normally appear. The contents
|
| 116 |
+
of the NOTICE file are for informational purposes only and
|
| 117 |
+
do not modify the License. You may add Your own attribution
|
| 118 |
+
notices within Derivative Works that You distribute, alongside
|
| 119 |
+
or as an addendum to the NOTICE text from the Work, provided
|
| 120 |
+
that such additional attribution notices cannot be construed
|
| 121 |
+
as modifying the License.
|
| 122 |
+
|
| 123 |
+
You may add Your own copyright statement to Your modifications and
|
| 124 |
+
may provide additional or different license terms and conditions
|
| 125 |
+
for use, reproduction, or distribution of Your modifications, or
|
| 126 |
+
for any such Derivative Works as a whole, provided Your use,
|
| 127 |
+
reproduction, and distribution of the Work otherwise complies with
|
| 128 |
+
the conditions stated in this License.
|
| 129 |
+
|
| 130 |
+
5. Submission of Contributions. Unless You explicitly state otherwise,
|
| 131 |
+
any Contribution intentionally submitted for inclusion in the Work
|
| 132 |
+
by You to the Licensor shall be under the terms and conditions of
|
| 133 |
+
this License, without any additional terms or conditions.
|
| 134 |
+
Notwithstanding the above, nothing herein shall supersede or modify
|
| 135 |
+
the terms of any separate license agreement you may have executed
|
| 136 |
+
with Licensor regarding such Contributions.
|
| 137 |
+
|
| 138 |
+
6. Trademarks. This License does not grant permission to use the trade
|
| 139 |
+
names, trademarks, service marks, or product names of the Licensor,
|
| 140 |
+
except as required for reasonable and customary use in describing the
|
| 141 |
+
origin of the Work and reproducing the content of the NOTICE file.
|
| 142 |
+
|
| 143 |
+
7. Disclaimer of Warranty. Unless required by applicable law or
|
| 144 |
+
agreed to in writing, Licensor provides the Work (and each
|
| 145 |
+
Contributor provides its Contributions) on an "AS IS" BASIS,
|
| 146 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
| 147 |
+
implied, including, without limitation, any warranties or conditions
|
| 148 |
+
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
| 149 |
+
PARTICULAR PURPOSE. You are solely responsible for determining the
|
| 150 |
+
appropriateness of using or redistributing the Work and assume any
|
| 151 |
+
risks associated with Your exercise of permissions under this License.
|
| 152 |
+
|
| 153 |
+
8. Limitation of Liability. In no event and under no legal theory,
|
| 154 |
+
whether in tort (including negligence), contract, or otherwise,
|
| 155 |
+
unless required by applicable law (such as deliberate and grossly
|
| 156 |
+
negligent acts) or agreed to in writing, shall any Contributor be
|
| 157 |
+
liable to You for damages, including any direct, indirect, special,
|
| 158 |
+
incidental, or consequential damages of any character arising as a
|
| 159 |
+
result of this License or out of the use or inability to use the
|
| 160 |
+
Work (including but not limited to damages for loss of goodwill,
|
| 161 |
+
work stoppage, computer failure or malfunction, or any and all
|
| 162 |
+
other commercial damages or losses), even if such Contributor
|
| 163 |
+
has been advised of the possibility of such damages.
|
| 164 |
+
|
| 165 |
+
9. Accepting Warranty or Additional Liability. While redistributing
|
| 166 |
+
the Work or Derivative Works thereof, You may choose to offer,
|
| 167 |
+
and charge a fee for, acceptance of support, warranty, indemnity,
|
| 168 |
+
or other liability obligations and/or rights consistent with this
|
| 169 |
+
License. However, in accepting such obligations, You may act only
|
| 170 |
+
on Your own behalf and on Your sole responsibility, not on behalf
|
| 171 |
+
of any other Contributor, and only if You agree to indemnify,
|
| 172 |
+
defend, and hold each Contributor harmless for any liability
|
| 173 |
+
incurred by, or claims asserted against, such Contributor by reason
|
| 174 |
+
of your accepting any such warranty or additional liability.
|
| 175 |
+
|
| 176 |
+
END OF TERMS AND CONDITIONS
|
| 177 |
+
|
| 178 |
+
APPENDIX: How to apply the Apache License to your work.
|
| 179 |
+
|
| 180 |
+
To apply the Apache License to your work, attach the following
|
| 181 |
+
boilerplate notice, with the fields enclosed by brackets "[]"
|
| 182 |
+
replaced with your own identifying information. (Don't include
|
| 183 |
+
the brackets!) The text should be enclosed in the appropriate
|
| 184 |
+
comment syntax for the file format. We also recommend that a
|
| 185 |
+
file or class name and description of purpose be included on the
|
| 186 |
+
same "printed page" as the copyright notice for easier
|
| 187 |
+
identification within third-party archives.
|
| 188 |
+
|
| 189 |
+
Copyright 2020 - present, Facebook, Inc
|
| 190 |
+
|
| 191 |
+
Licensed under the Apache License, Version 2.0 (the "License");
|
| 192 |
+
you may not use this file except in compliance with the License.
|
| 193 |
+
You may obtain a copy of the License at
|
| 194 |
+
|
| 195 |
+
http://www.apache.org/licenses/LICENSE-2.0
|
| 196 |
+
|
| 197 |
+
Unless required by applicable law or agreed to in writing, software
|
| 198 |
+
distributed under the License is distributed on an "AS IS" BASIS,
|
| 199 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 200 |
+
See the License for the specific language governing permissions and
|
| 201 |
+
limitations under the License.
|
ArtiAgent - DefectDiffu/src/GroundingDINO/README.md
ADDED
|
@@ -0,0 +1,163 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Grounding DINO
|
| 2 |
+
|
| 3 |
+
---
|
| 4 |
+
|
| 5 |
+
[](https://arxiv.org/abs/2303.05499)
|
| 6 |
+
[](https://youtu.be/wxWDt5UiwY8)
|
| 7 |
+
[](https://colab.research.google.com/github/roboflow-ai/notebooks/blob/main/notebooks/zero-shot-object-detection-with-grounding-dino.ipynb)
|
| 8 |
+
[](https://youtu.be/cMa77r3YrDk)
|
| 9 |
+
[](https://huggingface.co/spaces/ShilongLiu/Grounding_DINO_demo)
|
| 10 |
+
|
| 11 |
+
[](https://paperswithcode.com/sota/zero-shot-object-detection-on-mscoco?p=grounding-dino-marrying-dino-with-grounded) \
|
| 12 |
+
[](https://paperswithcode.com/sota/zero-shot-object-detection-on-odinw?p=grounding-dino-marrying-dino-with-grounded) \
|
| 13 |
+
[](https://paperswithcode.com/sota/object-detection-on-coco-minival?p=grounding-dino-marrying-dino-with-grounded) \
|
| 14 |
+
[](https://paperswithcode.com/sota/object-detection-on-coco?p=grounding-dino-marrying-dino-with-grounded)
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
Official PyTorch implementation of [Grounding DINO](https://arxiv.org/abs/2303.05499), a stronger open-set object detector. Code is available now!
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
## Highlight
|
| 22 |
+
|
| 23 |
+
- **Open-Set Detection.** Detect **everything** with language!
|
| 24 |
+
- **High Performancce.** COCO zero-shot **52.5 AP** (training without COCO data!). COCO fine-tune **63.0 AP**.
|
| 25 |
+
- **Flexible.** Collaboration with Stable Diffusion for Image Editting.
|
| 26 |
+
|
| 27 |
+
## News
|
| 28 |
+
[2023/03/28] A YouTube [video](https://youtu.be/cMa77r3YrDk) about Grounding DINO and basic object detection prompt engineering. [[SkalskiP](https://github.com/SkalskiP)] \
|
| 29 |
+
[2023/03/28] Add a [demo](https://huggingface.co/spaces/ShilongLiu/Grounding_DINO_demo) on Hugging Face Space! \
|
| 30 |
+
[2023/03/27] Support CPU-only mode. Now the model can run on machines without GPUs.\
|
| 31 |
+
[2023/03/25] A [demo](https://colab.research.google.com/github/roboflow-ai/notebooks/blob/main/notebooks/zero-shot-object-detection-with-grounding-dino.ipynb) for Grounding DINO is available at Colab. [[SkalskiP](https://github.com/SkalskiP)] \
|
| 32 |
+
[2023/03/22] Code is available Now!
|
| 33 |
+
|
| 34 |
+
<details open>
|
| 35 |
+
<summary><font size="4">
|
| 36 |
+
Description
|
| 37 |
+
</font></summary>
|
| 38 |
+
<img src=".asset/hero_figure.png" alt="ODinW" width="100%">
|
| 39 |
+
</details>
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
## TODO
|
| 44 |
+
|
| 45 |
+
- [x] Release inference code and demo.
|
| 46 |
+
- [x] Release checkpoints.
|
| 47 |
+
- [ ] Grounding DINO with Stable Diffusion and GLIGEN demos.
|
| 48 |
+
- [ ] Release training codes.
|
| 49 |
+
|
| 50 |
+
## Install
|
| 51 |
+
|
| 52 |
+
If you have a CUDA environment, please make sure the environment variable `CUDA_HOME` is set. It will be compiled under CPU-only mode if no CUDA available.
|
| 53 |
+
|
| 54 |
+
```bash
|
| 55 |
+
pip install -e .
|
| 56 |
+
```
|
| 57 |
+
|
| 58 |
+
## Demo
|
| 59 |
+
|
| 60 |
+
```bash
|
| 61 |
+
CUDA_VISIBLE_DEVICES=6 python demo/inference_on_a_image.py \
|
| 62 |
+
-c /path/to/config \
|
| 63 |
+
-p /path/to/checkpoint \
|
| 64 |
+
-i .asset/cats.png \
|
| 65 |
+
-o "outputs/0" \
|
| 66 |
+
-t "cat ear." \
|
| 67 |
+
[--cpu-only] # open it for cpu mode
|
| 68 |
+
```
|
| 69 |
+
See the `demo/inference_on_a_image.py` for more details.
|
| 70 |
+
|
| 71 |
+
**Web UI**
|
| 72 |
+
|
| 73 |
+
We also provide a demo code to integrate Grounding DINO with Gradio Web UI. See the file `demo/gradio_app.py` for more details.
|
| 74 |
+
|
| 75 |
+
## Checkpoints
|
| 76 |
+
|
| 77 |
+
<!-- insert a table -->
|
| 78 |
+
<table>
|
| 79 |
+
<thead>
|
| 80 |
+
<tr style="text-align: right;">
|
| 81 |
+
<th></th>
|
| 82 |
+
<th>name</th>
|
| 83 |
+
<th>backbone</th>
|
| 84 |
+
<th>Data</th>
|
| 85 |
+
<th>box AP on COCO</th>
|
| 86 |
+
<th>Checkpoint</th>
|
| 87 |
+
<th>Config</th>
|
| 88 |
+
</tr>
|
| 89 |
+
</thead>
|
| 90 |
+
<tbody>
|
| 91 |
+
<tr>
|
| 92 |
+
<th>1</th>
|
| 93 |
+
<td>GroundingDINO-T</td>
|
| 94 |
+
<td>Swin-T</td>
|
| 95 |
+
<td>O365,GoldG,Cap4M</td>
|
| 96 |
+
<td>48.4 (zero-shot) / 57.2 (fine-tune)</td>
|
| 97 |
+
<td><a href="https://github.com/IDEA-Research/GroundingDINO/releases/download/v0.1.0-alpha/groundingdino_swint_ogc.pth">Github link</a> | <a href="https://huggingface.co/ShilongLiu/GroundingDINO/resolve/main/groundingdino_swint_ogc.pth">HF link</a></td>
|
| 98 |
+
<td><a href="https://github.com/IDEA-Research/GroundingDINO/blob/main/groundingdino/config/GroundingDINO_SwinT_OGC.py">link</a></td>
|
| 99 |
+
</tr>
|
| 100 |
+
</tbody>
|
| 101 |
+
</table>
|
| 102 |
+
|
| 103 |
+
## Results
|
| 104 |
+
|
| 105 |
+
<details open>
|
| 106 |
+
<summary><font size="4">
|
| 107 |
+
COCO Object Detection Results
|
| 108 |
+
</font></summary>
|
| 109 |
+
<img src=".asset/COCO.png" alt="COCO" width="100%">
|
| 110 |
+
</details>
|
| 111 |
+
|
| 112 |
+
<details open>
|
| 113 |
+
<summary><font size="4">
|
| 114 |
+
ODinW Object Detection Results
|
| 115 |
+
</font></summary>
|
| 116 |
+
<img src=".asset/ODinW.png" alt="ODinW" width="100%">
|
| 117 |
+
</details>
|
| 118 |
+
|
| 119 |
+
<details open>
|
| 120 |
+
<summary><font size="4">
|
| 121 |
+
Marrying Grounding DINO with <a href="https://github.com/Stability-AI/StableDiffusion">Stable Diffusion</a> for Image Editing
|
| 122 |
+
</font></summary>
|
| 123 |
+
<img src=".asset/GD_SD.png" alt="GD_SD" width="100%">
|
| 124 |
+
</details>
|
| 125 |
+
|
| 126 |
+
<details open>
|
| 127 |
+
<summary><font size="4">
|
| 128 |
+
Marrying Grounding DINO with <a href="https://github.com/gligen/GLIGEN">GLIGEN</a> for more Detailed Image Editing
|
| 129 |
+
</font></summary>
|
| 130 |
+
<img src=".asset/GD_GLIGEN.png" alt="GD_GLIGEN" width="100%">
|
| 131 |
+
</details>
|
| 132 |
+
|
| 133 |
+
## Model
|
| 134 |
+
|
| 135 |
+
Includes: a text backbone, an image backbone, a feature enhancer, a language-guided query selection, and a cross-modality decoder.
|
| 136 |
+
|
| 137 |
+

|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
## Acknowledgement
|
| 141 |
+
|
| 142 |
+
Our model is related to [DINO](https://github.com/IDEA-Research/DINO) and [GLIP](https://github.com/microsoft/GLIP). Thanks for their great work!
|
| 143 |
+
|
| 144 |
+
We also thank great previous work including DETR, Deformable DETR, SMCA, Conditional DETR, Anchor DETR, Dynamic DETR, DAB-DETR, DN-DETR, etc. More related work are available at [Awesome Detection Transformer](https://github.com/IDEACVR/awesome-detection-transformer). A new toolbox [detrex](https://github.com/IDEA-Research/detrex) is available as well.
|
| 145 |
+
|
| 146 |
+
Thanks [Stable Diffusion](https://github.com/Stability-AI/StableDiffusion) and [GLIGEN](https://github.com/gligen/GLIGEN) for their awesome models.
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
## Citation
|
| 150 |
+
|
| 151 |
+
If you find our work helpful for your research, please consider citing the following BibTeX entry.
|
| 152 |
+
|
| 153 |
+
```bibtex
|
| 154 |
+
@inproceedings{ShilongLiu2023GroundingDM,
|
| 155 |
+
title={Grounding DINO: Marrying DINO with Grounded Pre-Training for Open-Set Object Detection},
|
| 156 |
+
author={Shilong Liu and Zhaoyang Zeng and Tianhe Ren and Feng Li and Hao Zhang and Jie Yang and Chunyuan Li and Jianwei Yang and Hang Su and Jun Zhu and Lei Zhang},
|
| 157 |
+
year={2023}
|
| 158 |
+
}
|
| 159 |
+
```
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
|
| 163 |
+
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/__init__.py
ADDED
|
File without changes
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/__pycache__/__init__.cpython-310.pyc
ADDED
|
Binary file (192 Bytes). View file
|
|
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/config/GroundingDINO_SwinB.py
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
batch_size = 1
|
| 2 |
+
modelname = "groundingdino"
|
| 3 |
+
backbone = "swin_B_384_22k"
|
| 4 |
+
position_embedding = "sine"
|
| 5 |
+
pe_temperatureH = 20
|
| 6 |
+
pe_temperatureW = 20
|
| 7 |
+
return_interm_indices = [1, 2, 3]
|
| 8 |
+
backbone_freeze_keywords = None
|
| 9 |
+
enc_layers = 6
|
| 10 |
+
dec_layers = 6
|
| 11 |
+
pre_norm = False
|
| 12 |
+
dim_feedforward = 2048
|
| 13 |
+
hidden_dim = 256
|
| 14 |
+
dropout = 0.0
|
| 15 |
+
nheads = 8
|
| 16 |
+
num_queries = 900
|
| 17 |
+
query_dim = 4
|
| 18 |
+
num_patterns = 0
|
| 19 |
+
num_feature_levels = 4
|
| 20 |
+
enc_n_points = 4
|
| 21 |
+
dec_n_points = 4
|
| 22 |
+
two_stage_type = "standard"
|
| 23 |
+
two_stage_bbox_embed_share = False
|
| 24 |
+
two_stage_class_embed_share = False
|
| 25 |
+
transformer_activation = "relu"
|
| 26 |
+
dec_pred_bbox_embed_share = True
|
| 27 |
+
dn_box_noise_scale = 1.0
|
| 28 |
+
dn_label_noise_ratio = 0.5
|
| 29 |
+
dn_label_coef = 1.0
|
| 30 |
+
dn_bbox_coef = 1.0
|
| 31 |
+
embed_init_tgt = True
|
| 32 |
+
dn_labelbook_size = 2000
|
| 33 |
+
max_text_len = 256
|
| 34 |
+
text_encoder_type = "bert-base-uncased"
|
| 35 |
+
use_text_enhancer = True
|
| 36 |
+
use_fusion_layer = True
|
| 37 |
+
use_checkpoint = True
|
| 38 |
+
use_transformer_ckpt = True
|
| 39 |
+
use_text_cross_attention = True
|
| 40 |
+
text_dropout = 0.0
|
| 41 |
+
fusion_dropout = 0.0
|
| 42 |
+
fusion_droppath = 0.1
|
| 43 |
+
sub_sentence_present = True
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/config/GroundingDINO_SwinT_OGC.py
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
batch_size = 1
|
| 2 |
+
modelname = "groundingdino"
|
| 3 |
+
backbone = "swin_T_224_1k"
|
| 4 |
+
position_embedding = "sine"
|
| 5 |
+
pe_temperatureH = 20
|
| 6 |
+
pe_temperatureW = 20
|
| 7 |
+
return_interm_indices = [1, 2, 3]
|
| 8 |
+
backbone_freeze_keywords = None
|
| 9 |
+
enc_layers = 6
|
| 10 |
+
dec_layers = 6
|
| 11 |
+
pre_norm = False
|
| 12 |
+
dim_feedforward = 2048
|
| 13 |
+
hidden_dim = 256
|
| 14 |
+
dropout = 0.0
|
| 15 |
+
nheads = 8
|
| 16 |
+
num_queries = 900
|
| 17 |
+
query_dim = 4
|
| 18 |
+
num_patterns = 0
|
| 19 |
+
num_feature_levels = 4
|
| 20 |
+
enc_n_points = 4
|
| 21 |
+
dec_n_points = 4
|
| 22 |
+
two_stage_type = "standard"
|
| 23 |
+
two_stage_bbox_embed_share = False
|
| 24 |
+
two_stage_class_embed_share = False
|
| 25 |
+
transformer_activation = "relu"
|
| 26 |
+
dec_pred_bbox_embed_share = True
|
| 27 |
+
dn_box_noise_scale = 1.0
|
| 28 |
+
dn_label_noise_ratio = 0.5
|
| 29 |
+
dn_label_coef = 1.0
|
| 30 |
+
dn_bbox_coef = 1.0
|
| 31 |
+
embed_init_tgt = True
|
| 32 |
+
dn_labelbook_size = 2000
|
| 33 |
+
max_text_len = 256
|
| 34 |
+
text_encoder_type = "bert-base-uncased"
|
| 35 |
+
use_text_enhancer = True
|
| 36 |
+
use_fusion_layer = True
|
| 37 |
+
use_checkpoint = True
|
| 38 |
+
use_transformer_ckpt = True
|
| 39 |
+
use_text_cross_attention = True
|
| 40 |
+
text_dropout = 0.0
|
| 41 |
+
fusion_dropout = 0.0
|
| 42 |
+
fusion_droppath = 0.1
|
| 43 |
+
sub_sentence_present = True
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/datasets/__init__.py
ADDED
|
File without changes
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/datasets/__pycache__/__init__.cpython-310.pyc
ADDED
|
Binary file (201 Bytes). View file
|
|
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/datasets/__pycache__/transforms.cpython-310.pyc
ADDED
|
Binary file (10.2 kB). View file
|
|
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/datasets/transforms.py
ADDED
|
@@ -0,0 +1,311 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved
|
| 2 |
+
"""
|
| 3 |
+
Transforms and data augmentation for both image + bbox.
|
| 4 |
+
"""
|
| 5 |
+
import os
|
| 6 |
+
import random
|
| 7 |
+
|
| 8 |
+
import PIL
|
| 9 |
+
import torch
|
| 10 |
+
import torchvision.transforms as T
|
| 11 |
+
import torchvision.transforms.functional as F
|
| 12 |
+
|
| 13 |
+
from groundingdino.util.box_ops import box_xyxy_to_cxcywh
|
| 14 |
+
from groundingdino.util.misc import interpolate
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def crop(image, target, region):
|
| 18 |
+
cropped_image = F.crop(image, *region)
|
| 19 |
+
|
| 20 |
+
target = target.copy()
|
| 21 |
+
i, j, h, w = region
|
| 22 |
+
|
| 23 |
+
# should we do something wrt the original size?
|
| 24 |
+
target["size"] = torch.tensor([h, w])
|
| 25 |
+
|
| 26 |
+
fields = ["labels", "area", "iscrowd", "positive_map"]
|
| 27 |
+
|
| 28 |
+
if "boxes" in target:
|
| 29 |
+
boxes = target["boxes"]
|
| 30 |
+
max_size = torch.as_tensor([w, h], dtype=torch.float32)
|
| 31 |
+
cropped_boxes = boxes - torch.as_tensor([j, i, j, i])
|
| 32 |
+
cropped_boxes = torch.min(cropped_boxes.reshape(-1, 2, 2), max_size)
|
| 33 |
+
cropped_boxes = cropped_boxes.clamp(min=0)
|
| 34 |
+
area = (cropped_boxes[:, 1, :] - cropped_boxes[:, 0, :]).prod(dim=1)
|
| 35 |
+
target["boxes"] = cropped_boxes.reshape(-1, 4)
|
| 36 |
+
target["area"] = area
|
| 37 |
+
fields.append("boxes")
|
| 38 |
+
|
| 39 |
+
if "masks" in target:
|
| 40 |
+
# FIXME should we update the area here if there are no boxes?
|
| 41 |
+
target["masks"] = target["masks"][:, i : i + h, j : j + w]
|
| 42 |
+
fields.append("masks")
|
| 43 |
+
|
| 44 |
+
# remove elements for which the boxes or masks that have zero area
|
| 45 |
+
if "boxes" in target or "masks" in target:
|
| 46 |
+
# favor boxes selection when defining which elements to keep
|
| 47 |
+
# this is compatible with previous implementation
|
| 48 |
+
if "boxes" in target:
|
| 49 |
+
cropped_boxes = target["boxes"].reshape(-1, 2, 2)
|
| 50 |
+
keep = torch.all(cropped_boxes[:, 1, :] > cropped_boxes[:, 0, :], dim=1)
|
| 51 |
+
else:
|
| 52 |
+
keep = target["masks"].flatten(1).any(1)
|
| 53 |
+
|
| 54 |
+
for field in fields:
|
| 55 |
+
if field in target:
|
| 56 |
+
target[field] = target[field][keep]
|
| 57 |
+
|
| 58 |
+
if os.environ.get("IPDB_SHILONG_DEBUG", None) == "INFO":
|
| 59 |
+
# for debug and visualization only.
|
| 60 |
+
if "strings_positive" in target:
|
| 61 |
+
target["strings_positive"] = [
|
| 62 |
+
_i for _i, _j in zip(target["strings_positive"], keep) if _j
|
| 63 |
+
]
|
| 64 |
+
|
| 65 |
+
return cropped_image, target
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def hflip(image, target):
|
| 69 |
+
flipped_image = F.hflip(image)
|
| 70 |
+
|
| 71 |
+
w, h = image.size
|
| 72 |
+
|
| 73 |
+
target = target.copy()
|
| 74 |
+
if "boxes" in target:
|
| 75 |
+
boxes = target["boxes"]
|
| 76 |
+
boxes = boxes[:, [2, 1, 0, 3]] * torch.as_tensor([-1, 1, -1, 1]) + torch.as_tensor(
|
| 77 |
+
[w, 0, w, 0]
|
| 78 |
+
)
|
| 79 |
+
target["boxes"] = boxes
|
| 80 |
+
|
| 81 |
+
if "masks" in target:
|
| 82 |
+
target["masks"] = target["masks"].flip(-1)
|
| 83 |
+
|
| 84 |
+
return flipped_image, target
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def resize(image, target, size, max_size=None):
|
| 88 |
+
# size can be min_size (scalar) or (w, h) tuple
|
| 89 |
+
|
| 90 |
+
def get_size_with_aspect_ratio(image_size, size, max_size=None):
|
| 91 |
+
w, h = image_size
|
| 92 |
+
if max_size is not None:
|
| 93 |
+
min_original_size = float(min((w, h)))
|
| 94 |
+
max_original_size = float(max((w, h)))
|
| 95 |
+
if max_original_size / min_original_size * size > max_size:
|
| 96 |
+
size = int(round(max_size * min_original_size / max_original_size))
|
| 97 |
+
|
| 98 |
+
if (w <= h and w == size) or (h <= w and h == size):
|
| 99 |
+
return (h, w)
|
| 100 |
+
|
| 101 |
+
if w < h:
|
| 102 |
+
ow = size
|
| 103 |
+
oh = int(size * h / w)
|
| 104 |
+
else:
|
| 105 |
+
oh = size
|
| 106 |
+
ow = int(size * w / h)
|
| 107 |
+
|
| 108 |
+
return (oh, ow)
|
| 109 |
+
|
| 110 |
+
def get_size(image_size, size, max_size=None):
|
| 111 |
+
if isinstance(size, (list, tuple)):
|
| 112 |
+
return size[::-1]
|
| 113 |
+
else:
|
| 114 |
+
return get_size_with_aspect_ratio(image_size, size, max_size)
|
| 115 |
+
|
| 116 |
+
size = get_size(image.size, size, max_size)
|
| 117 |
+
rescaled_image = F.resize(image, size)
|
| 118 |
+
|
| 119 |
+
if target is None:
|
| 120 |
+
return rescaled_image, None
|
| 121 |
+
|
| 122 |
+
ratios = tuple(float(s) / float(s_orig) for s, s_orig in zip(rescaled_image.size, image.size))
|
| 123 |
+
ratio_width, ratio_height = ratios
|
| 124 |
+
|
| 125 |
+
target = target.copy()
|
| 126 |
+
if "boxes" in target:
|
| 127 |
+
boxes = target["boxes"]
|
| 128 |
+
scaled_boxes = boxes * torch.as_tensor(
|
| 129 |
+
[ratio_width, ratio_height, ratio_width, ratio_height]
|
| 130 |
+
)
|
| 131 |
+
target["boxes"] = scaled_boxes
|
| 132 |
+
|
| 133 |
+
if "area" in target:
|
| 134 |
+
area = target["area"]
|
| 135 |
+
scaled_area = area * (ratio_width * ratio_height)
|
| 136 |
+
target["area"] = scaled_area
|
| 137 |
+
|
| 138 |
+
h, w = size
|
| 139 |
+
target["size"] = torch.tensor([h, w])
|
| 140 |
+
|
| 141 |
+
if "masks" in target:
|
| 142 |
+
target["masks"] = (
|
| 143 |
+
interpolate(target["masks"][:, None].float(), size, mode="nearest")[:, 0] > 0.5
|
| 144 |
+
)
|
| 145 |
+
|
| 146 |
+
return rescaled_image, target
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
def pad(image, target, padding):
|
| 150 |
+
# assumes that we only pad on the bottom right corners
|
| 151 |
+
padded_image = F.pad(image, (0, 0, padding[0], padding[1]))
|
| 152 |
+
if target is None:
|
| 153 |
+
return padded_image, None
|
| 154 |
+
target = target.copy()
|
| 155 |
+
# should we do something wrt the original size?
|
| 156 |
+
target["size"] = torch.tensor(padded_image.size[::-1])
|
| 157 |
+
if "masks" in target:
|
| 158 |
+
target["masks"] = torch.nn.functional.pad(target["masks"], (0, padding[0], 0, padding[1]))
|
| 159 |
+
return padded_image, target
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
class ResizeDebug(object):
|
| 163 |
+
def __init__(self, size):
|
| 164 |
+
self.size = size
|
| 165 |
+
|
| 166 |
+
def __call__(self, img, target):
|
| 167 |
+
return resize(img, target, self.size)
|
| 168 |
+
|
| 169 |
+
|
| 170 |
+
class RandomCrop(object):
|
| 171 |
+
def __init__(self, size):
|
| 172 |
+
self.size = size
|
| 173 |
+
|
| 174 |
+
def __call__(self, img, target):
|
| 175 |
+
region = T.RandomCrop.get_params(img, self.size)
|
| 176 |
+
return crop(img, target, region)
|
| 177 |
+
|
| 178 |
+
|
| 179 |
+
class RandomSizeCrop(object):
|
| 180 |
+
def __init__(self, min_size: int, max_size: int, respect_boxes: bool = False):
|
| 181 |
+
# respect_boxes: True to keep all boxes
|
| 182 |
+
# False to tolerence box filter
|
| 183 |
+
self.min_size = min_size
|
| 184 |
+
self.max_size = max_size
|
| 185 |
+
self.respect_boxes = respect_boxes
|
| 186 |
+
|
| 187 |
+
def __call__(self, img: PIL.Image.Image, target: dict):
|
| 188 |
+
init_boxes = len(target["boxes"])
|
| 189 |
+
max_patience = 10
|
| 190 |
+
for i in range(max_patience):
|
| 191 |
+
w = random.randint(self.min_size, min(img.width, self.max_size))
|
| 192 |
+
h = random.randint(self.min_size, min(img.height, self.max_size))
|
| 193 |
+
region = T.RandomCrop.get_params(img, [h, w])
|
| 194 |
+
result_img, result_target = crop(img, target, region)
|
| 195 |
+
if (
|
| 196 |
+
not self.respect_boxes
|
| 197 |
+
or len(result_target["boxes"]) == init_boxes
|
| 198 |
+
or i == max_patience - 1
|
| 199 |
+
):
|
| 200 |
+
return result_img, result_target
|
| 201 |
+
return result_img, result_target
|
| 202 |
+
|
| 203 |
+
|
| 204 |
+
class CenterCrop(object):
|
| 205 |
+
def __init__(self, size):
|
| 206 |
+
self.size = size
|
| 207 |
+
|
| 208 |
+
def __call__(self, img, target):
|
| 209 |
+
image_width, image_height = img.size
|
| 210 |
+
crop_height, crop_width = self.size
|
| 211 |
+
crop_top = int(round((image_height - crop_height) / 2.0))
|
| 212 |
+
crop_left = int(round((image_width - crop_width) / 2.0))
|
| 213 |
+
return crop(img, target, (crop_top, crop_left, crop_height, crop_width))
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
class RandomHorizontalFlip(object):
|
| 217 |
+
def __init__(self, p=0.5):
|
| 218 |
+
self.p = p
|
| 219 |
+
|
| 220 |
+
def __call__(self, img, target):
|
| 221 |
+
if random.random() < self.p:
|
| 222 |
+
return hflip(img, target)
|
| 223 |
+
return img, target
|
| 224 |
+
|
| 225 |
+
|
| 226 |
+
class RandomResize(object):
|
| 227 |
+
def __init__(self, sizes, max_size=None):
|
| 228 |
+
assert isinstance(sizes, (list, tuple))
|
| 229 |
+
self.sizes = sizes
|
| 230 |
+
self.max_size = max_size
|
| 231 |
+
|
| 232 |
+
def __call__(self, img, target=None):
|
| 233 |
+
size = random.choice(self.sizes)
|
| 234 |
+
return resize(img, target, size, self.max_size)
|
| 235 |
+
|
| 236 |
+
|
| 237 |
+
class RandomPad(object):
|
| 238 |
+
def __init__(self, max_pad):
|
| 239 |
+
self.max_pad = max_pad
|
| 240 |
+
|
| 241 |
+
def __call__(self, img, target):
|
| 242 |
+
pad_x = random.randint(0, self.max_pad)
|
| 243 |
+
pad_y = random.randint(0, self.max_pad)
|
| 244 |
+
return pad(img, target, (pad_x, pad_y))
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
class RandomSelect(object):
|
| 248 |
+
"""
|
| 249 |
+
Randomly selects between transforms1 and transforms2,
|
| 250 |
+
with probability p for transforms1 and (1 - p) for transforms2
|
| 251 |
+
"""
|
| 252 |
+
|
| 253 |
+
def __init__(self, transforms1, transforms2, p=0.5):
|
| 254 |
+
self.transforms1 = transforms1
|
| 255 |
+
self.transforms2 = transforms2
|
| 256 |
+
self.p = p
|
| 257 |
+
|
| 258 |
+
def __call__(self, img, target):
|
| 259 |
+
if random.random() < self.p:
|
| 260 |
+
return self.transforms1(img, target)
|
| 261 |
+
return self.transforms2(img, target)
|
| 262 |
+
|
| 263 |
+
|
| 264 |
+
class ToTensor(object):
|
| 265 |
+
def __call__(self, img, target):
|
| 266 |
+
return F.to_tensor(img), target
|
| 267 |
+
|
| 268 |
+
|
| 269 |
+
class RandomErasing(object):
|
| 270 |
+
def __init__(self, *args, **kwargs):
|
| 271 |
+
self.eraser = T.RandomErasing(*args, **kwargs)
|
| 272 |
+
|
| 273 |
+
def __call__(self, img, target):
|
| 274 |
+
return self.eraser(img), target
|
| 275 |
+
|
| 276 |
+
|
| 277 |
+
class Normalize(object):
|
| 278 |
+
def __init__(self, mean, std):
|
| 279 |
+
self.mean = mean
|
| 280 |
+
self.std = std
|
| 281 |
+
|
| 282 |
+
def __call__(self, image, target=None):
|
| 283 |
+
image = F.normalize(image, mean=self.mean, std=self.std)
|
| 284 |
+
if target is None:
|
| 285 |
+
return image, None
|
| 286 |
+
target = target.copy()
|
| 287 |
+
h, w = image.shape[-2:]
|
| 288 |
+
if "boxes" in target:
|
| 289 |
+
boxes = target["boxes"]
|
| 290 |
+
boxes = box_xyxy_to_cxcywh(boxes)
|
| 291 |
+
boxes = boxes / torch.tensor([w, h, w, h], dtype=torch.float32)
|
| 292 |
+
target["boxes"] = boxes
|
| 293 |
+
return image, target
|
| 294 |
+
|
| 295 |
+
|
| 296 |
+
class Compose(object):
|
| 297 |
+
def __init__(self, transforms):
|
| 298 |
+
self.transforms = transforms
|
| 299 |
+
|
| 300 |
+
def __call__(self, image, target):
|
| 301 |
+
for t in self.transforms:
|
| 302 |
+
image, target = t(image, target)
|
| 303 |
+
return image, target
|
| 304 |
+
|
| 305 |
+
def __repr__(self):
|
| 306 |
+
format_string = self.__class__.__name__ + "("
|
| 307 |
+
for t in self.transforms:
|
| 308 |
+
format_string += "\n"
|
| 309 |
+
format_string += " {0}".format(t)
|
| 310 |
+
format_string += "\n)"
|
| 311 |
+
return format_string
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/__init__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# ------------------------------------------------------------------------
|
| 2 |
+
# Grounding DINO
|
| 3 |
+
# url: https://github.com/IDEA-Research/GroundingDINO
|
| 4 |
+
# Copyright (c) 2023 IDEA. All Rights Reserved.
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
# ------------------------------------------------------------------------
|
| 7 |
+
# Conditional DETR
|
| 8 |
+
# Copyright (c) 2021 Microsoft. All Rights Reserved.
|
| 9 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 10 |
+
# ------------------------------------------------------------------------
|
| 11 |
+
# Copied from DETR (https://github.com/facebookresearch/detr)
|
| 12 |
+
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved.
|
| 13 |
+
# ------------------------------------------------------------------------
|
| 14 |
+
|
| 15 |
+
from .groundingdino import build_groundingdino
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/__pycache__/__init__.cpython-310.pyc
ADDED
|
Binary file (270 Bytes). View file
|
|
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/__pycache__/groundingdino.cpython-310.pyc
ADDED
|
Binary file (10.7 kB). View file
|
|
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/backbone/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
from .backbone import build_backbone
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/backbone/backbone.py
ADDED
|
@@ -0,0 +1,221 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# ------------------------------------------------------------------------
|
| 2 |
+
# Grounding DINO
|
| 3 |
+
# url: https://github.com/IDEA-Research/GroundingDINO
|
| 4 |
+
# Copyright (c) 2023 IDEA. All Rights Reserved.
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
# ------------------------------------------------------------------------
|
| 7 |
+
# Conditional DETR
|
| 8 |
+
# Copyright (c) 2021 Microsoft. All Rights Reserved.
|
| 9 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 10 |
+
# ------------------------------------------------------------------------
|
| 11 |
+
# Copied from DETR (https://github.com/facebookresearch/detr)
|
| 12 |
+
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved.
|
| 13 |
+
# ------------------------------------------------------------------------
|
| 14 |
+
|
| 15 |
+
"""
|
| 16 |
+
Backbone modules.
|
| 17 |
+
"""
|
| 18 |
+
|
| 19 |
+
from typing import Dict, List
|
| 20 |
+
|
| 21 |
+
import torch
|
| 22 |
+
import torch.nn.functional as F
|
| 23 |
+
import torchvision
|
| 24 |
+
from torch import nn
|
| 25 |
+
from torchvision.models._utils import IntermediateLayerGetter
|
| 26 |
+
|
| 27 |
+
from groundingdino.util.misc import NestedTensor, clean_state_dict, is_main_process
|
| 28 |
+
|
| 29 |
+
from .position_encoding import build_position_encoding
|
| 30 |
+
from .swin_transformer import build_swin_transformer
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
class FrozenBatchNorm2d(torch.nn.Module):
|
| 34 |
+
"""
|
| 35 |
+
BatchNorm2d where the batch statistics and the affine parameters are fixed.
|
| 36 |
+
|
| 37 |
+
Copy-paste from torchvision.misc.ops with added eps before rqsrt,
|
| 38 |
+
without which any other models than torchvision.models.resnet[18,34,50,101]
|
| 39 |
+
produce nans.
|
| 40 |
+
"""
|
| 41 |
+
|
| 42 |
+
def __init__(self, n):
|
| 43 |
+
super(FrozenBatchNorm2d, self).__init__()
|
| 44 |
+
self.register_buffer("weight", torch.ones(n))
|
| 45 |
+
self.register_buffer("bias", torch.zeros(n))
|
| 46 |
+
self.register_buffer("running_mean", torch.zeros(n))
|
| 47 |
+
self.register_buffer("running_var", torch.ones(n))
|
| 48 |
+
|
| 49 |
+
def _load_from_state_dict(
|
| 50 |
+
self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs
|
| 51 |
+
):
|
| 52 |
+
num_batches_tracked_key = prefix + "num_batches_tracked"
|
| 53 |
+
if num_batches_tracked_key in state_dict:
|
| 54 |
+
del state_dict[num_batches_tracked_key]
|
| 55 |
+
|
| 56 |
+
super(FrozenBatchNorm2d, self)._load_from_state_dict(
|
| 57 |
+
state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs
|
| 58 |
+
)
|
| 59 |
+
|
| 60 |
+
def forward(self, x):
|
| 61 |
+
# move reshapes to the beginning
|
| 62 |
+
# to make it fuser-friendly
|
| 63 |
+
w = self.weight.reshape(1, -1, 1, 1)
|
| 64 |
+
b = self.bias.reshape(1, -1, 1, 1)
|
| 65 |
+
rv = self.running_var.reshape(1, -1, 1, 1)
|
| 66 |
+
rm = self.running_mean.reshape(1, -1, 1, 1)
|
| 67 |
+
eps = 1e-5
|
| 68 |
+
scale = w * (rv + eps).rsqrt()
|
| 69 |
+
bias = b - rm * scale
|
| 70 |
+
return x * scale + bias
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
class BackboneBase(nn.Module):
|
| 74 |
+
def __init__(
|
| 75 |
+
self,
|
| 76 |
+
backbone: nn.Module,
|
| 77 |
+
train_backbone: bool,
|
| 78 |
+
num_channels: int,
|
| 79 |
+
return_interm_indices: list,
|
| 80 |
+
):
|
| 81 |
+
super().__init__()
|
| 82 |
+
for name, parameter in backbone.named_parameters():
|
| 83 |
+
if (
|
| 84 |
+
not train_backbone
|
| 85 |
+
or "layer2" not in name
|
| 86 |
+
and "layer3" not in name
|
| 87 |
+
and "layer4" not in name
|
| 88 |
+
):
|
| 89 |
+
parameter.requires_grad_(False)
|
| 90 |
+
|
| 91 |
+
return_layers = {}
|
| 92 |
+
for idx, layer_index in enumerate(return_interm_indices):
|
| 93 |
+
return_layers.update(
|
| 94 |
+
{"layer{}".format(5 - len(return_interm_indices) + idx): "{}".format(layer_index)}
|
| 95 |
+
)
|
| 96 |
+
|
| 97 |
+
# if len:
|
| 98 |
+
# if use_stage1_feature:
|
| 99 |
+
# return_layers = {"layer1": "0", "layer2": "1", "layer3": "2", "layer4": "3"}
|
| 100 |
+
# else:
|
| 101 |
+
# return_layers = {"layer2": "0", "layer3": "1", "layer4": "2"}
|
| 102 |
+
# else:
|
| 103 |
+
# return_layers = {'layer4': "0"}
|
| 104 |
+
self.body = IntermediateLayerGetter(backbone, return_layers=return_layers)
|
| 105 |
+
self.num_channels = num_channels
|
| 106 |
+
|
| 107 |
+
def forward(self, tensor_list: NestedTensor):
|
| 108 |
+
xs = self.body(tensor_list.tensors)
|
| 109 |
+
out: Dict[str, NestedTensor] = {}
|
| 110 |
+
for name, x in xs.items():
|
| 111 |
+
m = tensor_list.mask
|
| 112 |
+
assert m is not None
|
| 113 |
+
mask = F.interpolate(m[None].float(), size=x.shape[-2:]).to(torch.bool)[0]
|
| 114 |
+
out[name] = NestedTensor(x, mask)
|
| 115 |
+
# import ipdb; ipdb.set_trace()
|
| 116 |
+
return out
|
| 117 |
+
|
| 118 |
+
|
| 119 |
+
class Backbone(BackboneBase):
|
| 120 |
+
"""ResNet backbone with frozen BatchNorm."""
|
| 121 |
+
|
| 122 |
+
def __init__(
|
| 123 |
+
self,
|
| 124 |
+
name: str,
|
| 125 |
+
train_backbone: bool,
|
| 126 |
+
dilation: bool,
|
| 127 |
+
return_interm_indices: list,
|
| 128 |
+
batch_norm=FrozenBatchNorm2d,
|
| 129 |
+
):
|
| 130 |
+
if name in ["resnet18", "resnet34", "resnet50", "resnet101"]:
|
| 131 |
+
backbone = getattr(torchvision.models, name)(
|
| 132 |
+
replace_stride_with_dilation=[False, False, dilation],
|
| 133 |
+
pretrained=is_main_process(),
|
| 134 |
+
norm_layer=batch_norm,
|
| 135 |
+
)
|
| 136 |
+
else:
|
| 137 |
+
raise NotImplementedError("Why you can get here with name {}".format(name))
|
| 138 |
+
# num_channels = 512 if name in ('resnet18', 'resnet34') else 2048
|
| 139 |
+
assert name not in ("resnet18", "resnet34"), "Only resnet50 and resnet101 are available."
|
| 140 |
+
assert return_interm_indices in [[0, 1, 2, 3], [1, 2, 3], [3]]
|
| 141 |
+
num_channels_all = [256, 512, 1024, 2048]
|
| 142 |
+
num_channels = num_channels_all[4 - len(return_interm_indices) :]
|
| 143 |
+
super().__init__(backbone, train_backbone, num_channels, return_interm_indices)
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
class Joiner(nn.Sequential):
|
| 147 |
+
def __init__(self, backbone, position_embedding):
|
| 148 |
+
super().__init__(backbone, position_embedding)
|
| 149 |
+
|
| 150 |
+
def forward(self, tensor_list: NestedTensor):
|
| 151 |
+
xs = self[0](tensor_list)
|
| 152 |
+
out: List[NestedTensor] = []
|
| 153 |
+
pos = []
|
| 154 |
+
for name, x in xs.items():
|
| 155 |
+
out.append(x)
|
| 156 |
+
# position encoding
|
| 157 |
+
pos.append(self[1](x).to(x.tensors.dtype))
|
| 158 |
+
|
| 159 |
+
return out, pos
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
def build_backbone(args):
|
| 163 |
+
"""
|
| 164 |
+
Useful args:
|
| 165 |
+
- backbone: backbone name
|
| 166 |
+
- lr_backbone:
|
| 167 |
+
- dilation
|
| 168 |
+
- return_interm_indices: available: [0,1,2,3], [1,2,3], [3]
|
| 169 |
+
- backbone_freeze_keywords:
|
| 170 |
+
- use_checkpoint: for swin only for now
|
| 171 |
+
|
| 172 |
+
"""
|
| 173 |
+
position_embedding = build_position_encoding(args)
|
| 174 |
+
train_backbone = True
|
| 175 |
+
if not train_backbone:
|
| 176 |
+
raise ValueError("Please set lr_backbone > 0")
|
| 177 |
+
return_interm_indices = args.return_interm_indices
|
| 178 |
+
assert return_interm_indices in [[0, 1, 2, 3], [1, 2, 3], [3]]
|
| 179 |
+
args.backbone_freeze_keywords
|
| 180 |
+
use_checkpoint = getattr(args, "use_checkpoint", False)
|
| 181 |
+
|
| 182 |
+
if args.backbone in ["resnet50", "resnet101"]:
|
| 183 |
+
backbone = Backbone(
|
| 184 |
+
args.backbone,
|
| 185 |
+
train_backbone,
|
| 186 |
+
args.dilation,
|
| 187 |
+
return_interm_indices,
|
| 188 |
+
batch_norm=FrozenBatchNorm2d,
|
| 189 |
+
)
|
| 190 |
+
bb_num_channels = backbone.num_channels
|
| 191 |
+
elif args.backbone in [
|
| 192 |
+
"swin_T_224_1k",
|
| 193 |
+
"swin_B_224_22k",
|
| 194 |
+
"swin_B_384_22k",
|
| 195 |
+
"swin_L_224_22k",
|
| 196 |
+
"swin_L_384_22k",
|
| 197 |
+
]:
|
| 198 |
+
pretrain_img_size = int(args.backbone.split("_")[-2])
|
| 199 |
+
backbone = build_swin_transformer(
|
| 200 |
+
args.backbone,
|
| 201 |
+
pretrain_img_size=pretrain_img_size,
|
| 202 |
+
out_indices=tuple(return_interm_indices),
|
| 203 |
+
dilation=False,
|
| 204 |
+
use_checkpoint=use_checkpoint,
|
| 205 |
+
)
|
| 206 |
+
|
| 207 |
+
bb_num_channels = backbone.num_features[4 - len(return_interm_indices) :]
|
| 208 |
+
else:
|
| 209 |
+
raise NotImplementedError("Unknown backbone {}".format(args.backbone))
|
| 210 |
+
|
| 211 |
+
assert len(bb_num_channels) == len(
|
| 212 |
+
return_interm_indices
|
| 213 |
+
), f"len(bb_num_channels) {len(bb_num_channels)} != len(return_interm_indices) {len(return_interm_indices)}"
|
| 214 |
+
|
| 215 |
+
model = Joiner(backbone, position_embedding)
|
| 216 |
+
model.num_channels = bb_num_channels
|
| 217 |
+
assert isinstance(
|
| 218 |
+
bb_num_channels, List
|
| 219 |
+
), "bb_num_channels is expected to be a List but {}".format(type(bb_num_channels))
|
| 220 |
+
# import ipdb; ipdb.set_trace()
|
| 221 |
+
return model
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/backbone/position_encoding.py
ADDED
|
@@ -0,0 +1,186 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# ------------------------------------------------------------------------
|
| 2 |
+
# Grounding DINO
|
| 3 |
+
# url: https://github.com/IDEA-Research/GroundingDINO
|
| 4 |
+
# Copyright (c) 2023 IDEA. All Rights Reserved.
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
# ------------------------------------------------------------------------
|
| 7 |
+
# DINO
|
| 8 |
+
# Copyright (c) 2022 IDEA. All Rights Reserved.
|
| 9 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 10 |
+
# ------------------------------------------------------------------------
|
| 11 |
+
# Conditional DETR
|
| 12 |
+
# Copyright (c) 2021 Microsoft. All Rights Reserved.
|
| 13 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 14 |
+
# ------------------------------------------------------------------------
|
| 15 |
+
# Copied from DETR (https://github.com/facebookresearch/detr)
|
| 16 |
+
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved.
|
| 17 |
+
# ------------------------------------------------------------------------
|
| 18 |
+
|
| 19 |
+
"""
|
| 20 |
+
Various positional encodings for the transformer.
|
| 21 |
+
"""
|
| 22 |
+
import math
|
| 23 |
+
|
| 24 |
+
import torch
|
| 25 |
+
from torch import nn
|
| 26 |
+
|
| 27 |
+
from groundingdino.util.misc import NestedTensor
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
class PositionEmbeddingSine(nn.Module):
|
| 31 |
+
"""
|
| 32 |
+
This is a more standard version of the position embedding, very similar to the one
|
| 33 |
+
used by the Attention is all you need paper, generalized to work on images.
|
| 34 |
+
"""
|
| 35 |
+
|
| 36 |
+
def __init__(self, num_pos_feats=64, temperature=10000, normalize=False, scale=None):
|
| 37 |
+
super().__init__()
|
| 38 |
+
self.num_pos_feats = num_pos_feats
|
| 39 |
+
self.temperature = temperature
|
| 40 |
+
self.normalize = normalize
|
| 41 |
+
if scale is not None and normalize is False:
|
| 42 |
+
raise ValueError("normalize should be True if scale is passed")
|
| 43 |
+
if scale is None:
|
| 44 |
+
scale = 2 * math.pi
|
| 45 |
+
self.scale = scale
|
| 46 |
+
|
| 47 |
+
def forward(self, tensor_list: NestedTensor):
|
| 48 |
+
x = tensor_list.tensors
|
| 49 |
+
mask = tensor_list.mask
|
| 50 |
+
assert mask is not None
|
| 51 |
+
not_mask = ~mask
|
| 52 |
+
y_embed = not_mask.cumsum(1, dtype=torch.float32)
|
| 53 |
+
x_embed = not_mask.cumsum(2, dtype=torch.float32)
|
| 54 |
+
if self.normalize:
|
| 55 |
+
eps = 1e-6
|
| 56 |
+
# if os.environ.get("SHILONG_AMP", None) == '1':
|
| 57 |
+
# eps = 1e-4
|
| 58 |
+
# else:
|
| 59 |
+
# eps = 1e-6
|
| 60 |
+
y_embed = y_embed / (y_embed[:, -1:, :] + eps) * self.scale
|
| 61 |
+
x_embed = x_embed / (x_embed[:, :, -1:] + eps) * self.scale
|
| 62 |
+
|
| 63 |
+
dim_t = torch.arange(self.num_pos_feats, dtype=torch.float32, device=x.device)
|
| 64 |
+
dim_t = self.temperature ** (2 * (dim_t // 2) / self.num_pos_feats)
|
| 65 |
+
|
| 66 |
+
pos_x = x_embed[:, :, :, None] / dim_t
|
| 67 |
+
pos_y = y_embed[:, :, :, None] / dim_t
|
| 68 |
+
pos_x = torch.stack(
|
| 69 |
+
(pos_x[:, :, :, 0::2].sin(), pos_x[:, :, :, 1::2].cos()), dim=4
|
| 70 |
+
).flatten(3)
|
| 71 |
+
pos_y = torch.stack(
|
| 72 |
+
(pos_y[:, :, :, 0::2].sin(), pos_y[:, :, :, 1::2].cos()), dim=4
|
| 73 |
+
).flatten(3)
|
| 74 |
+
pos = torch.cat((pos_y, pos_x), dim=3).permute(0, 3, 1, 2)
|
| 75 |
+
return pos
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
class PositionEmbeddingSineHW(nn.Module):
|
| 79 |
+
"""
|
| 80 |
+
This is a more standard version of the position embedding, very similar to the one
|
| 81 |
+
used by the Attention is all you need paper, generalized to work on images.
|
| 82 |
+
"""
|
| 83 |
+
|
| 84 |
+
def __init__(
|
| 85 |
+
self, num_pos_feats=64, temperatureH=10000, temperatureW=10000, normalize=False, scale=None
|
| 86 |
+
):
|
| 87 |
+
super().__init__()
|
| 88 |
+
self.num_pos_feats = num_pos_feats
|
| 89 |
+
self.temperatureH = temperatureH
|
| 90 |
+
self.temperatureW = temperatureW
|
| 91 |
+
self.normalize = normalize
|
| 92 |
+
if scale is not None and normalize is False:
|
| 93 |
+
raise ValueError("normalize should be True if scale is passed")
|
| 94 |
+
if scale is None:
|
| 95 |
+
scale = 2 * math.pi
|
| 96 |
+
self.scale = scale
|
| 97 |
+
|
| 98 |
+
def forward(self, tensor_list: NestedTensor):
|
| 99 |
+
x = tensor_list.tensors
|
| 100 |
+
mask = tensor_list.mask
|
| 101 |
+
assert mask is not None
|
| 102 |
+
not_mask = ~mask
|
| 103 |
+
y_embed = not_mask.cumsum(1, dtype=torch.float32)
|
| 104 |
+
x_embed = not_mask.cumsum(2, dtype=torch.float32)
|
| 105 |
+
|
| 106 |
+
# import ipdb; ipdb.set_trace()
|
| 107 |
+
|
| 108 |
+
if self.normalize:
|
| 109 |
+
eps = 1e-6
|
| 110 |
+
y_embed = y_embed / (y_embed[:, -1:, :] + eps) * self.scale
|
| 111 |
+
x_embed = x_embed / (x_embed[:, :, -1:] + eps) * self.scale
|
| 112 |
+
|
| 113 |
+
dim_tx = torch.arange(self.num_pos_feats, dtype=torch.float32, device=x.device)
|
| 114 |
+
dim_tx = self.temperatureW ** (2 * (torch.div(dim_tx, 2, rounding_mode='floor')) / self.num_pos_feats)
|
| 115 |
+
pos_x = x_embed[:, :, :, None] / dim_tx
|
| 116 |
+
|
| 117 |
+
dim_ty = torch.arange(self.num_pos_feats, dtype=torch.float32, device=x.device)
|
| 118 |
+
dim_ty = self.temperatureH ** (2 * (torch.div(dim_ty, 2, rounding_mode='floor')) / self.num_pos_feats)
|
| 119 |
+
pos_y = y_embed[:, :, :, None] / dim_ty
|
| 120 |
+
|
| 121 |
+
pos_x = torch.stack(
|
| 122 |
+
(pos_x[:, :, :, 0::2].sin(), pos_x[:, :, :, 1::2].cos()), dim=4
|
| 123 |
+
).flatten(3)
|
| 124 |
+
pos_y = torch.stack(
|
| 125 |
+
(pos_y[:, :, :, 0::2].sin(), pos_y[:, :, :, 1::2].cos()), dim=4
|
| 126 |
+
).flatten(3)
|
| 127 |
+
pos = torch.cat((pos_y, pos_x), dim=3).permute(0, 3, 1, 2)
|
| 128 |
+
|
| 129 |
+
# import ipdb; ipdb.set_trace()
|
| 130 |
+
|
| 131 |
+
return pos
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
class PositionEmbeddingLearned(nn.Module):
|
| 135 |
+
"""
|
| 136 |
+
Absolute pos embedding, learned.
|
| 137 |
+
"""
|
| 138 |
+
|
| 139 |
+
def __init__(self, num_pos_feats=256):
|
| 140 |
+
super().__init__()
|
| 141 |
+
self.row_embed = nn.Embedding(50, num_pos_feats)
|
| 142 |
+
self.col_embed = nn.Embedding(50, num_pos_feats)
|
| 143 |
+
self.reset_parameters()
|
| 144 |
+
|
| 145 |
+
def reset_parameters(self):
|
| 146 |
+
nn.init.uniform_(self.row_embed.weight)
|
| 147 |
+
nn.init.uniform_(self.col_embed.weight)
|
| 148 |
+
|
| 149 |
+
def forward(self, tensor_list: NestedTensor):
|
| 150 |
+
x = tensor_list.tensors
|
| 151 |
+
h, w = x.shape[-2:]
|
| 152 |
+
i = torch.arange(w, device=x.device)
|
| 153 |
+
j = torch.arange(h, device=x.device)
|
| 154 |
+
x_emb = self.col_embed(i)
|
| 155 |
+
y_emb = self.row_embed(j)
|
| 156 |
+
pos = (
|
| 157 |
+
torch.cat(
|
| 158 |
+
[
|
| 159 |
+
x_emb.unsqueeze(0).repeat(h, 1, 1),
|
| 160 |
+
y_emb.unsqueeze(1).repeat(1, w, 1),
|
| 161 |
+
],
|
| 162 |
+
dim=-1,
|
| 163 |
+
)
|
| 164 |
+
.permute(2, 0, 1)
|
| 165 |
+
.unsqueeze(0)
|
| 166 |
+
.repeat(x.shape[0], 1, 1, 1)
|
| 167 |
+
)
|
| 168 |
+
return pos
|
| 169 |
+
|
| 170 |
+
|
| 171 |
+
def build_position_encoding(args):
|
| 172 |
+
N_steps = args.hidden_dim // 2
|
| 173 |
+
if args.position_embedding in ("v2", "sine"):
|
| 174 |
+
# TODO find a better way of exposing other arguments
|
| 175 |
+
position_embedding = PositionEmbeddingSineHW(
|
| 176 |
+
N_steps,
|
| 177 |
+
temperatureH=args.pe_temperatureH,
|
| 178 |
+
temperatureW=args.pe_temperatureW,
|
| 179 |
+
normalize=True,
|
| 180 |
+
)
|
| 181 |
+
elif args.position_embedding in ("v3", "learned"):
|
| 182 |
+
position_embedding = PositionEmbeddingLearned(N_steps)
|
| 183 |
+
else:
|
| 184 |
+
raise ValueError(f"not supported {args.position_embedding}")
|
| 185 |
+
|
| 186 |
+
return position_embedding
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/backbone/swin_transformer.py
ADDED
|
@@ -0,0 +1,802 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# ------------------------------------------------------------------------
|
| 2 |
+
# Grounding DINO
|
| 3 |
+
# url: https://github.com/IDEA-Research/GroundingDINO
|
| 4 |
+
# Copyright (c) 2023 IDEA. All Rights Reserved.
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
# ------------------------------------------------------------------------
|
| 7 |
+
# DINO
|
| 8 |
+
# Copyright (c) 2022 IDEA. All Rights Reserved.
|
| 9 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 10 |
+
# --------------------------------------------------------
|
| 11 |
+
# modified from https://github.com/SwinTransformer/Swin-Transformer-Object-Detection/blob/master/mmdet/models/backbones/swin_transformer.py
|
| 12 |
+
# --------------------------------------------------------
|
| 13 |
+
|
| 14 |
+
import numpy as np
|
| 15 |
+
import torch
|
| 16 |
+
import torch.nn as nn
|
| 17 |
+
import torch.nn.functional as F
|
| 18 |
+
import torch.utils.checkpoint as checkpoint
|
| 19 |
+
from timm.models.layers import DropPath, to_2tuple, trunc_normal_
|
| 20 |
+
|
| 21 |
+
from groundingdino.util.misc import NestedTensor
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class Mlp(nn.Module):
|
| 25 |
+
"""Multilayer perceptron."""
|
| 26 |
+
|
| 27 |
+
def __init__(
|
| 28 |
+
self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.0
|
| 29 |
+
):
|
| 30 |
+
super().__init__()
|
| 31 |
+
out_features = out_features or in_features
|
| 32 |
+
hidden_features = hidden_features or in_features
|
| 33 |
+
self.fc1 = nn.Linear(in_features, hidden_features)
|
| 34 |
+
self.act = act_layer()
|
| 35 |
+
self.fc2 = nn.Linear(hidden_features, out_features)
|
| 36 |
+
self.drop = nn.Dropout(drop)
|
| 37 |
+
|
| 38 |
+
def forward(self, x):
|
| 39 |
+
x = self.fc1(x)
|
| 40 |
+
x = self.act(x)
|
| 41 |
+
x = self.drop(x)
|
| 42 |
+
x = self.fc2(x)
|
| 43 |
+
x = self.drop(x)
|
| 44 |
+
return x
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def window_partition(x, window_size):
|
| 48 |
+
"""
|
| 49 |
+
Args:
|
| 50 |
+
x: (B, H, W, C)
|
| 51 |
+
window_size (int): window size
|
| 52 |
+
Returns:
|
| 53 |
+
windows: (num_windows*B, window_size, window_size, C)
|
| 54 |
+
"""
|
| 55 |
+
B, H, W, C = x.shape
|
| 56 |
+
x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)
|
| 57 |
+
windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
|
| 58 |
+
return windows
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def window_reverse(windows, window_size, H, W):
|
| 62 |
+
"""
|
| 63 |
+
Args:
|
| 64 |
+
windows: (num_windows*B, window_size, window_size, C)
|
| 65 |
+
window_size (int): Window size
|
| 66 |
+
H (int): Height of image
|
| 67 |
+
W (int): Width of image
|
| 68 |
+
Returns:
|
| 69 |
+
x: (B, H, W, C)
|
| 70 |
+
"""
|
| 71 |
+
B = int(windows.shape[0] / (H * W / window_size / window_size))
|
| 72 |
+
x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)
|
| 73 |
+
x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)
|
| 74 |
+
return x
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
class WindowAttention(nn.Module):
|
| 78 |
+
"""Window based multi-head self attention (W-MSA) module with relative position bias.
|
| 79 |
+
It supports both of shifted and non-shifted window.
|
| 80 |
+
Args:
|
| 81 |
+
dim (int): Number of input channels.
|
| 82 |
+
window_size (tuple[int]): The height and width of the window.
|
| 83 |
+
num_heads (int): Number of attention heads.
|
| 84 |
+
qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
|
| 85 |
+
qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set
|
| 86 |
+
attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0
|
| 87 |
+
proj_drop (float, optional): Dropout ratio of output. Default: 0.0
|
| 88 |
+
"""
|
| 89 |
+
|
| 90 |
+
def __init__(
|
| 91 |
+
self,
|
| 92 |
+
dim,
|
| 93 |
+
window_size,
|
| 94 |
+
num_heads,
|
| 95 |
+
qkv_bias=True,
|
| 96 |
+
qk_scale=None,
|
| 97 |
+
attn_drop=0.0,
|
| 98 |
+
proj_drop=0.0,
|
| 99 |
+
):
|
| 100 |
+
|
| 101 |
+
super().__init__()
|
| 102 |
+
self.dim = dim
|
| 103 |
+
self.window_size = window_size # Wh, Ww
|
| 104 |
+
self.num_heads = num_heads
|
| 105 |
+
head_dim = dim // num_heads
|
| 106 |
+
self.scale = qk_scale or head_dim**-0.5
|
| 107 |
+
|
| 108 |
+
# define a parameter table of relative position bias
|
| 109 |
+
self.relative_position_bias_table = nn.Parameter(
|
| 110 |
+
torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)
|
| 111 |
+
) # 2*Wh-1 * 2*Ww-1, nH
|
| 112 |
+
|
| 113 |
+
# get pair-wise relative position index for each token inside the window
|
| 114 |
+
coords_h = torch.arange(self.window_size[0])
|
| 115 |
+
coords_w = torch.arange(self.window_size[1])
|
| 116 |
+
coords = torch.stack(torch.meshgrid([coords_h, coords_w])) # 2, Wh, Ww
|
| 117 |
+
coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww
|
| 118 |
+
relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww
|
| 119 |
+
relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2
|
| 120 |
+
relative_coords[:, :, 0] += self.window_size[0] - 1 # shift to start from 0
|
| 121 |
+
relative_coords[:, :, 1] += self.window_size[1] - 1
|
| 122 |
+
relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1
|
| 123 |
+
relative_position_index = relative_coords.sum(-1) # Wh*Ww, Wh*Ww
|
| 124 |
+
self.register_buffer("relative_position_index", relative_position_index)
|
| 125 |
+
|
| 126 |
+
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
| 127 |
+
self.attn_drop = nn.Dropout(attn_drop)
|
| 128 |
+
self.proj = nn.Linear(dim, dim)
|
| 129 |
+
self.proj_drop = nn.Dropout(proj_drop)
|
| 130 |
+
|
| 131 |
+
trunc_normal_(self.relative_position_bias_table, std=0.02)
|
| 132 |
+
self.softmax = nn.Softmax(dim=-1)
|
| 133 |
+
|
| 134 |
+
def forward(self, x, mask=None):
|
| 135 |
+
"""Forward function.
|
| 136 |
+
Args:
|
| 137 |
+
x: input features with shape of (num_windows*B, N, C)
|
| 138 |
+
mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None
|
| 139 |
+
"""
|
| 140 |
+
B_, N, C = x.shape
|
| 141 |
+
qkv = (
|
| 142 |
+
self.qkv(x)
|
| 143 |
+
.reshape(B_, N, 3, self.num_heads, C // self.num_heads)
|
| 144 |
+
.permute(2, 0, 3, 1, 4)
|
| 145 |
+
)
|
| 146 |
+
q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
|
| 147 |
+
|
| 148 |
+
q = q * self.scale
|
| 149 |
+
attn = q @ k.transpose(-2, -1)
|
| 150 |
+
|
| 151 |
+
relative_position_bias = self.relative_position_bias_table[
|
| 152 |
+
self.relative_position_index.view(-1)
|
| 153 |
+
].view(
|
| 154 |
+
self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1
|
| 155 |
+
) # Wh*Ww,Wh*Ww,nH
|
| 156 |
+
relative_position_bias = relative_position_bias.permute(
|
| 157 |
+
2, 0, 1
|
| 158 |
+
).contiguous() # nH, Wh*Ww, Wh*Ww
|
| 159 |
+
attn = attn + relative_position_bias.unsqueeze(0)
|
| 160 |
+
|
| 161 |
+
if mask is not None:
|
| 162 |
+
nW = mask.shape[0]
|
| 163 |
+
attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)
|
| 164 |
+
attn = attn.view(-1, self.num_heads, N, N)
|
| 165 |
+
attn = self.softmax(attn)
|
| 166 |
+
else:
|
| 167 |
+
attn = self.softmax(attn)
|
| 168 |
+
|
| 169 |
+
attn = self.attn_drop(attn)
|
| 170 |
+
|
| 171 |
+
x = (attn @ v).transpose(1, 2).reshape(B_, N, C)
|
| 172 |
+
x = self.proj(x)
|
| 173 |
+
x = self.proj_drop(x)
|
| 174 |
+
return x
|
| 175 |
+
|
| 176 |
+
|
| 177 |
+
class SwinTransformerBlock(nn.Module):
|
| 178 |
+
"""Swin Transformer Block.
|
| 179 |
+
Args:
|
| 180 |
+
dim (int): Number of input channels.
|
| 181 |
+
num_heads (int): Number of attention heads.
|
| 182 |
+
window_size (int): Window size.
|
| 183 |
+
shift_size (int): Shift size for SW-MSA.
|
| 184 |
+
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
|
| 185 |
+
qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
|
| 186 |
+
qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
|
| 187 |
+
drop (float, optional): Dropout rate. Default: 0.0
|
| 188 |
+
attn_drop (float, optional): Attention dropout rate. Default: 0.0
|
| 189 |
+
drop_path (float, optional): Stochastic depth rate. Default: 0.0
|
| 190 |
+
act_layer (nn.Module, optional): Activation layer. Default: nn.GELU
|
| 191 |
+
norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
|
| 192 |
+
"""
|
| 193 |
+
|
| 194 |
+
def __init__(
|
| 195 |
+
self,
|
| 196 |
+
dim,
|
| 197 |
+
num_heads,
|
| 198 |
+
window_size=7,
|
| 199 |
+
shift_size=0,
|
| 200 |
+
mlp_ratio=4.0,
|
| 201 |
+
qkv_bias=True,
|
| 202 |
+
qk_scale=None,
|
| 203 |
+
drop=0.0,
|
| 204 |
+
attn_drop=0.0,
|
| 205 |
+
drop_path=0.0,
|
| 206 |
+
act_layer=nn.GELU,
|
| 207 |
+
norm_layer=nn.LayerNorm,
|
| 208 |
+
):
|
| 209 |
+
super().__init__()
|
| 210 |
+
self.dim = dim
|
| 211 |
+
self.num_heads = num_heads
|
| 212 |
+
self.window_size = window_size
|
| 213 |
+
self.shift_size = shift_size
|
| 214 |
+
self.mlp_ratio = mlp_ratio
|
| 215 |
+
assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size"
|
| 216 |
+
|
| 217 |
+
self.norm1 = norm_layer(dim)
|
| 218 |
+
self.attn = WindowAttention(
|
| 219 |
+
dim,
|
| 220 |
+
window_size=to_2tuple(self.window_size),
|
| 221 |
+
num_heads=num_heads,
|
| 222 |
+
qkv_bias=qkv_bias,
|
| 223 |
+
qk_scale=qk_scale,
|
| 224 |
+
attn_drop=attn_drop,
|
| 225 |
+
proj_drop=drop,
|
| 226 |
+
)
|
| 227 |
+
|
| 228 |
+
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
|
| 229 |
+
self.norm2 = norm_layer(dim)
|
| 230 |
+
mlp_hidden_dim = int(dim * mlp_ratio)
|
| 231 |
+
self.mlp = Mlp(
|
| 232 |
+
in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop
|
| 233 |
+
)
|
| 234 |
+
|
| 235 |
+
self.H = None
|
| 236 |
+
self.W = None
|
| 237 |
+
|
| 238 |
+
def forward(self, x, mask_matrix):
|
| 239 |
+
"""Forward function.
|
| 240 |
+
Args:
|
| 241 |
+
x: Input feature, tensor size (B, H*W, C).
|
| 242 |
+
H, W: Spatial resolution of the input feature.
|
| 243 |
+
mask_matrix: Attention mask for cyclic shift.
|
| 244 |
+
"""
|
| 245 |
+
B, L, C = x.shape
|
| 246 |
+
H, W = self.H, self.W
|
| 247 |
+
assert L == H * W, "input feature has wrong size"
|
| 248 |
+
|
| 249 |
+
shortcut = x
|
| 250 |
+
x = self.norm1(x)
|
| 251 |
+
x = x.view(B, H, W, C)
|
| 252 |
+
|
| 253 |
+
# pad feature maps to multiples of window size
|
| 254 |
+
pad_l = pad_t = 0
|
| 255 |
+
pad_r = (self.window_size - W % self.window_size) % self.window_size
|
| 256 |
+
pad_b = (self.window_size - H % self.window_size) % self.window_size
|
| 257 |
+
x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b))
|
| 258 |
+
_, Hp, Wp, _ = x.shape
|
| 259 |
+
|
| 260 |
+
# cyclic shift
|
| 261 |
+
if self.shift_size > 0:
|
| 262 |
+
shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))
|
| 263 |
+
attn_mask = mask_matrix
|
| 264 |
+
else:
|
| 265 |
+
shifted_x = x
|
| 266 |
+
attn_mask = None
|
| 267 |
+
|
| 268 |
+
# partition windows
|
| 269 |
+
x_windows = window_partition(
|
| 270 |
+
shifted_x, self.window_size
|
| 271 |
+
) # nW*B, window_size, window_size, C
|
| 272 |
+
x_windows = x_windows.view(
|
| 273 |
+
-1, self.window_size * self.window_size, C
|
| 274 |
+
) # nW*B, window_size*window_size, C
|
| 275 |
+
|
| 276 |
+
# W-MSA/SW-MSA
|
| 277 |
+
attn_windows = self.attn(x_windows, mask=attn_mask) # nW*B, window_size*window_size, C
|
| 278 |
+
|
| 279 |
+
# merge windows
|
| 280 |
+
attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)
|
| 281 |
+
shifted_x = window_reverse(attn_windows, self.window_size, Hp, Wp) # B H' W' C
|
| 282 |
+
|
| 283 |
+
# reverse cyclic shift
|
| 284 |
+
if self.shift_size > 0:
|
| 285 |
+
x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))
|
| 286 |
+
else:
|
| 287 |
+
x = shifted_x
|
| 288 |
+
|
| 289 |
+
if pad_r > 0 or pad_b > 0:
|
| 290 |
+
x = x[:, :H, :W, :].contiguous()
|
| 291 |
+
|
| 292 |
+
x = x.view(B, H * W, C)
|
| 293 |
+
|
| 294 |
+
# FFN
|
| 295 |
+
x = shortcut + self.drop_path(x)
|
| 296 |
+
x = x + self.drop_path(self.mlp(self.norm2(x)))
|
| 297 |
+
|
| 298 |
+
return x
|
| 299 |
+
|
| 300 |
+
|
| 301 |
+
class PatchMerging(nn.Module):
|
| 302 |
+
"""Patch Merging Layer
|
| 303 |
+
Args:
|
| 304 |
+
dim (int): Number of input channels.
|
| 305 |
+
norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
|
| 306 |
+
"""
|
| 307 |
+
|
| 308 |
+
def __init__(self, dim, norm_layer=nn.LayerNorm):
|
| 309 |
+
super().__init__()
|
| 310 |
+
self.dim = dim
|
| 311 |
+
self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)
|
| 312 |
+
self.norm = norm_layer(4 * dim)
|
| 313 |
+
|
| 314 |
+
def forward(self, x, H, W):
|
| 315 |
+
"""Forward function.
|
| 316 |
+
Args:
|
| 317 |
+
x: Input feature, tensor size (B, H*W, C).
|
| 318 |
+
H, W: Spatial resolution of the input feature.
|
| 319 |
+
"""
|
| 320 |
+
B, L, C = x.shape
|
| 321 |
+
assert L == H * W, "input feature has wrong size"
|
| 322 |
+
|
| 323 |
+
x = x.view(B, H, W, C)
|
| 324 |
+
|
| 325 |
+
# padding
|
| 326 |
+
pad_input = (H % 2 == 1) or (W % 2 == 1)
|
| 327 |
+
if pad_input:
|
| 328 |
+
x = F.pad(x, (0, 0, 0, W % 2, 0, H % 2))
|
| 329 |
+
|
| 330 |
+
x0 = x[:, 0::2, 0::2, :] # B H/2 W/2 C
|
| 331 |
+
x1 = x[:, 1::2, 0::2, :] # B H/2 W/2 C
|
| 332 |
+
x2 = x[:, 0::2, 1::2, :] # B H/2 W/2 C
|
| 333 |
+
x3 = x[:, 1::2, 1::2, :] # B H/2 W/2 C
|
| 334 |
+
x = torch.cat([x0, x1, x2, x3], -1) # B H/2 W/2 4*C
|
| 335 |
+
x = x.view(B, -1, 4 * C) # B H/2*W/2 4*C
|
| 336 |
+
|
| 337 |
+
x = self.norm(x)
|
| 338 |
+
x = self.reduction(x)
|
| 339 |
+
|
| 340 |
+
return x
|
| 341 |
+
|
| 342 |
+
|
| 343 |
+
class BasicLayer(nn.Module):
|
| 344 |
+
"""A basic Swin Transformer layer for one stage.
|
| 345 |
+
Args:
|
| 346 |
+
dim (int): Number of feature channels
|
| 347 |
+
depth (int): Depths of this stage.
|
| 348 |
+
num_heads (int): Number of attention head.
|
| 349 |
+
window_size (int): Local window size. Default: 7.
|
| 350 |
+
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
|
| 351 |
+
qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
|
| 352 |
+
qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
|
| 353 |
+
drop (float, optional): Dropout rate. Default: 0.0
|
| 354 |
+
attn_drop (float, optional): Attention dropout rate. Default: 0.0
|
| 355 |
+
drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0
|
| 356 |
+
norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
|
| 357 |
+
downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None
|
| 358 |
+
use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
|
| 359 |
+
"""
|
| 360 |
+
|
| 361 |
+
def __init__(
|
| 362 |
+
self,
|
| 363 |
+
dim,
|
| 364 |
+
depth,
|
| 365 |
+
num_heads,
|
| 366 |
+
window_size=7,
|
| 367 |
+
mlp_ratio=4.0,
|
| 368 |
+
qkv_bias=True,
|
| 369 |
+
qk_scale=None,
|
| 370 |
+
drop=0.0,
|
| 371 |
+
attn_drop=0.0,
|
| 372 |
+
drop_path=0.0,
|
| 373 |
+
norm_layer=nn.LayerNorm,
|
| 374 |
+
downsample=None,
|
| 375 |
+
use_checkpoint=False,
|
| 376 |
+
):
|
| 377 |
+
super().__init__()
|
| 378 |
+
self.window_size = window_size
|
| 379 |
+
self.shift_size = window_size // 2
|
| 380 |
+
self.depth = depth
|
| 381 |
+
self.use_checkpoint = use_checkpoint
|
| 382 |
+
|
| 383 |
+
# build blocks
|
| 384 |
+
self.blocks = nn.ModuleList(
|
| 385 |
+
[
|
| 386 |
+
SwinTransformerBlock(
|
| 387 |
+
dim=dim,
|
| 388 |
+
num_heads=num_heads,
|
| 389 |
+
window_size=window_size,
|
| 390 |
+
shift_size=0 if (i % 2 == 0) else window_size // 2,
|
| 391 |
+
mlp_ratio=mlp_ratio,
|
| 392 |
+
qkv_bias=qkv_bias,
|
| 393 |
+
qk_scale=qk_scale,
|
| 394 |
+
drop=drop,
|
| 395 |
+
attn_drop=attn_drop,
|
| 396 |
+
drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,
|
| 397 |
+
norm_layer=norm_layer,
|
| 398 |
+
)
|
| 399 |
+
for i in range(depth)
|
| 400 |
+
]
|
| 401 |
+
)
|
| 402 |
+
|
| 403 |
+
# patch merging layer
|
| 404 |
+
if downsample is not None:
|
| 405 |
+
self.downsample = downsample(dim=dim, norm_layer=norm_layer)
|
| 406 |
+
else:
|
| 407 |
+
self.downsample = None
|
| 408 |
+
|
| 409 |
+
def forward(self, x, H, W):
|
| 410 |
+
"""Forward function.
|
| 411 |
+
Args:
|
| 412 |
+
x: Input feature, tensor size (B, H*W, C).
|
| 413 |
+
H, W: Spatial resolution of the input feature.
|
| 414 |
+
"""
|
| 415 |
+
|
| 416 |
+
# calculate attention mask for SW-MSA
|
| 417 |
+
Hp = int(np.ceil(H / self.window_size)) * self.window_size
|
| 418 |
+
Wp = int(np.ceil(W / self.window_size)) * self.window_size
|
| 419 |
+
img_mask = torch.zeros((1, Hp, Wp, 1), device=x.device, dtype=x.dtype) # 1 Hp Wp 1
|
| 420 |
+
h_slices = (
|
| 421 |
+
slice(0, -self.window_size),
|
| 422 |
+
slice(-self.window_size, -self.shift_size),
|
| 423 |
+
slice(-self.shift_size, None),
|
| 424 |
+
)
|
| 425 |
+
w_slices = (
|
| 426 |
+
slice(0, -self.window_size),
|
| 427 |
+
slice(-self.window_size, -self.shift_size),
|
| 428 |
+
slice(-self.shift_size, None),
|
| 429 |
+
)
|
| 430 |
+
cnt = 0
|
| 431 |
+
for h in h_slices:
|
| 432 |
+
for w in w_slices:
|
| 433 |
+
img_mask[:, h, w, :] = cnt
|
| 434 |
+
cnt += 1
|
| 435 |
+
|
| 436 |
+
mask_windows = window_partition(
|
| 437 |
+
img_mask, self.window_size
|
| 438 |
+
) # nW, window_size, window_size, 1
|
| 439 |
+
mask_windows = mask_windows.view(-1, self.window_size * self.window_size)
|
| 440 |
+
attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
|
| 441 |
+
attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(
|
| 442 |
+
attn_mask == 0, float(0.0)
|
| 443 |
+
)
|
| 444 |
+
|
| 445 |
+
for blk in self.blocks:
|
| 446 |
+
blk.H, blk.W = H, W
|
| 447 |
+
if self.use_checkpoint:
|
| 448 |
+
x = checkpoint.checkpoint(blk, x, attn_mask)
|
| 449 |
+
else:
|
| 450 |
+
x = blk(x, attn_mask)
|
| 451 |
+
if self.downsample is not None:
|
| 452 |
+
x_down = self.downsample(x, H, W)
|
| 453 |
+
Wh, Ww = (H + 1) // 2, (W + 1) // 2
|
| 454 |
+
return x, H, W, x_down, Wh, Ww
|
| 455 |
+
else:
|
| 456 |
+
return x, H, W, x, H, W
|
| 457 |
+
|
| 458 |
+
|
| 459 |
+
class PatchEmbed(nn.Module):
|
| 460 |
+
"""Image to Patch Embedding
|
| 461 |
+
Args:
|
| 462 |
+
patch_size (int): Patch token size. Default: 4.
|
| 463 |
+
in_chans (int): Number of input image channels. Default: 3.
|
| 464 |
+
embed_dim (int): Number of linear projection output channels. Default: 96.
|
| 465 |
+
norm_layer (nn.Module, optional): Normalization layer. Default: None
|
| 466 |
+
"""
|
| 467 |
+
|
| 468 |
+
def __init__(self, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None):
|
| 469 |
+
super().__init__()
|
| 470 |
+
patch_size = to_2tuple(patch_size)
|
| 471 |
+
self.patch_size = patch_size
|
| 472 |
+
|
| 473 |
+
self.in_chans = in_chans
|
| 474 |
+
self.embed_dim = embed_dim
|
| 475 |
+
|
| 476 |
+
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
|
| 477 |
+
if norm_layer is not None:
|
| 478 |
+
self.norm = norm_layer(embed_dim)
|
| 479 |
+
else:
|
| 480 |
+
self.norm = None
|
| 481 |
+
|
| 482 |
+
def forward(self, x):
|
| 483 |
+
"""Forward function."""
|
| 484 |
+
# padding
|
| 485 |
+
_, _, H, W = x.size()
|
| 486 |
+
if W % self.patch_size[1] != 0:
|
| 487 |
+
x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1]))
|
| 488 |
+
if H % self.patch_size[0] != 0:
|
| 489 |
+
x = F.pad(x, (0, 0, 0, self.patch_size[0] - H % self.patch_size[0]))
|
| 490 |
+
|
| 491 |
+
x = self.proj(x) # B C Wh Ww
|
| 492 |
+
if self.norm is not None:
|
| 493 |
+
Wh, Ww = x.size(2), x.size(3)
|
| 494 |
+
x = x.flatten(2).transpose(1, 2)
|
| 495 |
+
x = self.norm(x)
|
| 496 |
+
x = x.transpose(1, 2).view(-1, self.embed_dim, Wh, Ww)
|
| 497 |
+
|
| 498 |
+
return x
|
| 499 |
+
|
| 500 |
+
|
| 501 |
+
class SwinTransformer(nn.Module):
|
| 502 |
+
"""Swin Transformer backbone.
|
| 503 |
+
A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows` -
|
| 504 |
+
https://arxiv.org/pdf/2103.14030
|
| 505 |
+
Args:
|
| 506 |
+
pretrain_img_size (int): Input image size for training the pretrained model,
|
| 507 |
+
used in absolute postion embedding. Default 224.
|
| 508 |
+
patch_size (int | tuple(int)): Patch size. Default: 4.
|
| 509 |
+
in_chans (int): Number of input image channels. Default: 3.
|
| 510 |
+
embed_dim (int): Number of linear projection output channels. Default: 96.
|
| 511 |
+
depths (tuple[int]): Depths of each Swin Transformer stage.
|
| 512 |
+
num_heads (tuple[int]): Number of attention head of each stage.
|
| 513 |
+
window_size (int): Window size. Default: 7.
|
| 514 |
+
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
|
| 515 |
+
qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True
|
| 516 |
+
qk_scale (float): Override default qk scale of head_dim ** -0.5 if set.
|
| 517 |
+
drop_rate (float): Dropout rate.
|
| 518 |
+
attn_drop_rate (float): Attention dropout rate. Default: 0.
|
| 519 |
+
drop_path_rate (float): Stochastic depth rate. Default: 0.2.
|
| 520 |
+
norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.
|
| 521 |
+
ape (bool): If True, add absolute position embedding to the patch embedding. Default: False.
|
| 522 |
+
patch_norm (bool): If True, add normalization after patch embedding. Default: True.
|
| 523 |
+
out_indices (Sequence[int]): Output from which stages.
|
| 524 |
+
frozen_stages (int): Stages to be frozen (stop grad and set eval mode).
|
| 525 |
+
-1 means not freezing any parameters.
|
| 526 |
+
use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
|
| 527 |
+
dilation (bool): if True, the output size if 16x downsample, ow 32x downsample.
|
| 528 |
+
"""
|
| 529 |
+
|
| 530 |
+
def __init__(
|
| 531 |
+
self,
|
| 532 |
+
pretrain_img_size=224,
|
| 533 |
+
patch_size=4,
|
| 534 |
+
in_chans=3,
|
| 535 |
+
embed_dim=96,
|
| 536 |
+
depths=[2, 2, 6, 2],
|
| 537 |
+
num_heads=[3, 6, 12, 24],
|
| 538 |
+
window_size=7,
|
| 539 |
+
mlp_ratio=4.0,
|
| 540 |
+
qkv_bias=True,
|
| 541 |
+
qk_scale=None,
|
| 542 |
+
drop_rate=0.0,
|
| 543 |
+
attn_drop_rate=0.0,
|
| 544 |
+
drop_path_rate=0.2,
|
| 545 |
+
norm_layer=nn.LayerNorm,
|
| 546 |
+
ape=False,
|
| 547 |
+
patch_norm=True,
|
| 548 |
+
out_indices=(0, 1, 2, 3),
|
| 549 |
+
frozen_stages=-1,
|
| 550 |
+
dilation=False,
|
| 551 |
+
use_checkpoint=False,
|
| 552 |
+
):
|
| 553 |
+
super().__init__()
|
| 554 |
+
|
| 555 |
+
self.pretrain_img_size = pretrain_img_size
|
| 556 |
+
self.num_layers = len(depths)
|
| 557 |
+
self.embed_dim = embed_dim
|
| 558 |
+
self.ape = ape
|
| 559 |
+
self.patch_norm = patch_norm
|
| 560 |
+
self.out_indices = out_indices
|
| 561 |
+
self.frozen_stages = frozen_stages
|
| 562 |
+
self.dilation = dilation
|
| 563 |
+
|
| 564 |
+
# if use_checkpoint:
|
| 565 |
+
# print("use_checkpoint!!!!!!!!!!!!!!!!!!!!!!!!")
|
| 566 |
+
|
| 567 |
+
# split image into non-overlapping patches
|
| 568 |
+
self.patch_embed = PatchEmbed(
|
| 569 |
+
patch_size=patch_size,
|
| 570 |
+
in_chans=in_chans,
|
| 571 |
+
embed_dim=embed_dim,
|
| 572 |
+
norm_layer=norm_layer if self.patch_norm else None,
|
| 573 |
+
)
|
| 574 |
+
|
| 575 |
+
# absolute position embedding
|
| 576 |
+
if self.ape:
|
| 577 |
+
pretrain_img_size = to_2tuple(pretrain_img_size)
|
| 578 |
+
patch_size = to_2tuple(patch_size)
|
| 579 |
+
patches_resolution = [
|
| 580 |
+
pretrain_img_size[0] // patch_size[0],
|
| 581 |
+
pretrain_img_size[1] // patch_size[1],
|
| 582 |
+
]
|
| 583 |
+
|
| 584 |
+
self.absolute_pos_embed = nn.Parameter(
|
| 585 |
+
torch.zeros(1, embed_dim, patches_resolution[0], patches_resolution[1])
|
| 586 |
+
)
|
| 587 |
+
trunc_normal_(self.absolute_pos_embed, std=0.02)
|
| 588 |
+
|
| 589 |
+
self.pos_drop = nn.Dropout(p=drop_rate)
|
| 590 |
+
|
| 591 |
+
# stochastic depth
|
| 592 |
+
dpr = [
|
| 593 |
+
x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))
|
| 594 |
+
] # stochastic depth decay rule
|
| 595 |
+
|
| 596 |
+
# build layers
|
| 597 |
+
self.layers = nn.ModuleList()
|
| 598 |
+
# prepare downsample list
|
| 599 |
+
downsamplelist = [PatchMerging for i in range(self.num_layers)]
|
| 600 |
+
downsamplelist[-1] = None
|
| 601 |
+
num_features = [int(embed_dim * 2**i) for i in range(self.num_layers)]
|
| 602 |
+
if self.dilation:
|
| 603 |
+
downsamplelist[-2] = None
|
| 604 |
+
num_features[-1] = int(embed_dim * 2 ** (self.num_layers - 1)) // 2
|
| 605 |
+
for i_layer in range(self.num_layers):
|
| 606 |
+
layer = BasicLayer(
|
| 607 |
+
# dim=int(embed_dim * 2 ** i_layer),
|
| 608 |
+
dim=num_features[i_layer],
|
| 609 |
+
depth=depths[i_layer],
|
| 610 |
+
num_heads=num_heads[i_layer],
|
| 611 |
+
window_size=window_size,
|
| 612 |
+
mlp_ratio=mlp_ratio,
|
| 613 |
+
qkv_bias=qkv_bias,
|
| 614 |
+
qk_scale=qk_scale,
|
| 615 |
+
drop=drop_rate,
|
| 616 |
+
attn_drop=attn_drop_rate,
|
| 617 |
+
drop_path=dpr[sum(depths[:i_layer]) : sum(depths[: i_layer + 1])],
|
| 618 |
+
norm_layer=norm_layer,
|
| 619 |
+
# downsample=PatchMerging if (i_layer < self.num_layers - 1) else None,
|
| 620 |
+
downsample=downsamplelist[i_layer],
|
| 621 |
+
use_checkpoint=use_checkpoint,
|
| 622 |
+
)
|
| 623 |
+
self.layers.append(layer)
|
| 624 |
+
|
| 625 |
+
# num_features = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]
|
| 626 |
+
self.num_features = num_features
|
| 627 |
+
|
| 628 |
+
# add a norm layer for each output
|
| 629 |
+
for i_layer in out_indices:
|
| 630 |
+
layer = norm_layer(num_features[i_layer])
|
| 631 |
+
layer_name = f"norm{i_layer}"
|
| 632 |
+
self.add_module(layer_name, layer)
|
| 633 |
+
|
| 634 |
+
self._freeze_stages()
|
| 635 |
+
|
| 636 |
+
def _freeze_stages(self):
|
| 637 |
+
if self.frozen_stages >= 0:
|
| 638 |
+
self.patch_embed.eval()
|
| 639 |
+
for param in self.patch_embed.parameters():
|
| 640 |
+
param.requires_grad = False
|
| 641 |
+
|
| 642 |
+
if self.frozen_stages >= 1 and self.ape:
|
| 643 |
+
self.absolute_pos_embed.requires_grad = False
|
| 644 |
+
|
| 645 |
+
if self.frozen_stages >= 2:
|
| 646 |
+
self.pos_drop.eval()
|
| 647 |
+
for i in range(0, self.frozen_stages - 1):
|
| 648 |
+
m = self.layers[i]
|
| 649 |
+
m.eval()
|
| 650 |
+
for param in m.parameters():
|
| 651 |
+
param.requires_grad = False
|
| 652 |
+
|
| 653 |
+
# def init_weights(self, pretrained=None):
|
| 654 |
+
# """Initialize the weights in backbone.
|
| 655 |
+
# Args:
|
| 656 |
+
# pretrained (str, optional): Path to pre-trained weights.
|
| 657 |
+
# Defaults to None.
|
| 658 |
+
# """
|
| 659 |
+
|
| 660 |
+
# def _init_weights(m):
|
| 661 |
+
# if isinstance(m, nn.Linear):
|
| 662 |
+
# trunc_normal_(m.weight, std=.02)
|
| 663 |
+
# if isinstance(m, nn.Linear) and m.bias is not None:
|
| 664 |
+
# nn.init.constant_(m.bias, 0)
|
| 665 |
+
# elif isinstance(m, nn.LayerNorm):
|
| 666 |
+
# nn.init.constant_(m.bias, 0)
|
| 667 |
+
# nn.init.constant_(m.weight, 1.0)
|
| 668 |
+
|
| 669 |
+
# if isinstance(pretrained, str):
|
| 670 |
+
# self.apply(_init_weights)
|
| 671 |
+
# logger = get_root_logger()
|
| 672 |
+
# load_checkpoint(self, pretrained, strict=False, logger=logger)
|
| 673 |
+
# elif pretrained is None:
|
| 674 |
+
# self.apply(_init_weights)
|
| 675 |
+
# else:
|
| 676 |
+
# raise TypeError('pretrained must be a str or None')
|
| 677 |
+
|
| 678 |
+
def forward_raw(self, x):
|
| 679 |
+
"""Forward function."""
|
| 680 |
+
x = self.patch_embed(x)
|
| 681 |
+
|
| 682 |
+
Wh, Ww = x.size(2), x.size(3)
|
| 683 |
+
if self.ape:
|
| 684 |
+
# interpolate the position embedding to the corresponding size
|
| 685 |
+
absolute_pos_embed = F.interpolate(
|
| 686 |
+
self.absolute_pos_embed, size=(Wh, Ww), mode="bicubic"
|
| 687 |
+
)
|
| 688 |
+
x = (x + absolute_pos_embed).flatten(2).transpose(1, 2) # B Wh*Ww C
|
| 689 |
+
else:
|
| 690 |
+
x = x.flatten(2).transpose(1, 2)
|
| 691 |
+
x = self.pos_drop(x)
|
| 692 |
+
|
| 693 |
+
outs = []
|
| 694 |
+
for i in range(self.num_layers):
|
| 695 |
+
layer = self.layers[i]
|
| 696 |
+
x_out, H, W, x, Wh, Ww = layer(x, Wh, Ww)
|
| 697 |
+
# import ipdb; ipdb.set_trace()
|
| 698 |
+
|
| 699 |
+
if i in self.out_indices:
|
| 700 |
+
norm_layer = getattr(self, f"norm{i}")
|
| 701 |
+
x_out = norm_layer(x_out)
|
| 702 |
+
|
| 703 |
+
out = x_out.view(-1, H, W, self.num_features[i]).permute(0, 3, 1, 2).contiguous()
|
| 704 |
+
outs.append(out)
|
| 705 |
+
# in:
|
| 706 |
+
# torch.Size([2, 3, 1024, 1024])
|
| 707 |
+
# outs:
|
| 708 |
+
# [torch.Size([2, 192, 256, 256]), torch.Size([2, 384, 128, 128]), \
|
| 709 |
+
# torch.Size([2, 768, 64, 64]), torch.Size([2, 1536, 32, 32])]
|
| 710 |
+
return tuple(outs)
|
| 711 |
+
|
| 712 |
+
def forward(self, tensor_list: NestedTensor):
|
| 713 |
+
x = tensor_list.tensors
|
| 714 |
+
|
| 715 |
+
"""Forward function."""
|
| 716 |
+
x = self.patch_embed(x)
|
| 717 |
+
|
| 718 |
+
Wh, Ww = x.size(2), x.size(3)
|
| 719 |
+
if self.ape:
|
| 720 |
+
# interpolate the position embedding to the corresponding size
|
| 721 |
+
absolute_pos_embed = F.interpolate(
|
| 722 |
+
self.absolute_pos_embed, size=(Wh, Ww), mode="bicubic"
|
| 723 |
+
)
|
| 724 |
+
x = (x + absolute_pos_embed).flatten(2).transpose(1, 2) # B Wh*Ww C
|
| 725 |
+
else:
|
| 726 |
+
x = x.flatten(2).transpose(1, 2)
|
| 727 |
+
x = self.pos_drop(x)
|
| 728 |
+
|
| 729 |
+
outs = []
|
| 730 |
+
for i in range(self.num_layers):
|
| 731 |
+
layer = self.layers[i]
|
| 732 |
+
x_out, H, W, x, Wh, Ww = layer(x, Wh, Ww)
|
| 733 |
+
|
| 734 |
+
if i in self.out_indices:
|
| 735 |
+
norm_layer = getattr(self, f"norm{i}")
|
| 736 |
+
x_out = norm_layer(x_out)
|
| 737 |
+
|
| 738 |
+
out = x_out.view(-1, H, W, self.num_features[i]).permute(0, 3, 1, 2).contiguous()
|
| 739 |
+
outs.append(out)
|
| 740 |
+
# in:
|
| 741 |
+
# torch.Size([2, 3, 1024, 1024])
|
| 742 |
+
# out:
|
| 743 |
+
# [torch.Size([2, 192, 256, 256]), torch.Size([2, 384, 128, 128]), \
|
| 744 |
+
# torch.Size([2, 768, 64, 64]), torch.Size([2, 1536, 32, 32])]
|
| 745 |
+
|
| 746 |
+
# collect for nesttensors
|
| 747 |
+
outs_dict = {}
|
| 748 |
+
for idx, out_i in enumerate(outs):
|
| 749 |
+
m = tensor_list.mask
|
| 750 |
+
assert m is not None
|
| 751 |
+
mask = F.interpolate(m[None].float(), size=out_i.shape[-2:]).to(torch.bool)[0]
|
| 752 |
+
outs_dict[idx] = NestedTensor(out_i, mask)
|
| 753 |
+
|
| 754 |
+
return outs_dict
|
| 755 |
+
|
| 756 |
+
def train(self, mode=True):
|
| 757 |
+
"""Convert the model into training mode while keep layers freezed."""
|
| 758 |
+
super(SwinTransformer, self).train(mode)
|
| 759 |
+
self._freeze_stages()
|
| 760 |
+
|
| 761 |
+
|
| 762 |
+
def build_swin_transformer(modelname, pretrain_img_size, **kw):
|
| 763 |
+
assert modelname in [
|
| 764 |
+
"swin_T_224_1k",
|
| 765 |
+
"swin_B_224_22k",
|
| 766 |
+
"swin_B_384_22k",
|
| 767 |
+
"swin_L_224_22k",
|
| 768 |
+
"swin_L_384_22k",
|
| 769 |
+
]
|
| 770 |
+
|
| 771 |
+
model_para_dict = {
|
| 772 |
+
"swin_T_224_1k": dict(
|
| 773 |
+
embed_dim=96, depths=[2, 2, 6, 2], num_heads=[3, 6, 12, 24], window_size=7
|
| 774 |
+
),
|
| 775 |
+
"swin_B_224_22k": dict(
|
| 776 |
+
embed_dim=128, depths=[2, 2, 18, 2], num_heads=[4, 8, 16, 32], window_size=7
|
| 777 |
+
),
|
| 778 |
+
"swin_B_384_22k": dict(
|
| 779 |
+
embed_dim=128, depths=[2, 2, 18, 2], num_heads=[4, 8, 16, 32], window_size=12
|
| 780 |
+
),
|
| 781 |
+
"swin_L_224_22k": dict(
|
| 782 |
+
embed_dim=192, depths=[2, 2, 18, 2], num_heads=[6, 12, 24, 48], window_size=7
|
| 783 |
+
),
|
| 784 |
+
"swin_L_384_22k": dict(
|
| 785 |
+
embed_dim=192, depths=[2, 2, 18, 2], num_heads=[6, 12, 24, 48], window_size=12
|
| 786 |
+
),
|
| 787 |
+
}
|
| 788 |
+
kw_cgf = model_para_dict[modelname]
|
| 789 |
+
kw_cgf.update(kw)
|
| 790 |
+
model = SwinTransformer(pretrain_img_size=pretrain_img_size, **kw_cgf)
|
| 791 |
+
return model
|
| 792 |
+
|
| 793 |
+
|
| 794 |
+
if __name__ == "__main__":
|
| 795 |
+
model = build_swin_transformer("swin_L_384_22k", 384, dilation=True)
|
| 796 |
+
x = torch.rand(2, 3, 1024, 1024)
|
| 797 |
+
y = model.forward_raw(x)
|
| 798 |
+
import ipdb
|
| 799 |
+
|
| 800 |
+
ipdb.set_trace()
|
| 801 |
+
x = torch.rand(2, 3, 384, 384)
|
| 802 |
+
y = model.forward_raw(x)
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/bertwarper.py
ADDED
|
@@ -0,0 +1,273 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# ------------------------------------------------------------------------
|
| 2 |
+
# Grounding DINO
|
| 3 |
+
# url: https://github.com/IDEA-Research/GroundingDINO
|
| 4 |
+
# Copyright (c) 2023 IDEA. All Rights Reserved.
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
# ------------------------------------------------------------------------
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
import torch.nn.functional as F
|
| 10 |
+
import torch.utils.checkpoint as checkpoint
|
| 11 |
+
from torch import Tensor, nn
|
| 12 |
+
from torchvision.ops.boxes import nms
|
| 13 |
+
from transformers import BertConfig, BertModel, BertPreTrainedModel
|
| 14 |
+
from transformers.modeling_outputs import BaseModelOutputWithPoolingAndCrossAttentions
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class BertModelWarper(nn.Module):
|
| 18 |
+
def __init__(self, bert_model):
|
| 19 |
+
super().__init__()
|
| 20 |
+
# self.bert = bert_modelc
|
| 21 |
+
|
| 22 |
+
self.config = bert_model.config
|
| 23 |
+
self.embeddings = bert_model.embeddings
|
| 24 |
+
self.encoder = bert_model.encoder
|
| 25 |
+
self.pooler = bert_model.pooler
|
| 26 |
+
|
| 27 |
+
self.get_extended_attention_mask = bert_model.get_extended_attention_mask
|
| 28 |
+
self.invert_attention_mask = bert_model.invert_attention_mask
|
| 29 |
+
self.get_head_mask = bert_model.get_head_mask
|
| 30 |
+
|
| 31 |
+
def forward(
|
| 32 |
+
self,
|
| 33 |
+
input_ids=None,
|
| 34 |
+
attention_mask=None,
|
| 35 |
+
token_type_ids=None,
|
| 36 |
+
position_ids=None,
|
| 37 |
+
head_mask=None,
|
| 38 |
+
inputs_embeds=None,
|
| 39 |
+
encoder_hidden_states=None,
|
| 40 |
+
encoder_attention_mask=None,
|
| 41 |
+
past_key_values=None,
|
| 42 |
+
use_cache=None,
|
| 43 |
+
output_attentions=None,
|
| 44 |
+
output_hidden_states=None,
|
| 45 |
+
return_dict=None,
|
| 46 |
+
):
|
| 47 |
+
r"""
|
| 48 |
+
encoder_hidden_states (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, hidden_size)`, `optional`):
|
| 49 |
+
Sequence of hidden-states at the output of the last layer of the encoder. Used in the cross-attention if
|
| 50 |
+
the model is configured as a decoder.
|
| 51 |
+
encoder_attention_mask (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
|
| 52 |
+
Mask to avoid performing attention on the padding token indices of the encoder input. This mask is used in
|
| 53 |
+
the cross-attention if the model is configured as a decoder. Mask values selected in ``[0, 1]``:
|
| 54 |
+
|
| 55 |
+
- 1 for tokens that are **not masked**,
|
| 56 |
+
- 0 for tokens that are **masked**.
|
| 57 |
+
past_key_values (:obj:`tuple(tuple(torch.FloatTensor))` of length :obj:`config.n_layers` with each tuple having 4 tensors of shape :obj:`(batch_size, num_heads, sequence_length - 1, embed_size_per_head)`):
|
| 58 |
+
Contains precomputed key and value hidden states of the attention blocks. Can be used to speed up decoding.
|
| 59 |
+
|
| 60 |
+
If :obj:`past_key_values` are used, the user can optionally input only the last :obj:`decoder_input_ids`
|
| 61 |
+
(those that don't have their past key value states given to this model) of shape :obj:`(batch_size, 1)`
|
| 62 |
+
instead of all :obj:`decoder_input_ids` of shape :obj:`(batch_size, sequence_length)`.
|
| 63 |
+
use_cache (:obj:`bool`, `optional`):
|
| 64 |
+
If set to :obj:`True`, :obj:`past_key_values` key value states are returned and can be used to speed up
|
| 65 |
+
decoding (see :obj:`past_key_values`).
|
| 66 |
+
"""
|
| 67 |
+
output_attentions = (
|
| 68 |
+
output_attentions if output_attentions is not None else self.config.output_attentions
|
| 69 |
+
)
|
| 70 |
+
output_hidden_states = (
|
| 71 |
+
output_hidden_states
|
| 72 |
+
if output_hidden_states is not None
|
| 73 |
+
else self.config.output_hidden_states
|
| 74 |
+
)
|
| 75 |
+
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
| 76 |
+
|
| 77 |
+
if self.config.is_decoder:
|
| 78 |
+
use_cache = use_cache if use_cache is not None else self.config.use_cache
|
| 79 |
+
else:
|
| 80 |
+
use_cache = False
|
| 81 |
+
|
| 82 |
+
if input_ids is not None and inputs_embeds is not None:
|
| 83 |
+
raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
|
| 84 |
+
elif input_ids is not None:
|
| 85 |
+
input_shape = input_ids.size()
|
| 86 |
+
batch_size, seq_length = input_shape
|
| 87 |
+
elif inputs_embeds is not None:
|
| 88 |
+
input_shape = inputs_embeds.size()[:-1]
|
| 89 |
+
batch_size, seq_length = input_shape
|
| 90 |
+
else:
|
| 91 |
+
raise ValueError("You have to specify either input_ids or inputs_embeds")
|
| 92 |
+
|
| 93 |
+
device = input_ids.device if input_ids is not None else inputs_embeds.device
|
| 94 |
+
|
| 95 |
+
# past_key_values_length
|
| 96 |
+
past_key_values_length = (
|
| 97 |
+
past_key_values[0][0].shape[2] if past_key_values is not None else 0
|
| 98 |
+
)
|
| 99 |
+
|
| 100 |
+
if attention_mask is None:
|
| 101 |
+
attention_mask = torch.ones(
|
| 102 |
+
((batch_size, seq_length + past_key_values_length)), device=device
|
| 103 |
+
)
|
| 104 |
+
if token_type_ids is None:
|
| 105 |
+
token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=device)
|
| 106 |
+
|
| 107 |
+
# We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]
|
| 108 |
+
# ourselves in which case we just need to make it broadcastable to all heads.
|
| 109 |
+
extended_attention_mask: torch.Tensor = self.get_extended_attention_mask(
|
| 110 |
+
attention_mask, input_shape, device
|
| 111 |
+
)
|
| 112 |
+
|
| 113 |
+
# If a 2D or 3D attention mask is provided for the cross-attention
|
| 114 |
+
# we need to make broadcastable to [batch_size, num_heads, seq_length, seq_length]
|
| 115 |
+
if self.config.is_decoder and encoder_hidden_states is not None:
|
| 116 |
+
encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states.size()
|
| 117 |
+
encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)
|
| 118 |
+
if encoder_attention_mask is None:
|
| 119 |
+
encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)
|
| 120 |
+
encoder_extended_attention_mask = self.invert_attention_mask(encoder_attention_mask)
|
| 121 |
+
else:
|
| 122 |
+
encoder_extended_attention_mask = None
|
| 123 |
+
# if os.environ.get('IPDB_SHILONG_DEBUG', None) == 'INFO':
|
| 124 |
+
# import ipdb; ipdb.set_trace()
|
| 125 |
+
|
| 126 |
+
# Prepare head mask if needed
|
| 127 |
+
# 1.0 in head_mask indicate we keep the head
|
| 128 |
+
# attention_probs has shape bsz x n_heads x N x N
|
| 129 |
+
# input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]
|
| 130 |
+
# and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]
|
| 131 |
+
head_mask = self.get_head_mask(head_mask, self.config.num_hidden_layers)
|
| 132 |
+
|
| 133 |
+
embedding_output = self.embeddings(
|
| 134 |
+
input_ids=input_ids,
|
| 135 |
+
position_ids=position_ids,
|
| 136 |
+
token_type_ids=token_type_ids,
|
| 137 |
+
inputs_embeds=inputs_embeds,
|
| 138 |
+
past_key_values_length=past_key_values_length,
|
| 139 |
+
)
|
| 140 |
+
|
| 141 |
+
encoder_outputs = self.encoder(
|
| 142 |
+
embedding_output,
|
| 143 |
+
attention_mask=extended_attention_mask,
|
| 144 |
+
head_mask=head_mask,
|
| 145 |
+
encoder_hidden_states=encoder_hidden_states,
|
| 146 |
+
encoder_attention_mask=encoder_extended_attention_mask,
|
| 147 |
+
past_key_values=past_key_values,
|
| 148 |
+
use_cache=use_cache,
|
| 149 |
+
output_attentions=output_attentions,
|
| 150 |
+
output_hidden_states=output_hidden_states,
|
| 151 |
+
return_dict=return_dict,
|
| 152 |
+
)
|
| 153 |
+
sequence_output = encoder_outputs[0]
|
| 154 |
+
pooled_output = self.pooler(sequence_output) if self.pooler is not None else None
|
| 155 |
+
|
| 156 |
+
if not return_dict:
|
| 157 |
+
return (sequence_output, pooled_output) + encoder_outputs[1:]
|
| 158 |
+
|
| 159 |
+
return BaseModelOutputWithPoolingAndCrossAttentions(
|
| 160 |
+
last_hidden_state=sequence_output,
|
| 161 |
+
pooler_output=pooled_output,
|
| 162 |
+
past_key_values=encoder_outputs.past_key_values,
|
| 163 |
+
hidden_states=encoder_outputs.hidden_states,
|
| 164 |
+
attentions=encoder_outputs.attentions,
|
| 165 |
+
cross_attentions=encoder_outputs.cross_attentions,
|
| 166 |
+
)
|
| 167 |
+
|
| 168 |
+
|
| 169 |
+
class TextEncoderShell(nn.Module):
|
| 170 |
+
def __init__(self, text_encoder):
|
| 171 |
+
super().__init__()
|
| 172 |
+
self.text_encoder = text_encoder
|
| 173 |
+
self.config = self.text_encoder.config
|
| 174 |
+
|
| 175 |
+
def forward(self, **kw):
|
| 176 |
+
# feed into text encoder
|
| 177 |
+
return self.text_encoder(**kw)
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
def generate_masks_with_special_tokens(tokenized, special_tokens_list, tokenizer):
|
| 181 |
+
"""Generate attention mask between each pair of special tokens
|
| 182 |
+
Args:
|
| 183 |
+
input_ids (torch.Tensor): input ids. Shape: [bs, num_token]
|
| 184 |
+
special_tokens_mask (list): special tokens mask.
|
| 185 |
+
Returns:
|
| 186 |
+
torch.Tensor: attention mask between each special tokens.
|
| 187 |
+
"""
|
| 188 |
+
input_ids = tokenized["input_ids"]
|
| 189 |
+
bs, num_token = input_ids.shape
|
| 190 |
+
# special_tokens_mask: bs, num_token. 1 for special tokens. 0 for normal tokens
|
| 191 |
+
special_tokens_mask = torch.zeros((bs, num_token), device=input_ids.device).bool()
|
| 192 |
+
for special_token in special_tokens_list:
|
| 193 |
+
special_tokens_mask |= input_ids == special_token
|
| 194 |
+
|
| 195 |
+
# idxs: each row is a list of indices of special tokens
|
| 196 |
+
idxs = torch.nonzero(special_tokens_mask)
|
| 197 |
+
|
| 198 |
+
# generate attention mask and positional ids
|
| 199 |
+
attention_mask = (
|
| 200 |
+
torch.eye(num_token, device=input_ids.device).bool().unsqueeze(0).repeat(bs, 1, 1)
|
| 201 |
+
)
|
| 202 |
+
position_ids = torch.zeros((bs, num_token), device=input_ids.device)
|
| 203 |
+
previous_col = 0
|
| 204 |
+
for i in range(idxs.shape[0]):
|
| 205 |
+
row, col = idxs[i]
|
| 206 |
+
if (col == 0) or (col == num_token - 1):
|
| 207 |
+
attention_mask[row, col, col] = True
|
| 208 |
+
position_ids[row, col] = 0
|
| 209 |
+
else:
|
| 210 |
+
attention_mask[row, previous_col + 1 : col + 1, previous_col + 1 : col + 1] = True
|
| 211 |
+
position_ids[row, previous_col + 1 : col + 1] = torch.arange(
|
| 212 |
+
0, col - previous_col, device=input_ids.device
|
| 213 |
+
)
|
| 214 |
+
|
| 215 |
+
previous_col = col
|
| 216 |
+
|
| 217 |
+
# # padding mask
|
| 218 |
+
# padding_mask = tokenized['attention_mask']
|
| 219 |
+
# attention_mask = attention_mask & padding_mask.unsqueeze(1).bool() & padding_mask.unsqueeze(2).bool()
|
| 220 |
+
|
| 221 |
+
return attention_mask, position_ids.to(torch.long)
|
| 222 |
+
|
| 223 |
+
|
| 224 |
+
def generate_masks_with_special_tokens_and_transfer_map(tokenized, special_tokens_list, tokenizer):
|
| 225 |
+
"""Generate attention mask between each pair of special tokens
|
| 226 |
+
Args:
|
| 227 |
+
input_ids (torch.Tensor): input ids. Shape: [bs, num_token]
|
| 228 |
+
special_tokens_mask (list): special tokens mask.
|
| 229 |
+
Returns:
|
| 230 |
+
torch.Tensor: attention mask between each special tokens.
|
| 231 |
+
"""
|
| 232 |
+
input_ids = tokenized["input_ids"]
|
| 233 |
+
bs, num_token = input_ids.shape
|
| 234 |
+
# special_tokens_mask: bs, num_token. 1 for special tokens. 0 for normal tokens
|
| 235 |
+
special_tokens_mask = torch.zeros((bs, num_token), device=input_ids.device).bool()
|
| 236 |
+
for special_token in special_tokens_list:
|
| 237 |
+
special_tokens_mask |= input_ids == special_token
|
| 238 |
+
|
| 239 |
+
# idxs: each row is a list of indices of special tokens
|
| 240 |
+
idxs = torch.nonzero(special_tokens_mask)
|
| 241 |
+
|
| 242 |
+
# generate attention mask and positional ids
|
| 243 |
+
attention_mask = (
|
| 244 |
+
torch.eye(num_token, device=input_ids.device).bool().unsqueeze(0).repeat(bs, 1, 1)
|
| 245 |
+
)
|
| 246 |
+
position_ids = torch.zeros((bs, num_token), device=input_ids.device)
|
| 247 |
+
cate_to_token_mask_list = [[] for _ in range(bs)]
|
| 248 |
+
previous_col = 0
|
| 249 |
+
for i in range(idxs.shape[0]):
|
| 250 |
+
row, col = idxs[i]
|
| 251 |
+
if (col == 0) or (col == num_token - 1):
|
| 252 |
+
attention_mask[row, col, col] = True
|
| 253 |
+
position_ids[row, col] = 0
|
| 254 |
+
else:
|
| 255 |
+
attention_mask[row, previous_col + 1 : col + 1, previous_col + 1 : col + 1] = True
|
| 256 |
+
position_ids[row, previous_col + 1 : col + 1] = torch.arange(
|
| 257 |
+
0, col - previous_col, device=input_ids.device
|
| 258 |
+
)
|
| 259 |
+
c2t_maski = torch.zeros((num_token), device=input_ids.device).bool()
|
| 260 |
+
c2t_maski[previous_col + 1 : col] = True
|
| 261 |
+
cate_to_token_mask_list[row].append(c2t_maski)
|
| 262 |
+
previous_col = col
|
| 263 |
+
|
| 264 |
+
cate_to_token_mask_list = [
|
| 265 |
+
torch.stack(cate_to_token_mask_listi, dim=0)
|
| 266 |
+
for cate_to_token_mask_listi in cate_to_token_mask_list
|
| 267 |
+
]
|
| 268 |
+
|
| 269 |
+
# # padding mask
|
| 270 |
+
# padding_mask = tokenized['attention_mask']
|
| 271 |
+
# attention_mask = attention_mask & padding_mask.unsqueeze(1).bool() & padding_mask.unsqueeze(2).bool()
|
| 272 |
+
|
| 273 |
+
return attention_mask, position_ids.to(torch.long), cate_to_token_mask_list
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_attn.h
ADDED
|
@@ -0,0 +1,64 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
/*!
|
| 2 |
+
**************************************************************************************************
|
| 3 |
+
* Deformable DETR
|
| 4 |
+
* Copyright (c) 2020 SenseTime. All Rights Reserved.
|
| 5 |
+
* Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
**************************************************************************************************
|
| 7 |
+
* Modified from https://github.com/chengdazhi/Deformable-Convolution-V2-PyTorch/tree/pytorch_1.0.0
|
| 8 |
+
**************************************************************************************************
|
| 9 |
+
*/
|
| 10 |
+
|
| 11 |
+
#pragma once
|
| 12 |
+
|
| 13 |
+
#include "ms_deform_attn_cpu.h"
|
| 14 |
+
|
| 15 |
+
#ifdef WITH_CUDA
|
| 16 |
+
#include "ms_deform_attn_cuda.h"
|
| 17 |
+
#endif
|
| 18 |
+
|
| 19 |
+
namespace groundingdino {
|
| 20 |
+
|
| 21 |
+
at::Tensor
|
| 22 |
+
ms_deform_attn_forward(
|
| 23 |
+
const at::Tensor &value,
|
| 24 |
+
const at::Tensor &spatial_shapes,
|
| 25 |
+
const at::Tensor &level_start_index,
|
| 26 |
+
const at::Tensor &sampling_loc,
|
| 27 |
+
const at::Tensor &attn_weight,
|
| 28 |
+
const int im2col_step)
|
| 29 |
+
{
|
| 30 |
+
if (value.type().is_cuda())
|
| 31 |
+
{
|
| 32 |
+
#ifdef WITH_CUDA
|
| 33 |
+
return ms_deform_attn_cuda_forward(
|
| 34 |
+
value, spatial_shapes, level_start_index, sampling_loc, attn_weight, im2col_step);
|
| 35 |
+
#else
|
| 36 |
+
AT_ERROR("Not compiled with GPU support");
|
| 37 |
+
#endif
|
| 38 |
+
}
|
| 39 |
+
AT_ERROR("Not implemented on the CPU");
|
| 40 |
+
}
|
| 41 |
+
|
| 42 |
+
std::vector<at::Tensor>
|
| 43 |
+
ms_deform_attn_backward(
|
| 44 |
+
const at::Tensor &value,
|
| 45 |
+
const at::Tensor &spatial_shapes,
|
| 46 |
+
const at::Tensor &level_start_index,
|
| 47 |
+
const at::Tensor &sampling_loc,
|
| 48 |
+
const at::Tensor &attn_weight,
|
| 49 |
+
const at::Tensor &grad_output,
|
| 50 |
+
const int im2col_step)
|
| 51 |
+
{
|
| 52 |
+
if (value.type().is_cuda())
|
| 53 |
+
{
|
| 54 |
+
#ifdef WITH_CUDA
|
| 55 |
+
return ms_deform_attn_cuda_backward(
|
| 56 |
+
value, spatial_shapes, level_start_index, sampling_loc, attn_weight, grad_output, im2col_step);
|
| 57 |
+
#else
|
| 58 |
+
AT_ERROR("Not compiled with GPU support");
|
| 59 |
+
#endif
|
| 60 |
+
}
|
| 61 |
+
AT_ERROR("Not implemented on the CPU");
|
| 62 |
+
}
|
| 63 |
+
|
| 64 |
+
} // namespace groundingdino
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_attn_cpu.cpp
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
/*!
|
| 2 |
+
**************************************************************************************************
|
| 3 |
+
* Deformable DETR
|
| 4 |
+
* Copyright (c) 2020 SenseTime. All Rights Reserved.
|
| 5 |
+
* Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
**************************************************************************************************
|
| 7 |
+
* Modified from https://github.com/chengdazhi/Deformable-Convolution-V2-PyTorch/tree/pytorch_1.0.0
|
| 8 |
+
**************************************************************************************************
|
| 9 |
+
*/
|
| 10 |
+
|
| 11 |
+
#include <vector>
|
| 12 |
+
|
| 13 |
+
#include <ATen/ATen.h>
|
| 14 |
+
#include <ATen/cuda/CUDAContext.h>
|
| 15 |
+
|
| 16 |
+
namespace groundingdino {
|
| 17 |
+
|
| 18 |
+
at::Tensor
|
| 19 |
+
ms_deform_attn_cpu_forward(
|
| 20 |
+
const at::Tensor &value,
|
| 21 |
+
const at::Tensor &spatial_shapes,
|
| 22 |
+
const at::Tensor &level_start_index,
|
| 23 |
+
const at::Tensor &sampling_loc,
|
| 24 |
+
const at::Tensor &attn_weight,
|
| 25 |
+
const int im2col_step)
|
| 26 |
+
{
|
| 27 |
+
AT_ERROR("Not implement on cpu");
|
| 28 |
+
}
|
| 29 |
+
|
| 30 |
+
std::vector<at::Tensor>
|
| 31 |
+
ms_deform_attn_cpu_backward(
|
| 32 |
+
const at::Tensor &value,
|
| 33 |
+
const at::Tensor &spatial_shapes,
|
| 34 |
+
const at::Tensor &level_start_index,
|
| 35 |
+
const at::Tensor &sampling_loc,
|
| 36 |
+
const at::Tensor &attn_weight,
|
| 37 |
+
const at::Tensor &grad_output,
|
| 38 |
+
const int im2col_step)
|
| 39 |
+
{
|
| 40 |
+
AT_ERROR("Not implement on cpu");
|
| 41 |
+
}
|
| 42 |
+
|
| 43 |
+
} // namespace groundingdino
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_attn_cpu.h
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
/*!
|
| 2 |
+
**************************************************************************************************
|
| 3 |
+
* Deformable DETR
|
| 4 |
+
* Copyright (c) 2020 SenseTime. All Rights Reserved.
|
| 5 |
+
* Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
**************************************************************************************************
|
| 7 |
+
* Modified from https://github.com/chengdazhi/Deformable-Convolution-V2-PyTorch/tree/pytorch_1.0.0
|
| 8 |
+
**************************************************************************************************
|
| 9 |
+
*/
|
| 10 |
+
|
| 11 |
+
#pragma once
|
| 12 |
+
#include <torch/extension.h>
|
| 13 |
+
|
| 14 |
+
namespace groundingdino {
|
| 15 |
+
|
| 16 |
+
at::Tensor
|
| 17 |
+
ms_deform_attn_cpu_forward(
|
| 18 |
+
const at::Tensor &value,
|
| 19 |
+
const at::Tensor &spatial_shapes,
|
| 20 |
+
const at::Tensor &level_start_index,
|
| 21 |
+
const at::Tensor &sampling_loc,
|
| 22 |
+
const at::Tensor &attn_weight,
|
| 23 |
+
const int im2col_step);
|
| 24 |
+
|
| 25 |
+
std::vector<at::Tensor>
|
| 26 |
+
ms_deform_attn_cpu_backward(
|
| 27 |
+
const at::Tensor &value,
|
| 28 |
+
const at::Tensor &spatial_shapes,
|
| 29 |
+
const at::Tensor &level_start_index,
|
| 30 |
+
const at::Tensor &sampling_loc,
|
| 31 |
+
const at::Tensor &attn_weight,
|
| 32 |
+
const at::Tensor &grad_output,
|
| 33 |
+
const int im2col_step);
|
| 34 |
+
|
| 35 |
+
} // namespace groundingdino
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_attn_cuda.cu
ADDED
|
@@ -0,0 +1,156 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
/*!
|
| 2 |
+
**************************************************************************************************
|
| 3 |
+
* Deformable DETR
|
| 4 |
+
* Copyright (c) 2020 SenseTime. All Rights Reserved.
|
| 5 |
+
* Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
**************************************************************************************************
|
| 7 |
+
* Modified from https://github.com/chengdazhi/Deformable-Convolution-V2-PyTorch/tree/pytorch_1.0.0
|
| 8 |
+
**************************************************************************************************
|
| 9 |
+
*/
|
| 10 |
+
|
| 11 |
+
#include <vector>
|
| 12 |
+
#include "ms_deform_im2col_cuda.cuh"
|
| 13 |
+
|
| 14 |
+
#include <ATen/ATen.h>
|
| 15 |
+
#include <ATen/cuda/CUDAContext.h>
|
| 16 |
+
#include <cuda.h>
|
| 17 |
+
#include <cuda_runtime.h>
|
| 18 |
+
|
| 19 |
+
namespace groundingdino {
|
| 20 |
+
|
| 21 |
+
at::Tensor ms_deform_attn_cuda_forward(
|
| 22 |
+
const at::Tensor &value,
|
| 23 |
+
const at::Tensor &spatial_shapes,
|
| 24 |
+
const at::Tensor &level_start_index,
|
| 25 |
+
const at::Tensor &sampling_loc,
|
| 26 |
+
const at::Tensor &attn_weight,
|
| 27 |
+
const int im2col_step)
|
| 28 |
+
{
|
| 29 |
+
AT_ASSERTM(value.is_contiguous(), "value tensor has to be contiguous");
|
| 30 |
+
AT_ASSERTM(spatial_shapes.is_contiguous(), "spatial_shapes tensor has to be contiguous");
|
| 31 |
+
AT_ASSERTM(level_start_index.is_contiguous(), "level_start_index tensor has to be contiguous");
|
| 32 |
+
AT_ASSERTM(sampling_loc.is_contiguous(), "sampling_loc tensor has to be contiguous");
|
| 33 |
+
AT_ASSERTM(attn_weight.is_contiguous(), "attn_weight tensor has to be contiguous");
|
| 34 |
+
|
| 35 |
+
AT_ASSERTM(value.type().is_cuda(), "value must be a CUDA tensor");
|
| 36 |
+
AT_ASSERTM(spatial_shapes.type().is_cuda(), "spatial_shapes must be a CUDA tensor");
|
| 37 |
+
AT_ASSERTM(level_start_index.type().is_cuda(), "level_start_index must be a CUDA tensor");
|
| 38 |
+
AT_ASSERTM(sampling_loc.type().is_cuda(), "sampling_loc must be a CUDA tensor");
|
| 39 |
+
AT_ASSERTM(attn_weight.type().is_cuda(), "attn_weight must be a CUDA tensor");
|
| 40 |
+
|
| 41 |
+
const int batch = value.size(0);
|
| 42 |
+
const int spatial_size = value.size(1);
|
| 43 |
+
const int num_heads = value.size(2);
|
| 44 |
+
const int channels = value.size(3);
|
| 45 |
+
|
| 46 |
+
const int num_levels = spatial_shapes.size(0);
|
| 47 |
+
|
| 48 |
+
const int num_query = sampling_loc.size(1);
|
| 49 |
+
const int num_point = sampling_loc.size(4);
|
| 50 |
+
|
| 51 |
+
const int im2col_step_ = std::min(batch, im2col_step);
|
| 52 |
+
|
| 53 |
+
AT_ASSERTM(batch % im2col_step_ == 0, "batch(%d) must divide im2col_step(%d)", batch, im2col_step_);
|
| 54 |
+
|
| 55 |
+
auto output = at::zeros({batch, num_query, num_heads, channels}, value.options());
|
| 56 |
+
|
| 57 |
+
const int batch_n = im2col_step_;
|
| 58 |
+
auto output_n = output.view({batch/im2col_step_, batch_n, num_query, num_heads, channels});
|
| 59 |
+
auto per_value_size = spatial_size * num_heads * channels;
|
| 60 |
+
auto per_sample_loc_size = num_query * num_heads * num_levels * num_point * 2;
|
| 61 |
+
auto per_attn_weight_size = num_query * num_heads * num_levels * num_point;
|
| 62 |
+
for (int n = 0; n < batch/im2col_step_; ++n)
|
| 63 |
+
{
|
| 64 |
+
auto columns = output_n.select(0, n);
|
| 65 |
+
AT_DISPATCH_FLOATING_TYPES(value.type(), "ms_deform_attn_forward_cuda", ([&] {
|
| 66 |
+
ms_deformable_im2col_cuda(at::cuda::getCurrentCUDAStream(),
|
| 67 |
+
value.data<scalar_t>() + n * im2col_step_ * per_value_size,
|
| 68 |
+
spatial_shapes.data<int64_t>(),
|
| 69 |
+
level_start_index.data<int64_t>(),
|
| 70 |
+
sampling_loc.data<scalar_t>() + n * im2col_step_ * per_sample_loc_size,
|
| 71 |
+
attn_weight.data<scalar_t>() + n * im2col_step_ * per_attn_weight_size,
|
| 72 |
+
batch_n, spatial_size, num_heads, channels, num_levels, num_query, num_point,
|
| 73 |
+
columns.data<scalar_t>());
|
| 74 |
+
|
| 75 |
+
}));
|
| 76 |
+
}
|
| 77 |
+
|
| 78 |
+
output = output.view({batch, num_query, num_heads*channels});
|
| 79 |
+
|
| 80 |
+
return output;
|
| 81 |
+
}
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
std::vector<at::Tensor> ms_deform_attn_cuda_backward(
|
| 85 |
+
const at::Tensor &value,
|
| 86 |
+
const at::Tensor &spatial_shapes,
|
| 87 |
+
const at::Tensor &level_start_index,
|
| 88 |
+
const at::Tensor &sampling_loc,
|
| 89 |
+
const at::Tensor &attn_weight,
|
| 90 |
+
const at::Tensor &grad_output,
|
| 91 |
+
const int im2col_step)
|
| 92 |
+
{
|
| 93 |
+
|
| 94 |
+
AT_ASSERTM(value.is_contiguous(), "value tensor has to be contiguous");
|
| 95 |
+
AT_ASSERTM(spatial_shapes.is_contiguous(), "spatial_shapes tensor has to be contiguous");
|
| 96 |
+
AT_ASSERTM(level_start_index.is_contiguous(), "level_start_index tensor has to be contiguous");
|
| 97 |
+
AT_ASSERTM(sampling_loc.is_contiguous(), "sampling_loc tensor has to be contiguous");
|
| 98 |
+
AT_ASSERTM(attn_weight.is_contiguous(), "attn_weight tensor has to be contiguous");
|
| 99 |
+
AT_ASSERTM(grad_output.is_contiguous(), "grad_output tensor has to be contiguous");
|
| 100 |
+
|
| 101 |
+
AT_ASSERTM(value.type().is_cuda(), "value must be a CUDA tensor");
|
| 102 |
+
AT_ASSERTM(spatial_shapes.type().is_cuda(), "spatial_shapes must be a CUDA tensor");
|
| 103 |
+
AT_ASSERTM(level_start_index.type().is_cuda(), "level_start_index must be a CUDA tensor");
|
| 104 |
+
AT_ASSERTM(sampling_loc.type().is_cuda(), "sampling_loc must be a CUDA tensor");
|
| 105 |
+
AT_ASSERTM(attn_weight.type().is_cuda(), "attn_weight must be a CUDA tensor");
|
| 106 |
+
AT_ASSERTM(grad_output.type().is_cuda(), "grad_output must be a CUDA tensor");
|
| 107 |
+
|
| 108 |
+
const int batch = value.size(0);
|
| 109 |
+
const int spatial_size = value.size(1);
|
| 110 |
+
const int num_heads = value.size(2);
|
| 111 |
+
const int channels = value.size(3);
|
| 112 |
+
|
| 113 |
+
const int num_levels = spatial_shapes.size(0);
|
| 114 |
+
|
| 115 |
+
const int num_query = sampling_loc.size(1);
|
| 116 |
+
const int num_point = sampling_loc.size(4);
|
| 117 |
+
|
| 118 |
+
const int im2col_step_ = std::min(batch, im2col_step);
|
| 119 |
+
|
| 120 |
+
AT_ASSERTM(batch % im2col_step_ == 0, "batch(%d) must divide im2col_step(%d)", batch, im2col_step_);
|
| 121 |
+
|
| 122 |
+
auto grad_value = at::zeros_like(value);
|
| 123 |
+
auto grad_sampling_loc = at::zeros_like(sampling_loc);
|
| 124 |
+
auto grad_attn_weight = at::zeros_like(attn_weight);
|
| 125 |
+
|
| 126 |
+
const int batch_n = im2col_step_;
|
| 127 |
+
auto per_value_size = spatial_size * num_heads * channels;
|
| 128 |
+
auto per_sample_loc_size = num_query * num_heads * num_levels * num_point * 2;
|
| 129 |
+
auto per_attn_weight_size = num_query * num_heads * num_levels * num_point;
|
| 130 |
+
auto grad_output_n = grad_output.view({batch/im2col_step_, batch_n, num_query, num_heads, channels});
|
| 131 |
+
|
| 132 |
+
for (int n = 0; n < batch/im2col_step_; ++n)
|
| 133 |
+
{
|
| 134 |
+
auto grad_output_g = grad_output_n.select(0, n);
|
| 135 |
+
AT_DISPATCH_FLOATING_TYPES(value.type(), "ms_deform_attn_backward_cuda", ([&] {
|
| 136 |
+
ms_deformable_col2im_cuda(at::cuda::getCurrentCUDAStream(),
|
| 137 |
+
grad_output_g.data<scalar_t>(),
|
| 138 |
+
value.data<scalar_t>() + n * im2col_step_ * per_value_size,
|
| 139 |
+
spatial_shapes.data<int64_t>(),
|
| 140 |
+
level_start_index.data<int64_t>(),
|
| 141 |
+
sampling_loc.data<scalar_t>() + n * im2col_step_ * per_sample_loc_size,
|
| 142 |
+
attn_weight.data<scalar_t>() + n * im2col_step_ * per_attn_weight_size,
|
| 143 |
+
batch_n, spatial_size, num_heads, channels, num_levels, num_query, num_point,
|
| 144 |
+
grad_value.data<scalar_t>() + n * im2col_step_ * per_value_size,
|
| 145 |
+
grad_sampling_loc.data<scalar_t>() + n * im2col_step_ * per_sample_loc_size,
|
| 146 |
+
grad_attn_weight.data<scalar_t>() + n * im2col_step_ * per_attn_weight_size);
|
| 147 |
+
|
| 148 |
+
}));
|
| 149 |
+
}
|
| 150 |
+
|
| 151 |
+
return {
|
| 152 |
+
grad_value, grad_sampling_loc, grad_attn_weight
|
| 153 |
+
};
|
| 154 |
+
}
|
| 155 |
+
|
| 156 |
+
} // namespace groundingdino
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_attn_cuda.h
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
/*!
|
| 2 |
+
**************************************************************************************************
|
| 3 |
+
* Deformable DETR
|
| 4 |
+
* Copyright (c) 2020 SenseTime. All Rights Reserved.
|
| 5 |
+
* Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
**************************************************************************************************
|
| 7 |
+
* Modified from https://github.com/chengdazhi/Deformable-Convolution-V2-PyTorch/tree/pytorch_1.0.0
|
| 8 |
+
**************************************************************************************************
|
| 9 |
+
*/
|
| 10 |
+
|
| 11 |
+
#pragma once
|
| 12 |
+
#include <torch/extension.h>
|
| 13 |
+
|
| 14 |
+
namespace groundingdino {
|
| 15 |
+
|
| 16 |
+
at::Tensor ms_deform_attn_cuda_forward(
|
| 17 |
+
const at::Tensor &value,
|
| 18 |
+
const at::Tensor &spatial_shapes,
|
| 19 |
+
const at::Tensor &level_start_index,
|
| 20 |
+
const at::Tensor &sampling_loc,
|
| 21 |
+
const at::Tensor &attn_weight,
|
| 22 |
+
const int im2col_step);
|
| 23 |
+
|
| 24 |
+
std::vector<at::Tensor> ms_deform_attn_cuda_backward(
|
| 25 |
+
const at::Tensor &value,
|
| 26 |
+
const at::Tensor &spatial_shapes,
|
| 27 |
+
const at::Tensor &level_start_index,
|
| 28 |
+
const at::Tensor &sampling_loc,
|
| 29 |
+
const at::Tensor &attn_weight,
|
| 30 |
+
const at::Tensor &grad_output,
|
| 31 |
+
const int im2col_step);
|
| 32 |
+
|
| 33 |
+
} // namespace groundingdino
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/MsDeformAttn/ms_deform_im2col_cuda.cuh
ADDED
|
@@ -0,0 +1,1327 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
/*!
|
| 2 |
+
**************************************************************************
|
| 3 |
+
* Deformable DETR
|
| 4 |
+
* Copyright (c) 2020 SenseTime. All Rights Reserved.
|
| 5 |
+
* Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
**************************************************************************
|
| 7 |
+
* Modified from DCN (https://github.com/msracver/Deformable-ConvNets)
|
| 8 |
+
* Copyright (c) 2018 Microsoft
|
| 9 |
+
**************************************************************************
|
| 10 |
+
*/
|
| 11 |
+
|
| 12 |
+
#include <cstdio>
|
| 13 |
+
#include <algorithm>
|
| 14 |
+
#include <cstring>
|
| 15 |
+
|
| 16 |
+
#include <ATen/ATen.h>
|
| 17 |
+
#include <ATen/cuda/CUDAContext.h>
|
| 18 |
+
|
| 19 |
+
#include <THC/THCAtomics.cuh>
|
| 20 |
+
|
| 21 |
+
#define CUDA_KERNEL_LOOP(i, n) \
|
| 22 |
+
for (int i = blockIdx.x * blockDim.x + threadIdx.x; \
|
| 23 |
+
i < (n); \
|
| 24 |
+
i += blockDim.x * gridDim.x)
|
| 25 |
+
|
| 26 |
+
const int CUDA_NUM_THREADS = 1024;
|
| 27 |
+
inline int GET_BLOCKS(const int N, const int num_threads)
|
| 28 |
+
{
|
| 29 |
+
return (N + num_threads - 1) / num_threads;
|
| 30 |
+
}
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
template <typename scalar_t>
|
| 34 |
+
__device__ scalar_t ms_deform_attn_im2col_bilinear(const scalar_t* &bottom_data,
|
| 35 |
+
const int &height, const int &width, const int &nheads, const int &channels,
|
| 36 |
+
const scalar_t &h, const scalar_t &w, const int &m, const int &c)
|
| 37 |
+
{
|
| 38 |
+
const int h_low = floor(h);
|
| 39 |
+
const int w_low = floor(w);
|
| 40 |
+
const int h_high = h_low + 1;
|
| 41 |
+
const int w_high = w_low + 1;
|
| 42 |
+
|
| 43 |
+
const scalar_t lh = h - h_low;
|
| 44 |
+
const scalar_t lw = w - w_low;
|
| 45 |
+
const scalar_t hh = 1 - lh, hw = 1 - lw;
|
| 46 |
+
|
| 47 |
+
const int w_stride = nheads * channels;
|
| 48 |
+
const int h_stride = width * w_stride;
|
| 49 |
+
const int h_low_ptr_offset = h_low * h_stride;
|
| 50 |
+
const int h_high_ptr_offset = h_low_ptr_offset + h_stride;
|
| 51 |
+
const int w_low_ptr_offset = w_low * w_stride;
|
| 52 |
+
const int w_high_ptr_offset = w_low_ptr_offset + w_stride;
|
| 53 |
+
const int base_ptr = m * channels + c;
|
| 54 |
+
|
| 55 |
+
scalar_t v1 = 0;
|
| 56 |
+
if (h_low >= 0 && w_low >= 0)
|
| 57 |
+
{
|
| 58 |
+
const int ptr1 = h_low_ptr_offset + w_low_ptr_offset + base_ptr;
|
| 59 |
+
v1 = bottom_data[ptr1];
|
| 60 |
+
}
|
| 61 |
+
scalar_t v2 = 0;
|
| 62 |
+
if (h_low >= 0 && w_high <= width - 1)
|
| 63 |
+
{
|
| 64 |
+
const int ptr2 = h_low_ptr_offset + w_high_ptr_offset + base_ptr;
|
| 65 |
+
v2 = bottom_data[ptr2];
|
| 66 |
+
}
|
| 67 |
+
scalar_t v3 = 0;
|
| 68 |
+
if (h_high <= height - 1 && w_low >= 0)
|
| 69 |
+
{
|
| 70 |
+
const int ptr3 = h_high_ptr_offset + w_low_ptr_offset + base_ptr;
|
| 71 |
+
v3 = bottom_data[ptr3];
|
| 72 |
+
}
|
| 73 |
+
scalar_t v4 = 0;
|
| 74 |
+
if (h_high <= height - 1 && w_high <= width - 1)
|
| 75 |
+
{
|
| 76 |
+
const int ptr4 = h_high_ptr_offset + w_high_ptr_offset + base_ptr;
|
| 77 |
+
v4 = bottom_data[ptr4];
|
| 78 |
+
}
|
| 79 |
+
|
| 80 |
+
const scalar_t w1 = hh * hw, w2 = hh * lw, w3 = lh * hw, w4 = lh * lw;
|
| 81 |
+
|
| 82 |
+
const scalar_t val = (w1 * v1 + w2 * v2 + w3 * v3 + w4 * v4);
|
| 83 |
+
return val;
|
| 84 |
+
}
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
template <typename scalar_t>
|
| 88 |
+
__device__ void ms_deform_attn_col2im_bilinear(const scalar_t* &bottom_data,
|
| 89 |
+
const int &height, const int &width, const int &nheads, const int &channels,
|
| 90 |
+
const scalar_t &h, const scalar_t &w, const int &m, const int &c,
|
| 91 |
+
const scalar_t &top_grad,
|
| 92 |
+
const scalar_t &attn_weight,
|
| 93 |
+
scalar_t* &grad_value,
|
| 94 |
+
scalar_t* grad_sampling_loc,
|
| 95 |
+
scalar_t* grad_attn_weight)
|
| 96 |
+
{
|
| 97 |
+
const int h_low = floor(h);
|
| 98 |
+
const int w_low = floor(w);
|
| 99 |
+
const int h_high = h_low + 1;
|
| 100 |
+
const int w_high = w_low + 1;
|
| 101 |
+
|
| 102 |
+
const scalar_t lh = h - h_low;
|
| 103 |
+
const scalar_t lw = w - w_low;
|
| 104 |
+
const scalar_t hh = 1 - lh, hw = 1 - lw;
|
| 105 |
+
|
| 106 |
+
const int w_stride = nheads * channels;
|
| 107 |
+
const int h_stride = width * w_stride;
|
| 108 |
+
const int h_low_ptr_offset = h_low * h_stride;
|
| 109 |
+
const int h_high_ptr_offset = h_low_ptr_offset + h_stride;
|
| 110 |
+
const int w_low_ptr_offset = w_low * w_stride;
|
| 111 |
+
const int w_high_ptr_offset = w_low_ptr_offset + w_stride;
|
| 112 |
+
const int base_ptr = m * channels + c;
|
| 113 |
+
|
| 114 |
+
const scalar_t w1 = hh * hw, w2 = hh * lw, w3 = lh * hw, w4 = lh * lw;
|
| 115 |
+
const scalar_t top_grad_value = top_grad * attn_weight;
|
| 116 |
+
scalar_t grad_h_weight = 0, grad_w_weight = 0;
|
| 117 |
+
|
| 118 |
+
scalar_t v1 = 0;
|
| 119 |
+
if (h_low >= 0 && w_low >= 0)
|
| 120 |
+
{
|
| 121 |
+
const int ptr1 = h_low_ptr_offset + w_low_ptr_offset + base_ptr;
|
| 122 |
+
v1 = bottom_data[ptr1];
|
| 123 |
+
grad_h_weight -= hw * v1;
|
| 124 |
+
grad_w_weight -= hh * v1;
|
| 125 |
+
atomicAdd(grad_value+ptr1, w1*top_grad_value);
|
| 126 |
+
}
|
| 127 |
+
scalar_t v2 = 0;
|
| 128 |
+
if (h_low >= 0 && w_high <= width - 1)
|
| 129 |
+
{
|
| 130 |
+
const int ptr2 = h_low_ptr_offset + w_high_ptr_offset + base_ptr;
|
| 131 |
+
v2 = bottom_data[ptr2];
|
| 132 |
+
grad_h_weight -= lw * v2;
|
| 133 |
+
grad_w_weight += hh * v2;
|
| 134 |
+
atomicAdd(grad_value+ptr2, w2*top_grad_value);
|
| 135 |
+
}
|
| 136 |
+
scalar_t v3 = 0;
|
| 137 |
+
if (h_high <= height - 1 && w_low >= 0)
|
| 138 |
+
{
|
| 139 |
+
const int ptr3 = h_high_ptr_offset + w_low_ptr_offset + base_ptr;
|
| 140 |
+
v3 = bottom_data[ptr3];
|
| 141 |
+
grad_h_weight += hw * v3;
|
| 142 |
+
grad_w_weight -= lh * v3;
|
| 143 |
+
atomicAdd(grad_value+ptr3, w3*top_grad_value);
|
| 144 |
+
}
|
| 145 |
+
scalar_t v4 = 0;
|
| 146 |
+
if (h_high <= height - 1 && w_high <= width - 1)
|
| 147 |
+
{
|
| 148 |
+
const int ptr4 = h_high_ptr_offset + w_high_ptr_offset + base_ptr;
|
| 149 |
+
v4 = bottom_data[ptr4];
|
| 150 |
+
grad_h_weight += lw * v4;
|
| 151 |
+
grad_w_weight += lh * v4;
|
| 152 |
+
atomicAdd(grad_value+ptr4, w4*top_grad_value);
|
| 153 |
+
}
|
| 154 |
+
|
| 155 |
+
const scalar_t val = (w1 * v1 + w2 * v2 + w3 * v3 + w4 * v4);
|
| 156 |
+
*grad_attn_weight = top_grad * val;
|
| 157 |
+
*grad_sampling_loc = width * grad_w_weight * top_grad_value;
|
| 158 |
+
*(grad_sampling_loc + 1) = height * grad_h_weight * top_grad_value;
|
| 159 |
+
}
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
template <typename scalar_t>
|
| 163 |
+
__device__ void ms_deform_attn_col2im_bilinear_gm(const scalar_t* &bottom_data,
|
| 164 |
+
const int &height, const int &width, const int &nheads, const int &channels,
|
| 165 |
+
const scalar_t &h, const scalar_t &w, const int &m, const int &c,
|
| 166 |
+
const scalar_t &top_grad,
|
| 167 |
+
const scalar_t &attn_weight,
|
| 168 |
+
scalar_t* &grad_value,
|
| 169 |
+
scalar_t* grad_sampling_loc,
|
| 170 |
+
scalar_t* grad_attn_weight)
|
| 171 |
+
{
|
| 172 |
+
const int h_low = floor(h);
|
| 173 |
+
const int w_low = floor(w);
|
| 174 |
+
const int h_high = h_low + 1;
|
| 175 |
+
const int w_high = w_low + 1;
|
| 176 |
+
|
| 177 |
+
const scalar_t lh = h - h_low;
|
| 178 |
+
const scalar_t lw = w - w_low;
|
| 179 |
+
const scalar_t hh = 1 - lh, hw = 1 - lw;
|
| 180 |
+
|
| 181 |
+
const int w_stride = nheads * channels;
|
| 182 |
+
const int h_stride = width * w_stride;
|
| 183 |
+
const int h_low_ptr_offset = h_low * h_stride;
|
| 184 |
+
const int h_high_ptr_offset = h_low_ptr_offset + h_stride;
|
| 185 |
+
const int w_low_ptr_offset = w_low * w_stride;
|
| 186 |
+
const int w_high_ptr_offset = w_low_ptr_offset + w_stride;
|
| 187 |
+
const int base_ptr = m * channels + c;
|
| 188 |
+
|
| 189 |
+
const scalar_t w1 = hh * hw, w2 = hh * lw, w3 = lh * hw, w4 = lh * lw;
|
| 190 |
+
const scalar_t top_grad_value = top_grad * attn_weight;
|
| 191 |
+
scalar_t grad_h_weight = 0, grad_w_weight = 0;
|
| 192 |
+
|
| 193 |
+
scalar_t v1 = 0;
|
| 194 |
+
if (h_low >= 0 && w_low >= 0)
|
| 195 |
+
{
|
| 196 |
+
const int ptr1 = h_low_ptr_offset + w_low_ptr_offset + base_ptr;
|
| 197 |
+
v1 = bottom_data[ptr1];
|
| 198 |
+
grad_h_weight -= hw * v1;
|
| 199 |
+
grad_w_weight -= hh * v1;
|
| 200 |
+
atomicAdd(grad_value+ptr1, w1*top_grad_value);
|
| 201 |
+
}
|
| 202 |
+
scalar_t v2 = 0;
|
| 203 |
+
if (h_low >= 0 && w_high <= width - 1)
|
| 204 |
+
{
|
| 205 |
+
const int ptr2 = h_low_ptr_offset + w_high_ptr_offset + base_ptr;
|
| 206 |
+
v2 = bottom_data[ptr2];
|
| 207 |
+
grad_h_weight -= lw * v2;
|
| 208 |
+
grad_w_weight += hh * v2;
|
| 209 |
+
atomicAdd(grad_value+ptr2, w2*top_grad_value);
|
| 210 |
+
}
|
| 211 |
+
scalar_t v3 = 0;
|
| 212 |
+
if (h_high <= height - 1 && w_low >= 0)
|
| 213 |
+
{
|
| 214 |
+
const int ptr3 = h_high_ptr_offset + w_low_ptr_offset + base_ptr;
|
| 215 |
+
v3 = bottom_data[ptr3];
|
| 216 |
+
grad_h_weight += hw * v3;
|
| 217 |
+
grad_w_weight -= lh * v3;
|
| 218 |
+
atomicAdd(grad_value+ptr3, w3*top_grad_value);
|
| 219 |
+
}
|
| 220 |
+
scalar_t v4 = 0;
|
| 221 |
+
if (h_high <= height - 1 && w_high <= width - 1)
|
| 222 |
+
{
|
| 223 |
+
const int ptr4 = h_high_ptr_offset + w_high_ptr_offset + base_ptr;
|
| 224 |
+
v4 = bottom_data[ptr4];
|
| 225 |
+
grad_h_weight += lw * v4;
|
| 226 |
+
grad_w_weight += lh * v4;
|
| 227 |
+
atomicAdd(grad_value+ptr4, w4*top_grad_value);
|
| 228 |
+
}
|
| 229 |
+
|
| 230 |
+
const scalar_t val = (w1 * v1 + w2 * v2 + w3 * v3 + w4 * v4);
|
| 231 |
+
atomicAdd(grad_attn_weight, top_grad * val);
|
| 232 |
+
atomicAdd(grad_sampling_loc, width * grad_w_weight * top_grad_value);
|
| 233 |
+
atomicAdd(grad_sampling_loc + 1, height * grad_h_weight * top_grad_value);
|
| 234 |
+
}
|
| 235 |
+
|
| 236 |
+
|
| 237 |
+
template <typename scalar_t>
|
| 238 |
+
__global__ void ms_deformable_im2col_gpu_kernel(const int n,
|
| 239 |
+
const scalar_t *data_value,
|
| 240 |
+
const int64_t *data_spatial_shapes,
|
| 241 |
+
const int64_t *data_level_start_index,
|
| 242 |
+
const scalar_t *data_sampling_loc,
|
| 243 |
+
const scalar_t *data_attn_weight,
|
| 244 |
+
const int batch_size,
|
| 245 |
+
const int spatial_size,
|
| 246 |
+
const int num_heads,
|
| 247 |
+
const int channels,
|
| 248 |
+
const int num_levels,
|
| 249 |
+
const int num_query,
|
| 250 |
+
const int num_point,
|
| 251 |
+
scalar_t *data_col)
|
| 252 |
+
{
|
| 253 |
+
CUDA_KERNEL_LOOP(index, n)
|
| 254 |
+
{
|
| 255 |
+
int _temp = index;
|
| 256 |
+
const int c_col = _temp % channels;
|
| 257 |
+
_temp /= channels;
|
| 258 |
+
const int sampling_index = _temp;
|
| 259 |
+
const int m_col = _temp % num_heads;
|
| 260 |
+
_temp /= num_heads;
|
| 261 |
+
const int q_col = _temp % num_query;
|
| 262 |
+
_temp /= num_query;
|
| 263 |
+
const int b_col = _temp;
|
| 264 |
+
|
| 265 |
+
scalar_t *data_col_ptr = data_col + index;
|
| 266 |
+
int data_weight_ptr = sampling_index * num_levels * num_point;
|
| 267 |
+
int data_loc_w_ptr = data_weight_ptr << 1;
|
| 268 |
+
const int qid_stride = num_heads * channels;
|
| 269 |
+
const int data_value_ptr_init_offset = b_col * spatial_size * qid_stride;
|
| 270 |
+
scalar_t col = 0;
|
| 271 |
+
|
| 272 |
+
for (int l_col=0; l_col < num_levels; ++l_col)
|
| 273 |
+
{
|
| 274 |
+
const int level_start_id = data_level_start_index[l_col];
|
| 275 |
+
const int spatial_h_ptr = l_col << 1;
|
| 276 |
+
const int spatial_h = data_spatial_shapes[spatial_h_ptr];
|
| 277 |
+
const int spatial_w = data_spatial_shapes[spatial_h_ptr + 1];
|
| 278 |
+
const scalar_t *data_value_ptr = data_value + (data_value_ptr_init_offset + level_start_id * qid_stride);
|
| 279 |
+
for (int p_col=0; p_col < num_point; ++p_col)
|
| 280 |
+
{
|
| 281 |
+
const scalar_t loc_w = data_sampling_loc[data_loc_w_ptr];
|
| 282 |
+
const scalar_t loc_h = data_sampling_loc[data_loc_w_ptr + 1];
|
| 283 |
+
const scalar_t weight = data_attn_weight[data_weight_ptr];
|
| 284 |
+
|
| 285 |
+
const scalar_t h_im = loc_h * spatial_h - 0.5;
|
| 286 |
+
const scalar_t w_im = loc_w * spatial_w - 0.5;
|
| 287 |
+
|
| 288 |
+
if (h_im > -1 && w_im > -1 && h_im < spatial_h && w_im < spatial_w)
|
| 289 |
+
{
|
| 290 |
+
col += ms_deform_attn_im2col_bilinear(data_value_ptr, spatial_h, spatial_w, num_heads, channels, h_im, w_im, m_col, c_col) * weight;
|
| 291 |
+
}
|
| 292 |
+
|
| 293 |
+
data_weight_ptr += 1;
|
| 294 |
+
data_loc_w_ptr += 2;
|
| 295 |
+
}
|
| 296 |
+
}
|
| 297 |
+
*data_col_ptr = col;
|
| 298 |
+
}
|
| 299 |
+
}
|
| 300 |
+
|
| 301 |
+
template <typename scalar_t, unsigned int blockSize>
|
| 302 |
+
__global__ void ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v1(const int n,
|
| 303 |
+
const scalar_t *grad_col,
|
| 304 |
+
const scalar_t *data_value,
|
| 305 |
+
const int64_t *data_spatial_shapes,
|
| 306 |
+
const int64_t *data_level_start_index,
|
| 307 |
+
const scalar_t *data_sampling_loc,
|
| 308 |
+
const scalar_t *data_attn_weight,
|
| 309 |
+
const int batch_size,
|
| 310 |
+
const int spatial_size,
|
| 311 |
+
const int num_heads,
|
| 312 |
+
const int channels,
|
| 313 |
+
const int num_levels,
|
| 314 |
+
const int num_query,
|
| 315 |
+
const int num_point,
|
| 316 |
+
scalar_t *grad_value,
|
| 317 |
+
scalar_t *grad_sampling_loc,
|
| 318 |
+
scalar_t *grad_attn_weight)
|
| 319 |
+
{
|
| 320 |
+
CUDA_KERNEL_LOOP(index, n)
|
| 321 |
+
{
|
| 322 |
+
__shared__ scalar_t cache_grad_sampling_loc[blockSize * 2];
|
| 323 |
+
__shared__ scalar_t cache_grad_attn_weight[blockSize];
|
| 324 |
+
unsigned int tid = threadIdx.x;
|
| 325 |
+
int _temp = index;
|
| 326 |
+
const int c_col = _temp % channels;
|
| 327 |
+
_temp /= channels;
|
| 328 |
+
const int sampling_index = _temp;
|
| 329 |
+
const int m_col = _temp % num_heads;
|
| 330 |
+
_temp /= num_heads;
|
| 331 |
+
const int q_col = _temp % num_query;
|
| 332 |
+
_temp /= num_query;
|
| 333 |
+
const int b_col = _temp;
|
| 334 |
+
|
| 335 |
+
const scalar_t top_grad = grad_col[index];
|
| 336 |
+
|
| 337 |
+
int data_weight_ptr = sampling_index * num_levels * num_point;
|
| 338 |
+
int data_loc_w_ptr = data_weight_ptr << 1;
|
| 339 |
+
const int grad_sampling_ptr = data_weight_ptr;
|
| 340 |
+
grad_sampling_loc += grad_sampling_ptr << 1;
|
| 341 |
+
grad_attn_weight += grad_sampling_ptr;
|
| 342 |
+
const int grad_weight_stride = 1;
|
| 343 |
+
const int grad_loc_stride = 2;
|
| 344 |
+
const int qid_stride = num_heads * channels;
|
| 345 |
+
const int data_value_ptr_init_offset = b_col * spatial_size * qid_stride;
|
| 346 |
+
|
| 347 |
+
for (int l_col=0; l_col < num_levels; ++l_col)
|
| 348 |
+
{
|
| 349 |
+
const int level_start_id = data_level_start_index[l_col];
|
| 350 |
+
const int spatial_h_ptr = l_col << 1;
|
| 351 |
+
const int spatial_h = data_spatial_shapes[spatial_h_ptr];
|
| 352 |
+
const int spatial_w = data_spatial_shapes[spatial_h_ptr + 1];
|
| 353 |
+
const int value_ptr_offset = data_value_ptr_init_offset + level_start_id * qid_stride;
|
| 354 |
+
const scalar_t *data_value_ptr = data_value + value_ptr_offset;
|
| 355 |
+
scalar_t *grad_value_ptr = grad_value + value_ptr_offset;
|
| 356 |
+
|
| 357 |
+
for (int p_col=0; p_col < num_point; ++p_col)
|
| 358 |
+
{
|
| 359 |
+
const scalar_t loc_w = data_sampling_loc[data_loc_w_ptr];
|
| 360 |
+
const scalar_t loc_h = data_sampling_loc[data_loc_w_ptr + 1];
|
| 361 |
+
const scalar_t weight = data_attn_weight[data_weight_ptr];
|
| 362 |
+
|
| 363 |
+
const scalar_t h_im = loc_h * spatial_h - 0.5;
|
| 364 |
+
const scalar_t w_im = loc_w * spatial_w - 0.5;
|
| 365 |
+
*(cache_grad_sampling_loc+(threadIdx.x << 1)) = 0;
|
| 366 |
+
*(cache_grad_sampling_loc+((threadIdx.x << 1) + 1)) = 0;
|
| 367 |
+
*(cache_grad_attn_weight+threadIdx.x)=0;
|
| 368 |
+
if (h_im > -1 && w_im > -1 && h_im < spatial_h && w_im < spatial_w)
|
| 369 |
+
{
|
| 370 |
+
ms_deform_attn_col2im_bilinear(
|
| 371 |
+
data_value_ptr, spatial_h, spatial_w, num_heads, channels, h_im, w_im, m_col, c_col,
|
| 372 |
+
top_grad, weight, grad_value_ptr,
|
| 373 |
+
cache_grad_sampling_loc+(threadIdx.x << 1), cache_grad_attn_weight+threadIdx.x);
|
| 374 |
+
}
|
| 375 |
+
|
| 376 |
+
__syncthreads();
|
| 377 |
+
if (tid == 0)
|
| 378 |
+
{
|
| 379 |
+
scalar_t _grad_w=cache_grad_sampling_loc[0], _grad_h=cache_grad_sampling_loc[1], _grad_a=cache_grad_attn_weight[0];
|
| 380 |
+
int sid=2;
|
| 381 |
+
for (unsigned int tid = 1; tid < blockSize; ++tid)
|
| 382 |
+
{
|
| 383 |
+
_grad_w += cache_grad_sampling_loc[sid];
|
| 384 |
+
_grad_h += cache_grad_sampling_loc[sid + 1];
|
| 385 |
+
_grad_a += cache_grad_attn_weight[tid];
|
| 386 |
+
sid += 2;
|
| 387 |
+
}
|
| 388 |
+
|
| 389 |
+
|
| 390 |
+
*grad_sampling_loc = _grad_w;
|
| 391 |
+
*(grad_sampling_loc + 1) = _grad_h;
|
| 392 |
+
*grad_attn_weight = _grad_a;
|
| 393 |
+
}
|
| 394 |
+
__syncthreads();
|
| 395 |
+
|
| 396 |
+
data_weight_ptr += 1;
|
| 397 |
+
data_loc_w_ptr += 2;
|
| 398 |
+
grad_attn_weight += grad_weight_stride;
|
| 399 |
+
grad_sampling_loc += grad_loc_stride;
|
| 400 |
+
}
|
| 401 |
+
}
|
| 402 |
+
}
|
| 403 |
+
}
|
| 404 |
+
|
| 405 |
+
|
| 406 |
+
template <typename scalar_t, unsigned int blockSize>
|
| 407 |
+
__global__ void ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v2(const int n,
|
| 408 |
+
const scalar_t *grad_col,
|
| 409 |
+
const scalar_t *data_value,
|
| 410 |
+
const int64_t *data_spatial_shapes,
|
| 411 |
+
const int64_t *data_level_start_index,
|
| 412 |
+
const scalar_t *data_sampling_loc,
|
| 413 |
+
const scalar_t *data_attn_weight,
|
| 414 |
+
const int batch_size,
|
| 415 |
+
const int spatial_size,
|
| 416 |
+
const int num_heads,
|
| 417 |
+
const int channels,
|
| 418 |
+
const int num_levels,
|
| 419 |
+
const int num_query,
|
| 420 |
+
const int num_point,
|
| 421 |
+
scalar_t *grad_value,
|
| 422 |
+
scalar_t *grad_sampling_loc,
|
| 423 |
+
scalar_t *grad_attn_weight)
|
| 424 |
+
{
|
| 425 |
+
CUDA_KERNEL_LOOP(index, n)
|
| 426 |
+
{
|
| 427 |
+
__shared__ scalar_t cache_grad_sampling_loc[blockSize * 2];
|
| 428 |
+
__shared__ scalar_t cache_grad_attn_weight[blockSize];
|
| 429 |
+
unsigned int tid = threadIdx.x;
|
| 430 |
+
int _temp = index;
|
| 431 |
+
const int c_col = _temp % channels;
|
| 432 |
+
_temp /= channels;
|
| 433 |
+
const int sampling_index = _temp;
|
| 434 |
+
const int m_col = _temp % num_heads;
|
| 435 |
+
_temp /= num_heads;
|
| 436 |
+
const int q_col = _temp % num_query;
|
| 437 |
+
_temp /= num_query;
|
| 438 |
+
const int b_col = _temp;
|
| 439 |
+
|
| 440 |
+
const scalar_t top_grad = grad_col[index];
|
| 441 |
+
|
| 442 |
+
int data_weight_ptr = sampling_index * num_levels * num_point;
|
| 443 |
+
int data_loc_w_ptr = data_weight_ptr << 1;
|
| 444 |
+
const int grad_sampling_ptr = data_weight_ptr;
|
| 445 |
+
grad_sampling_loc += grad_sampling_ptr << 1;
|
| 446 |
+
grad_attn_weight += grad_sampling_ptr;
|
| 447 |
+
const int grad_weight_stride = 1;
|
| 448 |
+
const int grad_loc_stride = 2;
|
| 449 |
+
const int qid_stride = num_heads * channels;
|
| 450 |
+
const int data_value_ptr_init_offset = b_col * spatial_size * qid_stride;
|
| 451 |
+
|
| 452 |
+
for (int l_col=0; l_col < num_levels; ++l_col)
|
| 453 |
+
{
|
| 454 |
+
const int level_start_id = data_level_start_index[l_col];
|
| 455 |
+
const int spatial_h_ptr = l_col << 1;
|
| 456 |
+
const int spatial_h = data_spatial_shapes[spatial_h_ptr];
|
| 457 |
+
const int spatial_w = data_spatial_shapes[spatial_h_ptr + 1];
|
| 458 |
+
const int value_ptr_offset = data_value_ptr_init_offset + level_start_id * qid_stride;
|
| 459 |
+
const scalar_t *data_value_ptr = data_value + value_ptr_offset;
|
| 460 |
+
scalar_t *grad_value_ptr = grad_value + value_ptr_offset;
|
| 461 |
+
|
| 462 |
+
for (int p_col=0; p_col < num_point; ++p_col)
|
| 463 |
+
{
|
| 464 |
+
const scalar_t loc_w = data_sampling_loc[data_loc_w_ptr];
|
| 465 |
+
const scalar_t loc_h = data_sampling_loc[data_loc_w_ptr + 1];
|
| 466 |
+
const scalar_t weight = data_attn_weight[data_weight_ptr];
|
| 467 |
+
|
| 468 |
+
const scalar_t h_im = loc_h * spatial_h - 0.5;
|
| 469 |
+
const scalar_t w_im = loc_w * spatial_w - 0.5;
|
| 470 |
+
*(cache_grad_sampling_loc+(threadIdx.x << 1)) = 0;
|
| 471 |
+
*(cache_grad_sampling_loc+((threadIdx.x << 1) + 1)) = 0;
|
| 472 |
+
*(cache_grad_attn_weight+threadIdx.x)=0;
|
| 473 |
+
if (h_im > -1 && w_im > -1 && h_im < spatial_h && w_im < spatial_w)
|
| 474 |
+
{
|
| 475 |
+
ms_deform_attn_col2im_bilinear(
|
| 476 |
+
data_value_ptr, spatial_h, spatial_w, num_heads, channels, h_im, w_im, m_col, c_col,
|
| 477 |
+
top_grad, weight, grad_value_ptr,
|
| 478 |
+
cache_grad_sampling_loc+(threadIdx.x << 1), cache_grad_attn_weight+threadIdx.x);
|
| 479 |
+
}
|
| 480 |
+
|
| 481 |
+
__syncthreads();
|
| 482 |
+
|
| 483 |
+
for (unsigned int s=blockSize/2; s>0; s>>=1)
|
| 484 |
+
{
|
| 485 |
+
if (tid < s) {
|
| 486 |
+
const unsigned int xid1 = tid << 1;
|
| 487 |
+
const unsigned int xid2 = (tid + s) << 1;
|
| 488 |
+
cache_grad_attn_weight[tid] += cache_grad_attn_weight[tid + s];
|
| 489 |
+
cache_grad_sampling_loc[xid1] += cache_grad_sampling_loc[xid2];
|
| 490 |
+
cache_grad_sampling_loc[xid1 + 1] += cache_grad_sampling_loc[xid2 + 1];
|
| 491 |
+
}
|
| 492 |
+
__syncthreads();
|
| 493 |
+
}
|
| 494 |
+
|
| 495 |
+
if (tid == 0)
|
| 496 |
+
{
|
| 497 |
+
*grad_sampling_loc = cache_grad_sampling_loc[0];
|
| 498 |
+
*(grad_sampling_loc + 1) = cache_grad_sampling_loc[1];
|
| 499 |
+
*grad_attn_weight = cache_grad_attn_weight[0];
|
| 500 |
+
}
|
| 501 |
+
__syncthreads();
|
| 502 |
+
|
| 503 |
+
data_weight_ptr += 1;
|
| 504 |
+
data_loc_w_ptr += 2;
|
| 505 |
+
grad_attn_weight += grad_weight_stride;
|
| 506 |
+
grad_sampling_loc += grad_loc_stride;
|
| 507 |
+
}
|
| 508 |
+
}
|
| 509 |
+
}
|
| 510 |
+
}
|
| 511 |
+
|
| 512 |
+
|
| 513 |
+
template <typename scalar_t>
|
| 514 |
+
__global__ void ms_deformable_col2im_gpu_kernel_shm_reduce_v1(const int n,
|
| 515 |
+
const scalar_t *grad_col,
|
| 516 |
+
const scalar_t *data_value,
|
| 517 |
+
const int64_t *data_spatial_shapes,
|
| 518 |
+
const int64_t *data_level_start_index,
|
| 519 |
+
const scalar_t *data_sampling_loc,
|
| 520 |
+
const scalar_t *data_attn_weight,
|
| 521 |
+
const int batch_size,
|
| 522 |
+
const int spatial_size,
|
| 523 |
+
const int num_heads,
|
| 524 |
+
const int channels,
|
| 525 |
+
const int num_levels,
|
| 526 |
+
const int num_query,
|
| 527 |
+
const int num_point,
|
| 528 |
+
scalar_t *grad_value,
|
| 529 |
+
scalar_t *grad_sampling_loc,
|
| 530 |
+
scalar_t *grad_attn_weight)
|
| 531 |
+
{
|
| 532 |
+
CUDA_KERNEL_LOOP(index, n)
|
| 533 |
+
{
|
| 534 |
+
extern __shared__ int _s[];
|
| 535 |
+
scalar_t* cache_grad_sampling_loc = (scalar_t*)_s;
|
| 536 |
+
scalar_t* cache_grad_attn_weight = cache_grad_sampling_loc + 2 * blockDim.x;
|
| 537 |
+
unsigned int tid = threadIdx.x;
|
| 538 |
+
int _temp = index;
|
| 539 |
+
const int c_col = _temp % channels;
|
| 540 |
+
_temp /= channels;
|
| 541 |
+
const int sampling_index = _temp;
|
| 542 |
+
const int m_col = _temp % num_heads;
|
| 543 |
+
_temp /= num_heads;
|
| 544 |
+
const int q_col = _temp % num_query;
|
| 545 |
+
_temp /= num_query;
|
| 546 |
+
const int b_col = _temp;
|
| 547 |
+
|
| 548 |
+
const scalar_t top_grad = grad_col[index];
|
| 549 |
+
|
| 550 |
+
int data_weight_ptr = sampling_index * num_levels * num_point;
|
| 551 |
+
int data_loc_w_ptr = data_weight_ptr << 1;
|
| 552 |
+
const int grad_sampling_ptr = data_weight_ptr;
|
| 553 |
+
grad_sampling_loc += grad_sampling_ptr << 1;
|
| 554 |
+
grad_attn_weight += grad_sampling_ptr;
|
| 555 |
+
const int grad_weight_stride = 1;
|
| 556 |
+
const int grad_loc_stride = 2;
|
| 557 |
+
const int qid_stride = num_heads * channels;
|
| 558 |
+
const int data_value_ptr_init_offset = b_col * spatial_size * qid_stride;
|
| 559 |
+
|
| 560 |
+
for (int l_col=0; l_col < num_levels; ++l_col)
|
| 561 |
+
{
|
| 562 |
+
const int level_start_id = data_level_start_index[l_col];
|
| 563 |
+
const int spatial_h_ptr = l_col << 1;
|
| 564 |
+
const int spatial_h = data_spatial_shapes[spatial_h_ptr];
|
| 565 |
+
const int spatial_w = data_spatial_shapes[spatial_h_ptr + 1];
|
| 566 |
+
const int value_ptr_offset = data_value_ptr_init_offset + level_start_id * qid_stride;
|
| 567 |
+
const scalar_t *data_value_ptr = data_value + value_ptr_offset;
|
| 568 |
+
scalar_t *grad_value_ptr = grad_value + value_ptr_offset;
|
| 569 |
+
|
| 570 |
+
for (int p_col=0; p_col < num_point; ++p_col)
|
| 571 |
+
{
|
| 572 |
+
const scalar_t loc_w = data_sampling_loc[data_loc_w_ptr];
|
| 573 |
+
const scalar_t loc_h = data_sampling_loc[data_loc_w_ptr + 1];
|
| 574 |
+
const scalar_t weight = data_attn_weight[data_weight_ptr];
|
| 575 |
+
|
| 576 |
+
const scalar_t h_im = loc_h * spatial_h - 0.5;
|
| 577 |
+
const scalar_t w_im = loc_w * spatial_w - 0.5;
|
| 578 |
+
*(cache_grad_sampling_loc+(threadIdx.x << 1)) = 0;
|
| 579 |
+
*(cache_grad_sampling_loc+((threadIdx.x << 1) + 1)) = 0;
|
| 580 |
+
*(cache_grad_attn_weight+threadIdx.x)=0;
|
| 581 |
+
if (h_im > -1 && w_im > -1 && h_im < spatial_h && w_im < spatial_w)
|
| 582 |
+
{
|
| 583 |
+
ms_deform_attn_col2im_bilinear(
|
| 584 |
+
data_value_ptr, spatial_h, spatial_w, num_heads, channels, h_im, w_im, m_col, c_col,
|
| 585 |
+
top_grad, weight, grad_value_ptr,
|
| 586 |
+
cache_grad_sampling_loc+(threadIdx.x << 1), cache_grad_attn_weight+threadIdx.x);
|
| 587 |
+
}
|
| 588 |
+
|
| 589 |
+
__syncthreads();
|
| 590 |
+
if (tid == 0)
|
| 591 |
+
{
|
| 592 |
+
scalar_t _grad_w=cache_grad_sampling_loc[0], _grad_h=cache_grad_sampling_loc[1], _grad_a=cache_grad_attn_weight[0];
|
| 593 |
+
int sid=2;
|
| 594 |
+
for (unsigned int tid = 1; tid < blockDim.x; ++tid)
|
| 595 |
+
{
|
| 596 |
+
_grad_w += cache_grad_sampling_loc[sid];
|
| 597 |
+
_grad_h += cache_grad_sampling_loc[sid + 1];
|
| 598 |
+
_grad_a += cache_grad_attn_weight[tid];
|
| 599 |
+
sid += 2;
|
| 600 |
+
}
|
| 601 |
+
|
| 602 |
+
|
| 603 |
+
*grad_sampling_loc = _grad_w;
|
| 604 |
+
*(grad_sampling_loc + 1) = _grad_h;
|
| 605 |
+
*grad_attn_weight = _grad_a;
|
| 606 |
+
}
|
| 607 |
+
__syncthreads();
|
| 608 |
+
|
| 609 |
+
data_weight_ptr += 1;
|
| 610 |
+
data_loc_w_ptr += 2;
|
| 611 |
+
grad_attn_weight += grad_weight_stride;
|
| 612 |
+
grad_sampling_loc += grad_loc_stride;
|
| 613 |
+
}
|
| 614 |
+
}
|
| 615 |
+
}
|
| 616 |
+
}
|
| 617 |
+
|
| 618 |
+
template <typename scalar_t>
|
| 619 |
+
__global__ void ms_deformable_col2im_gpu_kernel_shm_reduce_v2(const int n,
|
| 620 |
+
const scalar_t *grad_col,
|
| 621 |
+
const scalar_t *data_value,
|
| 622 |
+
const int64_t *data_spatial_shapes,
|
| 623 |
+
const int64_t *data_level_start_index,
|
| 624 |
+
const scalar_t *data_sampling_loc,
|
| 625 |
+
const scalar_t *data_attn_weight,
|
| 626 |
+
const int batch_size,
|
| 627 |
+
const int spatial_size,
|
| 628 |
+
const int num_heads,
|
| 629 |
+
const int channels,
|
| 630 |
+
const int num_levels,
|
| 631 |
+
const int num_query,
|
| 632 |
+
const int num_point,
|
| 633 |
+
scalar_t *grad_value,
|
| 634 |
+
scalar_t *grad_sampling_loc,
|
| 635 |
+
scalar_t *grad_attn_weight)
|
| 636 |
+
{
|
| 637 |
+
CUDA_KERNEL_LOOP(index, n)
|
| 638 |
+
{
|
| 639 |
+
extern __shared__ int _s[];
|
| 640 |
+
scalar_t* cache_grad_sampling_loc = (scalar_t*)_s;
|
| 641 |
+
scalar_t* cache_grad_attn_weight = cache_grad_sampling_loc + 2 * blockDim.x;
|
| 642 |
+
unsigned int tid = threadIdx.x;
|
| 643 |
+
int _temp = index;
|
| 644 |
+
const int c_col = _temp % channels;
|
| 645 |
+
_temp /= channels;
|
| 646 |
+
const int sampling_index = _temp;
|
| 647 |
+
const int m_col = _temp % num_heads;
|
| 648 |
+
_temp /= num_heads;
|
| 649 |
+
const int q_col = _temp % num_query;
|
| 650 |
+
_temp /= num_query;
|
| 651 |
+
const int b_col = _temp;
|
| 652 |
+
|
| 653 |
+
const scalar_t top_grad = grad_col[index];
|
| 654 |
+
|
| 655 |
+
int data_weight_ptr = sampling_index * num_levels * num_point;
|
| 656 |
+
int data_loc_w_ptr = data_weight_ptr << 1;
|
| 657 |
+
const int grad_sampling_ptr = data_weight_ptr;
|
| 658 |
+
grad_sampling_loc += grad_sampling_ptr << 1;
|
| 659 |
+
grad_attn_weight += grad_sampling_ptr;
|
| 660 |
+
const int grad_weight_stride = 1;
|
| 661 |
+
const int grad_loc_stride = 2;
|
| 662 |
+
const int qid_stride = num_heads * channels;
|
| 663 |
+
const int data_value_ptr_init_offset = b_col * spatial_size * qid_stride;
|
| 664 |
+
|
| 665 |
+
for (int l_col=0; l_col < num_levels; ++l_col)
|
| 666 |
+
{
|
| 667 |
+
const int level_start_id = data_level_start_index[l_col];
|
| 668 |
+
const int spatial_h_ptr = l_col << 1;
|
| 669 |
+
const int spatial_h = data_spatial_shapes[spatial_h_ptr];
|
| 670 |
+
const int spatial_w = data_spatial_shapes[spatial_h_ptr + 1];
|
| 671 |
+
const int value_ptr_offset = data_value_ptr_init_offset + level_start_id * qid_stride;
|
| 672 |
+
const scalar_t *data_value_ptr = data_value + value_ptr_offset;
|
| 673 |
+
scalar_t *grad_value_ptr = grad_value + value_ptr_offset;
|
| 674 |
+
|
| 675 |
+
for (int p_col=0; p_col < num_point; ++p_col)
|
| 676 |
+
{
|
| 677 |
+
const scalar_t loc_w = data_sampling_loc[data_loc_w_ptr];
|
| 678 |
+
const scalar_t loc_h = data_sampling_loc[data_loc_w_ptr + 1];
|
| 679 |
+
const scalar_t weight = data_attn_weight[data_weight_ptr];
|
| 680 |
+
|
| 681 |
+
const scalar_t h_im = loc_h * spatial_h - 0.5;
|
| 682 |
+
const scalar_t w_im = loc_w * spatial_w - 0.5;
|
| 683 |
+
*(cache_grad_sampling_loc+(threadIdx.x << 1)) = 0;
|
| 684 |
+
*(cache_grad_sampling_loc+((threadIdx.x << 1) + 1)) = 0;
|
| 685 |
+
*(cache_grad_attn_weight+threadIdx.x)=0;
|
| 686 |
+
if (h_im > -1 && w_im > -1 && h_im < spatial_h && w_im < spatial_w)
|
| 687 |
+
{
|
| 688 |
+
ms_deform_attn_col2im_bilinear(
|
| 689 |
+
data_value_ptr, spatial_h, spatial_w, num_heads, channels, h_im, w_im, m_col, c_col,
|
| 690 |
+
top_grad, weight, grad_value_ptr,
|
| 691 |
+
cache_grad_sampling_loc+(threadIdx.x << 1), cache_grad_attn_weight+threadIdx.x);
|
| 692 |
+
}
|
| 693 |
+
|
| 694 |
+
__syncthreads();
|
| 695 |
+
|
| 696 |
+
for (unsigned int s=blockDim.x/2, spre=blockDim.x; s>0; s>>=1, spre>>=1)
|
| 697 |
+
{
|
| 698 |
+
if (tid < s) {
|
| 699 |
+
const unsigned int xid1 = tid << 1;
|
| 700 |
+
const unsigned int xid2 = (tid + s) << 1;
|
| 701 |
+
cache_grad_attn_weight[tid] += cache_grad_attn_weight[tid + s];
|
| 702 |
+
cache_grad_sampling_loc[xid1] += cache_grad_sampling_loc[xid2];
|
| 703 |
+
cache_grad_sampling_loc[xid1 + 1] += cache_grad_sampling_loc[xid2 + 1];
|
| 704 |
+
if (tid + (s << 1) < spre)
|
| 705 |
+
{
|
| 706 |
+
cache_grad_attn_weight[tid] += cache_grad_attn_weight[tid + (s << 1)];
|
| 707 |
+
cache_grad_sampling_loc[xid1] += cache_grad_sampling_loc[xid2 + (s << 1)];
|
| 708 |
+
cache_grad_sampling_loc[xid1 + 1] += cache_grad_sampling_loc[xid2 + 1 + (s << 1)];
|
| 709 |
+
}
|
| 710 |
+
}
|
| 711 |
+
__syncthreads();
|
| 712 |
+
}
|
| 713 |
+
|
| 714 |
+
if (tid == 0)
|
| 715 |
+
{
|
| 716 |
+
*grad_sampling_loc = cache_grad_sampling_loc[0];
|
| 717 |
+
*(grad_sampling_loc + 1) = cache_grad_sampling_loc[1];
|
| 718 |
+
*grad_attn_weight = cache_grad_attn_weight[0];
|
| 719 |
+
}
|
| 720 |
+
__syncthreads();
|
| 721 |
+
|
| 722 |
+
data_weight_ptr += 1;
|
| 723 |
+
data_loc_w_ptr += 2;
|
| 724 |
+
grad_attn_weight += grad_weight_stride;
|
| 725 |
+
grad_sampling_loc += grad_loc_stride;
|
| 726 |
+
}
|
| 727 |
+
}
|
| 728 |
+
}
|
| 729 |
+
}
|
| 730 |
+
|
| 731 |
+
template <typename scalar_t>
|
| 732 |
+
__global__ void ms_deformable_col2im_gpu_kernel_shm_reduce_v2_multi_blocks(const int n,
|
| 733 |
+
const scalar_t *grad_col,
|
| 734 |
+
const scalar_t *data_value,
|
| 735 |
+
const int64_t *data_spatial_shapes,
|
| 736 |
+
const int64_t *data_level_start_index,
|
| 737 |
+
const scalar_t *data_sampling_loc,
|
| 738 |
+
const scalar_t *data_attn_weight,
|
| 739 |
+
const int batch_size,
|
| 740 |
+
const int spatial_size,
|
| 741 |
+
const int num_heads,
|
| 742 |
+
const int channels,
|
| 743 |
+
const int num_levels,
|
| 744 |
+
const int num_query,
|
| 745 |
+
const int num_point,
|
| 746 |
+
scalar_t *grad_value,
|
| 747 |
+
scalar_t *grad_sampling_loc,
|
| 748 |
+
scalar_t *grad_attn_weight)
|
| 749 |
+
{
|
| 750 |
+
CUDA_KERNEL_LOOP(index, n)
|
| 751 |
+
{
|
| 752 |
+
extern __shared__ int _s[];
|
| 753 |
+
scalar_t* cache_grad_sampling_loc = (scalar_t*)_s;
|
| 754 |
+
scalar_t* cache_grad_attn_weight = cache_grad_sampling_loc + 2 * blockDim.x;
|
| 755 |
+
unsigned int tid = threadIdx.x;
|
| 756 |
+
int _temp = index;
|
| 757 |
+
const int c_col = _temp % channels;
|
| 758 |
+
_temp /= channels;
|
| 759 |
+
const int sampling_index = _temp;
|
| 760 |
+
const int m_col = _temp % num_heads;
|
| 761 |
+
_temp /= num_heads;
|
| 762 |
+
const int q_col = _temp % num_query;
|
| 763 |
+
_temp /= num_query;
|
| 764 |
+
const int b_col = _temp;
|
| 765 |
+
|
| 766 |
+
const scalar_t top_grad = grad_col[index];
|
| 767 |
+
|
| 768 |
+
int data_weight_ptr = sampling_index * num_levels * num_point;
|
| 769 |
+
int data_loc_w_ptr = data_weight_ptr << 1;
|
| 770 |
+
const int grad_sampling_ptr = data_weight_ptr;
|
| 771 |
+
grad_sampling_loc += grad_sampling_ptr << 1;
|
| 772 |
+
grad_attn_weight += grad_sampling_ptr;
|
| 773 |
+
const int grad_weight_stride = 1;
|
| 774 |
+
const int grad_loc_stride = 2;
|
| 775 |
+
const int qid_stride = num_heads * channels;
|
| 776 |
+
const int data_value_ptr_init_offset = b_col * spatial_size * qid_stride;
|
| 777 |
+
|
| 778 |
+
for (int l_col=0; l_col < num_levels; ++l_col)
|
| 779 |
+
{
|
| 780 |
+
const int level_start_id = data_level_start_index[l_col];
|
| 781 |
+
const int spatial_h_ptr = l_col << 1;
|
| 782 |
+
const int spatial_h = data_spatial_shapes[spatial_h_ptr];
|
| 783 |
+
const int spatial_w = data_spatial_shapes[spatial_h_ptr + 1];
|
| 784 |
+
const int value_ptr_offset = data_value_ptr_init_offset + level_start_id * qid_stride;
|
| 785 |
+
const scalar_t *data_value_ptr = data_value + value_ptr_offset;
|
| 786 |
+
scalar_t *grad_value_ptr = grad_value + value_ptr_offset;
|
| 787 |
+
|
| 788 |
+
for (int p_col=0; p_col < num_point; ++p_col)
|
| 789 |
+
{
|
| 790 |
+
const scalar_t loc_w = data_sampling_loc[data_loc_w_ptr];
|
| 791 |
+
const scalar_t loc_h = data_sampling_loc[data_loc_w_ptr + 1];
|
| 792 |
+
const scalar_t weight = data_attn_weight[data_weight_ptr];
|
| 793 |
+
|
| 794 |
+
const scalar_t h_im = loc_h * spatial_h - 0.5;
|
| 795 |
+
const scalar_t w_im = loc_w * spatial_w - 0.5;
|
| 796 |
+
*(cache_grad_sampling_loc+(threadIdx.x << 1)) = 0;
|
| 797 |
+
*(cache_grad_sampling_loc+((threadIdx.x << 1) + 1)) = 0;
|
| 798 |
+
*(cache_grad_attn_weight+threadIdx.x)=0;
|
| 799 |
+
if (h_im > -1 && w_im > -1 && h_im < spatial_h && w_im < spatial_w)
|
| 800 |
+
{
|
| 801 |
+
ms_deform_attn_col2im_bilinear(
|
| 802 |
+
data_value_ptr, spatial_h, spatial_w, num_heads, channels, h_im, w_im, m_col, c_col,
|
| 803 |
+
top_grad, weight, grad_value_ptr,
|
| 804 |
+
cache_grad_sampling_loc+(threadIdx.x << 1), cache_grad_attn_weight+threadIdx.x);
|
| 805 |
+
}
|
| 806 |
+
|
| 807 |
+
__syncthreads();
|
| 808 |
+
|
| 809 |
+
for (unsigned int s=blockDim.x/2, spre=blockDim.x; s>0; s>>=1, spre>>=1)
|
| 810 |
+
{
|
| 811 |
+
if (tid < s) {
|
| 812 |
+
const unsigned int xid1 = tid << 1;
|
| 813 |
+
const unsigned int xid2 = (tid + s) << 1;
|
| 814 |
+
cache_grad_attn_weight[tid] += cache_grad_attn_weight[tid + s];
|
| 815 |
+
cache_grad_sampling_loc[xid1] += cache_grad_sampling_loc[xid2];
|
| 816 |
+
cache_grad_sampling_loc[xid1 + 1] += cache_grad_sampling_loc[xid2 + 1];
|
| 817 |
+
if (tid + (s << 1) < spre)
|
| 818 |
+
{
|
| 819 |
+
cache_grad_attn_weight[tid] += cache_grad_attn_weight[tid + (s << 1)];
|
| 820 |
+
cache_grad_sampling_loc[xid1] += cache_grad_sampling_loc[xid2 + (s << 1)];
|
| 821 |
+
cache_grad_sampling_loc[xid1 + 1] += cache_grad_sampling_loc[xid2 + 1 + (s << 1)];
|
| 822 |
+
}
|
| 823 |
+
}
|
| 824 |
+
__syncthreads();
|
| 825 |
+
}
|
| 826 |
+
|
| 827 |
+
if (tid == 0)
|
| 828 |
+
{
|
| 829 |
+
atomicAdd(grad_sampling_loc, cache_grad_sampling_loc[0]);
|
| 830 |
+
atomicAdd(grad_sampling_loc + 1, cache_grad_sampling_loc[1]);
|
| 831 |
+
atomicAdd(grad_attn_weight, cache_grad_attn_weight[0]);
|
| 832 |
+
}
|
| 833 |
+
__syncthreads();
|
| 834 |
+
|
| 835 |
+
data_weight_ptr += 1;
|
| 836 |
+
data_loc_w_ptr += 2;
|
| 837 |
+
grad_attn_weight += grad_weight_stride;
|
| 838 |
+
grad_sampling_loc += grad_loc_stride;
|
| 839 |
+
}
|
| 840 |
+
}
|
| 841 |
+
}
|
| 842 |
+
}
|
| 843 |
+
|
| 844 |
+
|
| 845 |
+
template <typename scalar_t>
|
| 846 |
+
__global__ void ms_deformable_col2im_gpu_kernel_gm(const int n,
|
| 847 |
+
const scalar_t *grad_col,
|
| 848 |
+
const scalar_t *data_value,
|
| 849 |
+
const int64_t *data_spatial_shapes,
|
| 850 |
+
const int64_t *data_level_start_index,
|
| 851 |
+
const scalar_t *data_sampling_loc,
|
| 852 |
+
const scalar_t *data_attn_weight,
|
| 853 |
+
const int batch_size,
|
| 854 |
+
const int spatial_size,
|
| 855 |
+
const int num_heads,
|
| 856 |
+
const int channels,
|
| 857 |
+
const int num_levels,
|
| 858 |
+
const int num_query,
|
| 859 |
+
const int num_point,
|
| 860 |
+
scalar_t *grad_value,
|
| 861 |
+
scalar_t *grad_sampling_loc,
|
| 862 |
+
scalar_t *grad_attn_weight)
|
| 863 |
+
{
|
| 864 |
+
CUDA_KERNEL_LOOP(index, n)
|
| 865 |
+
{
|
| 866 |
+
int _temp = index;
|
| 867 |
+
const int c_col = _temp % channels;
|
| 868 |
+
_temp /= channels;
|
| 869 |
+
const int sampling_index = _temp;
|
| 870 |
+
const int m_col = _temp % num_heads;
|
| 871 |
+
_temp /= num_heads;
|
| 872 |
+
const int q_col = _temp % num_query;
|
| 873 |
+
_temp /= num_query;
|
| 874 |
+
const int b_col = _temp;
|
| 875 |
+
|
| 876 |
+
const scalar_t top_grad = grad_col[index];
|
| 877 |
+
|
| 878 |
+
int data_weight_ptr = sampling_index * num_levels * num_point;
|
| 879 |
+
int data_loc_w_ptr = data_weight_ptr << 1;
|
| 880 |
+
const int grad_sampling_ptr = data_weight_ptr;
|
| 881 |
+
grad_sampling_loc += grad_sampling_ptr << 1;
|
| 882 |
+
grad_attn_weight += grad_sampling_ptr;
|
| 883 |
+
const int grad_weight_stride = 1;
|
| 884 |
+
const int grad_loc_stride = 2;
|
| 885 |
+
const int qid_stride = num_heads * channels;
|
| 886 |
+
const int data_value_ptr_init_offset = b_col * spatial_size * qid_stride;
|
| 887 |
+
|
| 888 |
+
for (int l_col=0; l_col < num_levels; ++l_col)
|
| 889 |
+
{
|
| 890 |
+
const int level_start_id = data_level_start_index[l_col];
|
| 891 |
+
const int spatial_h_ptr = l_col << 1;
|
| 892 |
+
const int spatial_h = data_spatial_shapes[spatial_h_ptr];
|
| 893 |
+
const int spatial_w = data_spatial_shapes[spatial_h_ptr + 1];
|
| 894 |
+
const int value_ptr_offset = data_value_ptr_init_offset + level_start_id * qid_stride;
|
| 895 |
+
const scalar_t *data_value_ptr = data_value + value_ptr_offset;
|
| 896 |
+
scalar_t *grad_value_ptr = grad_value + value_ptr_offset;
|
| 897 |
+
|
| 898 |
+
for (int p_col=0; p_col < num_point; ++p_col)
|
| 899 |
+
{
|
| 900 |
+
const scalar_t loc_w = data_sampling_loc[data_loc_w_ptr];
|
| 901 |
+
const scalar_t loc_h = data_sampling_loc[data_loc_w_ptr + 1];
|
| 902 |
+
const scalar_t weight = data_attn_weight[data_weight_ptr];
|
| 903 |
+
|
| 904 |
+
const scalar_t h_im = loc_h * spatial_h - 0.5;
|
| 905 |
+
const scalar_t w_im = loc_w * spatial_w - 0.5;
|
| 906 |
+
if (h_im > -1 && w_im > -1 && h_im < spatial_h && w_im < spatial_w)
|
| 907 |
+
{
|
| 908 |
+
ms_deform_attn_col2im_bilinear_gm(
|
| 909 |
+
data_value_ptr, spatial_h, spatial_w, num_heads, channels, h_im, w_im, m_col, c_col,
|
| 910 |
+
top_grad, weight, grad_value_ptr,
|
| 911 |
+
grad_sampling_loc, grad_attn_weight);
|
| 912 |
+
}
|
| 913 |
+
data_weight_ptr += 1;
|
| 914 |
+
data_loc_w_ptr += 2;
|
| 915 |
+
grad_attn_weight += grad_weight_stride;
|
| 916 |
+
grad_sampling_loc += grad_loc_stride;
|
| 917 |
+
}
|
| 918 |
+
}
|
| 919 |
+
}
|
| 920 |
+
}
|
| 921 |
+
|
| 922 |
+
|
| 923 |
+
template <typename scalar_t>
|
| 924 |
+
void ms_deformable_im2col_cuda(cudaStream_t stream,
|
| 925 |
+
const scalar_t* data_value,
|
| 926 |
+
const int64_t* data_spatial_shapes,
|
| 927 |
+
const int64_t* data_level_start_index,
|
| 928 |
+
const scalar_t* data_sampling_loc,
|
| 929 |
+
const scalar_t* data_attn_weight,
|
| 930 |
+
const int batch_size,
|
| 931 |
+
const int spatial_size,
|
| 932 |
+
const int num_heads,
|
| 933 |
+
const int channels,
|
| 934 |
+
const int num_levels,
|
| 935 |
+
const int num_query,
|
| 936 |
+
const int num_point,
|
| 937 |
+
scalar_t* data_col)
|
| 938 |
+
{
|
| 939 |
+
const int num_kernels = batch_size * num_query * num_heads * channels;
|
| 940 |
+
const int num_actual_kernels = batch_size * num_query * num_heads * channels;
|
| 941 |
+
const int num_threads = CUDA_NUM_THREADS;
|
| 942 |
+
ms_deformable_im2col_gpu_kernel<scalar_t>
|
| 943 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 944 |
+
0, stream>>>(
|
| 945 |
+
num_kernels, data_value, data_spatial_shapes, data_level_start_index, data_sampling_loc, data_attn_weight,
|
| 946 |
+
batch_size, spatial_size, num_heads, channels, num_levels, num_query, num_point, data_col);
|
| 947 |
+
|
| 948 |
+
cudaError_t err = cudaGetLastError();
|
| 949 |
+
if (err != cudaSuccess)
|
| 950 |
+
{
|
| 951 |
+
printf("error in ms_deformable_im2col_cuda: %s\n", cudaGetErrorString(err));
|
| 952 |
+
}
|
| 953 |
+
|
| 954 |
+
}
|
| 955 |
+
|
| 956 |
+
template <typename scalar_t>
|
| 957 |
+
void ms_deformable_col2im_cuda(cudaStream_t stream,
|
| 958 |
+
const scalar_t* grad_col,
|
| 959 |
+
const scalar_t* data_value,
|
| 960 |
+
const int64_t * data_spatial_shapes,
|
| 961 |
+
const int64_t * data_level_start_index,
|
| 962 |
+
const scalar_t * data_sampling_loc,
|
| 963 |
+
const scalar_t * data_attn_weight,
|
| 964 |
+
const int batch_size,
|
| 965 |
+
const int spatial_size,
|
| 966 |
+
const int num_heads,
|
| 967 |
+
const int channels,
|
| 968 |
+
const int num_levels,
|
| 969 |
+
const int num_query,
|
| 970 |
+
const int num_point,
|
| 971 |
+
scalar_t* grad_value,
|
| 972 |
+
scalar_t* grad_sampling_loc,
|
| 973 |
+
scalar_t* grad_attn_weight)
|
| 974 |
+
{
|
| 975 |
+
const int num_threads = (channels > CUDA_NUM_THREADS)?CUDA_NUM_THREADS:channels;
|
| 976 |
+
const int num_kernels = batch_size * num_query * num_heads * channels;
|
| 977 |
+
const int num_actual_kernels = batch_size * num_query * num_heads * channels;
|
| 978 |
+
if (channels > 1024)
|
| 979 |
+
{
|
| 980 |
+
if ((channels & 1023) == 0)
|
| 981 |
+
{
|
| 982 |
+
ms_deformable_col2im_gpu_kernel_shm_reduce_v2_multi_blocks<scalar_t>
|
| 983 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 984 |
+
num_threads*3*sizeof(scalar_t), stream>>>(
|
| 985 |
+
num_kernels,
|
| 986 |
+
grad_col,
|
| 987 |
+
data_value,
|
| 988 |
+
data_spatial_shapes,
|
| 989 |
+
data_level_start_index,
|
| 990 |
+
data_sampling_loc,
|
| 991 |
+
data_attn_weight,
|
| 992 |
+
batch_size,
|
| 993 |
+
spatial_size,
|
| 994 |
+
num_heads,
|
| 995 |
+
channels,
|
| 996 |
+
num_levels,
|
| 997 |
+
num_query,
|
| 998 |
+
num_point,
|
| 999 |
+
grad_value,
|
| 1000 |
+
grad_sampling_loc,
|
| 1001 |
+
grad_attn_weight);
|
| 1002 |
+
}
|
| 1003 |
+
else
|
| 1004 |
+
{
|
| 1005 |
+
ms_deformable_col2im_gpu_kernel_gm<scalar_t>
|
| 1006 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1007 |
+
0, stream>>>(
|
| 1008 |
+
num_kernels,
|
| 1009 |
+
grad_col,
|
| 1010 |
+
data_value,
|
| 1011 |
+
data_spatial_shapes,
|
| 1012 |
+
data_level_start_index,
|
| 1013 |
+
data_sampling_loc,
|
| 1014 |
+
data_attn_weight,
|
| 1015 |
+
batch_size,
|
| 1016 |
+
spatial_size,
|
| 1017 |
+
num_heads,
|
| 1018 |
+
channels,
|
| 1019 |
+
num_levels,
|
| 1020 |
+
num_query,
|
| 1021 |
+
num_point,
|
| 1022 |
+
grad_value,
|
| 1023 |
+
grad_sampling_loc,
|
| 1024 |
+
grad_attn_weight);
|
| 1025 |
+
}
|
| 1026 |
+
}
|
| 1027 |
+
else{
|
| 1028 |
+
switch(channels)
|
| 1029 |
+
{
|
| 1030 |
+
case 1:
|
| 1031 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v1<scalar_t, 1>
|
| 1032 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1033 |
+
0, stream>>>(
|
| 1034 |
+
num_kernels,
|
| 1035 |
+
grad_col,
|
| 1036 |
+
data_value,
|
| 1037 |
+
data_spatial_shapes,
|
| 1038 |
+
data_level_start_index,
|
| 1039 |
+
data_sampling_loc,
|
| 1040 |
+
data_attn_weight,
|
| 1041 |
+
batch_size,
|
| 1042 |
+
spatial_size,
|
| 1043 |
+
num_heads,
|
| 1044 |
+
channels,
|
| 1045 |
+
num_levels,
|
| 1046 |
+
num_query,
|
| 1047 |
+
num_point,
|
| 1048 |
+
grad_value,
|
| 1049 |
+
grad_sampling_loc,
|
| 1050 |
+
grad_attn_weight);
|
| 1051 |
+
break;
|
| 1052 |
+
case 2:
|
| 1053 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v1<scalar_t, 2>
|
| 1054 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1055 |
+
0, stream>>>(
|
| 1056 |
+
num_kernels,
|
| 1057 |
+
grad_col,
|
| 1058 |
+
data_value,
|
| 1059 |
+
data_spatial_shapes,
|
| 1060 |
+
data_level_start_index,
|
| 1061 |
+
data_sampling_loc,
|
| 1062 |
+
data_attn_weight,
|
| 1063 |
+
batch_size,
|
| 1064 |
+
spatial_size,
|
| 1065 |
+
num_heads,
|
| 1066 |
+
channels,
|
| 1067 |
+
num_levels,
|
| 1068 |
+
num_query,
|
| 1069 |
+
num_point,
|
| 1070 |
+
grad_value,
|
| 1071 |
+
grad_sampling_loc,
|
| 1072 |
+
grad_attn_weight);
|
| 1073 |
+
break;
|
| 1074 |
+
case 4:
|
| 1075 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v1<scalar_t, 4>
|
| 1076 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1077 |
+
0, stream>>>(
|
| 1078 |
+
num_kernels,
|
| 1079 |
+
grad_col,
|
| 1080 |
+
data_value,
|
| 1081 |
+
data_spatial_shapes,
|
| 1082 |
+
data_level_start_index,
|
| 1083 |
+
data_sampling_loc,
|
| 1084 |
+
data_attn_weight,
|
| 1085 |
+
batch_size,
|
| 1086 |
+
spatial_size,
|
| 1087 |
+
num_heads,
|
| 1088 |
+
channels,
|
| 1089 |
+
num_levels,
|
| 1090 |
+
num_query,
|
| 1091 |
+
num_point,
|
| 1092 |
+
grad_value,
|
| 1093 |
+
grad_sampling_loc,
|
| 1094 |
+
grad_attn_weight);
|
| 1095 |
+
break;
|
| 1096 |
+
case 8:
|
| 1097 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v1<scalar_t, 8>
|
| 1098 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1099 |
+
0, stream>>>(
|
| 1100 |
+
num_kernels,
|
| 1101 |
+
grad_col,
|
| 1102 |
+
data_value,
|
| 1103 |
+
data_spatial_shapes,
|
| 1104 |
+
data_level_start_index,
|
| 1105 |
+
data_sampling_loc,
|
| 1106 |
+
data_attn_weight,
|
| 1107 |
+
batch_size,
|
| 1108 |
+
spatial_size,
|
| 1109 |
+
num_heads,
|
| 1110 |
+
channels,
|
| 1111 |
+
num_levels,
|
| 1112 |
+
num_query,
|
| 1113 |
+
num_point,
|
| 1114 |
+
grad_value,
|
| 1115 |
+
grad_sampling_loc,
|
| 1116 |
+
grad_attn_weight);
|
| 1117 |
+
break;
|
| 1118 |
+
case 16:
|
| 1119 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v1<scalar_t, 16>
|
| 1120 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1121 |
+
0, stream>>>(
|
| 1122 |
+
num_kernels,
|
| 1123 |
+
grad_col,
|
| 1124 |
+
data_value,
|
| 1125 |
+
data_spatial_shapes,
|
| 1126 |
+
data_level_start_index,
|
| 1127 |
+
data_sampling_loc,
|
| 1128 |
+
data_attn_weight,
|
| 1129 |
+
batch_size,
|
| 1130 |
+
spatial_size,
|
| 1131 |
+
num_heads,
|
| 1132 |
+
channels,
|
| 1133 |
+
num_levels,
|
| 1134 |
+
num_query,
|
| 1135 |
+
num_point,
|
| 1136 |
+
grad_value,
|
| 1137 |
+
grad_sampling_loc,
|
| 1138 |
+
grad_attn_weight);
|
| 1139 |
+
break;
|
| 1140 |
+
case 32:
|
| 1141 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v1<scalar_t, 32>
|
| 1142 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1143 |
+
0, stream>>>(
|
| 1144 |
+
num_kernels,
|
| 1145 |
+
grad_col,
|
| 1146 |
+
data_value,
|
| 1147 |
+
data_spatial_shapes,
|
| 1148 |
+
data_level_start_index,
|
| 1149 |
+
data_sampling_loc,
|
| 1150 |
+
data_attn_weight,
|
| 1151 |
+
batch_size,
|
| 1152 |
+
spatial_size,
|
| 1153 |
+
num_heads,
|
| 1154 |
+
channels,
|
| 1155 |
+
num_levels,
|
| 1156 |
+
num_query,
|
| 1157 |
+
num_point,
|
| 1158 |
+
grad_value,
|
| 1159 |
+
grad_sampling_loc,
|
| 1160 |
+
grad_attn_weight);
|
| 1161 |
+
break;
|
| 1162 |
+
case 64:
|
| 1163 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v2<scalar_t, 64>
|
| 1164 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1165 |
+
0, stream>>>(
|
| 1166 |
+
num_kernels,
|
| 1167 |
+
grad_col,
|
| 1168 |
+
data_value,
|
| 1169 |
+
data_spatial_shapes,
|
| 1170 |
+
data_level_start_index,
|
| 1171 |
+
data_sampling_loc,
|
| 1172 |
+
data_attn_weight,
|
| 1173 |
+
batch_size,
|
| 1174 |
+
spatial_size,
|
| 1175 |
+
num_heads,
|
| 1176 |
+
channels,
|
| 1177 |
+
num_levels,
|
| 1178 |
+
num_query,
|
| 1179 |
+
num_point,
|
| 1180 |
+
grad_value,
|
| 1181 |
+
grad_sampling_loc,
|
| 1182 |
+
grad_attn_weight);
|
| 1183 |
+
break;
|
| 1184 |
+
case 128:
|
| 1185 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v2<scalar_t, 128>
|
| 1186 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1187 |
+
0, stream>>>(
|
| 1188 |
+
num_kernels,
|
| 1189 |
+
grad_col,
|
| 1190 |
+
data_value,
|
| 1191 |
+
data_spatial_shapes,
|
| 1192 |
+
data_level_start_index,
|
| 1193 |
+
data_sampling_loc,
|
| 1194 |
+
data_attn_weight,
|
| 1195 |
+
batch_size,
|
| 1196 |
+
spatial_size,
|
| 1197 |
+
num_heads,
|
| 1198 |
+
channels,
|
| 1199 |
+
num_levels,
|
| 1200 |
+
num_query,
|
| 1201 |
+
num_point,
|
| 1202 |
+
grad_value,
|
| 1203 |
+
grad_sampling_loc,
|
| 1204 |
+
grad_attn_weight);
|
| 1205 |
+
break;
|
| 1206 |
+
case 256:
|
| 1207 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v2<scalar_t, 256>
|
| 1208 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1209 |
+
0, stream>>>(
|
| 1210 |
+
num_kernels,
|
| 1211 |
+
grad_col,
|
| 1212 |
+
data_value,
|
| 1213 |
+
data_spatial_shapes,
|
| 1214 |
+
data_level_start_index,
|
| 1215 |
+
data_sampling_loc,
|
| 1216 |
+
data_attn_weight,
|
| 1217 |
+
batch_size,
|
| 1218 |
+
spatial_size,
|
| 1219 |
+
num_heads,
|
| 1220 |
+
channels,
|
| 1221 |
+
num_levels,
|
| 1222 |
+
num_query,
|
| 1223 |
+
num_point,
|
| 1224 |
+
grad_value,
|
| 1225 |
+
grad_sampling_loc,
|
| 1226 |
+
grad_attn_weight);
|
| 1227 |
+
break;
|
| 1228 |
+
case 512:
|
| 1229 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v2<scalar_t, 512>
|
| 1230 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1231 |
+
0, stream>>>(
|
| 1232 |
+
num_kernels,
|
| 1233 |
+
grad_col,
|
| 1234 |
+
data_value,
|
| 1235 |
+
data_spatial_shapes,
|
| 1236 |
+
data_level_start_index,
|
| 1237 |
+
data_sampling_loc,
|
| 1238 |
+
data_attn_weight,
|
| 1239 |
+
batch_size,
|
| 1240 |
+
spatial_size,
|
| 1241 |
+
num_heads,
|
| 1242 |
+
channels,
|
| 1243 |
+
num_levels,
|
| 1244 |
+
num_query,
|
| 1245 |
+
num_point,
|
| 1246 |
+
grad_value,
|
| 1247 |
+
grad_sampling_loc,
|
| 1248 |
+
grad_attn_weight);
|
| 1249 |
+
break;
|
| 1250 |
+
case 1024:
|
| 1251 |
+
ms_deformable_col2im_gpu_kernel_shm_blocksize_aware_reduce_v2<scalar_t, 1024>
|
| 1252 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1253 |
+
0, stream>>>(
|
| 1254 |
+
num_kernels,
|
| 1255 |
+
grad_col,
|
| 1256 |
+
data_value,
|
| 1257 |
+
data_spatial_shapes,
|
| 1258 |
+
data_level_start_index,
|
| 1259 |
+
data_sampling_loc,
|
| 1260 |
+
data_attn_weight,
|
| 1261 |
+
batch_size,
|
| 1262 |
+
spatial_size,
|
| 1263 |
+
num_heads,
|
| 1264 |
+
channels,
|
| 1265 |
+
num_levels,
|
| 1266 |
+
num_query,
|
| 1267 |
+
num_point,
|
| 1268 |
+
grad_value,
|
| 1269 |
+
grad_sampling_loc,
|
| 1270 |
+
grad_attn_weight);
|
| 1271 |
+
break;
|
| 1272 |
+
default:
|
| 1273 |
+
if (channels < 64)
|
| 1274 |
+
{
|
| 1275 |
+
ms_deformable_col2im_gpu_kernel_shm_reduce_v1<scalar_t>
|
| 1276 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1277 |
+
num_threads*3*sizeof(scalar_t), stream>>>(
|
| 1278 |
+
num_kernels,
|
| 1279 |
+
grad_col,
|
| 1280 |
+
data_value,
|
| 1281 |
+
data_spatial_shapes,
|
| 1282 |
+
data_level_start_index,
|
| 1283 |
+
data_sampling_loc,
|
| 1284 |
+
data_attn_weight,
|
| 1285 |
+
batch_size,
|
| 1286 |
+
spatial_size,
|
| 1287 |
+
num_heads,
|
| 1288 |
+
channels,
|
| 1289 |
+
num_levels,
|
| 1290 |
+
num_query,
|
| 1291 |
+
num_point,
|
| 1292 |
+
grad_value,
|
| 1293 |
+
grad_sampling_loc,
|
| 1294 |
+
grad_attn_weight);
|
| 1295 |
+
}
|
| 1296 |
+
else
|
| 1297 |
+
{
|
| 1298 |
+
ms_deformable_col2im_gpu_kernel_shm_reduce_v2<scalar_t>
|
| 1299 |
+
<<<GET_BLOCKS(num_actual_kernels, num_threads), num_threads,
|
| 1300 |
+
num_threads*3*sizeof(scalar_t), stream>>>(
|
| 1301 |
+
num_kernels,
|
| 1302 |
+
grad_col,
|
| 1303 |
+
data_value,
|
| 1304 |
+
data_spatial_shapes,
|
| 1305 |
+
data_level_start_index,
|
| 1306 |
+
data_sampling_loc,
|
| 1307 |
+
data_attn_weight,
|
| 1308 |
+
batch_size,
|
| 1309 |
+
spatial_size,
|
| 1310 |
+
num_heads,
|
| 1311 |
+
channels,
|
| 1312 |
+
num_levels,
|
| 1313 |
+
num_query,
|
| 1314 |
+
num_point,
|
| 1315 |
+
grad_value,
|
| 1316 |
+
grad_sampling_loc,
|
| 1317 |
+
grad_attn_weight);
|
| 1318 |
+
}
|
| 1319 |
+
}
|
| 1320 |
+
}
|
| 1321 |
+
cudaError_t err = cudaGetLastError();
|
| 1322 |
+
if (err != cudaSuccess)
|
| 1323 |
+
{
|
| 1324 |
+
printf("error in ms_deformable_col2im_cuda: %s\n", cudaGetErrorString(err));
|
| 1325 |
+
}
|
| 1326 |
+
|
| 1327 |
+
}
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/cuda_version.cu
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#include <cuda_runtime_api.h>
|
| 2 |
+
|
| 3 |
+
namespace groundingdino {
|
| 4 |
+
int get_cudart_version() {
|
| 5 |
+
return CUDART_VERSION;
|
| 6 |
+
}
|
| 7 |
+
} // namespace groundingdino
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/csrc/vision.cpp
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved
|
| 2 |
+
|
| 3 |
+
#include "MsDeformAttn/ms_deform_attn.h"
|
| 4 |
+
|
| 5 |
+
namespace groundingdino {
|
| 6 |
+
|
| 7 |
+
#ifdef WITH_CUDA
|
| 8 |
+
extern int get_cudart_version();
|
| 9 |
+
#endif
|
| 10 |
+
|
| 11 |
+
std::string get_cuda_version() {
|
| 12 |
+
#ifdef WITH_CUDA
|
| 13 |
+
std::ostringstream oss;
|
| 14 |
+
|
| 15 |
+
// copied from
|
| 16 |
+
// https://github.com/pytorch/pytorch/blob/master/aten/src/ATen/cuda/detail/CUDAHooks.cpp#L231
|
| 17 |
+
auto printCudaStyleVersion = [&](int v) {
|
| 18 |
+
oss << (v / 1000) << "." << (v / 10 % 100);
|
| 19 |
+
if (v % 10 != 0) {
|
| 20 |
+
oss << "." << (v % 10);
|
| 21 |
+
}
|
| 22 |
+
};
|
| 23 |
+
printCudaStyleVersion(get_cudart_version());
|
| 24 |
+
return oss.str();
|
| 25 |
+
#else
|
| 26 |
+
return std::string("not available");
|
| 27 |
+
#endif
|
| 28 |
+
}
|
| 29 |
+
|
| 30 |
+
// similar to
|
| 31 |
+
// https://github.com/pytorch/pytorch/blob/master/aten/src/ATen/Version.cpp
|
| 32 |
+
std::string get_compiler_version() {
|
| 33 |
+
std::ostringstream ss;
|
| 34 |
+
#if defined(__GNUC__)
|
| 35 |
+
#ifndef __clang__
|
| 36 |
+
{ ss << "GCC " << __GNUC__ << "." << __GNUC_MINOR__; }
|
| 37 |
+
#endif
|
| 38 |
+
#endif
|
| 39 |
+
|
| 40 |
+
#if defined(__clang_major__)
|
| 41 |
+
{
|
| 42 |
+
ss << "clang " << __clang_major__ << "." << __clang_minor__ << "."
|
| 43 |
+
<< __clang_patchlevel__;
|
| 44 |
+
}
|
| 45 |
+
#endif
|
| 46 |
+
|
| 47 |
+
#if defined(_MSC_VER)
|
| 48 |
+
{ ss << "MSVC " << _MSC_FULL_VER; }
|
| 49 |
+
#endif
|
| 50 |
+
return ss.str();
|
| 51 |
+
}
|
| 52 |
+
|
| 53 |
+
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
| 54 |
+
m.def("ms_deform_attn_forward", &ms_deform_attn_forward, "ms_deform_attn_forward");
|
| 55 |
+
m.def("ms_deform_attn_backward", &ms_deform_attn_backward, "ms_deform_attn_backward");
|
| 56 |
+
}
|
| 57 |
+
|
| 58 |
+
} // namespace groundingdino
|
ArtiAgent - DefectDiffu/src/GroundingDINO/groundingdino/models/GroundingDINO/fuse_modules.py
ADDED
|
@@ -0,0 +1,297 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# ------------------------------------------------------------------------
|
| 2 |
+
# Grounding DINO
|
| 3 |
+
# url: https://github.com/IDEA-Research/GroundingDINO
|
| 4 |
+
# Copyright (c) 2023 IDEA. All Rights Reserved.
|
| 5 |
+
# Licensed under the Apache License, Version 2.0 [see LICENSE for details]
|
| 6 |
+
# ------------------------------------------------------------------------
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
import torch.nn as nn
|
| 10 |
+
import torch.nn.functional as F
|
| 11 |
+
from timm.models.layers import DropPath
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class FeatureResizer(nn.Module):
|
| 15 |
+
"""
|
| 16 |
+
This class takes as input a set of embeddings of dimension C1 and outputs a set of
|
| 17 |
+
embedding of dimension C2, after a linear transformation, dropout and normalization (LN).
|
| 18 |
+
"""
|
| 19 |
+
|
| 20 |
+
def __init__(self, input_feat_size, output_feat_size, dropout, do_ln=True):
|
| 21 |
+
super().__init__()
|
| 22 |
+
self.do_ln = do_ln
|
| 23 |
+
# Object feature encoding
|
| 24 |
+
self.fc = nn.Linear(input_feat_size, output_feat_size, bias=True)
|
| 25 |
+
self.layer_norm = nn.LayerNorm(output_feat_size, eps=1e-12)
|
| 26 |
+
self.dropout = nn.Dropout(dropout)
|
| 27 |
+
|
| 28 |
+
def forward(self, encoder_features):
|
| 29 |
+
x = self.fc(encoder_features)
|
| 30 |
+
if self.do_ln:
|
| 31 |
+
x = self.layer_norm(x)
|
| 32 |
+
output = self.dropout(x)
|
| 33 |
+
return output
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def l1norm(X, dim, eps=1e-8):
|
| 37 |
+
"""L1-normalize columns of X"""
|
| 38 |
+
norm = torch.abs(X).sum(dim=dim, keepdim=True) + eps
|
| 39 |
+
X = torch.div(X, norm)
|
| 40 |
+
return X
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def l2norm(X, dim, eps=1e-8):
|
| 44 |
+
"""L2-normalize columns of X"""
|
| 45 |
+
norm = torch.pow(X, 2).sum(dim=dim, keepdim=True).sqrt() + eps
|
| 46 |
+
X = torch.div(X, norm)
|
| 47 |
+
return X
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def func_attention(query, context, smooth=1, raw_feature_norm="softmax", eps=1e-8):
|
| 51 |
+
"""
|
| 52 |
+
query: (n_context, queryL, d)
|
| 53 |
+
context: (n_context, sourceL, d)
|
| 54 |
+
"""
|
| 55 |
+
batch_size_q, queryL = query.size(0), query.size(1)
|
| 56 |
+
batch_size, sourceL = context.size(0), context.size(1)
|
| 57 |
+
|
| 58 |
+
# Get attention
|
| 59 |
+
# --> (batch, d, queryL)
|
| 60 |
+
queryT = torch.transpose(query, 1, 2)
|
| 61 |
+
|
| 62 |
+
# (batch, sourceL, d)(batch, d, queryL)
|
| 63 |
+
# --> (batch, sourceL, queryL)
|
| 64 |
+
attn = torch.bmm(context, queryT)
|
| 65 |
+
if raw_feature_norm == "softmax":
|
| 66 |
+
# --> (batch*sourceL, queryL)
|
| 67 |
+
attn = attn.view(batch_size * sourceL, queryL)
|
| 68 |
+
attn = nn.Softmax()(attn)
|
| 69 |
+
# --> (batch, sourceL, queryL)
|
| 70 |
+
attn = attn.view(batch_size, sourceL, queryL)
|
| 71 |
+
elif raw_feature_norm == "l2norm":
|
| 72 |
+
attn = l2norm(attn, 2)
|
| 73 |
+
elif raw_feature_norm == "clipped_l2norm":
|
| 74 |
+
attn = nn.LeakyReLU(0.1)(attn)
|
| 75 |
+
attn = l2norm(attn, 2)
|
| 76 |
+
else:
|
| 77 |
+
raise ValueError("unknown first norm type:", raw_feature_norm)
|
| 78 |
+
# --> (batch, queryL, sourceL)
|
| 79 |
+
attn = torch.transpose(attn, 1, 2).contiguous()
|
| 80 |
+
# --> (batch*queryL, sourceL)
|
| 81 |
+
attn = attn.view(batch_size * queryL, sourceL)
|
| 82 |
+
attn = nn.Softmax()(attn * smooth)
|
| 83 |
+
# --> (batch, queryL, sourceL)
|
| 84 |
+
attn = attn.view(batch_size, queryL, sourceL)
|
| 85 |
+
# --> (batch, sourceL, queryL)
|
| 86 |
+
attnT = torch.transpose(attn, 1, 2).contiguous()
|
| 87 |
+
|
| 88 |
+
# --> (batch, d, sourceL)
|
| 89 |
+
contextT = torch.transpose(context, 1, 2)
|
| 90 |
+
# (batch x d x sourceL)(batch x sourceL x queryL)
|
| 91 |
+
# --> (batch, d, queryL)
|
| 92 |
+
weightedContext = torch.bmm(contextT, attnT)
|
| 93 |
+
# --> (batch, queryL, d)
|
| 94 |
+
weightedContext = torch.transpose(weightedContext, 1, 2)
|
| 95 |
+
|
| 96 |
+
return weightedContext, attnT
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
class BiMultiHeadAttention(nn.Module):
|
| 100 |
+
def __init__(self, v_dim, l_dim, embed_dim, num_heads, dropout=0.1, cfg=None):
|
| 101 |
+
super(BiMultiHeadAttention, self).__init__()
|
| 102 |
+
|
| 103 |
+
self.embed_dim = embed_dim
|
| 104 |
+
self.num_heads = num_heads
|
| 105 |
+
self.head_dim = embed_dim // num_heads
|
| 106 |
+
self.v_dim = v_dim
|
| 107 |
+
self.l_dim = l_dim
|
| 108 |
+
|
| 109 |
+
assert (
|
| 110 |
+
self.head_dim * self.num_heads == self.embed_dim
|
| 111 |
+
), f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`: {self.num_heads})."
|
| 112 |
+
self.scale = self.head_dim ** (-0.5)
|
| 113 |
+
self.dropout = dropout
|
| 114 |
+
|
| 115 |
+
self.v_proj = nn.Linear(self.v_dim, self.embed_dim)
|
| 116 |
+
self.l_proj = nn.Linear(self.l_dim, self.embed_dim)
|
| 117 |
+
self.values_v_proj = nn.Linear(self.v_dim, self.embed_dim)
|
| 118 |
+
self.values_l_proj = nn.Linear(self.l_dim, self.embed_dim)
|
| 119 |
+
|
| 120 |
+
self.out_v_proj = nn.Linear(self.embed_dim, self.v_dim)
|
| 121 |
+
self.out_l_proj = nn.Linear(self.embed_dim, self.l_dim)
|
| 122 |
+
|
| 123 |
+
self.stable_softmax_2d = True
|
| 124 |
+
self.clamp_min_for_underflow = True
|
| 125 |
+
self.clamp_max_for_overflow = True
|
| 126 |
+
|
| 127 |
+
self._reset_parameters()
|
| 128 |
+
|
| 129 |
+
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
|
| 130 |
+
return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()
|
| 131 |
+
|
| 132 |
+
def _reset_parameters(self):
|
| 133 |
+
nn.init.xavier_uniform_(self.v_proj.weight)
|
| 134 |
+
self.v_proj.bias.data.fill_(0)
|
| 135 |
+
nn.init.xavier_uniform_(self.l_proj.weight)
|
| 136 |
+
self.l_proj.bias.data.fill_(0)
|
| 137 |
+
nn.init.xavier_uniform_(self.values_v_proj.weight)
|
| 138 |
+
self.values_v_proj.bias.data.fill_(0)
|
| 139 |
+
nn.init.xavier_uniform_(self.values_l_proj.weight)
|
| 140 |
+
self.values_l_proj.bias.data.fill_(0)
|
| 141 |
+
nn.init.xavier_uniform_(self.out_v_proj.weight)
|
| 142 |
+
self.out_v_proj.bias.data.fill_(0)
|
| 143 |
+
nn.init.xavier_uniform_(self.out_l_proj.weight)
|
| 144 |
+
self.out_l_proj.bias.data.fill_(0)
|
| 145 |
+
|
| 146 |
+
def forward(self, v, l, attention_mask_v=None, attention_mask_l=None):
|
| 147 |
+
"""_summary_
|
| 148 |
+
|
| 149 |
+
Args:
|
| 150 |
+
v (_type_): bs, n_img, dim
|
| 151 |
+
l (_type_): bs, n_text, dim
|
| 152 |
+
attention_mask_v (_type_, optional): _description_. bs, n_img
|
| 153 |
+
attention_mask_l (_type_, optional): _description_. bs, n_text
|
| 154 |
+
|
| 155 |
+
Returns:
|
| 156 |
+
_type_: _description_
|
| 157 |
+
"""
|
| 158 |
+
# if os.environ.get('IPDB_SHILONG_DEBUG', None) == 'INFO':
|
| 159 |
+
# import ipdb; ipdb.set_trace()
|
| 160 |
+
bsz, tgt_len, _ = v.size()
|
| 161 |
+
|
| 162 |
+
query_states = self.v_proj(v) * self.scale
|
| 163 |
+
key_states = self._shape(self.l_proj(l), -1, bsz)
|
| 164 |
+
value_v_states = self._shape(self.values_v_proj(v), -1, bsz)
|
| 165 |
+
value_l_states = self._shape(self.values_l_proj(l), -1, bsz)
|
| 166 |
+
|
| 167 |
+
proj_shape = (bsz * self.num_heads, -1, self.head_dim)
|
| 168 |
+
query_states = self._shape(query_states, tgt_len, bsz).view(*proj_shape)
|
| 169 |
+
key_states = key_states.view(*proj_shape)
|
| 170 |
+
value_v_states = value_v_states.view(*proj_shape)
|
| 171 |
+
value_l_states = value_l_states.view(*proj_shape)
|
| 172 |
+
|
| 173 |
+
src_len = key_states.size(1)
|
| 174 |
+
attn_weights = torch.bmm(query_states, key_states.transpose(1, 2)) # bs*nhead, nimg, ntxt
|
| 175 |
+
|
| 176 |
+
if attn_weights.size() != (bsz * self.num_heads, tgt_len, src_len):
|
| 177 |
+
raise ValueError(
|
| 178 |
+
f"Attention weights should be of size {(bsz * self.num_heads, tgt_len, src_len)}, but is {attn_weights.size()}"
|
| 179 |
+
)
|
| 180 |
+
|
| 181 |
+
if self.stable_softmax_2d:
|
| 182 |
+
attn_weights = attn_weights - attn_weights.max()
|
| 183 |
+
|
| 184 |
+
if self.clamp_min_for_underflow:
|
| 185 |
+
attn_weights = torch.clamp(
|
| 186 |
+
attn_weights, min=-50000
|
| 187 |
+
) # Do not increase -50000, data type half has quite limited range
|
| 188 |
+
if self.clamp_max_for_overflow:
|
| 189 |
+
attn_weights = torch.clamp(
|
| 190 |
+
attn_weights, max=50000
|
| 191 |
+
) # Do not increase 50000, data type half has quite limited range
|
| 192 |
+
|
| 193 |
+
attn_weights_T = attn_weights.transpose(1, 2)
|
| 194 |
+
attn_weights_l = attn_weights_T - torch.max(attn_weights_T, dim=-1, keepdim=True)[0]
|
| 195 |
+
if self.clamp_min_for_underflow:
|
| 196 |
+
attn_weights_l = torch.clamp(
|
| 197 |
+
attn_weights_l, min=-50000
|
| 198 |
+
) # Do not increase -50000, data type half has quite limited range
|
| 199 |
+
if self.clamp_max_for_overflow:
|
| 200 |
+
attn_weights_l = torch.clamp(
|
| 201 |
+
attn_weights_l, max=50000
|
| 202 |
+
) # Do not increase 50000, data type half has quite limited range
|
| 203 |
+
|
| 204 |
+
# mask vison for language
|
| 205 |
+
if attention_mask_v is not None:
|
| 206 |
+
attention_mask_v = (
|
| 207 |
+
attention_mask_v[:, None, None, :].repeat(1, self.num_heads, 1, 1).flatten(0, 1)
|
| 208 |
+
)
|
| 209 |
+
attn_weights_l.masked_fill_(attention_mask_v, float("-inf"))
|
| 210 |
+
|
| 211 |
+
attn_weights_l = attn_weights_l.softmax(dim=-1)
|
| 212 |
+
|
| 213 |
+
# mask language for vision
|
| 214 |
+
if attention_mask_l is not None:
|
| 215 |
+
attention_mask_l = (
|
| 216 |
+
attention_mask_l[:, None, None, :].repeat(1, self.num_heads, 1, 1).flatten(0, 1)
|
| 217 |
+
)
|
| 218 |
+
attn_weights.masked_fill_(attention_mask_l, float("-inf"))
|
| 219 |
+
attn_weights_v = attn_weights.softmax(dim=-1)
|
| 220 |
+
|
| 221 |
+
attn_probs_v = F.dropout(attn_weights_v, p=self.dropout, training=self.training)
|
| 222 |
+
attn_probs_l = F.dropout(attn_weights_l, p=self.dropout, training=self.training)
|
| 223 |
+
|
| 224 |
+
attn_output_v = torch.bmm(attn_probs_v, value_l_states)
|
| 225 |
+
attn_output_l = torch.bmm(attn_probs_l, value_v_states)
|
| 226 |
+
|
| 227 |
+
if attn_output_v.size() != (bsz * self.num_heads, tgt_len, self.head_dim):
|
| 228 |
+
raise ValueError(
|
| 229 |
+
f"`attn_output_v` should be of size {(bsz, self.num_heads, tgt_len, self.head_dim)}, but is {attn_output_v.size()}"
|
| 230 |
+
)
|
| 231 |
+
|
| 232 |
+
if attn_output_l.size() != (bsz * self.num_heads, src_len, self.head_dim):
|
| 233 |
+
raise ValueError(
|
| 234 |
+
f"`attn_output_l` should be of size {(bsz, self.num_heads, src_len, self.head_dim)}, but is {attn_output_l.size()}"
|
| 235 |
+
)
|
| 236 |
+
|
| 237 |
+
attn_output_v = attn_output_v.view(bsz, self.num_heads, tgt_len, self.head_dim)
|
| 238 |
+
attn_output_v = attn_output_v.transpose(1, 2)
|
| 239 |
+
attn_output_v = attn_output_v.reshape(bsz, tgt_len, self.embed_dim)
|
| 240 |
+
|
| 241 |
+
attn_output_l = attn_output_l.view(bsz, self.num_heads, src_len, self.head_dim)
|
| 242 |
+
attn_output_l = attn_output_l.transpose(1, 2)
|
| 243 |
+
attn_output_l = attn_output_l.reshape(bsz, src_len, self.embed_dim)
|
| 244 |
+
|
| 245 |
+
attn_output_v = self.out_v_proj(attn_output_v)
|
| 246 |
+
attn_output_l = self.out_l_proj(attn_output_l)
|
| 247 |
+
|
| 248 |
+
return attn_output_v, attn_output_l
|
| 249 |
+
|
| 250 |
+
|
| 251 |
+
# Bi-Direction MHA (text->image, image->text)
|
| 252 |
+
class BiAttentionBlock(nn.Module):
|
| 253 |
+
def __init__(
|
| 254 |
+
self,
|
| 255 |
+
v_dim,
|
| 256 |
+
l_dim,
|
| 257 |
+
embed_dim,
|
| 258 |
+
num_heads,
|
| 259 |
+
dropout=0.1,
|
| 260 |
+
drop_path=0.0,
|
| 261 |
+
init_values=1e-4,
|
| 262 |
+
cfg=None,
|
| 263 |
+
):
|
| 264 |
+
"""
|
| 265 |
+
Inputs:
|
| 266 |
+
embed_dim - Dimensionality of input and attention feature vectors
|
| 267 |
+
hidden_dim - Dimensionality of hidden layer in feed-forward network
|
| 268 |
+
(usually 2-4x larger than embed_dim)
|
| 269 |
+
num_heads - Number of heads to use in the Multi-Head Attention block
|
| 270 |
+
dropout - Amount of dropout to apply in the feed-forward network
|
| 271 |
+
"""
|
| 272 |
+
super(BiAttentionBlock, self).__init__()
|
| 273 |
+
|
| 274 |
+
# pre layer norm
|
| 275 |
+
self.layer_norm_v = nn.LayerNorm(v_dim)
|
| 276 |
+
self.layer_norm_l = nn.LayerNorm(l_dim)
|
| 277 |
+
self.attn = BiMultiHeadAttention(
|
| 278 |
+
v_dim=v_dim, l_dim=l_dim, embed_dim=embed_dim, num_heads=num_heads, dropout=dropout
|
| 279 |
+
)
|
| 280 |
+
|
| 281 |
+
# add layer scale for training stability
|
| 282 |
+
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
|
| 283 |
+
self.gamma_v = nn.Parameter(init_values * torch.ones((v_dim)), requires_grad=True)
|
| 284 |
+
self.gamma_l = nn.Parameter(init_values * torch.ones((l_dim)), requires_grad=True)
|
| 285 |
+
|
| 286 |
+
def forward(self, v, l, attention_mask_v=None, attention_mask_l=None):
|
| 287 |
+
v = self.layer_norm_v(v)
|
| 288 |
+
l = self.layer_norm_l(l)
|
| 289 |
+
delta_v, delta_l = self.attn(
|
| 290 |
+
v, l, attention_mask_v=attention_mask_v, attention_mask_l=attention_mask_l
|
| 291 |
+
)
|
| 292 |
+
# v, l = v + delta_v, l + delta_l
|
| 293 |
+
v = v + self.drop_path(self.gamma_v * delta_v)
|
| 294 |
+
l = l + self.drop_path(self.gamma_l * delta_l)
|
| 295 |
+
return v, l
|
| 296 |
+
|
| 297 |
+
# def forward(self, v:List[torch.Tensor], l, attention_mask_v=None, attention_mask_l=None)
|