建26M參數(shù)微型GPT:2小時(shí)掌握Transformer核心原理與實(shí)戰(zhàn))
1. 項(xiàng)目概述為什么26M參數(shù)的GPT值得你花2小時(shí)看到“26M參數(shù)”和“GPT”這兩個(gè)詞放在一起很多人的第一反應(yīng)可能是這能干什么現(xiàn)在動(dòng)輒百億、千億參數(shù)的大模型滿天飛一個(gè)區(qū)區(qū)兩千六百萬參數(shù)的“小玩意兒”有什么訓(xùn)練的必要這正是這個(gè)教學(xué)項(xiàng)目的精妙之處——它剝離了所有關(guān)于算力的神話和資源的焦慮直指大語言模型LLM最核心的運(yùn)作原理。這個(gè)項(xiàng)目的目標(biāo)不是讓你復(fù)現(xiàn)一個(gè)能寫詩、編程、聊天的ChatGPT而是讓你在短短兩小時(shí)內(nèi)親手“捏”出一個(gè)能理解字符序列、并基于此生成新文本的微型GPT。這26M參數(shù)就像一個(gè)精密的鐘表機(jī)芯雖然體積小但齒輪注意力機(jī)制、發(fā)條前饋網(wǎng)絡(luò)、擒縱機(jī)構(gòu)層歸一化一應(yīng)俱全。通過訓(xùn)練它你將透徹理解Token是如何被嵌入成向量的自注意力機(jī)制到底在“注意”什么模型是如何通過概率預(yù)測下一個(gè)詞的這些問題的答案遠(yuǎn)比盲目調(diào)用API來得深刻。它適合所有對(duì)AI底層原理抱有好奇心但被海量數(shù)學(xué)公式和龐大工程嚇退的開發(fā)者、學(xué)生甚至產(chǎn)品經(jīng)理。你不需要八卡A100一臺(tái)有GPU的消費(fèi)級(jí)電腦甚至用CPU也能跑只是慢點(diǎn)就足夠了。這個(gè)項(xiàng)目的價(jià)值在于“教學(xué)”在于“體驗(yàn)”在于讓你獲得對(duì)Transformer架構(gòu)最直觀的、肌肉記憶般的理解。當(dāng)你看著自己從零搭建的模型從輸出亂碼到逐漸能拼湊出有意義的單詞和短句時(shí)那種成就感是無可替代的。接下來我們就拆開這個(gè)“鐘表”看看每一個(gè)零件是怎么工作的。2. 核心架構(gòu)拆解微型GPT的“五臟六腑”一個(gè)完整的GPT模型無論參數(shù)大小其核心架構(gòu)都是Transformer的解碼器Decoder堆疊。我們的26M參數(shù)版本可以看作是一個(gè)高度精簡但功能完備的“教學(xué)模型”。我們來逐一拆解它的核心組件并解釋為什么在這個(gè)規(guī)模下我們?nèi)绱嗽O(shè)計(jì)。2.1 詞表與嵌入層從字符到數(shù)字世界的橋梁首先模型不認(rèn)識(shí)單詞它只認(rèn)識(shí)數(shù)字。我們需要一個(gè)“詞典”把輸入的文本比如“hello world”轉(zhuǎn)換成一串?dāng)?shù)字ID這個(gè)過程叫Tokenization分詞。對(duì)于教學(xué)項(xiàng)目為了極致簡單我們通常采用字符級(jí)Character-level分詞。也就是說我們的詞表Vocabulary就是所有可能出現(xiàn)的字符集合例如英文小寫字母a-z、數(shù)字0-9、空格、標(biāo)點(diǎn)等。假設(shè)我們有100個(gè)字符那么詞表大小vocab_size就是100。為什么用字符級(jí)而不是更先進(jìn)的子詞Subword分詞如BPEByte Pair Encoding原因很簡單簡化。字符級(jí)分詞無需復(fù)雜的合并算法詞表極小實(shí)現(xiàn)直觀。雖然它會(huì)讓模型學(xué)習(xí)更長距離的依賴關(guān)系變得更難因?yàn)椤癶ello”需要5個(gè)token而不是1個(gè)但對(duì)于理解原理和在小數(shù)據(jù)集上快速驗(yàn)證它是完美的選擇。在26M參數(shù)規(guī)模下模型有能力學(xué)習(xí)字符間的組合規(guī)律。嵌入層Embedding Layer就是一個(gè)簡單的查找表。每個(gè)字符ID一個(gè)整數(shù)通過這個(gè)查找表被映射為一個(gè)固定長度的稠密向量比如128維。這個(gè)向量就是該字符的“分布式表示”它會(huì)在訓(xùn)練過程中被不斷調(diào)整使得語義相近的字符如‘a(chǎn)’和‘A’在向量空間中的位置也接近。2.2 核心引擎Transformer解碼器塊這是模型的心臟。一個(gè)解碼器塊主要由以下部分組成我們的微型GPT可能會(huì)堆疊4到6個(gè)這樣的塊自注意力機(jī)制Causal Self-Attention這是Transformer的靈魂。它允許序列中的每個(gè)“位置”去查看序列中所有之前的位置因果掩碼確保它不能“偷看”未來并計(jì)算一個(gè)加權(quán)和。簡單來說模型在預(yù)測下一個(gè)字符時(shí)會(huì)問自己“根據(jù)我已經(jīng)看到的這些字符哪一個(gè)或哪幾個(gè)對(duì)預(yù)測下一個(gè)字符最重要” 權(quán)重就是通過查詢Query、鍵Key、值Value三組向量計(jì)算得出的。在微型GPT中我們通常使用多頭注意力比如4個(gè)頭每個(gè)頭學(xué)習(xí)不同方面的依賴關(guān)系例如一個(gè)頭關(guān)注語法結(jié)構(gòu)一個(gè)頭關(guān)注詞性搭配。前饋神經(jīng)網(wǎng)絡(luò)Feed-Forward Network注意力層的輸出會(huì)經(jīng)過一個(gè)簡單的全連接網(wǎng)絡(luò)通常包含一個(gè)放大和縮小的過程例如從128維放大到512維再縮回128維。它的作用是為每個(gè)位置的特征提供一次非線性變換和特征混合增加模型的表達(dá)能力。層歸一化LayerNorm與殘差連接Residual Connection這是訓(xùn)練深層網(wǎng)絡(luò)穩(wěn)定的關(guān)鍵。每個(gè)子層注意力、前饋之前或之后都會(huì)應(yīng)用層歸一化將數(shù)據(jù)分布拉回穩(wěn)定狀態(tài)。殘差連接則是將子層的輸入直接加到其輸出上輸出 子層(輸入) 輸入。這有效地解決了深度網(wǎng)絡(luò)中的梯度消失問題讓信息可以暢通無阻地穿越很多層。2.3 輸出層從特征到概率經(jīng)過多個(gè)解碼器塊處理后我們得到了每個(gè)位置的一個(gè)高級(jí)特征向量。最后我們需要將這個(gè)向量映射回詞表空間。我們使用一個(gè)線性層Linear Layer將特征向量的維度如128投影到詞表大小100。這個(gè)操作會(huì)為詞表中的每個(gè)字符生成一個(gè)“分?jǐn)?shù)”logits。然后我們使用Softmax函數(shù)將這些分?jǐn)?shù)轉(zhuǎn)換為概率分布。模型預(yù)測的下一個(gè)字符就是從這個(gè)概率分布中采樣或取概率最大的那個(gè)得到的。參數(shù)估算26M參數(shù)從哪里來我們來粗略算一下。假設(shè)詞表大小100嵌入維度128那么嵌入層參數(shù)約100 * 128 12.8K。一個(gè)解碼器塊的主要參數(shù)在注意力層和前饋層注意力層的QKV投影矩陣和前饋層的兩個(gè)線性層。如果堆疊6個(gè)塊每個(gè)塊參數(shù)約4M總共就在24M左右加上最后的輸出層總數(shù)就接近26M。這是一個(gè)非常緊湊但足以演示Transformer核心機(jī)制的設(shè)計(jì)。3. 實(shí)戰(zhàn)兩小時(shí)訓(xùn)練流水線全解析理論清晰后我們進(jìn)入實(shí)戰(zhàn)環(huán)節(jié)。這兩小時(shí)需要高效利用每一步都有其目的和技巧。3.1 環(huán)境準(zhǔn)備與數(shù)據(jù)加載10分鐘環(huán)境推薦使用Python和PyTorch。安裝命令極其簡單pip install torch。如果你有NVIDIA GPU確保安裝了對(duì)應(yīng)版本的CUDA和cuDNNPyTorch安裝時(shí)會(huì)自動(dòng)匹配。數(shù)據(jù)選擇一個(gè)小而經(jīng)典的數(shù)據(jù)集。莎士比亞全集、維基百科的某個(gè)小條目、甚至是幾篇新聞文章都可以。數(shù)據(jù)量在1MB到10MB之間為宜。太大的數(shù)據(jù)兩小時(shí)處理不完太小則模型學(xué)不到模式。這里我們以“莎士比亞作品”文本為例。import torch import torch.nn as nn import torch.nn.functional as F import requests # 下載數(shù)據(jù) url https://raw.githubusercontent.com/karpathy/char-rnn/master/data/tinyshakespeare/input.txt text requests.get(url).text print(f數(shù)據(jù)長度: {len(text)} 字符) print(text[:500]) # 預(yù)覽前500個(gè)字符數(shù)據(jù)預(yù)處理構(gòu)建字符級(jí)詞表。# 創(chuàng)建字符到索引和索引到字符的映射 chars sorted(list(set(text))) vocab_size len(chars) print(f詞表大小: {vocab_size}) print(.join(chars)) stoi {ch:i for i,ch in enumerate(chars)} # 字符 - 索引 itos {i:ch for i,ch in enumerate(chars)} # 索引 - 字符 encode lambda s: [stoi[c] for c in s] # 編碼函數(shù) decode lambda l: .join([itos[i] for i in l]) # 解碼函數(shù) # 將整個(gè)文本編碼為張量 data torch.tensor(encode(text), dtypetorch.long) print(data.shape, data.dtype)3.2 模型定義與初始化20分鐘現(xiàn)在我們根據(jù)第二部分的設(shè)計(jì)用PyTorch定義模型。這里給出一個(gè)高度精簡但結(jié)構(gòu)清晰的實(shí)現(xiàn)框架。import torch.nn as nn import math class CausalSelfAttention(nn.Module): 帶因果掩碼的多頭自注意力 def __init__(self, embed_dim, num_heads): super().__init__() assert embed_dim % num_heads 0 self.num_heads num_heads self.head_dim embed_dim // num_heads # 通常將Q,K,V投影合并到一個(gè)線性層中提升效率 self.c_attn nn.Linear(embed_dim, 3 * embed_dim) # 輸出Q, K, V self.c_proj nn.Linear(embed_dim, embed_dim) # 輸出投影 # 因果掩碼確保位置i只能看到i的位置 self.register_buffer(bias, torch.tril(torch.ones(block_size, block_size)) .view(1, 1, block_size, block_size)) def forward(self, x): B, T, C x.size() # 批大小序列長度特征維度 # 計(jì)算Q, K, V qkv self.c_attn(x) q, k, v qkv.split(self.embed_dim, dim2) # 重塑為多頭 k k.view(B, T, self.num_heads, self.head_dim).transpose(1, 2) q q.view(B, T, self.num_heads, self.head_dim).transpose(1, 2) v v.view(B, T, self.num_heads, self.head_dim).transpose(1, 2) # 注意力計(jì)算 (縮放點(diǎn)積注意力) att (q k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1))) att att.masked_fill(self.bias[:,:,:T,:T] 0, float(-inf)) att F.softmax(att, dim-1) y att v # 合并多頭輸出 y y.transpose(1, 2).contiguous().view(B, T, C) y self.c_proj(y) return y class Block(nn.Module): 一個(gè)Transformer解碼器塊 def __init__(self, embed_dim, num_heads): super().__init__() self.ln1 nn.LayerNorm(embed_dim) self.attn CausalSelfAttention(embed_dim, num_heads) self.ln2 nn.LayerNorm(embed_dim) self.mlp nn.Sequential( nn.Linear(embed_dim, 4 * embed_dim), # 放大 nn.GELU(), # 激活函數(shù) nn.Linear(4 * embed_dim, embed_dim), # 縮小 ) def forward(self, x): # 殘差連接 層歸一化Pre-Norm結(jié)構(gòu)更穩(wěn)定 x x self.attn(self.ln1(x)) x x self.mlp(self.ln2(x)) return x class MiniGPT(nn.Module): 我們的26M參數(shù)微型GPT def __init__(self, vocab_size, embed_dim256, block_size256, num_layers6, num_heads8): super().__init__() self.block_size block_size self.token_embedding nn.Embedding(vocab_size, embed_dim) self.position_embedding nn.Embedding(block_size, embed_dim) # 位置編碼 self.blocks nn.Sequential(*[Block(embed_dim, num_heads) for _ in range(num_layers)]) self.ln_f nn.LayerNorm(embed_dim) self.lm_head nn.Linear(embed_dim, vocab_size) # 語言模型頭 # 參數(shù)初始化很重要 self.apply(self._init_weights) def _init_weights(self, module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean0.0, std0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, std0.02) def forward(self, idx, targetsNone): B, T idx.shape # 詞嵌入 位置嵌入 tok_emb self.token_embedding(idx) # (B,T,embed_dim) pos torch.arange(0, T, deviceidx.device) # (T) pos_emb self.position_embedding(pos) # (T, embed_dim) x tok_emb pos_emb # (B,T,embed_dim) x self.blocks(x) x self.ln_f(x) logits self.lm_head(x) # (B, T, vocab_size) loss None if targets is not None: B, T, C logits.shape logits logits.view(B*T, C) targets targets.view(B*T) loss F.cross_entropy(logits, targets) return logits, loss def generate(self, idx, max_new_tokens): 自回歸生成文本 for _ in range(max_new_tokens): # 裁剪上下文到block_size idx_cond idx[:, -self.block_size:] # 前向傳播 logits, _ self(idx_cond) # 聚焦最后一個(gè)時(shí)間步 logits logits[:, -1, :] # (B, C) # 用溫度采樣增加隨機(jī)性 probs F.softmax(logits, dim-1) idx_next torch.multinomial(probs, num_samples1) # (B, 1) # 拼接生成結(jié)果 idx torch.cat((idx, idx_next), dim1) return idx初始化技巧注意代碼中的_init_weights方法。用較小的正態(tài)分布std0.02初始化權(quán)重是訓(xùn)練Transformer模型的標(biāo)準(zhǔn)做法這有助于在訓(xùn)練初期保持激活值的穩(wěn)定性。將偏置bias初始化為0也是常見操作。3.3 訓(xùn)練循環(huán)與超參數(shù)設(shè)置80分鐘這是最耗時(shí)的部分但代碼結(jié)構(gòu)很清晰。我們將數(shù)據(jù)分割成訓(xùn)練集和驗(yàn)證集90%/10%并創(chuàng)建數(shù)據(jù)加載器。# 分割數(shù)據(jù) n int(0.9 * len(data)) train_data data[:n] val_data data[n:] def get_batch(split): 隨機(jī)獲取一個(gè)小批量的數(shù)據(jù) data train_data if split train else val_data ix torch.randint(len(data) - block_size, (batch_size,)) x torch.stack([data[i:iblock_size] for i in ix]) y torch.stack([data[i1:iblock_size1] for i in ix]) return x, y # 超參數(shù)設(shè)置這是關(guān)鍵 batch_size 32 # 每次訓(xùn)練輸入的樣本數(shù) block_size 256 # 模型能處理的最大上下文長度 learning_rate 3e-4 # 學(xué)習(xí)率Adam優(yōu)化器的黃金標(biāo)準(zhǔn) max_iters 5000 # 最大迭代步數(shù)控制訓(xùn)練時(shí)間 eval_interval 500 # 每多少步評(píng)估一次 eval_iters 200 # 評(píng)估時(shí)使用的迭代次數(shù)用于估算平均損失 # 初始化模型、優(yōu)化器 model MiniGPT(vocab_sizevocab_size, embed_dim256, block_sizeblock_size, num_layers6, num_heads8) print(f模型參數(shù)量: {sum(p.numel() for p in model.parameters())/1e6:.2f}M) model model.to(device) # 如果有GPU移到GPU上 optimizer torch.optim.AdamW(model.parameters(), lrlearning_rate) # AdamW是Adam的改進(jìn)版帶權(quán)重衰減 torch.no_grad() def estimate_loss(): 估算訓(xùn)練集和驗(yàn)證集的損失 out {} model.eval() for split in [train, val]: losses torch.zeros(eval_iters) for k in range(eval_iters): X, Y get_batch(split) X, Y X.to(device), Y.to(device) _, loss model(X, Y) losses[k] loss.item() out[split] losses.mean() model.train() return out # 訓(xùn)練循環(huán) for iter in range(max_iters): # 每隔一段時(shí)間評(píng)估一次 if iter % eval_interval 0 or iter max_iters - 1: losses estimate_loss() print(f第{iter}步: 訓(xùn)練損失 {losses[train]:.4f}, 驗(yàn)證損失 {losses[val]:.4f}) # 獲取一個(gè)批量數(shù)據(jù) xb, yb get_batch(train) xb, yb xb.to(device), yb.to(device) # 前向傳播計(jì)算損失 _, loss model(xb, yb) # 反向傳播更新參數(shù) optimizer.zero_grad(set_to_noneTrue) # 清零梯度set_to_noneTrue可以節(jié)省內(nèi)存 loss.backward() optimizer.step() print(訓(xùn)練完成)超參數(shù)設(shè)置心得學(xué)習(xí)率3e-4對(duì)于Adam優(yōu)化器這是一個(gè)經(jīng)過大量實(shí)踐驗(yàn)證的、近乎“萬能”的起始學(xué)習(xí)率。對(duì)于我們的微型模型這個(gè)值非常安全。批量大小32在GPU內(nèi)存允許的范圍內(nèi)批量大小越大梯度估計(jì)越準(zhǔn)訓(xùn)練越穩(wěn)定。但太大也可能導(dǎo)致泛化能力下降。32是一個(gè)兼顧速度和穩(wěn)定性的常見值。上下文長度256這限制了模型能“看到”多遠(yuǎn)的過去。對(duì)于字符級(jí)模型256個(gè)字符大約是一段話的長度足以讓模型學(xué)習(xí)到基本的單詞拼寫和短句結(jié)構(gòu)。最大迭代步數(shù)5000在兩小時(shí)的限制下我們需要估算每一步的時(shí)間。在消費(fèi)級(jí)GPU上5000步大約需要60-80分鐘留出評(píng)估和生成的時(shí)間。3.4 文本生成與效果評(píng)估10分鐘訓(xùn)練結(jié)束后最激動(dòng)人心的時(shí)刻到了讓模型“開口說話”。我們提供一個(gè)起始字符串context讓模型自回歸地生成后續(xù)文本。# 將模型設(shè)置為評(píng)估模式 model.eval() # 生成文本 context torch.tensor([encode(KING: )], dtypetorch.long, devicedevice) # 以“KING: ”開頭 generated_ids model.generate(context, max_new_tokens500)[0].tolist() generated_text decode(generated_ids) print(generated_text)生成策略解析代碼中使用了torch.multinomial進(jìn)行采樣。這意味著模型不是永遠(yuǎn)選擇概率最高的那個(gè)字符貪婪搜索而是根據(jù)概率分布隨機(jī)采樣。這能帶來更多樣化、更有趣的生成結(jié)果但有時(shí)也會(huì)產(chǎn)生不合邏輯的內(nèi)容。你可以嘗試“溫度”Temperature采樣。在Softmax之前將logits除以一個(gè)溫度系數(shù)T。T 1如1.2會(huì)使分布更平滑生成更隨機(jī)、更有創(chuàng)造性的文本T 1如0.8會(huì)使分布更尖銳生成更確定、更保守的文本。代碼中可以這樣修改temperature 0.8 logits logits / temperature probs F.softmax(logits, dim-1)另一種高級(jí)策略是Top-k 或 Top-p核采樣即只從概率最高的k個(gè)候選詞中采樣或從累積概率達(dá)到p的最小候選詞集合中采樣。這能有效避免采樣到概率極低的奇怪字符。評(píng)估生成質(zhì)量沒有絕對(duì)標(biāo)準(zhǔn)但你可以觀察字符級(jí)連貫性生成的單詞看起來像英文單詞嗎如“helllo”是錯(cuò)的“hello”是對(duì)的。語法結(jié)構(gòu)有沒有出現(xiàn)大寫字母開頭、句號(hào)結(jié)尾的短句上下文一致性如果輸入是“KING: ”生成的內(nèi)容是否像戲劇臺(tái)詞模型是否學(xué)到了訓(xùn)練數(shù)據(jù)莎士比亞的風(fēng)格4. 避坑指南與性能優(yōu)化實(shí)錄在實(shí)際操作中你幾乎一定會(huì)遇到下面這些問題。這里記錄了我的踩坑經(jīng)驗(yàn)和解決方案。4.1 訓(xùn)練不收斂或損失為NaN這是新手最常見的問題。檢查初始化確保你按照示例代碼進(jìn)行了正確的權(quán)重初始化std0.02。錯(cuò)誤的初始化如std過大會(huì)導(dǎo)致激活值爆炸梯度變成NaN。檢查學(xué)習(xí)率3e-4對(duì)AdamW通常是安全的。如果你手動(dòng)調(diào)整了模型架構(gòu)如大幅增加embed_dim可能需要微調(diào)學(xué)習(xí)率。一個(gè)簡單的策略是使用學(xué)習(xí)率預(yù)熱Warmup在訓(xùn)練的前幾百步將學(xué)習(xí)率從0線性增加到設(shè)定值這有助于訓(xùn)練初期穩(wěn)定。檢查梯度裁剪Gradient Clipping在loss.backward()之后optimizer.step()之前加入一行代碼torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)。這可以防止梯度爆炸將梯度向量的范數(shù)norm限制在1.0以內(nèi)是訓(xùn)練RNN和Transformer的常用穩(wěn)定技巧。檢查輸入數(shù)據(jù)確保你的輸入張量idx和targets的 dtype 是torch.long整數(shù)類型而不是浮點(diǎn)數(shù)。交叉熵?fù)p失函數(shù)要求索引是整數(shù)。4.2 模型過擬合與欠擬合過擬合表現(xiàn)訓(xùn)練損失持續(xù)下降但驗(yàn)證損失在某個(gè)點(diǎn)后開始上升。模型“死記硬背”了訓(xùn)練數(shù)據(jù)而無法泛化到新數(shù)據(jù)。解決方案增加數(shù)據(jù)量是最根本的。此外可以嘗試Dropout在注意力層和前饋層之后添加Dropout。例如在Block的forward函數(shù)中x x F.dropout(self.attn(self.ln1(x)), p0.1)。權(quán)重衰減Weight Decay我們使用的AdamW優(yōu)化器已經(jīng)內(nèi)置了權(quán)重衰減通過weight_decay參數(shù)設(shè)置通常為0.01或0.1這相當(dāng)于L2正則化能有效防止過擬合。早停Early Stopping監(jiān)控驗(yàn)證損失當(dāng)其在連續(xù)多個(gè)評(píng)估周期內(nèi)不再下降時(shí)停止訓(xùn)練。欠擬合表現(xiàn)訓(xùn)練損失和驗(yàn)證損失都很高且下降緩慢。模型能力不足無法捕捉數(shù)據(jù)中的模式。解決方案增加模型容量更多層、更大的embed_dim、延長訓(xùn)練時(shí)間增加max_iters或者檢查模型架構(gòu)是否有錯(cuò)誤例如注意力掩碼是否正確殘差連接是否生效。4.3 生成文本質(zhì)量差輸出重復(fù)或陷入循環(huán)這是采樣策略的問題。貪婪搜索總是選最高概率極易導(dǎo)致循環(huán)。務(wù)必使用采樣sampling而非貪婪搜索。同時(shí)可以嘗試降低溫度如0.7或使用Top-p采樣如p0.9來平衡生成的質(zhì)量和多樣性。生成亂碼或非字符檢查你的詞表itos和解碼函數(shù)decode。確保模型輸出的索引在詞表范圍內(nèi)。有時(shí)在生成時(shí)模型可能輸出超出范圍的索引這通常是因?yàn)镾oftmax前的logits有問題或者采樣函數(shù)出錯(cuò)。4.4 訓(xùn)練速度慢兩小時(shí)是目標(biāo)但如果你的機(jī)器只有CPU可能會(huì)超時(shí)。使用GPU這是最大的加速手段。確保你的PyTorch安裝了CUDA版本并使用.to(device)將模型和數(shù)據(jù)移到GPU上。降低精度使用混合精度訓(xùn)練Mixed Precision Training。這可以顯著減少GPU顯存占用并加快計(jì)算。PyTorch中可以使用torch.cuda.amp自動(dòng)混合精度模塊。調(diào)整批量大小在GPU顯存允許的前提下盡可能增大batch_size。更大的批次意味著更少的迭代步數(shù)就能看完一遍數(shù)據(jù)并且梯度估計(jì)更準(zhǔn)確。減少評(píng)估頻率將eval_interval設(shè)得大一些如1000步減少驗(yàn)證集上的前向傳播次數(shù)這些是不更新梯度的純耗時(shí)間。5. 從教學(xué)模型到實(shí)用化的思考完成這個(gè)26M參數(shù)GPT的訓(xùn)練后你已經(jīng)掌握了Transformer語言模型最核心的構(gòu)建、訓(xùn)練和生成流程。但這只是一個(gè)起點(diǎn)。如果你想走向更實(shí)用、更強(qiáng)大的模型以下方向值得深入1. 分詞器的升級(jí)將字符級(jí)分詞換成子詞分詞如BPE。這能極大提升模型處理常見單詞和未知詞的效率。你可以使用Hugging Face的tokenizers庫在更大的語料上訓(xùn)練一個(gè)BPE分詞器然后替換掉項(xiàng)目中的簡單字符詞表。2. 數(shù)據(jù)與規(guī)模的擴(kuò)展嘗試用更大的數(shù)據(jù)集如幾十MB的文本訓(xùn)練一個(gè)參數(shù)稍多如100M的模型。你會(huì)發(fā)現(xiàn)模型開始能生成更長的、語法更正確的段落甚至表現(xiàn)出初步的“主題”一致性。這就是“規(guī)模定律”Scaling Law的直觀體現(xiàn)更多的數(shù)據(jù)和參數(shù)會(huì)涌現(xiàn)出更復(fù)雜的能力。3. 引入更先進(jìn)的架構(gòu)細(xì)節(jié) -旋轉(zhuǎn)位置編碼RoPE替換掉簡單的可學(xué)習(xí)位置嵌入RoPE能更好地處理長序列也是LLaMA、GPT-4等主流模型的選擇。 -SwiGLU/RMSNorm嘗試使用SwiGLU激活函數(shù)和RMSNorm層歸一化這些是近年來被證明更有效的變體。 -Flash Attention如果你的GPU支持使用Flash Attention實(shí)現(xiàn)可以大幅加速注意力計(jì)算并降低內(nèi)存占用讓你能處理更長的序列。4. 指令微調(diào)Instruction Tuning與對(duì)齊我們的模型現(xiàn)在只是一個(gè)“續(xù)寫模型”。要讓它能回答問題、遵循指令你需要進(jìn)行指令微調(diào)。這需要收集或構(gòu)造大量的(指令, 輸入, 輸出)三元組數(shù)據(jù)在預(yù)訓(xùn)練好的模型基礎(chǔ)上進(jìn)行有監(jiān)督微調(diào)SFT。這之后還可以通過人類反饋強(qiáng)化學(xué)習(xí)RLHF進(jìn)一步對(duì)齊模型的輸出與人類偏好。這個(gè)2小時(shí)的項(xiàng)目就像給你一張地圖和一把鑰匙。地圖是Transformer的架構(gòu)圖鑰匙是親手運(yùn)行代碼的體驗(yàn)?,F(xiàn)在你已經(jīng)站在了大語言模型世界的大門口門后的廣闊天地等待你去探索。真正的挑戰(zhàn)和樂趣始于你開始根據(jù)自己的想法修改架構(gòu)處理新數(shù)據(jù)解決新問題的那一刻。