vivek0-chat / model.py
vivekkopthsd's picture
publish vivek0-chat (professional card v2)
406cb5e verified
Raw
History Blame Contribute Delete
1.58 kB
"""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})