AbstractPhil commited on
Commit
f6c03ba
·
verified ·
1 Parent(s): d73e959

mini-beatrix-2s automodel: embedding.py (mission final 16.101B, alephllm 0.8.6)

Browse files
Files changed (1) hide show
  1. embedding.py +44 -0
embedding.py ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Embeddings.
2
+
3
+ TrigramByteEmbedding — the validated composed byte embedding:
4
+ e_t = E0[x_t] + E1[x_{t-1}] + E2[x_{t-2}] + P[t]
5
+ with the PAD LAW built in permanently: the shift tables carry a dedicated
6
+ pad row (index 256). Padding trigram shifts with a legal byte conflates
7
+ real history with sequence starts and starves address consumption
8
+ (measured +.05..+.11 on repair) — the fix ships on, not opt-in.
9
+
10
+ TokenEmbedding — plain table + positions for BPE crafts.
11
+ """
12
+ from __future__ import annotations
13
+
14
+ import torch
15
+ import torch.nn as nn
16
+ import torch.nn.functional as F
17
+
18
+ BYTE_VOCAB = 256
19
+ PAD_ROW = 256 # dedicated pad index in the shift tables (size 257)
20
+
21
+
22
+ class TrigramByteEmbedding(nn.Module):
23
+ def __init__(self, d: int, context: int):
24
+ super().__init__()
25
+ self.emb0 = nn.Embedding(BYTE_VOCAB, d)
26
+ self.emb1 = nn.Embedding(BYTE_VOCAB + 1, d) # + pad row
27
+ self.emb2 = nn.Embedding(BYTE_VOCAB + 1, d)
28
+ self.pos = nn.Parameter(0.01 * torch.randn(1, context, d))
29
+
30
+ def forward(self, idx):
31
+ x = self.emb0(idx) \
32
+ + self.emb1(F.pad(idx, (1, 0), value=PAD_ROW)[:, :-1]) \
33
+ + self.emb2(F.pad(idx, (2, 0), value=PAD_ROW)[:, :-2])
34
+ return x + self.pos[:, : idx.shape[1]]
35
+
36
+
37
+ class TokenEmbedding(nn.Module):
38
+ def __init__(self, vocab: int, d: int, context: int):
39
+ super().__init__()
40
+ self.emb = nn.Embedding(vocab, d)
41
+ self.pos = nn.Parameter(0.01 * torch.randn(1, context, d))
42
+
43
+ def forward(self, idx):
44
+ return self.emb(idx) + self.pos[:, : idx.shape[1]]