File size: 1,578 Bytes
8b64ba0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 | """TinyLM — from-scratch transformer used for the VIVEK series.
Character-level language model: embedding + positional encoding +
causal Transformer encoder blocks + output head."""
import torch
import torch.nn as nn
class TinyLM(nn.Module):
def __init__(self, vocab_size, d_model=128, n_layers=2, n_heads=4,
ff=256, max_len=256, dropout=0.1):
super().__init__()
self.config = dict(vocab_size=vocab_size, d_model=d_model,
n_layers=n_layers, n_heads=n_heads, ff=ff,
max_len=max_len, dropout=dropout)
self.tok = nn.Embedding(vocab_size, d_model)
self.pos = nn.Embedding(max_len, d_model)
layer = nn.TransformerEncoderLayer(
d_model=d_model, nhead=n_heads, dim_feedforward=ff,
dropout=dropout, activation="gelu", batch_first=True,
norm_first=True)
self.enc = nn.TransformerEncoder(layer, num_layers=n_layers)
self.head = nn.Linear(d_model, vocab_size)
def forward(self, x):
B, L = x.shape
h = self.tok(x) + self.pos(torch.arange(L, device=x.device))
mask = torch.triu(torch.ones(L, L, dtype=torch.bool), diagonal=1)
h = self.enc(h, mask=mask)
return self.head(h)
@staticmethod
def from_config(path="config.json"):
import json
cfg = json.load(open(path))
return TinyLM(**{k: cfg[k] for k in
("vocab_size", "d_model", "n_layers", "n_heads",
"ff", "max_len", "dropout") if k in cfg})
|