BlueSourceJY commited on
Commit
e6a2ef1
·
verified ·
1 Parent(s): 57e491d

Upload models/lightningdit_rot.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. models/lightningdit_rot.py +665 -0
models/lightningdit_rot.py ADDED
@@ -0,0 +1,665 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Lightning DiT's codes are built from original DiT & SiT.
3
+ (https://github.com/facebookresearch/DiT; https://github.com/willisma/SiT)
4
+ It demonstrates that a advanced DiT together with advanced diffusion skills
5
+ could also achieve a very promising result with 1.35 FID on ImageNet 256 generation.
6
+
7
+ Enjoy everyone, DiT strikes back!
8
+
9
+ by Maple (Jingfeng Yao) from HUST-VL
10
+ """
11
+
12
+ import os
13
+ import math
14
+ import numpy as np
15
+
16
+ import torch
17
+ import torch.nn as nn
18
+ import torch.nn.functional as F
19
+ from torch.utils.checkpoint import checkpoint
20
+
21
+ from timm.models.vision_transformer import PatchEmbed, Mlp
22
+ from models.swiglu_ffn import SwiGLUFFN
23
+ from models.pos_embed import VisionRotaryEmbeddingFast
24
+ from models.rmsnorm import RMSNorm
25
+ from visualize_attention import visualize_attention_matrix
26
+
27
+ def rot90(x):
28
+ # print("x.shape:", x.shape) #(128,256,1152)
29
+ B, N, C = x.shape
30
+ H = W = int(N ** 0.5)
31
+ x = x.reshape(B, H, W, C)
32
+ x = torch.rot90(x, k=1, dims=(2,1))
33
+ x = x.reshape(B, N, C)
34
+ return x
35
+
36
+ def rot180(x):
37
+ # print("x.shape:", x.shape) #(128,256,1152)
38
+ B, N, C = x.shape
39
+ H = W = int(N ** 0.5)
40
+ x = x.reshape(B, H, W, C)
41
+ x = torch.rot90(x, k=2, dims=(2,1))
42
+ x = x.reshape(B, N, C)
43
+ return x
44
+
45
+
46
+ @torch.compile
47
+ def modulate(x, shift, scale):
48
+ if shift is None:
49
+ return x * (1 + scale.unsqueeze(1))
50
+ return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
51
+
52
+ class Attention(nn.Module):
53
+ """
54
+ Attention module of LightningDiT.
55
+ """
56
+ def __init__(
57
+ self,
58
+ dim: int,
59
+ num_heads: int = 8,
60
+ qkv_bias: bool = False,
61
+ qk_norm: bool = False,
62
+ attn_drop: float = 0.,
63
+ proj_drop: float = 0.,
64
+ norm_layer: nn.Module = nn.LayerNorm,
65
+ fused_attn: bool = True,
66
+ use_rmsnorm: bool = False,
67
+ is_causal: bool = False
68
+ ) -> None:
69
+ super().__init__()
70
+ assert dim % num_heads == 0, 'dim should be divisible by num_heads'
71
+
72
+ self.num_heads = num_heads
73
+ self.head_dim = dim // num_heads
74
+ self.scale = self.head_dim ** -0.5
75
+ self.fused_attn = fused_attn
76
+
77
+ if use_rmsnorm:
78
+ norm_layer = RMSNorm
79
+
80
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
81
+ self.q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
82
+ self.k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
83
+ self.attn_drop = nn.Dropout(attn_drop)
84
+ self.proj = nn.Linear(dim, dim)
85
+ self.proj_drop = nn.Dropout(proj_drop)
86
+
87
+ self.is_causal = is_causal
88
+
89
+ def forward(self, x: torch.Tensor, rope=None) -> torch.Tensor:
90
+ B, N, C = x.shape
91
+ qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4)
92
+ q, k, v = qkv.unbind(0)
93
+ q, k = self.q_norm(q), self.k_norm(k)
94
+
95
+ if rope is not None:
96
+ q = rope(q)
97
+ k = rope(k) #(B, self.num_heads, N, self.head_dim)
98
+
99
+ if self.fused_attn:
100
+ # Use PyTorch's fused scaled dot-product attention with causal mask.
101
+ # Set is_causal=True to apply causal (autoregressive) masking along the sequence dimension.
102
+ # If you need a custom mask instead, build it like:
103
+ # L = q.size(-2)
104
+ # attn_mask = torch.triu(torch.ones(L, L, device=q.device), diagonal=1).bool()
105
+ # and pass attn_mask=attn_mask to scaled_dot_product_attention.
106
+ x = F.scaled_dot_product_attention(
107
+ q, k, v,
108
+ attn_mask=None,
109
+ dropout_p=self.attn_drop.p if self.training else 0.,
110
+ is_causal=self.is_causal
111
+ )
112
+ else:
113
+ q = q * self.scale
114
+ attn = q @ k.transpose(-2, -1)
115
+ attn = attn.softmax(dim=-1)
116
+ attn = self.attn_drop(attn)
117
+ x = attn @ v
118
+
119
+ #------Store attention map for visualization------#
120
+ # q_vis = q * self.scale
121
+ # attn_vis = q_vis @ k.transpose(-2, -1)
122
+ # L, S = q.size(-2), k.size(-2)
123
+ # assert L == S
124
+ # temp_mask = torch.ones(L, S, dtype=torch.bool, device=attn_vis.device).tril(diagonal=0)
125
+ # attn_bias = torch.zeros_like(attn_vis)
126
+ # attn_bias.masked_fill_(temp_mask.logical_not(), float("-inf"))
127
+ # attn_vis = attn_vis + attn_bias
128
+ # attn_vis = attn_vis.softmax(dim=-1)
129
+
130
+ # #before mask:vmin: 9.168302e-07 vmax: 0.3206765
131
+
132
+
133
+ # # Save attention weights as class attribute for visualization
134
+ # self.attn_weights = attn_vis.detach()
135
+ #-------------------------------------------------#
136
+
137
+
138
+ x = x.transpose(1, 2).reshape(B, N, C)
139
+ x = self.proj(x)
140
+ x = self.proj_drop(x)
141
+ return x
142
+
143
+
144
+ class TimestepEmbedder(nn.Module):
145
+ """
146
+ Embeds scalar timesteps into vector representations.
147
+ Same as DiT.
148
+ """
149
+ def __init__(self, hidden_size: int, frequency_embedding_size: int = 256) -> None:
150
+ super().__init__()
151
+ self.frequency_embedding_size = frequency_embedding_size
152
+ self.mlp = nn.Sequential(
153
+ nn.Linear(frequency_embedding_size, hidden_size, bias=True),
154
+ nn.SiLU(),
155
+ nn.Linear(hidden_size, hidden_size, bias=True),
156
+ )
157
+
158
+ @staticmethod
159
+ def timestep_embedding(t: torch.Tensor, dim: int, max_period: int = 10000) -> torch.Tensor:
160
+ """
161
+ Create sinusoidal timestep embeddings.
162
+ Args:
163
+ t: A 1-D Tensor of N indices, one per batch element. These may be fractional.
164
+ dim: The dimension of the output.
165
+ max_period: Controls the minimum frequency of the embeddings.
166
+ Returns:
167
+ An (N, D) Tensor of positional embeddings.
168
+ """
169
+ # https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
170
+ half = dim // 2
171
+ freqs = torch.exp(
172
+ -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
173
+ ).to(device=t.device)
174
+
175
+ args = t[:, None].float() * freqs[None]
176
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
177
+
178
+ if dim % 2:
179
+ embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
180
+
181
+ return embedding
182
+
183
+ @torch.compile
184
+ def forward(self, t: torch.Tensor) -> torch.Tensor:
185
+ t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
186
+ t_emb = self.mlp(t_freq)
187
+ return t_emb
188
+
189
+
190
+ class LabelEmbedder(nn.Module):
191
+ """
192
+ Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
193
+ Same as DiT.
194
+ """
195
+ def __init__(self, num_classes, hidden_size, dropout_prob):
196
+ super().__init__()
197
+ use_cfg_embedding = dropout_prob > 0
198
+ self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding, hidden_size)
199
+ self.num_classes = num_classes
200
+ self.dropout_prob = dropout_prob
201
+
202
+ def token_drop(self, labels, force_drop_ids=None):
203
+ """
204
+ Drops labels to enable classifier-free guidance.
205
+ """
206
+ if force_drop_ids is None:
207
+ drop_ids = torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob
208
+ else:
209
+ drop_ids = force_drop_ids == 1
210
+ labels = torch.where(drop_ids, self.num_classes, labels)
211
+ return labels
212
+
213
+ @torch.compile
214
+ def forward(self, labels, train, force_drop_ids=None):
215
+ use_dropout = self.dropout_prob > 0
216
+ if (train and use_dropout) or (force_drop_ids is not None):
217
+ labels = self.token_drop(labels, force_drop_ids)
218
+ embeddings = self.embedding_table(labels)
219
+ return embeddings
220
+
221
+ class LightningDiTBlock(nn.Module):
222
+ """
223
+ Lightning DiT Block. We add features including:
224
+ - ROPE
225
+ - QKNorm
226
+ - RMSNorm
227
+ - SwiGLU
228
+ - No shift AdaLN.
229
+ Not all of them are used in the final model, please refer to the paper for more details.
230
+ """
231
+ def __init__(
232
+ self,
233
+ hidden_size,
234
+ num_heads,
235
+ mlp_ratio=4.0,
236
+ use_qknorm=False,
237
+ use_swiglu=False,
238
+ use_rmsnorm=False,
239
+ wo_shift=False,
240
+ is_causal=False,
241
+ **block_kwargs
242
+ ):
243
+ super().__init__()
244
+
245
+ # Initialize normalization layers
246
+ if not use_rmsnorm:
247
+ self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
248
+ self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
249
+ else:
250
+ self.norm1 = RMSNorm(hidden_size)
251
+ self.norm2 = RMSNorm(hidden_size)
252
+
253
+ self.is_causal = is_causal
254
+
255
+ # Initialize attention layer
256
+ self.attn = Attention(
257
+ hidden_size,
258
+ num_heads=num_heads,
259
+ qkv_bias=True,
260
+ qk_norm=use_qknorm,
261
+ use_rmsnorm=use_rmsnorm,
262
+ is_causal=self.is_causal,
263
+ **block_kwargs
264
+ )
265
+
266
+ # Initialize MLP layer
267
+ mlp_hidden_dim = int(hidden_size * mlp_ratio)
268
+ approx_gelu = lambda: nn.GELU(approximate="tanh")
269
+ if use_swiglu:
270
+ # here we did not use SwiGLU from xformers because it is not compatible with torch.compile for now.
271
+ self.mlp = SwiGLUFFN(hidden_size, int(2/3 * mlp_hidden_dim))
272
+ else:
273
+ self.mlp = Mlp(
274
+ in_features=hidden_size,
275
+ hidden_features=mlp_hidden_dim,
276
+ act_layer=approx_gelu,
277
+ drop=0
278
+ )
279
+
280
+ # Initialize AdaLN modulation
281
+ if wo_shift:
282
+ self.adaLN_modulation = nn.Sequential(
283
+ nn.SiLU(),
284
+ nn.Linear(hidden_size, 4 * hidden_size, bias=True)
285
+ )
286
+ else:
287
+ self.adaLN_modulation = nn.Sequential(
288
+ nn.SiLU(),
289
+ nn.Linear(hidden_size, 6 * hidden_size, bias=True)
290
+ )
291
+ self.wo_shift = wo_shift
292
+
293
+ @torch.compile
294
+ def forward(self, x, c, feat_rope=None):
295
+ if self.wo_shift:
296
+ scale_msa, gate_msa, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(4, dim=1)
297
+ shift_msa = None
298
+ shift_mlp = None
299
+ else:
300
+ shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1)
301
+
302
+ x = x + gate_msa.unsqueeze(1) * self.attn(modulate(self.norm1(x), shift_msa, scale_msa), rope=feat_rope)
303
+ x = x + gate_mlp.unsqueeze(1) * self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))
304
+ return x
305
+
306
+ class FinalLayer(nn.Module):
307
+ """
308
+ The final layer of LightningDiT.
309
+ """
310
+ def __init__(self, hidden_size, patch_size, out_channels, use_rmsnorm=False):
311
+ super().__init__()
312
+ if not use_rmsnorm:
313
+ self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
314
+ else:
315
+ self.norm_final = RMSNorm(hidden_size)
316
+ self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
317
+ self.adaLN_modulation = nn.Sequential(
318
+ nn.SiLU(),
319
+ nn.Linear(hidden_size, 2 * hidden_size, bias=True)
320
+ )
321
+ @torch.compile
322
+ def forward(self, x, c):
323
+ shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
324
+ x = modulate(self.norm_final(x), shift, scale)
325
+ x = self.linear(x)
326
+ return x
327
+
328
+
329
+ class LightningDiT(nn.Module):
330
+ """
331
+ Diffusion model with a Transformer backbone.
332
+ """
333
+ def __init__(
334
+ self,
335
+ input_size=32,
336
+ patch_size=2,
337
+ in_channels=32,
338
+ hidden_size=1152,
339
+ depth=28,
340
+ num_heads=16,
341
+ mlp_ratio=4.0,
342
+ class_dropout_prob=0.1,
343
+ num_classes=1000,
344
+ learn_sigma=False,
345
+ use_qknorm=False,
346
+ use_swiglu=False,
347
+ use_rope=False,
348
+ use_rmsnorm=False,
349
+ wo_shift=False,
350
+ degree='180',
351
+ use_checkpoint=False,
352
+ ):
353
+ super().__init__()
354
+ self.learn_sigma = learn_sigma
355
+ self.in_channels = in_channels
356
+ self.out_channels = in_channels if not learn_sigma else in_channels * 2
357
+ self.patch_size = patch_size
358
+ self.num_heads = num_heads
359
+ self.use_rope = use_rope
360
+ self.use_rmsnorm = use_rmsnorm
361
+ self.depth = depth
362
+ self.hidden_size = hidden_size
363
+ self.use_checkpoint = use_checkpoint
364
+ self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size, bias=True)
365
+ self.t_embedder = TimestepEmbedder(hidden_size)
366
+ self.y_embedder = LabelEmbedder(num_classes, hidden_size, class_dropout_prob)
367
+ num_patches = self.x_embedder.num_patches
368
+ # Will use fixed sin-cos embedding:
369
+ self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, hidden_size), requires_grad=False)
370
+
371
+ # use rotary position encoding, borrow from EVA
372
+ if self.use_rope:
373
+ half_head_dim = hidden_size // num_heads // 2
374
+ hw_seq_len = input_size // patch_size
375
+ self.feat_rope = VisionRotaryEmbeddingFast(
376
+ dim=half_head_dim,
377
+ pt_seq_len=hw_seq_len,
378
+ )
379
+ else:
380
+ self.feat_rope = None
381
+
382
+ # Set rotation function based on degree parameter
383
+ if degree == '180':
384
+ self.rot_func = rot180
385
+ elif degree == '90':
386
+ self.rot_func = rot90
387
+ else:
388
+ raise ValueError(f"Unsupported degree value: {degree}. Only '90' and '180' are supported.")
389
+
390
+ # self.blocks = nn.ModuleList([
391
+ # LightningDiTBlock(hidden_size,
392
+ # num_heads,
393
+ # mlp_ratio=mlp_ratio,
394
+ # use_qknorm=use_qknorm,
395
+ # use_swiglu=use_swiglu,
396
+ # use_rmsnorm=use_rmsnorm,
397
+ # wo_shift=wo_shift,
398
+ # ) for _ in range(depth)
399
+ # ])
400
+ # self.blocks = nn.ModuleList()
401
+ # group_size = 2
402
+ # rot_per_group = 2
403
+ # normal_per_group = 0
404
+ # num_groups = depth // group_size
405
+ # num_res_layer = depth % group_size
406
+
407
+ # print("*********")
408
+ # print("num_groups:", num_groups)
409
+ # print("group_size:", group_size)
410
+ # print("total depth:", depth)
411
+ # print("rot_per_group:", rot_per_group)
412
+ # print("normal_per_group:", normal_per_group)
413
+ # print("res_blocks:", num_res_layer)
414
+ # print("*********")
415
+
416
+ # for _ in range(num_groups):
417
+ # for i in range(group_size):
418
+ # if i < rot_per_group:
419
+ # self.blocks.append(LightningDiTBlock(
420
+ # hidden_size,
421
+ # num_heads,
422
+ # mlp_ratio=mlp_ratio,
423
+ # use_qknorm=use_qknorm,
424
+ # use_swiglu=use_swiglu,
425
+ # use_rmsnorm=use_rmsnorm,
426
+ # wo_shift=wo_shift,
427
+ # is_causal=True
428
+ # ))
429
+ # print("add causal block")
430
+ # else:
431
+ # self.blocks.append(LightningDiTBlock(
432
+ # hidden_size,
433
+ # num_heads,
434
+ # mlp_ratio=mlp_ratio,
435
+ # use_qknorm=use_qknorm,
436
+ # use_swiglu=use_swiglu,
437
+ # use_rmsnorm=use_rmsnorm,
438
+ # wo_shift=wo_shift,
439
+ # is_causal=False
440
+ # ))
441
+ # print("add full block")
442
+ # for _ in range(num_res_layer):
443
+ # self.blocks.append(LightningDiTBlock(
444
+ # hidden_size,
445
+ # num_heads,
446
+ # mlp_ratio=mlp_ratio,
447
+ # use_qknorm=use_qknorm,
448
+ # use_swiglu=use_swiglu,
449
+ # use_rmsnorm=use_rmsnorm,
450
+ # wo_shift=wo_shift,
451
+ # is_causal=False
452
+ # ))
453
+ self.blocks = nn.ModuleList([
454
+ LightningDiTBlock(hidden_size,
455
+ num_heads,
456
+ mlp_ratio=mlp_ratio,
457
+ use_qknorm=use_qknorm,
458
+ use_swiglu=use_swiglu,
459
+ use_rmsnorm=use_rmsnorm,
460
+ wo_shift=wo_shift,
461
+ is_causal=True
462
+ ) for _ in range(depth)
463
+ ])
464
+ assert len(self.blocks) == depth, f"Total blocks {len(self.blocks)} not equal to depth {depth}"
465
+
466
+ self.final_layer = FinalLayer(hidden_size, patch_size, self.out_channels, use_rmsnorm=use_rmsnorm)
467
+ self.initialize_weights()
468
+
469
+ def initialize_weights(self):
470
+ # Initialize transformer layers:
471
+ def _basic_init(module):
472
+ if isinstance(module, nn.Linear):
473
+ torch.nn.init.xavier_uniform_(module.weight)
474
+ if module.bias is not None:
475
+ nn.init.constant_(module.bias, 0)
476
+ self.apply(_basic_init)
477
+
478
+ # Initialize (and freeze) pos_embed by sin-cos embedding:
479
+ pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], int(self.x_embedder.num_patches ** 0.5))
480
+ self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0))
481
+
482
+ # Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
483
+ w = self.x_embedder.proj.weight.data
484
+ nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
485
+ nn.init.constant_(self.x_embedder.proj.bias, 0)
486
+
487
+ # Initialize label embedding table:
488
+ nn.init.normal_(self.y_embedder.embedding_table.weight, std=0.02)
489
+
490
+ # Initialize timestep embedding MLP:
491
+ nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
492
+ nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
493
+
494
+ # Zero-out adaLN modulation layers in LightningDiT blocks:
495
+ for block in self.blocks:
496
+ nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
497
+ nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
498
+
499
+ # Zero-out output layers:
500
+ nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)
501
+ nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
502
+ nn.init.constant_(self.final_layer.linear.weight, 0)
503
+ nn.init.constant_(self.final_layer.linear.bias, 0)
504
+
505
+ def unpatchify(self, x):
506
+ """
507
+ x: (N, T, patch_size**2 * C)
508
+ imgs: (N, H, W, C)
509
+ """
510
+ c = self.out_channels
511
+ p = self.x_embedder.patch_size[0]
512
+ h = w = int(x.shape[1] ** 0.5)
513
+ assert h * w == x.shape[1]
514
+
515
+ x = x.reshape(shape=(x.shape[0], h, w, p, p, c))
516
+ x = torch.einsum('nhwpqc->nchpwq', x)
517
+ imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p))
518
+ return imgs
519
+
520
+ def forward(self, x, t=None, y=None):
521
+ """
522
+ Forward pass of LightningDiT.
523
+ x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
524
+ t: (N,) tensor of diffusion timesteps
525
+ y: (N,) tensor of class labels
526
+ use_checkpoint: boolean to toggle checkpointing
527
+ """
528
+
529
+ use_checkpoint = self.use_checkpoint
530
+
531
+ x = self.x_embedder(x) + self.pos_embed # (N, T, D), where T = H * W / patch_size ** 2
532
+ t = self.t_embedder(t) # (N, D)
533
+ y = self.y_embedder(y, self.training) # (N, D)
534
+ c = t + y # (N, D)
535
+
536
+ for block in self.blocks:
537
+ assert block.is_causal, "All blocks should be causal for rot180 setting"
538
+ # x = rot90(x) # for 3+1, rot 90, 180, 270
539
+ if use_checkpoint:
540
+ x = checkpoint(block, x, c, self.feat_rope, use_reentrant=True)
541
+ else:
542
+ x = block(x, c, self.feat_rope)
543
+ # if block.is_causal:
544
+ # x = rot90(x) # for 4+1/4+0, rot 0, 90, 180, 270
545
+ x = self.rot_func(x) # Use the rotation function based on degree parameter
546
+
547
+ x = self.final_layer(x, c) # (N, T, patch_size ** 2 * out_channels)
548
+ x = self.unpatchify(x) # (N, out_channels, H, W)
549
+
550
+ if self.learn_sigma:
551
+ x, _ = x.chunk(2, dim=1)
552
+ return x
553
+
554
+ def forward_with_cfg(self, x, t, y, cfg_scale, cfg_interval=None, cfg_interval_start=None):
555
+ """
556
+ Forward pass of LightningDiT, but also batches the unconditional forward pass for classifier-free guidance.
557
+ """
558
+ # https://github.com/openai/glide-text2im/blob/main/notebooks/text2im.ipynb
559
+ half = x[: len(x) // 2]
560
+ combined = torch.cat([half, half], dim=0)
561
+ model_out = self.forward(combined, t, y)
562
+ # For exact reproducibility reasons, we apply classifier-free guidance on only
563
+ # three channels by default. The standard approach to cfg applies it to all channels.
564
+ # This can be done by uncommenting the following line and commenting-out the line following that.
565
+ # eps, rest = model_out[:, :self.in_channels], model_out[:, self.in_channels:]
566
+ eps, rest = model_out[:, :3], model_out[:, 3:]
567
+ cond_eps, uncond_eps = torch.split(eps, len(eps) // 2, dim=0)
568
+ half_eps = uncond_eps + cfg_scale * (cond_eps - uncond_eps)
569
+
570
+ if cfg_interval is True:
571
+ timestep = t[0]
572
+ if timestep < cfg_interval_start:
573
+ half_eps = cond_eps
574
+
575
+ eps = torch.cat([half_eps, half_eps], dim=0)
576
+ return torch.cat([eps, rest], dim=1)
577
+
578
+ def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0):
579
+ """
580
+ grid_size: int of the grid height and width
581
+ return:
582
+ pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)
583
+ """
584
+ grid_h = np.arange(grid_size, dtype=np.float32)
585
+ grid_w = np.arange(grid_size, dtype=np.float32)
586
+ grid = np.meshgrid(grid_w, grid_h) # here w goes first
587
+ grid = np.stack(grid, axis=0)
588
+
589
+ grid = grid.reshape([2, 1, grid_size, grid_size])
590
+ pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
591
+ if cls_token and extra_tokens > 0:
592
+ pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0)
593
+ return pos_embed
594
+
595
+
596
+ def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
597
+ assert embed_dim % 2 == 0
598
+
599
+ # use half of dimensions to encode grid_h
600
+ emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
601
+ emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
602
+
603
+ emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
604
+ return emb
605
+
606
+
607
+ def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
608
+ """
609
+ embed_dim: output dimension for each position
610
+ pos: a list of positions to be encoded: size (M,)
611
+ out: (M, D)
612
+ """
613
+ assert embed_dim % 2 == 0
614
+ omega = np.arange(embed_dim // 2, dtype=np.float64)
615
+ omega /= embed_dim / 2.
616
+ omega = 1. / 10000**omega # (D/2,)
617
+
618
+ pos = pos.reshape(-1) # (M,)
619
+ out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
620
+
621
+ emb_sin = np.sin(out) # (M, D/2)
622
+ emb_cos = np.cos(out) # (M, D/2)
623
+
624
+ emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
625
+ return emb
626
+
627
+
628
+ #################################################################################
629
+ # LightningDiT Configs #
630
+ #################################################################################
631
+
632
+ def LightningDiT_XL_1(**kwargs):
633
+ return LightningDiT(depth=28, hidden_size=1152, patch_size=1, num_heads=16, **kwargs)
634
+
635
+ def LightningDiT_XL_2(**kwargs):
636
+ return LightningDiT(depth=28, hidden_size=1152, patch_size=2, num_heads=16, **kwargs)
637
+
638
+ def LightningDiT_L_2(**kwargs):
639
+ return LightningDiT(depth=24, hidden_size=1024, patch_size=2, num_heads=16, **kwargs)
640
+
641
+ def LightningDiT_B_1(**kwargs):
642
+ return LightningDiT(depth=12, hidden_size=768, patch_size=1, num_heads=12, **kwargs)
643
+
644
+ def LightningDiT_B_2(**kwargs):
645
+ return LightningDiT(depth=12, hidden_size=768, patch_size=2, num_heads=12, **kwargs)
646
+
647
+ def LightningDiT_1p0B_1(**kwargs):
648
+ return LightningDiT(depth=24, hidden_size=1536, patch_size=1, num_heads=24, **kwargs)
649
+
650
+ def LightningDiT_1p0B_2(**kwargs):
651
+ return LightningDiT(depth=24, hidden_size=1536, patch_size=2, num_heads=24, **kwargs)
652
+
653
+ def LightningDiT_1p6B_1(**kwargs):
654
+ return LightningDiT(depth=28, hidden_size=1792, patch_size=1, num_heads=28, **kwargs)
655
+
656
+ def LightningDiT_1p6B_2(**kwargs):
657
+ return LightningDiT(depth=28, hidden_size=1792, patch_size=2, num_heads=28, **kwargs)
658
+
659
+ LightningDiT_models = {
660
+ 'LightningDiT-B/1': LightningDiT_B_1, 'LightningDiT-B/2': LightningDiT_B_2,
661
+ 'LightningDiT-L/2': LightningDiT_L_2,
662
+ 'LightningDiT-XL/1': LightningDiT_XL_1, 'LightningDiT-XL/2': LightningDiT_XL_2,
663
+ 'LightningDiT-1p0B/1': LightningDiT_1p0B_1, 'LightningDiT-1p0B/2': LightningDiT_1p0B_2,
664
+ 'LightningDiT-1p6B/1': LightningDiT_1p6B_1, 'LightningDiT-1p6B/2': LightningDiT_1p6B_2,
665
+ }