pandeyps commited on
Commit
bd4f849
·
verified ·
1 Parent(s): a740356

Fela 1.6M

Browse files
README.md ADDED
@@ -0,0 +1,49 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ tags:
4
+ - protein
5
+ - biology
6
+ - language-model
7
+ - causal-lm
8
+ - hyena
9
+ - pytorch
10
+ pipeline_tag: feature extraction
11
+ ---
12
+
13
+ # Fela
14
+
15
+ PyTorch written protein language model on the hyena operator (1.6M params)
16
+
17
+ - Architecture: long conv + MLP blocks, pre-norm, LM head
18
+ - Tokenizer: char level over `ACDEFGHIKLMNPQRSTVWYX`, `<pad>`=0, `<eos>`=22, `<unk>`=23
19
+ - Data: Pfam-A (filtered to 20–512 residues, standard alphabet only), ~9.5B tokens
20
+ - Training: 40k steps, batch 256, bf16, AdamW (wd 0.1), cosine LR 6e-4 → 6e-5
21
+
22
+ ## Config
23
+
24
+ | Parameter | Value |
25
+ |---|---|
26
+ | d_model | 256 |
27
+ | n_layer | 2 |
28
+ | d_inner | 1024 |
29
+ | vocab_size | 32 |
30
+ | l_max | 514 |
31
+ | order | 2 |
32
+ | filter_order | 64 |
33
+ | short_filter_order | 3 |
34
+ | emb_dim | 5 |
35
+ | w | 10 |
36
+ | num_inner_mlps | 2 |
37
+ | residual_in_fp32 | true |
38
+
39
+ ## Usage
40
+
41
+ ```python
42
+ from transformers import AutoModelForCausalLM, AutoTokenizer
43
+
44
+ model = AutoModelForCausalLM.from_pretrained("pandeyps/fela", trust_remote_code=True)
45
+ tok = AutoTokenizer.from_pretrained("pandeyps/fela", trust_remote_code=True)
46
+
47
+ ids = tok.encode("MSDKIIEYDETARRAIEAGVNTLADAV", return_tensors="pt")
48
+ gen = model.generate(ids, max_new_tokens=64, do_sample=True, temperature=0.7)
49
+ print(tok.decode(gen[0]))
config.json ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "d_inner": 1024,
3
+ "d_model": 256,
4
+ "embed_dropout": 0.1,
5
+ "eos_token_id": 22,
6
+ "eps": 1e-05,
7
+ "l_max": 514,
8
+ "model_type": "fela",
9
+ "n_layer": 2,
10
+ "pad_token_id": 0,
11
+ "resid_dropout": 0.0,
12
+ "residual_in_fp32": true,
13
+ "transformers_version": "4.57.6",
14
+ "vocab_size": 32,
15
+ "auto_map": {
16
+ "AutoConfig": "modeling_fela.FelaConfig",
17
+ "AutoModelForCausalLM": "modeling_fela.FelaForCausalLM",
18
+ "AutoTokenizer": "tokenization_fela.FelaTokenizer"
19
+ }
20
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:fd95f80847e19fc12494f869447d77803204bb460421fb56def031bca5e59546
3
+ size 6613656
modeling_fela.py ADDED
@@ -0,0 +1,224 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ import torch
3
+ import torch.nn as nn
4
+ import torch.nn.functional as F
5
+ from einops import rearrange
6
+ from transformers import PreTrainedModel, PretrainedConfig
7
+ from transformers.modeling_outputs import CausalLMOutput
8
+
9
+
10
+ class Sin(nn.Module):
11
+ def __init__(self, dim, w=10, train_freq=True):
12
+ super().__init__()
13
+ self.freq = nn.Parameter(w * torch.ones(1, dim)) if train_freq else w * torch.ones(1, dim)
14
+ def forward(self, x):
15
+ return torch.sin(self.freq * x)
16
+
17
+ class PositionalEmbedding(nn.Module):
18
+ def __init__(self, emb_dim, seq_len):
19
+ super().__init__()
20
+ t = torch.linspace(0, 1, seq_len)[None, :, None]
21
+ bands = (emb_dim - 1) // 2
22
+ t_rescaled = torch.linspace(0, seq_len - 1, seq_len)[None, :, None]
23
+ w = 2 * math.pi * t_rescaled / seq_len
24
+ f = torch.linspace(1e-4, bands - 1, bands)[None, None]
25
+ z = torch.exp(-1j * f * w)
26
+ z = torch.cat([t, z.real, z.imag], dim=-1)
27
+ self.register_buffer("z", z)
28
+ self.register_buffer("t", t)
29
+ def forward(self, L):
30
+ return self.z[:, :L], self.t[:, :L]
31
+
32
+ class ExponentialModulation(nn.Module):
33
+ def __init__(self, d_model, fast_decay_pct=0.3, slow_decay_pct=1.5, target=1e-2, shift=0.0):
34
+ super().__init__()
35
+ self.shift = shift
36
+ max_decay = math.log(target) / fast_decay_pct
37
+ min_decay = math.log(target) / slow_decay_pct
38
+ deltas = torch.linspace(min_decay, max_decay, d_model)[None, None]
39
+ self.register_buffer("deltas", deltas)
40
+ def forward(self, t, x):
41
+ return x * (torch.exp(-t * self.deltas.abs()) + self.shift)
42
+
43
+ class HyenaFilter(nn.Module):
44
+ def __init__(self, d_model=256, emb_dim=5, order=64, seq_len=514,
45
+ num_inner_mlps=2, w=10, modulate=True):
46
+ super().__init__()
47
+ self.modulate = modulate
48
+ self.bias = nn.Parameter(torch.randn(d_model))
49
+ act = Sin(dim=order, w=w)
50
+ self.pos_emb = PositionalEmbedding(emb_dim, seq_len)
51
+ self.implicit_filter = nn.Sequential(
52
+ nn.Linear(emb_dim, order), act,
53
+ *[m for _ in range(num_inner_mlps) for m in (nn.Linear(order, order), act)],
54
+ nn.Linear(order, d_model, bias=False),
55
+ )
56
+ self.modulation = ExponentialModulation(d_model)
57
+ def filter(self, L):
58
+ z, t = self.pos_emb(L)
59
+ h = self.implicit_filter(z)
60
+ if self.modulate:
61
+ h = self.modulation(t, h)
62
+ return h
63
+
64
+ class ShortConv(nn.Module):
65
+ def __init__(self, d_model=256, order=2, short_filter_order=3):
66
+ super().__init__()
67
+ total_width = d_model * (order + 1)
68
+ self.in_proj = nn.Linear(d_model, total_width)
69
+ self.conv = nn.Conv1d(total_width, total_width, short_filter_order,
70
+ groups=total_width, padding=short_filter_order - 1)
71
+ def forward(self, u):
72
+ u = self.in_proj(u)
73
+ u = u.transpose(1, 2)
74
+ u = self.conv(u)[..., :u.shape[-1]]
75
+ return u
76
+
77
+ def fft_conv(u, k, bias=None):
78
+ seqlen = u.shape[-1]
79
+ fft_size = 2 * seqlen
80
+ k_f = torch.fft.rfft(k, n=fft_size) / fft_size
81
+ if len(u.shape) > 3:
82
+ k_f = k_f.unsqueeze(1)
83
+ u_f = torch.fft.rfft(u.to(dtype=k.dtype), n=fft_size)
84
+ y = torch.fft.irfft(u_f * k_f, n=fft_size, norm="forward")[..., :seqlen]
85
+ return y + u * bias.unsqueeze(-1) if bias is not None else y
86
+
87
+ class HyenaOperator(nn.Module):
88
+ def __init__(self, d_model=256, l_max=514, order=2, filter_order=64,
89
+ short_filter_order=3, drop_rate=0.0):
90
+ super().__init__()
91
+ self.d_model, self.order, self.l_max = d_model, order, l_max
92
+ self.in_proj = nn.Linear(d_model, (order + 1) * d_model)
93
+ self.out_proj = nn.Linear(d_model, d_model)
94
+ total_width = d_model * (order + 1)
95
+ self.short_filter = nn.Conv1d(total_width, total_width, short_filter_order,
96
+ groups=total_width, padding=short_filter_order - 1)
97
+ self.filter_fn = HyenaFilter(d_model=d_model, order=filter_order, seq_len=l_max)
98
+ self.dropout = nn.Dropout(drop_rate)
99
+ def forward(self, u):
100
+ l_filter = min(u.size(-2), self.l_max)
101
+ u = rearrange(self.in_proj(u), "b l d -> b d l")
102
+ uc = self.short_filter(u)[..., :l_filter]
103
+ *x, v = uc.split(self.d_model, dim=1)
104
+ k = self.filter_fn.filter(l_filter)
105
+ k = rearrange(k, "c l (v o) -> c o v l", v=self.d_model, o=self.order - 1)
106
+ bias = rearrange(self.filter_fn.bias, "(v o) -> o v", o=self.order - 1)
107
+ for o, x_i in enumerate(reversed(x[1:])):
108
+ v = self.dropout(v * x_i)
109
+ v = fft_conv(v, k[o], bias[o])
110
+ return self.out_proj(rearrange(v * x[0], "b v l -> b l v"))
111
+
112
+ class _Block(nn.Module):
113
+ def __init__(self, d_model, d_inner, l_max, drop1_p, drop2_p, eps=1e-5, residual_in_fp32=True):
114
+ super().__init__()
115
+ self.drop1 = nn.Dropout(drop1_p)
116
+ self.norm1 = nn.LayerNorm(d_model, eps=eps)
117
+ self.mixer = HyenaOperator(d_model=d_model, l_max=l_max)
118
+ self.drop2 = nn.Dropout(drop2_p)
119
+ self.norm2 = nn.LayerNorm(d_model, eps=eps)
120
+ self.mlp = nn.Sequential(nn.Linear(d_model, d_inner),
121
+ nn.GELU(approximate="tanh"),
122
+ nn.Linear(d_inner, d_model))
123
+ self.residual_in_fp32 = residual_in_fp32
124
+ def forward(self, hidden, residual):
125
+ dropped = self.drop1(hidden)
126
+ residual = (dropped + residual) if residual is not None else dropped
127
+ hidden = self.mixer(self.norm1(residual.to(dtype=self.norm1.weight.dtype)))
128
+ if self.residual_in_fp32:
129
+ residual = residual.float()
130
+ dropped = self.drop2(hidden)
131
+ residual = (dropped + residual) if residual is not None else dropped
132
+ hidden = self.mlp(self.norm2(residual.to(dtype=self.norm2.weight.dtype)))
133
+ if self.residual_in_fp32:
134
+ residual = residual.float()
135
+ return hidden, residual
136
+
137
+ class Fela(nn.Module):
138
+ def __init__(self, d_model=256, n_layer=2, d_inner=1024, vocab_size=32, l_max=514,
139
+ embed_dropout=0.1, resid_dropout=0.0, eps=1e-5, residual_in_fp32=True):
140
+ super().__init__()
141
+ torch.manual_seed(2222)
142
+ self.embed = nn.Embedding(vocab_size, d_model)
143
+ self.blocks = nn.ModuleList(
144
+ _Block(d_model, d_inner, l_max,
145
+ drop1_p=embed_dropout if i == 0 else resid_dropout,
146
+ drop2_p=resid_dropout, eps=eps,
147
+ residual_in_fp32=residual_in_fp32)
148
+ for i in range(n_layer)
149
+ )
150
+ self.drop_f = nn.Dropout(resid_dropout)
151
+ self.ln_f = nn.LayerNorm(d_model, eps=eps)
152
+ self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
153
+ self._init_weights(n_layer)
154
+ self.lm_head.weight = self.embed.weight
155
+ def _init_weights(self, n_layer):
156
+ for m in self.modules():
157
+ if isinstance(m, nn.Linear):
158
+ nn.init.normal_(m.weight, std=0.02)
159
+ if m.bias is not None:
160
+ nn.init.zeros_(m.bias)
161
+ elif isinstance(m, nn.Embedding):
162
+ nn.init.normal_(m.weight, std=0.02)
163
+ for name, p in self.named_parameters():
164
+ if name.endswith("out_proj.weight") or name.endswith("mlp.2.weight"):
165
+ nn.init.normal_(p, std=0.02 / math.sqrt(2 * n_layer))
166
+ def forward(self, input_ids):
167
+ hidden = self.embed(input_ids)
168
+ residual = None
169
+ for block in self.blocks:
170
+ hidden, residual = block(hidden, residual)
171
+ dropped = self.drop_f(hidden)
172
+ residual = (dropped + residual) if residual is not None else dropped
173
+ hidden = self.ln_f(residual.to(dtype=self.ln_f.weight.dtype))
174
+ return self.lm_head(hidden)
175
+
176
+
177
+ try:
178
+ from transformers.generation import GenerationMixin
179
+ except ImportError:
180
+ from transformers.generation_utils import GenerationMixin
181
+
182
+ class FelaConfig(PretrainedConfig):
183
+ model_type = "fela"
184
+ def __init__(self, d_model=256, n_layer=2, d_inner=1024, vocab_size=32, l_max=514,
185
+ embed_dropout=0.1, resid_dropout=0.0, eps=1e-5, residual_in_fp32=True, **kwargs):
186
+ super().__init__(**kwargs)
187
+ self.d_model, self.n_layer, self.d_inner = d_model, n_layer, d_inner
188
+ self.vocab_size, self.l_max = vocab_size, l_max
189
+ self.embed_dropout, self.resid_dropout = embed_dropout, resid_dropout
190
+ self.eps, self.residual_in_fp32 = eps, residual_in_fp32
191
+ self.pad_token_id = 0
192
+ self.eos_token_id = 22
193
+
194
+ class FelaPreTrainedModel(PreTrainedModel):
195
+ config_class = FelaConfig
196
+ base_model_prefix = "fela"
197
+ def _init_weights(self, module):
198
+ pass
199
+
200
+ class FelaForCausalLM(FelaPreTrainedModel, GenerationMixin):
201
+ config_class = FelaConfig
202
+ base_model_prefix = "fela"
203
+ def __init__(self, config):
204
+ super().__init__(config)
205
+ self.fela = Fela(
206
+ d_model=config.d_model, n_layer=config.n_layer,
207
+ d_inner=config.d_inner, vocab_size=config.vocab_size,
208
+ l_max=config.l_max, embed_dropout=config.embed_dropout,
209
+ resid_dropout=config.resid_dropout, eps=config.eps,
210
+ residual_in_fp32=config.residual_in_fp32,
211
+ )
212
+ def get_input_embeddings(self):
213
+ return self.fela.embed
214
+ def set_input_embeddings(self, v):
215
+ self.fela.embed = v
216
+ self.fela.lm_head.weight = v.weight
217
+ def prepare_inputs_for_generation(self, input_ids, **kwargs):
218
+ return {"input_ids": input_ids}
219
+ def forward(self, input_ids=None, labels=None, attention_mask=None, **kwargs):
220
+ logits = self.fela(input_ids)
221
+ loss = None
222
+ if labels is not None:
223
+ loss = F.cross_entropy(logits.reshape(-1, self.config.vocab_size), labels.reshape(-1))
224
+ return CausalLMOutput(logits=logits, loss=loss)
special_tokens_map.json ADDED
@@ -0,0 +1 @@
 
 
1
+ {}
tokenization_fela.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import json, os
3
+ from transformers import PreTrainedTokenizer
4
+
5
+ class FelaTokenizer(PreTrainedTokenizer):
6
+ vocab_files_names = {"vocab_file": "vocab.json"}
7
+ def __init__(self, vocab_file=None, **kwargs):
8
+ self.vocab = json.load(open(vocab_file)) if vocab_file else {}
9
+ self.itos = {v: k for k, v in self.vocab.items()}
10
+ super().__init__(**kwargs)
11
+ def _tokenize(self, text, **kwargs):
12
+ return list(text)
13
+ def _convert_token_to_id(self, token):
14
+ return self.vocab.get(token, self.vocab.get("<unk>", 23))
15
+ def _convert_id_to_token(self, index):
16
+ return self.itos.get(index, "<unk>")
17
+ def get_vocab(self):
18
+ return dict(self.vocab)
19
+ def vocab_size(self):
20
+ return len(self.vocab)
21
+ def save_vocabulary(self, save_directory, filename_prefix=None):
22
+ fname = os.path.join(save_directory, (filename_prefix + "-" if filename_prefix else "") + "vocab.json")
23
+ with open(fname, "w") as f:
24
+ json.dump(self.get_vocab(), f)
25
+ return (fname,)
tokenizer_config.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "added_tokens_decoder": {},
3
+ "clean_up_tokenization_spaces": false,
4
+ "extra_special_tokens": {},
5
+ "model_max_length": 1000000000000000019884624838656,
6
+ "tokenizer_class": "FelaTokenizer"
7
+ }
vocab.json ADDED
@@ -0,0 +1 @@
 
 
1
+ {"A": 1, "C": 2, "D": 3, "E": 4, "F": 5, "G": 6, "H": 7, "I": 8, "K": 9, "L": 10, "M": 11, "N": 12, "P": 13, "Q": 14, "R": 15, "S": 16, "T": 17, "V": 18, "W": 19, "Y": 20, "X": 21, "<pad>": 0, "<eos>": 22, "<unk>": 23}