comdoleger commited on
Commit
e3e3fdf
·
verified ·
1 Parent(s): 2b1f42f

Upload extensions_built_in/diffusion_models/omnigen2/src/models/transformers/repo.py with huggingface_hub

Browse files
extensions_built_in/diffusion_models/omnigen2/src/models/transformers/repo.py ADDED
@@ -0,0 +1,135 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import List, Tuple
2
+
3
+ import torch
4
+ import torch.nn as nn
5
+
6
+ from einops import repeat
7
+ from diffusers.models.embeddings import get_1d_rotary_pos_embed
8
+
9
+ class OmniGen2RotaryPosEmbed(nn.Module):
10
+ def __init__(self, theta: int,
11
+ axes_dim: Tuple[int, int, int],
12
+ axes_lens: Tuple[int, int, int] = (300, 512, 512),
13
+ patch_size: int = 2):
14
+ super().__init__()
15
+ self.theta = theta
16
+ self.axes_dim = axes_dim
17
+ self.axes_lens = axes_lens
18
+ self.patch_size = patch_size
19
+
20
+ @staticmethod
21
+ def get_freqs_cis(axes_dim: Tuple[int, int, int],
22
+ axes_lens: Tuple[int, int, int],
23
+ theta: int) -> List[torch.Tensor]:
24
+ freqs_cis = []
25
+ freqs_dtype = torch.float32 if torch.backends.mps.is_available() else torch.float64
26
+ for i, (d, e) in enumerate(zip(axes_dim, axes_lens)):
27
+ emb = get_1d_rotary_pos_embed(d, e, theta=theta, freqs_dtype=freqs_dtype)
28
+ freqs_cis.append(emb)
29
+ return freqs_cis
30
+
31
+ def _get_freqs_cis(self, freqs_cis, ids: torch.Tensor) -> torch.Tensor:
32
+ device = ids.device
33
+ if ids.device.type == "mps":
34
+ ids = ids.to("cpu")
35
+
36
+ result = []
37
+ for i in range(len(self.axes_dim)):
38
+ freqs = freqs_cis[i].to(ids.device)
39
+ index = ids[:, :, i : i + 1].repeat(1, 1, freqs.shape[-1]).to(torch.int64)
40
+ result.append(torch.gather(freqs.unsqueeze(0).repeat(index.shape[0], 1, 1), dim=1, index=index))
41
+ return torch.cat(result, dim=-1).to(device)
42
+
43
+ def forward(
44
+ self,
45
+ freqs_cis,
46
+ attention_mask,
47
+ l_effective_ref_img_len,
48
+ l_effective_img_len,
49
+ ref_img_sizes,
50
+ img_sizes,
51
+ device
52
+ ):
53
+ batch_size = len(attention_mask)
54
+ p = self.patch_size
55
+
56
+ encoder_seq_len = attention_mask.shape[1]
57
+ l_effective_cap_len = attention_mask.sum(dim=1).tolist()
58
+
59
+ seq_lengths = [cap_len + sum(ref_img_len) + img_len for cap_len, ref_img_len, img_len in zip(l_effective_cap_len, l_effective_ref_img_len, l_effective_img_len)]
60
+
61
+ max_seq_len = int(max(seq_lengths))
62
+ max_ref_img_len = max([sum(ref_img_len) for ref_img_len in l_effective_ref_img_len])
63
+ max_img_len = max(l_effective_img_len)
64
+
65
+ # Create position IDs
66
+ position_ids = torch.zeros(batch_size, max_seq_len, 3, dtype=torch.int32, device=device)
67
+
68
+ for i, (cap_seq_len, seq_len) in enumerate(zip(l_effective_cap_len, seq_lengths)):
69
+ cap_seq_len = int(cap_seq_len)
70
+ seq_len = int(seq_len)
71
+ # add text position ids
72
+ position_ids[i, :cap_seq_len] = repeat(torch.arange(cap_seq_len, dtype=torch.int32, device=device), "l -> l 3")
73
+
74
+ pe_shift = cap_seq_len
75
+ pe_shift_len = cap_seq_len
76
+
77
+ if ref_img_sizes[i] is not None:
78
+ for ref_img_size, ref_img_len in zip(ref_img_sizes[i], l_effective_ref_img_len[i]):
79
+ H, W = ref_img_size
80
+ ref_H_tokens, ref_W_tokens = H // p, W // p
81
+ assert ref_H_tokens * ref_W_tokens == ref_img_len
82
+ # add image position ids
83
+
84
+ row_ids = repeat(torch.arange(ref_H_tokens, dtype=torch.int32, device=device), "h -> h w", w=ref_W_tokens).flatten()
85
+ col_ids = repeat(torch.arange(ref_W_tokens, dtype=torch.int32, device=device), "w -> h w", h=ref_H_tokens).flatten()
86
+ position_ids[i, pe_shift_len:pe_shift_len + ref_img_len, 0] = pe_shift
87
+ position_ids[i, pe_shift_len:pe_shift_len + ref_img_len, 1] = row_ids
88
+ position_ids[i, pe_shift_len:pe_shift_len + ref_img_len, 2] = col_ids
89
+
90
+ pe_shift += max(ref_H_tokens, ref_W_tokens)
91
+ pe_shift_len += ref_img_len
92
+
93
+ H, W = img_sizes[i]
94
+ H_tokens, W_tokens = H // p, W // p
95
+ assert H_tokens * W_tokens == l_effective_img_len[i]
96
+
97
+ row_ids = repeat(torch.arange(H_tokens, dtype=torch.int32, device=device), "h -> h w", w=W_tokens).flatten()
98
+ col_ids = repeat(torch.arange(W_tokens, dtype=torch.int32, device=device), "w -> h w", h=H_tokens).flatten()
99
+
100
+ assert pe_shift_len + l_effective_img_len[i] == seq_len
101
+ position_ids[i, pe_shift_len: seq_len, 0] = pe_shift
102
+ position_ids[i, pe_shift_len: seq_len, 1] = row_ids
103
+ position_ids[i, pe_shift_len: seq_len, 2] = col_ids
104
+
105
+ # Get combined rotary embeddings
106
+ freqs_cis = self._get_freqs_cis(freqs_cis, position_ids)
107
+
108
+ # create separate rotary embeddings for captions and images
109
+ cap_freqs_cis = torch.zeros(
110
+ batch_size, encoder_seq_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype
111
+ )
112
+ ref_img_freqs_cis = torch.zeros(
113
+ batch_size, max_ref_img_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype
114
+ )
115
+ img_freqs_cis = torch.zeros(
116
+ batch_size, max_img_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype
117
+ )
118
+
119
+ for i, (cap_seq_len, ref_img_len, img_len, seq_len) in enumerate(zip(l_effective_cap_len, l_effective_ref_img_len, l_effective_img_len, seq_lengths)):
120
+ cap_seq_len = int(cap_seq_len)
121
+ sum_ref_img_len = int(sum(ref_img_len))
122
+ img_len = int(img_len)
123
+ seq_len = int(seq_len)
124
+ cap_freqs_cis[i, :cap_seq_len] = freqs_cis[i, :cap_seq_len]
125
+ ref_img_freqs_cis[i, :sum_ref_img_len] = freqs_cis[i, cap_seq_len:cap_seq_len + sum_ref_img_len]
126
+ img_freqs_cis[i, :img_len] = freqs_cis[i, cap_seq_len + sum_ref_img_len:cap_seq_len + sum_ref_img_len + img_len]
127
+
128
+ return (
129
+ cap_freqs_cis,
130
+ ref_img_freqs_cis,
131
+ img_freqs_cis,
132
+ freqs_cis,
133
+ l_effective_cap_len,
134
+ seq_lengths,
135
+ )