XBFDI commited on
Commit
413caa6
·
verified ·
1 Parent(s): 252435f

Fix: all_tied_weights_keys

Browse files
Files changed (1) hide show
  1. modeling_ezfgraphic.py +93 -91
modeling_ezfgraphic.py CHANGED
@@ -1,91 +1,93 @@
1
- from .configuration_ezfgraphic import EZFGraphicConfig
2
-
3
- import torch, torch.nn as nn, torch.nn.functional as F
4
- from transformers import PreTrainedModel
5
-
6
- class RMSNorm(nn.Module):
7
- def __init__(self, d, eps=1e-5):
8
- super().__init__()
9
- self.w = nn.Parameter(torch.ones(d)); self.eps = eps
10
- def forward(self, x):
11
- return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.w
12
-
13
- class Attention(nn.Module):
14
- def __init__(self, cfg):
15
- super().__init__()
16
- self.qkv = nn.Linear(cfg.d_model, 3 * cfg.d_model, bias=False)
17
- self.proj = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
18
- self.drop = nn.Dropout(cfg.dropout)
19
- self.register_buffer("mask", torch.tril(torch.ones(512, 512)).bool())
20
- def forward(self, x):
21
- B, T, C = x.shape
22
- q, k, v = self.qkv(x).chunk(3, dim=-1)
23
- att = (q @ k.transpose(-2, -1)) / (C ** 0.5)
24
- att = att.masked_fill(~self.mask[:T, :T], float("-inf"))
25
- att = self.drop(F.softmax(att, dim=-1))
26
- return self.proj(att @ v)
27
-
28
- class FeedForward(nn.Module):
29
- def __init__(self, cfg):
30
- super().__init__()
31
- self.fc1 = nn.Linear(cfg.d_model, 2 * cfg.d_model)
32
- self.fc2 = nn.Linear(2 * cfg.d_model, cfg.d_model)
33
- self.drop = nn.Dropout(cfg.dropout)
34
- def forward(self, x):
35
- return self.drop(self.fc2(F.relu(self.fc1(x))))
36
-
37
- class Block(nn.Module):
38
- def __init__(self, cfg):
39
- super().__init__()
40
- self.ln1 = RMSNorm(cfg.d_model); self.attn = Attention(cfg)
41
- self.ln2 = RMSNorm(cfg.d_model); self.ffn = FeedForward(cfg)
42
- def forward(self, x):
43
- x = x + self.attn(self.ln1(x))
44
- x = x + self.ffn(self.ln2(x))
45
- return x
46
-
47
- class EZFGraphic(PreTrainedModel):
48
- config_class = EZFGraphicConfig
49
- def __init__(self, config):
50
- super().__init__(config)
51
- cfg = config
52
- self.cfg = cfg
53
- self.text_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)
54
- self.pos_emb_text = nn.Embedding(cfg.ctx, cfg.d_model)
55
- n_patches = (cfg.image_size // cfg.patch_size) ** 2
56
- self.n_patches = n_patches
57
- self.pos_emb_visual = nn.Embedding(n_patches, cfg.d_model)
58
- self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.n_layer)])
59
- self.ln_f = RMSNorm(cfg.d_model)
60
- self.patch_pixels = cfg.patch_size * cfg.patch_size * 3
61
- self.pixel_decoder = nn.Sequential(
62
- nn.Linear(cfg.d_model, 512),
63
- nn.ReLU(),
64
- nn.Linear(512, 1024),
65
- nn.ReLU(),
66
- nn.Linear(1024, self.patch_pixels),
67
- )
68
- def encode_text(self, text_ids):
69
- B, T = text_ids.shape
70
- pos = torch.arange(T, device=text_ids.device)
71
- x = self.text_emb(text_ids) + self.pos_emb_text(pos)
72
- for b in self.blocks:
73
- x = b(x)
74
- return self.ln_f(x)
75
- def decode_pixels(self, x):
76
- B, T, _ = x.shape
77
- patches = self.pixel_decoder(x)
78
- ps = self.cfg.patch_size
79
- h = w = self.cfg.image_size // ps
80
- patches = patches.view(B, h, w, ps, ps, 3)
81
- pixels = patches.permute(0, 5, 1, 3, 2, 4).reshape(B, 3, self.cfg.image_size, self.cfg.image_size)
82
- return pixels
83
- def forward(self, text_ids):
84
- x = self.encode_text(text_ids)
85
- last = x[:, -1:, :]
86
- visual_input = last.expand(-1, self.n_patches, -1).contiguous()
87
- visual_input = visual_input + self.pos_emb_visual(
88
- torch.arange(self.n_patches, device=text_ids.device)
89
- ).unsqueeze(0)
90
- pixels = self.decode_pixels(visual_input)
91
- return pixels
 
 
 
1
+ from .configuration_ezfgraphic import EZFGraphicConfig
2
+
3
+ import torch, torch.nn as nn, torch.nn.functional as F
4
+ from transformers import PreTrainedModel
5
+
6
+ class RMSNorm(nn.Module):
7
+ def __init__(self, d, eps=1e-5):
8
+ super().__init__()
9
+ self.w = nn.Parameter(torch.ones(d)); self.eps = eps
10
+ def forward(self, x):
11
+ return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.w
12
+
13
+ class Attention(nn.Module):
14
+ def __init__(self, cfg):
15
+ super().__init__()
16
+ self.qkv = nn.Linear(cfg.d_model, 3 * cfg.d_model, bias=False)
17
+ self.proj = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
18
+ self.drop = nn.Dropout(cfg.dropout)
19
+ self.register_buffer("mask", torch.tril(torch.ones(512, 512)).bool())
20
+ def forward(self, x):
21
+ B, T, C = x.shape
22
+ q, k, v = self.qkv(x).chunk(3, dim=-1)
23
+ att = (q @ k.transpose(-2, -1)) / (C ** 0.5)
24
+ att = att.masked_fill(~self.mask[:T, :T], float("-inf"))
25
+ att = self.drop(F.softmax(att, dim=-1))
26
+ return self.proj(att @ v)
27
+
28
+ class FeedForward(nn.Module):
29
+ def __init__(self, cfg):
30
+ super().__init__()
31
+ self.fc1 = nn.Linear(cfg.d_model, 2 * cfg.d_model)
32
+ self.fc2 = nn.Linear(2 * cfg.d_model, cfg.d_model)
33
+ self.drop = nn.Dropout(cfg.dropout)
34
+ def forward(self, x):
35
+ return self.drop(self.fc2(F.relu(self.fc1(x))))
36
+
37
+ class Block(nn.Module):
38
+ def __init__(self, cfg):
39
+ super().__init__()
40
+ self.ln1 = RMSNorm(cfg.d_model); self.attn = Attention(cfg)
41
+ self.ln2 = RMSNorm(cfg.d_model); self.ffn = FeedForward(cfg)
42
+ def forward(self, x):
43
+ x = x + self.attn(self.ln1(x))
44
+ x = x + self.ffn(self.ln2(x))
45
+ return x
46
+
47
+ class EZFGraphic(PreTrainedModel):
48
+ config_class = EZFGraphicConfig
49
+ _tied_weights_keys = []
50
+ all_tied_weights_keys = {}
51
+ def __init__(self, config):
52
+ super().__init__(config)
53
+ cfg = config
54
+ self.cfg = cfg
55
+ self.text_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)
56
+ self.pos_emb_text = nn.Embedding(cfg.ctx, cfg.d_model)
57
+ n_patches = (cfg.image_size // cfg.patch_size) ** 2
58
+ self.n_patches = n_patches
59
+ self.pos_emb_visual = nn.Embedding(n_patches, cfg.d_model)
60
+ self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.n_layer)])
61
+ self.ln_f = RMSNorm(cfg.d_model)
62
+ self.patch_pixels = cfg.patch_size * cfg.patch_size * 3
63
+ self.pixel_decoder = nn.Sequential(
64
+ nn.Linear(cfg.d_model, 512),
65
+ nn.ReLU(),
66
+ nn.Linear(512, 1024),
67
+ nn.ReLU(),
68
+ nn.Linear(1024, self.patch_pixels),
69
+ )
70
+ def encode_text(self, text_ids):
71
+ B, T = text_ids.shape
72
+ pos = torch.arange(T, device=text_ids.device)
73
+ x = self.text_emb(text_ids) + self.pos_emb_text(pos)
74
+ for b in self.blocks:
75
+ x = b(x)
76
+ return self.ln_f(x)
77
+ def decode_pixels(self, x):
78
+ B, T, _ = x.shape
79
+ patches = self.pixel_decoder(x)
80
+ ps = self.cfg.patch_size
81
+ h = w = self.cfg.image_size // ps
82
+ patches = patches.view(B, h, w, ps, ps, 3)
83
+ pixels = patches.permute(0, 5, 1, 3, 2, 4).reshape(B, 3, self.cfg.image_size, self.cfg.image_size)
84
+ return pixels
85
+ def forward(self, text_ids):
86
+ x = self.encode_text(text_ids)
87
+ last = x[:, -1:, :]
88
+ visual_input = last.expand(-1, self.n_patches, -1).contiguous()
89
+ visual_input = visual_input + self.pos_emb_visual(
90
+ torch.arange(self.n_patches, device=text_ids.device)
91
+ ).unsqueeze(0)
92
+ pixels = self.decode_pixels(visual_input)
93
+ return pixels