autoregressive-diffusion-pytorch 0.2.0__tar.gz → 0.2.2__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {autoregressive_diffusion_pytorch-0.2.0 → autoregressive_diffusion_pytorch-0.2.2}/PKG-INFO +71 -3
- {autoregressive_diffusion_pytorch-0.2.0 → autoregressive_diffusion_pytorch-0.2.2}/README.md +70 -2
- autoregressive_diffusion_pytorch-0.2.2/autoregressive_diffusion_pytorch/__init__.py +15 -0
- {autoregressive_diffusion_pytorch-0.2.0 → autoregressive_diffusion_pytorch-0.2.2}/autoregressive_diffusion_pytorch/autoregressive_diffusion.py +2 -2
- autoregressive_diffusion_pytorch-0.2.2/autoregressive_diffusion_pytorch/autoregressive_flow.py +335 -0
- autoregressive_diffusion_pytorch-0.2.2/autoregressive_diffusion_pytorch/image_trainer.py +194 -0
- autoregressive_diffusion_pytorch-0.2.2/images/results.96600.png +0 -0
- {autoregressive_diffusion_pytorch-0.2.0 → autoregressive_diffusion_pytorch-0.2.2}/pyproject.toml +1 -1
- autoregressive_diffusion_pytorch-0.2.0/autoregressive_diffusion_pytorch/__init__.py +0 -5
- autoregressive_diffusion_pytorch-0.2.0/images/sample.flowers.59000.png +0 -0
- {autoregressive_diffusion_pytorch-0.2.0 → autoregressive_diffusion_pytorch-0.2.2}/.github/workflows/python-publish.yml +0 -0
- {autoregressive_diffusion_pytorch-0.2.0 → autoregressive_diffusion_pytorch-0.2.2}/.gitignore +0 -0
- {autoregressive_diffusion_pytorch-0.2.0 → autoregressive_diffusion_pytorch-0.2.2}/LICENSE +0 -0
- {autoregressive_diffusion_pytorch-0.2.0 → autoregressive_diffusion_pytorch-0.2.2}/ar-diffusion.png +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: autoregressive-diffusion-pytorch
|
|
3
|
-
Version: 0.2.
|
|
3
|
+
Version: 0.2.2
|
|
4
4
|
Summary: Autoregressive Diffusion - Pytorch
|
|
5
5
|
Project-URL: Homepage, https://pypi.org/project/autoregressive-diffusion-pytorch/
|
|
6
6
|
Project-URL: Repository, https://github.com/lucidrains/autoregressive-diffusion-pytorch
|
|
@@ -52,9 +52,9 @@ Implementation of the architecture behind <a href="https://arxiv.org/abs/2406.11
|
|
|
52
52
|
|
|
53
53
|
Official repository has been released <a href="https://github.com/LTH14/mar">here</a>
|
|
54
54
|
|
|
55
|
-
<img src="./images/
|
|
55
|
+
<img src="./images/results.96600.png" width="400px"></img>
|
|
56
56
|
|
|
57
|
-
*oxford flowers at
|
|
57
|
+
*oxford flowers at 96k steps*
|
|
58
58
|
|
|
59
59
|
## Install
|
|
60
60
|
|
|
@@ -115,6 +115,74 @@ assert sampled.shape == images.shape
|
|
|
115
115
|
|
|
116
116
|
```
|
|
117
117
|
|
|
118
|
+
An images trainer
|
|
119
|
+
|
|
120
|
+
```python
|
|
121
|
+
import torch
|
|
122
|
+
|
|
123
|
+
from autoregressive_diffusion_pytorch import (
|
|
124
|
+
ImageDataset,
|
|
125
|
+
ImageAutoregressiveDiffusion,
|
|
126
|
+
ImageTrainer
|
|
127
|
+
)
|
|
128
|
+
|
|
129
|
+
dataset = ImageDataset(
|
|
130
|
+
'/path/to/your/images',
|
|
131
|
+
image_size = 128
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
model = ImageAutoregressiveDiffusion(
|
|
135
|
+
model = dict(
|
|
136
|
+
dim = 512
|
|
137
|
+
),
|
|
138
|
+
image_size = 128,
|
|
139
|
+
patch_size = 16
|
|
140
|
+
)
|
|
141
|
+
|
|
142
|
+
trainer = ImageTrainer(
|
|
143
|
+
model = model,
|
|
144
|
+
dataset = dataset
|
|
145
|
+
)
|
|
146
|
+
|
|
147
|
+
trainer()
|
|
148
|
+
```
|
|
149
|
+
|
|
150
|
+
For an improvised version using flow matching, just import `ImageAutoregressiveFlow` and `AutoregressiveFlow` instead
|
|
151
|
+
|
|
152
|
+
The rest is the same
|
|
153
|
+
|
|
154
|
+
ex.
|
|
155
|
+
|
|
156
|
+
```python
|
|
157
|
+
import torch
|
|
158
|
+
|
|
159
|
+
from autoregressive_diffusion_pytorch import (
|
|
160
|
+
ImageDataset,
|
|
161
|
+
ImageTrainer,
|
|
162
|
+
ImageAutoregressiveFlow,
|
|
163
|
+
)
|
|
164
|
+
|
|
165
|
+
dataset = ImageDataset(
|
|
166
|
+
'/path/to/your/images',
|
|
167
|
+
image_size = 128
|
|
168
|
+
)
|
|
169
|
+
|
|
170
|
+
model = ImageAutoregressiveFlow(
|
|
171
|
+
model = dict(
|
|
172
|
+
dim = 512
|
|
173
|
+
),
|
|
174
|
+
image_size = 128,
|
|
175
|
+
patch_size = 16
|
|
176
|
+
)
|
|
177
|
+
|
|
178
|
+
trainer = ImageTrainer(
|
|
179
|
+
model = model,
|
|
180
|
+
dataset = dataset
|
|
181
|
+
)
|
|
182
|
+
|
|
183
|
+
trainer()
|
|
184
|
+
```
|
|
185
|
+
|
|
118
186
|
## Citations
|
|
119
187
|
|
|
120
188
|
```bibtex
|
|
@@ -6,9 +6,9 @@ Implementation of the architecture behind <a href="https://arxiv.org/abs/2406.11
|
|
|
6
6
|
|
|
7
7
|
Official repository has been released <a href="https://github.com/LTH14/mar">here</a>
|
|
8
8
|
|
|
9
|
-
<img src="./images/
|
|
9
|
+
<img src="./images/results.96600.png" width="400px"></img>
|
|
10
10
|
|
|
11
|
-
*oxford flowers at
|
|
11
|
+
*oxford flowers at 96k steps*
|
|
12
12
|
|
|
13
13
|
## Install
|
|
14
14
|
|
|
@@ -69,6 +69,74 @@ assert sampled.shape == images.shape
|
|
|
69
69
|
|
|
70
70
|
```
|
|
71
71
|
|
|
72
|
+
An images trainer
|
|
73
|
+
|
|
74
|
+
```python
|
|
75
|
+
import torch
|
|
76
|
+
|
|
77
|
+
from autoregressive_diffusion_pytorch import (
|
|
78
|
+
ImageDataset,
|
|
79
|
+
ImageAutoregressiveDiffusion,
|
|
80
|
+
ImageTrainer
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
dataset = ImageDataset(
|
|
84
|
+
'/path/to/your/images',
|
|
85
|
+
image_size = 128
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
model = ImageAutoregressiveDiffusion(
|
|
89
|
+
model = dict(
|
|
90
|
+
dim = 512
|
|
91
|
+
),
|
|
92
|
+
image_size = 128,
|
|
93
|
+
patch_size = 16
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
trainer = ImageTrainer(
|
|
97
|
+
model = model,
|
|
98
|
+
dataset = dataset
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
trainer()
|
|
102
|
+
```
|
|
103
|
+
|
|
104
|
+
For an improvised version using flow matching, just import `ImageAutoregressiveFlow` and `AutoregressiveFlow` instead
|
|
105
|
+
|
|
106
|
+
The rest is the same
|
|
107
|
+
|
|
108
|
+
ex.
|
|
109
|
+
|
|
110
|
+
```python
|
|
111
|
+
import torch
|
|
112
|
+
|
|
113
|
+
from autoregressive_diffusion_pytorch import (
|
|
114
|
+
ImageDataset,
|
|
115
|
+
ImageTrainer,
|
|
116
|
+
ImageAutoregressiveFlow,
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
dataset = ImageDataset(
|
|
120
|
+
'/path/to/your/images',
|
|
121
|
+
image_size = 128
|
|
122
|
+
)
|
|
123
|
+
|
|
124
|
+
model = ImageAutoregressiveFlow(
|
|
125
|
+
model = dict(
|
|
126
|
+
dim = 512
|
|
127
|
+
),
|
|
128
|
+
image_size = 128,
|
|
129
|
+
patch_size = 16
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
trainer = ImageTrainer(
|
|
133
|
+
model = model,
|
|
134
|
+
dataset = dataset
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
trainer()
|
|
138
|
+
```
|
|
139
|
+
|
|
72
140
|
## Citations
|
|
73
141
|
|
|
74
142
|
```bibtex
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
from autoregressive_diffusion_pytorch.autoregressive_diffusion import (
|
|
2
|
+
MLP,
|
|
3
|
+
AutoregressiveDiffusion,
|
|
4
|
+
ImageAutoregressiveDiffusion
|
|
5
|
+
)
|
|
6
|
+
|
|
7
|
+
from autoregressive_diffusion_pytorch.autoregressive_flow import (
|
|
8
|
+
AutoregressiveFlow,
|
|
9
|
+
ImageAutoregressiveFlow
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
from autoregressive_diffusion_pytorch.image_trainer import (
|
|
13
|
+
ImageDataset,
|
|
14
|
+
ImageTrainer
|
|
15
|
+
)
|
|
@@ -231,7 +231,7 @@ class ElucidatedDiffusion(Module):
|
|
|
231
231
|
if isinstance(sigma, float):
|
|
232
232
|
sigma = torch.full((batch,), sigma, device = device)
|
|
233
233
|
|
|
234
|
-
padded_sigma =
|
|
234
|
+
padded_sigma = right_pad_dims_to(noised_seq, sigma)
|
|
235
235
|
|
|
236
236
|
net_out = self.net(
|
|
237
237
|
self.c_in(padded_sigma) * noised_seq,
|
|
@@ -331,7 +331,7 @@ class ElucidatedDiffusion(Module):
|
|
|
331
331
|
assert dim == self.dim, f'dimension of sequence being passed in must be {self.dim} but received {dim}'
|
|
332
332
|
|
|
333
333
|
sigmas = self.noise_distribution(batch_size)
|
|
334
|
-
padded_sigmas =
|
|
334
|
+
padded_sigmas = right_pad_dims_to(seq, sigmas)
|
|
335
335
|
|
|
336
336
|
noise = torch.randn_like(seq)
|
|
337
337
|
|
autoregressive_diffusion_pytorch-0.2.2/autoregressive_diffusion_pytorch/autoregressive_flow.py
ADDED
|
@@ -0,0 +1,335 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
from math import sqrt
|
|
5
|
+
from typing import Literal
|
|
6
|
+
from functools import partial
|
|
7
|
+
|
|
8
|
+
import torch
|
|
9
|
+
from torch import nn, pi
|
|
10
|
+
import torch.nn.functional as F
|
|
11
|
+
from torch.nn import Module, ModuleList
|
|
12
|
+
|
|
13
|
+
from torchdiffeq import odeint
|
|
14
|
+
|
|
15
|
+
import einx
|
|
16
|
+
from einops import rearrange, repeat, reduce, pack, unpack
|
|
17
|
+
from einops.layers.torch import Rearrange
|
|
18
|
+
|
|
19
|
+
from tqdm import tqdm
|
|
20
|
+
|
|
21
|
+
from x_transformers import Decoder
|
|
22
|
+
|
|
23
|
+
from autoregressive_diffusion_pytorch.autoregressive_diffusion import MLP
|
|
24
|
+
|
|
25
|
+
# helpers
|
|
26
|
+
|
|
27
|
+
def exists(v):
|
|
28
|
+
return v is not None
|
|
29
|
+
|
|
30
|
+
def default(v, d):
|
|
31
|
+
return v if exists(v) else d
|
|
32
|
+
|
|
33
|
+
def divisible_by(num, den):
|
|
34
|
+
return (num % den) == 0
|
|
35
|
+
|
|
36
|
+
# tensor helpers
|
|
37
|
+
|
|
38
|
+
def log(t, eps = 1e-20):
|
|
39
|
+
return torch.log(t.clamp(min = eps))
|
|
40
|
+
|
|
41
|
+
def safe_div(num, den, eps = 1e-5):
|
|
42
|
+
return num / den.clamp(min = eps)
|
|
43
|
+
|
|
44
|
+
def right_pad_dims_to(x, t):
|
|
45
|
+
padding_dims = x.ndim - t.ndim
|
|
46
|
+
|
|
47
|
+
if padding_dims <= 0:
|
|
48
|
+
return t
|
|
49
|
+
|
|
50
|
+
return t.view(*t.shape, *((1,) * padding_dims))
|
|
51
|
+
|
|
52
|
+
def pack_one(t, pattern):
|
|
53
|
+
packed, ps = pack([t], pattern)
|
|
54
|
+
|
|
55
|
+
def unpack_one(to_unpack, unpack_pattern = None):
|
|
56
|
+
unpacked, = unpack(to_unpack, ps, default(unpack_pattern, pattern))
|
|
57
|
+
return unpacked
|
|
58
|
+
|
|
59
|
+
return packed, unpack_one
|
|
60
|
+
|
|
61
|
+
# sinusoidal embedding
|
|
62
|
+
|
|
63
|
+
class AdaptiveLayerNorm(Module):
|
|
64
|
+
def __init__(
|
|
65
|
+
self,
|
|
66
|
+
dim,
|
|
67
|
+
dim_condition = None
|
|
68
|
+
):
|
|
69
|
+
super().__init__()
|
|
70
|
+
dim_condition = default(dim_condition, dim)
|
|
71
|
+
|
|
72
|
+
self.ln = nn.LayerNorm(dim, elementwise_affine = False)
|
|
73
|
+
self.to_gamma = nn.Linear(dim_condition, dim, bias = False)
|
|
74
|
+
nn.init.zeros_(self.to_gamma.weight)
|
|
75
|
+
|
|
76
|
+
def forward(self, x, *, condition):
|
|
77
|
+
normed = self.ln(x)
|
|
78
|
+
gamma = self.to_gamma(condition)
|
|
79
|
+
return normed * (gamma + 1.)
|
|
80
|
+
|
|
81
|
+
class LearnedSinusoidalPosEmb(Module):
|
|
82
|
+
def __init__(self, dim):
|
|
83
|
+
super().__init__()
|
|
84
|
+
assert divisible_by(dim, 2)
|
|
85
|
+
half_dim = dim // 2
|
|
86
|
+
self.weights = nn.Parameter(torch.randn(half_dim))
|
|
87
|
+
|
|
88
|
+
def forward(self, x):
|
|
89
|
+
x = rearrange(x, 'b -> b 1')
|
|
90
|
+
freqs = x * rearrange(self.weights, 'd -> 1 d') * 2 * pi
|
|
91
|
+
fouriered = torch.cat((freqs.sin(), freqs.cos()), dim = -1)
|
|
92
|
+
fouriered = torch.cat((x, fouriered), dim = -1)
|
|
93
|
+
return fouriered
|
|
94
|
+
|
|
95
|
+
# gaussian diffusion
|
|
96
|
+
|
|
97
|
+
class Flow(Module):
|
|
98
|
+
def __init__(
|
|
99
|
+
self,
|
|
100
|
+
dim: int,
|
|
101
|
+
net: MLP,
|
|
102
|
+
*,
|
|
103
|
+
atol = 1e-5,
|
|
104
|
+
rtol = 1e-5,
|
|
105
|
+
method = 'midpoint'
|
|
106
|
+
):
|
|
107
|
+
super().__init__()
|
|
108
|
+
self.net = net
|
|
109
|
+
self.dim = dim
|
|
110
|
+
|
|
111
|
+
self.odeint_kwargs = dict(
|
|
112
|
+
atol = atol,
|
|
113
|
+
rtol = rtol,
|
|
114
|
+
method = method
|
|
115
|
+
)
|
|
116
|
+
|
|
117
|
+
@property
|
|
118
|
+
def device(self):
|
|
119
|
+
return next(self.net.parameters()).device
|
|
120
|
+
|
|
121
|
+
@torch.no_grad()
|
|
122
|
+
def sample(
|
|
123
|
+
self,
|
|
124
|
+
cond,
|
|
125
|
+
num_sample_steps = 16
|
|
126
|
+
):
|
|
127
|
+
|
|
128
|
+
batch = cond.shape[0]
|
|
129
|
+
|
|
130
|
+
sampled_data_shape = (batch, self.dim)
|
|
131
|
+
|
|
132
|
+
# start with random gaussian noise - y0
|
|
133
|
+
|
|
134
|
+
noise = torch.randn(sampled_data_shape, device = self.device)
|
|
135
|
+
|
|
136
|
+
# time steps
|
|
137
|
+
|
|
138
|
+
times = torch.linspace(0., 1., num_sample_steps, device = self.device)
|
|
139
|
+
|
|
140
|
+
# ode
|
|
141
|
+
|
|
142
|
+
def ode_fn(t, x):
|
|
143
|
+
t = repeat(t, '-> b', b = batch)
|
|
144
|
+
flow = self.net(x, times = t, cond = cond)
|
|
145
|
+
return flow
|
|
146
|
+
|
|
147
|
+
trajectory = odeint(ode_fn, noise, times, **self.odeint_kwargs)
|
|
148
|
+
|
|
149
|
+
sampled = trajectory[-1]
|
|
150
|
+
|
|
151
|
+
return sampled
|
|
152
|
+
|
|
153
|
+
# training
|
|
154
|
+
|
|
155
|
+
def forward(self, seq, *, cond):
|
|
156
|
+
batch_size, dim, device = *seq.shape, self.device
|
|
157
|
+
|
|
158
|
+
assert dim == self.dim, f'dimension of sequence being passed in must be {self.dim} but received {dim}'
|
|
159
|
+
|
|
160
|
+
times = torch.rand(batch_size, device = device)
|
|
161
|
+
noise = torch.randn_like(seq)
|
|
162
|
+
padded_times = right_pad_dims_to(seq, times)
|
|
163
|
+
|
|
164
|
+
flow = seq - noise
|
|
165
|
+
|
|
166
|
+
noised = (1.- padded_times) * noise + padded_times * seq
|
|
167
|
+
|
|
168
|
+
pred_flow = self.net(noised, times = times, cond = cond)
|
|
169
|
+
|
|
170
|
+
return F.mse_loss(pred_flow, flow)
|
|
171
|
+
|
|
172
|
+
# main model, a decoder with continuous wrapper + small denoising mlp
|
|
173
|
+
|
|
174
|
+
class AutoregressiveFlow(Module):
|
|
175
|
+
def __init__(
|
|
176
|
+
self,
|
|
177
|
+
dim,
|
|
178
|
+
*,
|
|
179
|
+
max_seq_len,
|
|
180
|
+
depth = 8,
|
|
181
|
+
dim_head = 64,
|
|
182
|
+
heads = 8,
|
|
183
|
+
mlp_depth = 3,
|
|
184
|
+
mlp_width = None,
|
|
185
|
+
dim_input = None,
|
|
186
|
+
decoder_kwargs: dict = dict(),
|
|
187
|
+
mlp_kwargs: dict = dict(),
|
|
188
|
+
flow_kwargs: dict = dict()
|
|
189
|
+
):
|
|
190
|
+
super().__init__()
|
|
191
|
+
|
|
192
|
+
self.start_token = nn.Parameter(torch.zeros(dim))
|
|
193
|
+
self.max_seq_len = max_seq_len
|
|
194
|
+
self.abs_pos_emb = nn.Embedding(max_seq_len, dim)
|
|
195
|
+
|
|
196
|
+
dim_input = default(dim_input, dim)
|
|
197
|
+
self.dim_input = dim_input
|
|
198
|
+
self.proj_in = nn.Linear(dim_input, dim)
|
|
199
|
+
|
|
200
|
+
self.transformer = Decoder(
|
|
201
|
+
dim = dim,
|
|
202
|
+
depth = depth,
|
|
203
|
+
heads = heads,
|
|
204
|
+
attn_dim_head = dim_head,
|
|
205
|
+
**decoder_kwargs
|
|
206
|
+
)
|
|
207
|
+
|
|
208
|
+
self.denoiser = MLP(
|
|
209
|
+
dim_cond = dim,
|
|
210
|
+
dim_input = dim_input,
|
|
211
|
+
depth = mlp_depth,
|
|
212
|
+
width = default(mlp_width, dim),
|
|
213
|
+
**mlp_kwargs
|
|
214
|
+
)
|
|
215
|
+
|
|
216
|
+
self.flow = Flow(
|
|
217
|
+
dim_input,
|
|
218
|
+
self.denoiser,
|
|
219
|
+
**flow_kwargs
|
|
220
|
+
)
|
|
221
|
+
|
|
222
|
+
@property
|
|
223
|
+
def device(self):
|
|
224
|
+
return next(self.transformer.parameters()).device
|
|
225
|
+
|
|
226
|
+
@torch.no_grad()
|
|
227
|
+
def sample(
|
|
228
|
+
self,
|
|
229
|
+
batch_size = 1,
|
|
230
|
+
prompt = None
|
|
231
|
+
):
|
|
232
|
+
self.eval()
|
|
233
|
+
|
|
234
|
+
start_tokens = repeat(self.start_token, 'd -> b 1 d', b = batch_size)
|
|
235
|
+
|
|
236
|
+
if not exists(prompt):
|
|
237
|
+
out = torch.empty((batch_size, 0, self.dim_input), device = self.device, dtype = torch.float32)
|
|
238
|
+
else:
|
|
239
|
+
out = prompt
|
|
240
|
+
|
|
241
|
+
cache = None
|
|
242
|
+
|
|
243
|
+
for _ in tqdm(range(self.max_seq_len - out.shape[1]), desc = 'tokens'):
|
|
244
|
+
|
|
245
|
+
cond = self.proj_in(out)
|
|
246
|
+
|
|
247
|
+
cond = torch.cat((start_tokens, cond), dim = 1)
|
|
248
|
+
cond = cond + self.abs_pos_emb(torch.arange(cond.shape[1], device = self.device))
|
|
249
|
+
|
|
250
|
+
cond, cache = self.transformer(cond, cache = cache, return_hiddens = True)
|
|
251
|
+
|
|
252
|
+
last_cond = cond[:, -1]
|
|
253
|
+
|
|
254
|
+
denoised_pred = self.flow.sample(cond = last_cond)
|
|
255
|
+
|
|
256
|
+
denoised_pred = rearrange(denoised_pred, 'b d -> b 1 d')
|
|
257
|
+
out = torch.cat((out, denoised_pred), dim = 1)
|
|
258
|
+
|
|
259
|
+
return out
|
|
260
|
+
|
|
261
|
+
def forward(
|
|
262
|
+
self,
|
|
263
|
+
seq
|
|
264
|
+
):
|
|
265
|
+
b, seq_len, dim = seq.shape
|
|
266
|
+
|
|
267
|
+
assert dim == self.dim_input
|
|
268
|
+
assert seq_len == self.max_seq_len
|
|
269
|
+
|
|
270
|
+
# break into seq and the continuous targets to be predicted
|
|
271
|
+
|
|
272
|
+
seq, target = seq[:, :-1], seq
|
|
273
|
+
|
|
274
|
+
# append start tokens
|
|
275
|
+
|
|
276
|
+
seq = self.proj_in(seq)
|
|
277
|
+
start_token = repeat(self.start_token, 'd -> b 1 d', b = b)
|
|
278
|
+
|
|
279
|
+
seq = torch.cat((start_token, seq), dim = 1)
|
|
280
|
+
seq = seq + self.abs_pos_emb(torch.arange(seq_len, device = self.device))
|
|
281
|
+
|
|
282
|
+
cond = self.transformer(seq)
|
|
283
|
+
|
|
284
|
+
# pack batch and sequence dimensions, so to train each token with different noise levels
|
|
285
|
+
|
|
286
|
+
target, _ = pack_one(target, '* d')
|
|
287
|
+
cond, _ = pack_one(cond, '* d')
|
|
288
|
+
|
|
289
|
+
return self.flow(target, cond = cond)
|
|
290
|
+
|
|
291
|
+
# image wrapper
|
|
292
|
+
|
|
293
|
+
def normalize_to_neg_one_to_one(img):
|
|
294
|
+
return img * 2 - 1
|
|
295
|
+
|
|
296
|
+
def unnormalize_to_zero_to_one(t):
|
|
297
|
+
return (t + 1) * 0.5
|
|
298
|
+
|
|
299
|
+
class ImageAutoregressiveFlow(Module):
|
|
300
|
+
def __init__(
|
|
301
|
+
self,
|
|
302
|
+
*,
|
|
303
|
+
image_size,
|
|
304
|
+
patch_size,
|
|
305
|
+
channels = 3,
|
|
306
|
+
model: dict = dict(),
|
|
307
|
+
):
|
|
308
|
+
super().__init__()
|
|
309
|
+
assert divisible_by(image_size, patch_size)
|
|
310
|
+
|
|
311
|
+
num_patches = (image_size // patch_size) ** 2
|
|
312
|
+
dim_in = channels * patch_size ** 2
|
|
313
|
+
|
|
314
|
+
self.image_size = image_size
|
|
315
|
+
self.patch_size = patch_size
|
|
316
|
+
|
|
317
|
+
self.to_tokens = Rearrange('b c (h p1) (w p2) -> b (h w) (c p1 p2)', p1 = patch_size, p2 = patch_size)
|
|
318
|
+
|
|
319
|
+
self.model = AutoregressiveFlow(
|
|
320
|
+
**model,
|
|
321
|
+
dim_input = dim_in,
|
|
322
|
+
max_seq_len = num_patches
|
|
323
|
+
)
|
|
324
|
+
|
|
325
|
+
self.to_image = Rearrange('b (h w) (c p1 p2) -> b c (h p1) (w p2)', p1 = patch_size, p2 = patch_size, h = int(math.sqrt(num_patches)))
|
|
326
|
+
|
|
327
|
+
def sample(self, batch_size = 1):
|
|
328
|
+
tokens = self.model.sample(batch_size = batch_size)
|
|
329
|
+
images = self.to_image(tokens)
|
|
330
|
+
return unnormalize_to_zero_to_one(images)
|
|
331
|
+
|
|
332
|
+
def forward(self, images):
|
|
333
|
+
images = normalize_to_neg_one_to_one(images)
|
|
334
|
+
tokens = self.to_tokens(images)
|
|
335
|
+
return self.model(tokens)
|
|
@@ -0,0 +1,194 @@
|
|
|
1
|
+
from typing import List
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
|
|
6
|
+
from accelerate import Accelerator
|
|
7
|
+
from ema_pytorch import EMA
|
|
8
|
+
|
|
9
|
+
import torch
|
|
10
|
+
from torch import nn
|
|
11
|
+
from torch.optim import Adam
|
|
12
|
+
from torch.utils.data import DataLoader
|
|
13
|
+
from torch.nn import Module, ModuleList
|
|
14
|
+
from torch.utils.data import Dataset
|
|
15
|
+
|
|
16
|
+
from torchvision.utils import save_image
|
|
17
|
+
import torchvision.transforms as T
|
|
18
|
+
|
|
19
|
+
from PIL import Image
|
|
20
|
+
|
|
21
|
+
# functions
|
|
22
|
+
|
|
23
|
+
def exists(v):
|
|
24
|
+
return v is not None
|
|
25
|
+
|
|
26
|
+
def default(v, d):
|
|
27
|
+
return v if exists(v) else d
|
|
28
|
+
|
|
29
|
+
def divisible_by(num, den):
|
|
30
|
+
return (num % den) == 0
|
|
31
|
+
|
|
32
|
+
def cycle(dl):
|
|
33
|
+
while True:
|
|
34
|
+
for batch in dl:
|
|
35
|
+
yield batch
|
|
36
|
+
|
|
37
|
+
# dataset classes
|
|
38
|
+
|
|
39
|
+
class ImageDataset(Dataset):
|
|
40
|
+
def __init__(
|
|
41
|
+
self,
|
|
42
|
+
folder: str | Path,
|
|
43
|
+
image_size: int,
|
|
44
|
+
exts: List[str] = ['jpg', 'jpeg', 'png', 'tiff'],
|
|
45
|
+
augment_horizontal_flip = False,
|
|
46
|
+
convert_image_to = None
|
|
47
|
+
):
|
|
48
|
+
super().__init__()
|
|
49
|
+
if isinstance(folder, str):
|
|
50
|
+
folder = Path(folder)
|
|
51
|
+
|
|
52
|
+
assert folder.is_dir()
|
|
53
|
+
|
|
54
|
+
self.folder = folder
|
|
55
|
+
self.image_size = image_size
|
|
56
|
+
|
|
57
|
+
self.paths = [p for ext in exts for p in folder.glob(f'**/*.{ext}')]
|
|
58
|
+
|
|
59
|
+
def convert_image_to_fn(img_type, image):
|
|
60
|
+
if image.mode == img_type:
|
|
61
|
+
return image
|
|
62
|
+
|
|
63
|
+
return image.convert(img_type)
|
|
64
|
+
|
|
65
|
+
maybe_convert_fn = partial(convert_image_to_fn, convert_image_to) if exists(convert_image_to) else nn.Identity()
|
|
66
|
+
|
|
67
|
+
self.transform = T.Compose([
|
|
68
|
+
T.Lambda(maybe_convert_fn),
|
|
69
|
+
T.Resize(image_size),
|
|
70
|
+
T.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(),
|
|
71
|
+
T.CenterCrop(image_size),
|
|
72
|
+
T.ToTensor()
|
|
73
|
+
])
|
|
74
|
+
|
|
75
|
+
def __len__(self):
|
|
76
|
+
return len(self.paths)
|
|
77
|
+
|
|
78
|
+
def __getitem__(self, index):
|
|
79
|
+
path = self.paths[index]
|
|
80
|
+
img = Image.open(path)
|
|
81
|
+
return self.transform(img)
|
|
82
|
+
|
|
83
|
+
# trainer
|
|
84
|
+
|
|
85
|
+
class ImageTrainer(Module):
|
|
86
|
+
def __init__(
|
|
87
|
+
self,
|
|
88
|
+
model,
|
|
89
|
+
*,
|
|
90
|
+
dataset: Dataset,
|
|
91
|
+
num_train_steps = 70_000,
|
|
92
|
+
learning_rate = 3e-4,
|
|
93
|
+
batch_size = 16,
|
|
94
|
+
checkpoints_folder: str = './checkpoints',
|
|
95
|
+
results_folder: str = './results',
|
|
96
|
+
save_results_every: int = 100,
|
|
97
|
+
checkpoint_every: int = 1000,
|
|
98
|
+
num_samples: int = 16,
|
|
99
|
+
adam_kwargs: dict = dict(),
|
|
100
|
+
accelerate_kwargs: dict = dict(),
|
|
101
|
+
ema_kwargs: dict = dict()
|
|
102
|
+
):
|
|
103
|
+
super().__init__()
|
|
104
|
+
self.accelerator = Accelerator(**accelerate_kwargs)
|
|
105
|
+
|
|
106
|
+
self.model = model
|
|
107
|
+
|
|
108
|
+
if self.is_main:
|
|
109
|
+
self.ema_model = EMA(
|
|
110
|
+
self.model,
|
|
111
|
+
forward_method_names = ('sample',),
|
|
112
|
+
**ema_kwargs
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
self.ema_model.to(self.accelerator.device)
|
|
116
|
+
|
|
117
|
+
self.optimizer = Adam(model.parameters(), lr = learning_rate, **adam_kwargs)
|
|
118
|
+
self.dl = DataLoader(dataset, batch_size = batch_size, shuffle = True, drop_last = True)
|
|
119
|
+
|
|
120
|
+
self.model, self.optimizer, self.dl = self.accelerator.prepare(self.model, self.optimizer, self.dl)
|
|
121
|
+
|
|
122
|
+
self.num_train_steps = num_train_steps
|
|
123
|
+
|
|
124
|
+
self.checkpoints_folder = Path(checkpoints_folder)
|
|
125
|
+
self.results_folder = Path(results_folder)
|
|
126
|
+
|
|
127
|
+
self.checkpoints_folder.mkdir(exist_ok = True, parents = True)
|
|
128
|
+
self.results_folder.mkdir(exist_ok = True, parents = True)
|
|
129
|
+
|
|
130
|
+
self.checkpoint_every = checkpoint_every
|
|
131
|
+
self.save_results_every = save_results_every
|
|
132
|
+
|
|
133
|
+
self.num_sample_rows = int(math.sqrt(num_samples))
|
|
134
|
+
assert (self.num_sample_rows ** 2) == num_samples, f'{num_samples} must be a square'
|
|
135
|
+
self.num_samples = num_samples
|
|
136
|
+
|
|
137
|
+
assert self.checkpoints_folder.is_dir()
|
|
138
|
+
assert self.results_folder.is_dir()
|
|
139
|
+
|
|
140
|
+
@property
|
|
141
|
+
def is_main(self):
|
|
142
|
+
return self.accelerator.is_main_process
|
|
143
|
+
|
|
144
|
+
def save(self, path):
|
|
145
|
+
if not self.is_main:
|
|
146
|
+
return
|
|
147
|
+
|
|
148
|
+
save_package = dict(
|
|
149
|
+
model = self.accelerator.unwrap_model(self.model).state_dict(),
|
|
150
|
+
ema_model = self.ema_model.state_dict(),
|
|
151
|
+
optimizer = self.accelerator.unwrap_model(self.optimizer).state_dict(),
|
|
152
|
+
)
|
|
153
|
+
|
|
154
|
+
torch.save(save_package, str(self.checkpoints_folder / path))
|
|
155
|
+
|
|
156
|
+
def forward(self):
|
|
157
|
+
|
|
158
|
+
dl = cycle(self.dl)
|
|
159
|
+
|
|
160
|
+
for ind in range(self.num_train_steps):
|
|
161
|
+
step = ind + 1
|
|
162
|
+
|
|
163
|
+
self.model.train()
|
|
164
|
+
|
|
165
|
+
data = next(dl)
|
|
166
|
+
loss = self.model(data)
|
|
167
|
+
|
|
168
|
+
self.accelerator.print(f'[{step}] loss: {loss.item():.3f}')
|
|
169
|
+
self.accelerator.backward(loss)
|
|
170
|
+
|
|
171
|
+
self.optimizer.step()
|
|
172
|
+
self.optimizer.zero_grad()
|
|
173
|
+
|
|
174
|
+
if self.is_main:
|
|
175
|
+
self.ema_model.update()
|
|
176
|
+
|
|
177
|
+
self.accelerator.wait_for_everyone()
|
|
178
|
+
|
|
179
|
+
if self.is_main:
|
|
180
|
+
if divisible_by(step, self.save_results_every):
|
|
181
|
+
|
|
182
|
+
with torch.no_grad():
|
|
183
|
+
sampled = self.ema_model.sample(batch_size = self.num_samples)
|
|
184
|
+
|
|
185
|
+
sampled.clamp_(0., 1.)
|
|
186
|
+
save_image(sampled, str(self.results_folder / f'results.{step}.png'), nrow = self.num_sample_rows)
|
|
187
|
+
|
|
188
|
+
if divisible_by(step, self.checkpoint_every):
|
|
189
|
+
self.save(f'checkpoint.{step}.pt')
|
|
190
|
+
|
|
191
|
+
self.accelerator.wait_for_everyone()
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
print('training complete')
|
|
Binary file
|
|
Binary file
|
|
File without changes
|
{autoregressive_diffusion_pytorch-0.2.0 → autoregressive_diffusion_pytorch-0.2.2}/.gitignore
RENAMED
|
File without changes
|
|
File without changes
|
{autoregressive_diffusion_pytorch-0.2.0 → autoregressive_diffusion_pytorch-0.2.2}/ar-diffusion.png
RENAMED
|
File without changes
|