MiniGPT Phase 9 β€” 85.37M Parameter Byte-Level GPT + GPT-2 Transfer

Live training β€” updated automatically every 250 steps.

What's New in Phase 9

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.

Architecture

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)

Tokenizer

None. Raw UTF-8 bytes (vocab=256). No tokenizer files needed.

Training

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

Current Progress

Metric Value
Step 500 / 8,000 (6.2%)
Loss 2.5844
Perplexity 13.26
Initial Loss 9.0427

Phase History

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

Optimizations

  • GPT-2 transfer: Pretrained transformer blocks β€” language structure for free
  • SDPA: F.scaled_dot_product_attention β€” fused causal attention kernel
  • Fused AdamW: AMX-accelerated optimizer
  • AMX BF16: DNNL_DEFAULT_FPMATH_MODE=BF16 β€” Intel tiles on Sapphire Rapids
  • Intel Sapphire Rapids: 260MB L3 cache keeps per-layer weights hot

Loading

import 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"))

Checkpoint

  • phase9_latest.pt β€” model + optimizer state (resumes training exactly)
  • Saved every 250 steps automatically
Downloads last month
2,307
Safetensors
Model size
10.8M params
Tensor type
F32
Β·
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support

Dataset used to train xerxesxi/minigpt-phase7