Oxpose Martico2432 commited on
Commit
60b63b7
·
0 Parent(s):

Duplicate from Martico2432/srlm-1m

Browse files

Co-authored-by: Martí <Martico2432@users.noreply.huggingface.co>

.gitattributes ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ckpt filter=lfs diff=lfs merge=lfs -text
6
+ *.ftz filter=lfs diff=lfs merge=lfs -text
7
+ *.gz filter=lfs diff=lfs merge=lfs -text
8
+ *.h5 filter=lfs diff=lfs merge=lfs -text
9
+ *.joblib filter=lfs diff=lfs merge=lfs -text
10
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
+ *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
+ *.model filter=lfs diff=lfs merge=lfs -text
13
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
14
+ *.npy filter=lfs diff=lfs merge=lfs -text
15
+ *.npz filter=lfs diff=lfs merge=lfs -text
16
+ *.onnx filter=lfs diff=lfs merge=lfs -text
17
+ *.ot filter=lfs diff=lfs merge=lfs -text
18
+ *.parquet filter=lfs diff=lfs merge=lfs -text
19
+ *.pb filter=lfs diff=lfs merge=lfs -text
20
+ *.pickle filter=lfs diff=lfs merge=lfs -text
21
+ *.pkl filter=lfs diff=lfs merge=lfs -text
22
+ *.pt filter=lfs diff=lfs merge=lfs -text
23
+ *.pth filter=lfs diff=lfs merge=lfs -text
24
+ *.rar filter=lfs diff=lfs merge=lfs -text
25
+ *.safetensors filter=lfs diff=lfs merge=lfs -text
26
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
28
+ *.tar filter=lfs diff=lfs merge=lfs -text
29
+ *.tflite filter=lfs diff=lfs merge=lfs -text
30
+ *.tgz filter=lfs diff=lfs merge=lfs -text
31
+ *.wasm filter=lfs diff=lfs merge=lfs -text
32
+ *.xz filter=lfs diff=lfs merge=lfs -text
33
+ *.zip filter=lfs diff=lfs merge=lfs -text
34
+ *.zst filter=lfs diff=lfs merge=lfs -text
35
+ *tfevents* filter=lfs diff=lfs merge=lfs -text
README.md ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ datasets:
3
+ - HuggingFaceTB/smollm-corpus
4
+ language:
5
+ - en
6
+ pipeline_tag: text-generation
7
+ tags:
8
+ - SLM
9
+ - ROSA
10
+ license: mit
11
+ ---
12
+
13
+ # SRLM-1M
14
+
15
+ A **S**mall **R**OSA based **L**anguage **M**odel, of 900k parameters.
16
+
17
+ ## Training inforation
18
+
19
+ Trained on smollm-corpus fineweb-edu-dedup subset, over 16M tokens.
20
+
21
+ ## Evaluations
22
+
23
+ - Wikitext v2 byte_perplexity: 4.4553
24
+ - Arc_easy acc_norm: 28.66%
25
+ - Arc_easy acc: 26.39%
26
+ - Blimp acc: 53.1%
27
+ - Hellaswag acc: 26.65%
28
+ - Hellaswag acc_norm: 27.07%
29
+
30
+ ## How to run it
31
+
32
+ To run this model, you have to:
33
+ 1. Download the files
34
+ 2. Compile rosa.pyx using `python setup.py build_ext --inplace`
35
+ 3. Load the model with the correct configuration
36
+ 4. Call the model
config.json ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": ["SRLMForCausalLM"],
3
+ "model_type": "srlm",
4
+ "vocab_size": 5000,
5
+ "d_model": 256,
6
+ "rank_emb": 48,
7
+ "rank_rosa": 48,
8
+ "num_rosa_layers": 6,
9
+ "max_position_embeddings": 512,
10
+ "torch_dtype": "float32"
11
+ }
model.py ADDED
@@ -0,0 +1,232 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+
3
+ import torch
4
+ import torch.nn as nn
5
+ import torch.nn.functional as F
6
+ from rosa import rosa, rosa_batch
7
+
8
+ # def rosa(x):
9
+ # n = len(x)
10
+ # y = [-1] * n
11
+ # s = 2 * n + 1
12
+ # b = [None] * s
13
+ # c = [-1] * s
14
+ # d = [0] * s
15
+ # e = [-1] * s
16
+ # b[0] = {}
17
+ # g = 0
18
+ # z = 1
19
+ # for i, t in enumerate(x):
20
+ # r = z
21
+ # z += 1
22
+ # b[r] = {}
23
+ # d[r] = d[g] + 1
24
+ # p = g
25
+ # while p != -1 and t not in b[p]:
26
+ # b[p][t] = r
27
+ # p = c[p]
28
+ # if p == -1:
29
+ # c[r] = 0
30
+ # else:
31
+ # q = b[p][t]
32
+ # if d[p] + 1 == d[q]:
33
+ # c[r] = q
34
+ # else:
35
+ # u = z
36
+ # z += 1
37
+ # b[u] = b[q].copy()
38
+ # d[u] = d[p] + 1
39
+ # c[u] = c[q]
40
+ # e[u] = e[q]
41
+ # while p != -1 and b[p][t] == q:
42
+ # b[p][t] = u
43
+ # p = c[p]
44
+ # c[q] = c[r] = u
45
+ # v = g = r
46
+ # a = -1
47
+ # while v != -1:
48
+ # if d[v] > 0 and e[v] >= 0:
49
+ # a = x[e[v] + 1]
50
+ # break
51
+ # v = c[v]
52
+ # y[i] = a
53
+ # v = g
54
+ # while v != -1 and e[v] < i:
55
+ # e[v] = i
56
+ # v = c[v]
57
+ # return y
58
+
59
+
60
+ def rosa_batch_python_orig(z: torch.Tensor, alphabet: int) -> torch.Tensor:
61
+ assert z.dtype == torch.long and z.ndim == 2
62
+ zc = z.detach().contiguous().cpu().numpy()
63
+ out = rosa_batch(zc, alphabet)
64
+ return torch.from_numpy(out).to(z.device)
65
+
66
+
67
+ def rosa_batch_python(z: torch.Tensor) -> torch.Tensor:
68
+ assert z.dtype == torch.uint8 and z.ndim == 2
69
+ zc = z.detach().contiguous().cpu().to(torch.int64).numpy()
70
+ out = rosa_batch(zc, 16) # 4-bit alphabet
71
+ out = out.clip(min=0).astype("uint8")
72
+ return torch.from_numpy(out).to(z.device)
73
+
74
+
75
+ class FactorizedTiedEmbedding(nn.Module):
76
+ def __init__(self, vocab_size, d_model, rank):
77
+ super().__init__()
78
+ self.A = nn.Parameter(torch.randn(vocab_size, rank) * 0.02)
79
+ self.B = nn.Parameter(torch.randn(rank, d_model) * (1.0 / math.sqrt(rank)))
80
+
81
+ def embed(self, ids):
82
+ codes = F.embedding(ids, self.A) # (B, T, r) gather, not (V, d) matmul
83
+ return codes @ self.B # (B, T, d)
84
+
85
+ def logits(self, hidden):
86
+ r = hidden @ self.B.t() # (B, T, r)
87
+ return r @ self.A.t() # (B, T, vocab)
88
+
89
+
90
+ class rosa_emb_layer(nn.Module):
91
+ def __init__(self, V, C, rank):
92
+ super().__init__()
93
+ self.emb = FactorizedTiedEmbedding(V, C, rank)
94
+ self.V = V
95
+
96
+ def forward(self, idx):
97
+ idx = rosa_batch_python_orig(idx, self.V)
98
+ out = self.emb.embed(idx.clamp_min(0))
99
+ return out.masked_fill(idx.eq(-1).unsqueeze(-1), 0.0)
100
+
101
+
102
+ class rosa_4bit_layer(nn.Module):
103
+ def __init__(self, C: int, eps: float = 1e-5):
104
+ super().__init__()
105
+ assert C % 4 == 0
106
+ self.emb0 = nn.Parameter(torch.full((1, 1, C), -eps))
107
+ self.emb1 = nn.Parameter(torch.full((1, 1, C), eps))
108
+
109
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
110
+ B, T, C = x.shape
111
+ Cg = C // 4
112
+
113
+ b = (x.reshape(B, T, Cg, 4) > 0).to(torch.uint8)
114
+ tok2d = b[..., 0] | (b[..., 1] << 1) | (b[..., 2] << 2) | (b[..., 3] << 3)
115
+
116
+ # Orient to (B, Cg, T)
117
+ tok2d_oriented = tok2d.permute(0, 2, 1).contiguous()
118
+
119
+ tok2d_flat = tok2d_oriented.view(B * Cg, T)
120
+
121
+ idx_q_flat = rosa_batch_python(tok2d_flat)
122
+
123
+ # Reshape back to the 3D track orientation
124
+ idx_q = idx_q_flat.view(B, Cg, T)
125
+ idx_q = idx_q.transpose(1, 2).contiguous() # (B, T, Cg)
126
+
127
+ bit0 = (idx_q & 1).bool()
128
+ bit1 = ((idx_q >> 1) & 1).bool()
129
+ bit2 = ((idx_q >> 2) & 1).bool()
130
+ bit3 = ((idx_q >> 3) & 1).bool()
131
+ bits = torch.stack([bit0, bit1, bit2, bit3], dim=-1)
132
+
133
+ e0 = self.emb0.view(1, 1, Cg, 4).expand(B, T, -1, -1)
134
+ e1 = self.emb1.view(1, 1, Cg, 4).expand(B, T, -1, -1)
135
+
136
+ return torch.where(bits, e1, e0).reshape(B, T, C)
137
+
138
+
139
+ class Model(nn.Module):
140
+ def __init__(self, V, C, rank_emb, rank_rosa, num_rosa_layers):
141
+ super().__init__()
142
+ self.embedding = FactorizedTiedEmbedding(V, C, rank_emb)
143
+ self.emb_rosa = rosa_emb_layer(V, C, rank_rosa)
144
+ # Now a list of Rosa embeddings
145
+ self.emb_rosa_list = nn.ModuleList(
146
+ [rosa_4bit_layer(C) for _ in range(num_rosa_layers)]
147
+ )
148
+ self.num_rosa_layers = num_rosa_layers
149
+ self.linear_list = nn.ModuleList(
150
+ [nn.Linear(C, C) for _ in range(num_rosa_layers)]
151
+ )
152
+ self.norm_list = nn.ModuleList(
153
+ [nn.RMSNorm(C) for _ in range(num_rosa_layers)]
154
+ ) # me save params, me repeat
155
+
156
+ def forward(self, x):
157
+ x = self.embedding.embed(x) + self.emb_rosa(x)
158
+ for i in range(self.num_rosa_layers):
159
+ x = self.norm_list[i](x)
160
+ x = x + self.emb_rosa_list[i](x) # Really want to add RMSNorm here
161
+ x = x + self.linear_list[i](x)
162
+ return self.embedding.logits(x)
163
+
164
+
165
+ if __name__ == "__main__":
166
+ import time
167
+
168
+ device = "cuda" if torch.cuda.is_available() else "cpu"
169
+ print(f"Using device: {device.upper()}")
170
+
171
+ V = 5000 # Vocab size
172
+ C = 256 # Hidden Dimension
173
+ rank_emb = 48 # Factorization Rank
174
+ rank_rosa = 48
175
+ num_rosa_layers = 6 # Deeper structural depth
176
+
177
+ B, T = 16, 512 # Batch size and Context length for benchmark loop
178
+
179
+ print(f"Initializing model on {device.upper()}...")
180
+ model = Model(V, C, rank_emb, rank_rosa, num_rosa_layers).to(device)
181
+
182
+ total_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
183
+ print("\n" + "=" * 60)
184
+ print(f" TOTAL TRAINABLE FOOTPRINT: {total_params:,} parameters")
185
+ print("=" * 60)
186
+
187
+ def get_batch():
188
+ x = torch.randint(0, V, (B, T), device=device)
189
+ y = torch.roll(x, shifts=-1, dims=1)
190
+ y[:, -1] = 0
191
+ return x, y
192
+
193
+ optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
194
+ criterion = nn.CrossEntropyLoss()
195
+
196
+ print("\nStarting Benchmarking iterations with Loss tracking...")
197
+ print(
198
+ f"Config: Batch={B}, SeqLen={T}, Vocab={V}, Channels={C}, Layers={num_rosa_layers}"
199
+ )
200
+ print("-" * 60)
201
+
202
+ model.train()
203
+ total_time = 0.0
204
+ steps = 5
205
+
206
+ for step in range(1, steps + 1):
207
+ x, y = get_batch()
208
+
209
+ torch.cuda.synchronize() if device == "cuda" else None
210
+ start_time = time.perf_counter()
211
+
212
+ logits = model(x)
213
+ loss = criterion(logits.view(-1, V), y.view(-1))
214
+
215
+ optimizer.zero_grad(set_to_none=True)
216
+ loss.backward()
217
+ optimizer.step()
218
+
219
+ torch.cuda.synchronize() if device == "cuda" else None
220
+ end_time = time.perf_counter()
221
+
222
+ step_time = end_time - start_time
223
+ total_time += step_time
224
+
225
+ tokens_per_sec = (B * T) / step_time
226
+ print(
227
+ f"Step {step}/{steps} | Loss: {loss.item():.4f} | Time: {step_time * 1000:.2f}ms | Throughput: {tokens_per_sec:.2f} tok/sec"
228
+ )
229
+
230
+ print("-" * 60)
231
+ print(f"Average Benchmark Step Velocity: {(total_time / steps) * 1000:.2f} ms")
232
+ print("=" * 60)
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4174a6c5f8cfd05d552666051242362b58466c185b92112a4acbe9f3fe21271b
3
+ size 3618736
pytorch_model.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4e3cdccfe3f4dce1013150e76edab12c74f9b6ab1259d3a086bb24e6f52858c4
3
+ size 3627893
rosa.pyx ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # cython: language_level=3
2
+ # cython: boundscheck=False
3
+ # cython: wraparound=False
4
+ # cython: initializedcheck=False
5
+ # cython: nonecheck=False
6
+ # cython: cdivision=True
7
+
8
+ from libc.stdlib cimport malloc, free
9
+ from libc.string cimport memset, memcpy
10
+ from cython.parallel cimport prange
11
+ cimport cython
12
+ import numpy as np
13
+ cimport numpy as cnp
14
+
15
+ cnp.import_array()
16
+
17
+ @cython.boundscheck(False)
18
+ @cython.wraparound(False)
19
+ @cython.cdivision(True)
20
+ cdef void _rosa_row(const cnp.int64_t* x, int n, int alphabet, cnp.int64_t* y_out) noexcept nogil:
21
+ cdef int s = 2 * n + 1
22
+ cdef int *trans = <int*> malloc(<size_t>s * alphabet * sizeof(int))
23
+ cdef int *c = <int*> malloc(<size_t>s * sizeof(int))
24
+ cdef int *d = <int*> malloc(<size_t>s * sizeof(int))
25
+ cdef cnp.int64_t *e = <cnp.int64_t*> malloc(<size_t>s * sizeof(cnp.int64_t))
26
+
27
+ memset(trans, 0xFF, <size_t>s * alphabet * sizeof(int)) # all -1
28
+
29
+ cdef int g = 0, z = 1, i, t, r, p, q, u, v
30
+ cdef cnp.int64_t a
31
+
32
+ d[0] = 0
33
+ c[0] = -1
34
+ e[0] = -1
35
+
36
+ for i in range(n):
37
+ t = <int> x[i]
38
+ r = z
39
+ z += 1
40
+ d[r] = d[g] + 1
41
+ e[r] = -1
42
+ p = g
43
+ while p != -1 and trans[p * alphabet + t] == -1:
44
+ trans[p * alphabet + t] = r
45
+ p = c[p]
46
+ if p == -1:
47
+ c[r] = 0
48
+ else:
49
+ q = trans[p * alphabet + t]
50
+ if d[p] + 1 == d[q]:
51
+ c[r] = q
52
+ else:
53
+ u = z
54
+ z += 1
55
+ memcpy(trans + <size_t>u * alphabet, trans + <size_t>q * alphabet,
56
+ <size_t>alphabet * sizeof(int))
57
+ d[u] = d[p] + 1
58
+ c[u] = c[q]
59
+ e[u] = e[q]
60
+ while p != -1 and trans[p * alphabet + t] == q:
61
+ trans[p * alphabet + t] = u
62
+ p = c[p]
63
+ c[q] = u
64
+ c[r] = u
65
+ v = g = r
66
+ a = -1
67
+ while v != -1:
68
+ if d[v] > 0 and e[v] >= 0:
69
+ a = x[e[v] + 1]
70
+ break
71
+ v = c[v]
72
+ y_out[i] = a
73
+ v = g
74
+ while v != -1 and e[v] < i:
75
+ e[v] = i
76
+ v = c[v]
77
+
78
+ free(trans)
79
+ free(c)
80
+ free(d)
81
+ free(e)
82
+
83
+
84
+ def rosa_batch(cnp.int64_t[:, :] x not None, int alphabet):
85
+ """
86
+ x: (num_rows, n) int64, values in [0, alphabet)
87
+ returns: (num_rows, n) int64 numpy array, -1 where no next-distinct-symbol exists
88
+ """
89
+ cdef Py_ssize_t num_rows = x.shape[0]
90
+ cdef Py_ssize_t n = x.shape[1]
91
+ y_np = np.empty((num_rows, n), dtype=np.int64)
92
+ cdef cnp.int64_t[:, :] y = y_np
93
+ cdef Py_ssize_t i
94
+
95
+ for i in prange(num_rows, nogil=True, schedule='static'):
96
+ _rosa_row(&x[i, 0], <int>n, alphabet, &y[i, 0])
97
+
98
+ return y_np
99
+
100
+
101
+ cpdef list rosa(list x):
102
+ cdef int n = len(x)
103
+ cdef list y = [-1] * n
104
+ cdef int s = 2 * n + 1
105
+ cdef list b = [None] * s
106
+ cdef list c = [-1] * s
107
+ cdef list d = [0] * s
108
+ cdef list e = [-1] * s
109
+ b[0] = {}
110
+ cdef int g = 0
111
+ cdef int z = 1
112
+
113
+ cdef int i, t
114
+ cdef int r, p, q, u, v, a
115
+ for i in range(n):
116
+ t = x[i]
117
+ r = z
118
+ z += 1
119
+ b[r] = {}
120
+ d[r] = d[g] + 1
121
+ p = g
122
+ while p != -1 and t not in b[p]:
123
+ b[p][t] = r
124
+ p = c[p]
125
+ if p == -1:
126
+ c[r] = 0
127
+ else:
128
+ q = b[p][t]
129
+ if d[p] + 1 == d[q]:
130
+ c[r] = q
131
+ else:
132
+ u = z
133
+ z += 1
134
+ b[u] = b[q].copy()
135
+ d[u] = d[p] + 1
136
+ c[u] = c[q]
137
+ e[u] = e[q]
138
+ while p != -1 and b[p][t] == q:
139
+ b[p][t] = u
140
+ p = c[p]
141
+ c[q] = c[r] = u
142
+ v = g = r
143
+ a = -1
144
+ while v != -1:
145
+ if d[v] > 0 and e[v] >= 0:
146
+ a = x[e[v] + 1]
147
+ break
148
+ v = c[v]
149
+ y[i] = a
150
+ v = g
151
+ while v != -1 and e[v] < i:
152
+ e[v] = i
153
+ v = c[v]
154
+ return y
setup.py ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy
2
+ from Cython.Build import cythonize
3
+ from setuptools import Extension, setup
4
+
5
+ ext = Extension(
6
+ "rosa",
7
+ ["rosa.pyx"],
8
+ include_dirs=[numpy.get_include()],
9
+ extra_compile_args=["/O3", "/openmp"],
10
+ )
11
+
12
+ setup(ext_modules=cythonize([ext], language_level=3))
tiny_lm_tokenizer/tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tiny_lm_tokenizer/tokenizer_config.json ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "backend": "tokenizers",
3
+ "bos_token": "[BOS]",
4
+ "eos_token": "[EOS]",
5
+ "model_max_length": 512,
6
+ "pad_token": "[PAD]",
7
+ "tokenizer_class": "TokenizersBackend",
8
+ "unk_token": "[UNK]"
9
+ }