Salesforce/wikitext
Viewer β’ Updated β’ 3.71M β’ 1.88M β’ 814
Live training β updated automatically every 250 steps.
GPT-2 weight transfer: The 12 transformer blocks (attention + MLP + LayerNorm) are initialized from pretrained GPT-2 small weights. Only wte/wpe/lm_head are trained from scratch (byte-level vocabulary). This skips ~8000 steps of learning English structure β quality shortcut without cheating.
| Parameter | Value |
|---|---|
| Layers | 12 |
| Attention Heads | 12 |
| Embedding Dim | 768 |
| Feed-Forward Dim | 3072 |
| Context Length | 256 tokens |
| Vocabulary Size | 256 (raw bytes) |
| Total Parameters | 85.37M |
| Weight Tying | Yes (wte = lm_head) |
| Transfer Source | GPT-2 small (12 blocks + ln_f) |
None. Raw UTF-8 bytes (vocab=256). No tokenizer files needed.
| Field | Value |
|---|---|
| Dataset | WikiText-103-raw (~541M tokens via memmap) |
| Transfer | GPT-2 small transformer blocks |
| Optimizer | AdamW (fused=True) |
| Betas | (0.9, 0.95) |
| Weight Decay | 0.1 |
| LR Schedule | Cosine decay |
| LR Max | 2e-4 (lower than Phase 8 for fine-tuning) |
| LR Min | 5e-6 |
| Grad Clip | 1.0 |
| Total Steps | 8,000 |
| Batch Size | 1 |
| Context Length | 256 |
| Hardware | Intel Sapphire Rapids, AMX BF16, SDPA |
| Metric | Value |
|---|---|
| Step | 500 / 8,000 (6.2%) |
| Loss | 2.5844 |
| Perplexity | 13.26 |
| Initial Loss | 9.0427 |
| Phase | Architecture Change | Notes |
|---|---|---|
| Phase 1-6 | Baseline GPT | Progressive training |
| Phase 7 | Full 12L/12H/768D | Stable architecture |
| Phase 8 | +SDPA +fused AdamW | 1.26 sps, AMX BF16 |
| Phase 9 | +GPT-2 block transfer | Skips ~8k steps of language learning |
F.scaled_dot_product_attention β fused causal attention kernelDNNL_DEFAULT_FPMATH_MODE=BF16 β Intel tiles on Sapphire Rapidsimport torch, torch.nn as nn, torch.nn.functional as F
N_VOCAB=256; N_CTX=256; N_EMBD=768; N_HEAD=12; N_LAYER=12; N_FF=3072
POS_IDS=torch.arange(N_CTX)
class CausalSelfAttention(nn.Module):
def __init__(self):
super().__init__()
self.c_attn=nn.Linear(N_EMBD,3*N_EMBD,bias=False)
self.c_proj=nn.Linear(N_EMBD,N_EMBD,bias=False)
def forward(self,x):
B,T,C=x.shape; D=C//N_HEAD
q,k,v=self.c_attn(x).split(C,dim=2)
q=q.view(B,T,N_HEAD,D).transpose(1,2)
k=k.view(B,T,N_HEAD,D).transpose(1,2)
v=v.view(B,T,N_HEAD,D).transpose(1,2)
y=F.scaled_dot_product_attention(q,k,v,is_causal=True)
return self.c_proj(y.transpose(1,2).contiguous().view(B,T,C))
class MLP(nn.Module):
def __init__(self):
super().__init__()
self.c_fc=nn.Linear(N_EMBD,N_FF,bias=False)
self.c_proj=nn.Linear(N_FF,N_EMBD,bias=False)
def forward(self,x): return self.c_proj(F.gelu(self.c_fc(x)))
class Block(nn.Module):
def __init__(self):
super().__init__()
self.ln1=nn.LayerNorm(N_EMBD); self.attn=CausalSelfAttention()
self.ln2=nn.LayerNorm(N_EMBD); self.mlp=MLP()
def forward(self,x):
x=x+self.attn(self.ln1(x)); x=x+self.mlp(self.ln2(x)); return x
class GPT(nn.Module):
def __init__(self):
super().__init__()
self.wte=nn.Embedding(N_VOCAB,N_EMBD); self.wpe=nn.Embedding(N_CTX,N_EMBD)
self.blocks=nn.ModuleList([Block() for _ in range(N_LAYER)])
self.ln_f=nn.LayerNorm(N_EMBD); self.lm_head=nn.Linear(N_EMBD,N_VOCAB,bias=False)
self.lm_head.weight=self.wte.weight
def forward(self,idx,targets=None):
B,T=idx.shape
x=self.wte(idx)+self.wpe(POS_IDS[:T])
for block in self.blocks: x=block(x)
logits=self.lm_head(self.ln_f(x))
loss=F.cross_entropy(logits.view(-1,N_VOCAB),targets.view(-1)) if targets is not None else None
return logits,loss
model=GPT()
ckpt=torch.load('phase9_latest.pt',map_location='cpu',weights_only=True)
model.load_state_dict(ckpt['model'])
model.eval()
def generate(model, prompt_bytes, max_new=200, temperature=0.8):
ctx=torch.tensor([list(prompt_bytes)],dtype=torch.long)
for _ in range(max_new):
ctx_crop=ctx[:,-N_CTX:]
with torch.no_grad():
logits,_=model(ctx_crop)
logits=logits[0,-1,:]/temperature
probs=torch.softmax(logits,dim=-1)
next_tok=torch.multinomial(probs,1)
ctx=torch.cat([ctx,next_tok.unsqueeze(0)],dim=1)
return bytes(ctx[0].tolist()).decode('utf-8',errors='replace')
print(generate(model,b"The history of"))
phase9_latest.pt β model + optimizer state (resumes training exactly)