包如何掌控數(shù)據(jù)到服務全鏈路)
1. 從零手搓AI工程為什么我不建議你直接調(diào)包很多人一上來就想搞個大模型應用第一反應是找API、裝框架、跑通一個Demo然后覺得自己“入門AI工程”了。我剛開始也這么干過結(jié)果踩了一堆坑接口一改就崩、成本失控、延遲高得離譜、出了問題完全不知道從哪查。后來我才意識到AI工程的核心能力不是“會調(diào)包”而是理解從數(shù)據(jù)到模型再到服務的整條鏈路。這就是我決定從零開始手搓一遍的原因?!癮i-engineering-from-scratch”這個方向說白了就是不依賴高級封裝用最基礎(chǔ)的工具把AI系統(tǒng)的關(guān)鍵環(huán)節(jié)自己實現(xiàn)一遍。它解決的不是“能不能跑通”的問題而是“跑通之后你能不能掌控它”的問題。適合誰看如果你已經(jīng)會寫Python、懂一點線性代數(shù)和概率但每次遇到模型效果不好、推理太慢、顯存爆了就只能上網(wǎng)搜答案那這套東西就是給你準備的。我下面會按實際動手的順序把數(shù)據(jù)管線、模型訓練、推理優(yōu)化、服務部署這幾個環(huán)節(jié)拆開講每個環(huán)節(jié)都給出可復現(xiàn)的代碼思路和參數(shù)選擇的理由。2. 整體設(shè)計思路為什么我要把“調(diào)包”拆成“手搓”2.1 先想清楚手搓到底練的是什么很多人對手搓有誤解覺得是要自己寫一個PyTorch出來。不是的。手搓的目的是讓你對每一層抽象都有“掀開蓋子”的能力。比如你知道m(xù)odel.fit()背后發(fā)生了什么嗎梯度是怎么累積的學習率在哪個時刻衰減數(shù)據(jù)是怎么分批的這些細節(jié)在調(diào)包時全是黑盒但一旦你自己用NumPy實現(xiàn)一遍反向傳播再用PyTorch對照驗證你對訓練過程的理解會完全不一樣。我給自己定的原則是能用基礎(chǔ)庫就不用高級封裝能自己寫的模塊就不調(diào)現(xiàn)成函數(shù)。但也不是什么都從零寫像矩陣乘法這種底層算子直接用NumPy或者PyTorch的底層接口就行沒必要自己寫CUDA核。關(guān)鍵是理解每一層的輸入輸出和計算邏輯。2.2 技術(shù)選型為什么是Python NumPy PyTorch底層API選Python沒什么好說的生態(tài)最全。NumPy用來做數(shù)據(jù)預處理和手寫算法驗證因為它足夠底層你能看到每一個數(shù)組的形狀變化。PyTorch我只用torch.tensor、torch.autograd和torch.nn.functional這些底層接口不用nn.Module的高級封裝更不用Trainer。這樣做的代價是代碼量會多兩三倍但好處是每一個超參數(shù)、每一次前向傳播、每一次梯度更新都在你眼皮底下。有人會問為什么不直接用JAX或者TensorFlow。我的考慮是PyTorch的動態(tài)圖機制對調(diào)試最友好而且它的底層API和NumPy的思維模式最接近從NumPy過渡到PyTorch幾乎沒有認知負擔。JAX雖然快但函數(shù)式編程的風格對新手不太友好調(diào)試也麻煩。2.3 整體架構(gòu)從數(shù)據(jù)到服務的四層拆分我把整個系統(tǒng)拆成四層每一層都可以獨立測試和替換數(shù)據(jù)層負責原始數(shù)據(jù)的讀取、清洗、分詞、分批。這一層的關(guān)鍵是可復現(xiàn)同樣的隨機種子必須產(chǎn)生同樣的批次順序。模型層定義網(wǎng)絡(luò)結(jié)構(gòu)、損失函數(shù)、優(yōu)化器。這一層的關(guān)鍵是可解釋每一層的參數(shù)量、計算量、梯度流動都要能打印出來。訓練層控制訓練循環(huán)、學習率調(diào)度、梯度裁剪、模型保存。這一層的關(guān)鍵是可觀測loss曲線、梯度范數(shù)、學習率變化都要實時記錄。服務層把訓練好的模型封裝成API處理并發(fā)請求、批處理、超時。這一層的關(guān)鍵是可伸縮單機能跑多機也能擴。這四層之間的接口我全部用最樸素的Python字典和NumPy數(shù)組來傳遞不用任何框架特有的數(shù)據(jù)結(jié)構(gòu)。這樣做的目的是讓每一層都可以單獨拿出來測試比如我可以不啟動訓練直接用假數(shù)據(jù)測試服務層的吞吐量。3. 核心細節(jié)解析數(shù)據(jù)管線與模型訓練的實操要點3.1 數(shù)據(jù)管線為什么你的模型效果不好八成是數(shù)據(jù)沒弄對我見過太多人把精力全花在調(diào)模型結(jié)構(gòu)上結(jié)果數(shù)據(jù)管線里藏著一堆bug。數(shù)據(jù)管線的第一原則是任何一步都要能單獨驗證。比如分詞之后你要能隨機抽幾條出來看分批之后你要能打印出每個批次的形狀和標簽分布。具體怎么做我一般會寫一個DataPipeline類里面每個方法只做一件事class DataPipeline: def __init__(self, raw_texts, labels, tokenizer, batch_size, seed42): self.raw_texts raw_texts self.labels labels self.tokenizer tokenizer self.batch_size batch_size self.rng np.random.default_rng(seed) def clean(self): # 去重、去空、去異常字符 cleaned [] for text in self.raw_texts: text text.strip() if len(text) 2: continue cleaned.append(text) return cleaned def tokenize(self, texts): # 這里用最簡單的空格分詞實際項目可以換成BPE return [self.tokenizer.encode(t) for t in texts] def batch(self, token_ids, labels): # 先打亂再按batch_size切分 indices self.rng.permutation(len(token_ids)) for i in range(0, len(indices), self.batch_size): batch_idx indices[i:iself.batch_size] yield [token_ids[j] for j in batch_idx], [labels[j] for j in batch_idx]注意幾個細節(jié)隨機種子要固定不然每次跑的結(jié)果都不一樣沒法對比實驗。清洗規(guī)則要可配置不同數(shù)據(jù)集的最短長度要求不一樣。分批之前一定要打亂不然模型會學到順序信息這在很多任務里是致命的。還有一個坑padding的位置。如果你用固定長度的批次短句子后面補0那計算loss的時候一定要mask掉這些0不然模型會學著去預測padding。我一般會在batch方法里同時返回一個mask數(shù)組訓練時用loss (loss * mask).sum() / mask.sum()來算真實loss。3.2 模型層手寫一個Transformer的注意力機制既然是從零手搓那注意力機制肯定要自己寫一遍。很多人覺得Transformer很復雜其實拆開看就是幾個矩陣乘法和softmax。我用NumPy寫一個最基礎(chǔ)的單頭注意力def attention(Q, K, V, maskNone): # Q, K, V的形狀都是 (batch_size, seq_len, d_model) d_k Q.shape[-1] scores np.matmul(Q, K.transpose(0, 2, 1)) / np.sqrt(d_k) if mask is not None: scores np.where(mask 0, -1e9, scores) weights softmax(scores, axis-1) return np.matmul(weights, V), weights def softmax(x, axis-1): x_max np.max(x, axisaxis, keepdimsTrue) exp_x np.exp(x - x_max) return exp_x / np.sum(exp_x, axisaxis, keepdimsTrue)這里有幾個關(guān)鍵點除以sqrt(d_k)是為了防止點積過大導致softmax梯度消失這個縮放因子不是隨便選的是讓方差保持在1左右。mask要在softmax之前加而且要用一個很大的負數(shù)而不是0因為softmax(0)是有值的會污染注意力分布。softmax要減去最大值這是數(shù)值穩(wěn)定性的常規(guī)操作不然exp容易溢出。寫完之后我會用PyTorch的torch.nn.functional.scaled_dot_product_attention對照驗證確保輸出一致。這一步很重要手搓的代碼必須和成熟實現(xiàn)對齊不然你根本不知道是自己寫錯了還是模型本身效果不好。3.3 訓練層學習率調(diào)度和梯度裁剪的實操參數(shù)訓練循環(huán)看起來簡單但里面的坑最多。我一般會記錄四個東西訓練loss、驗證loss、梯度范數(shù)、學習率。這四個指標能覆蓋90%的訓練問題。學習率調(diào)度我用的是帶warmup的余弦退火參數(shù)是這樣選的warmup步數(shù)總步數(shù)的5%到10%。比如總共訓練10000步warmup設(shè)500到1000步。warmup的作用是讓模型在初期不要更新太猛避免梯度爆炸。最大學習率1e-4到3e-4之間。我一般從3e-4開始試如果loss震蕩就降到1e-4。最小學習率最大學習率的十分之一。余弦退火到最后會降到這個值讓模型在末期微調(diào)。梯度裁剪我設(shè)的是全局范數(shù)裁剪閾值1.0。具體做法是把所有參數(shù)的梯度拼成一個向量算它的L2范數(shù)如果超過1.0就按比例縮放。這個操作能防止個別批次的異常梯度把模型帶偏。def clip_gradients(parameters, max_norm1.0): total_norm 0.0 for p in parameters: total_norm np.sum(p.grad ** 2) total_norm np.sqrt(total_norm) clip_coef max_norm / (total_norm 1e-6) if clip_coef 1.0: for p in parameters: p.grad * clip_coef return total_norm注意1e-6這個epsilon不能省不然total_norm為0的時候會除零。返回的total_norm要記錄下來如果它一直很大說明學習率可能太高了。4. 實操過程從零搭建一個文本分類服務的完整記錄4.1 環(huán)境準備與依賴安裝我用的環(huán)境是Python 3.10依賴只有四個numpy、torch、flask、requests。不用transformers、不用datasets、不用accelerate。安裝命令很簡單pip install numpy torch flask requests有人會問不用transformers怎么加載預訓練模型我的做法是自己寫一個最小的模型加載器從HuggingFace的bin文件里讀權(quán)重然后映射到我手寫的網(wǎng)絡(luò)結(jié)構(gòu)上。這個過程很麻煩但能讓你徹底搞清楚預訓練模型的參數(shù)命名規(guī)則和結(jié)構(gòu)。如果只是想快速驗證也可以先用隨機初始化的模型跑通流程再替換成預訓練權(quán)重。4.2 數(shù)據(jù)準備用一個小數(shù)據(jù)集跑通全流程我用的是一個公開的中文情感分類數(shù)據(jù)集大概1萬條數(shù)據(jù)正負樣本各半。數(shù)據(jù)格式是每行一個JSON包含text和label兩個字段。讀取和清洗的代碼如下import json def load_data(path): texts, labels [], [] with open(path, r, encodingutf-8) as f: for line in f: item json.loads(line) texts.append(item[text]) labels.append(item[label]) return texts, labels texts, labels load_data(sentiment.jsonl) print(f總樣本數(shù): {len(texts)}) print(f正樣本比例: {sum(labels) / len(labels):.2f})打印正樣本比例這一步很重要如果比例嚴重失衡準確率這個指標就沒意義了得換F1或者AUC。我一般會先看一眼這個比例再決定用哪些評估指標。4.3 模型定義一個極簡的文本分類網(wǎng)絡(luò)我的模型結(jié)構(gòu)很簡單詞嵌入 平均池化 全連接。沒有用Transformer因為在這個數(shù)據(jù)量下簡單模型反而更穩(wěn)。結(jié)構(gòu)如下import torch import torch.nn as nn class TextClassifier(nn.Module): def __init__(self, vocab_size, embed_dim, num_classes): super().__init__() self.embedding nn.Embedding(vocab_size, embed_dim, padding_idx0) self.fc nn.Linear(embed_dim, num_classes) def forward(self, input_ids, mask): # input_ids: (batch, seq_len) embeds self.embedding(input_ids) # (batch, seq_len, embed_dim) # 用mask做加權(quán)平均池化 mask mask.unsqueeze(-1).float() pooled (embeds * mask).sum(dim1) / mask.sum(dim1).clamp(min1e-6) logits self.fc(pooled) return logits注意padding_idx0這個參數(shù)它讓padding位置的嵌入向量不參與梯度更新。池化的時候用mask加權(quán)平均而不是直接mean這樣padding不會影響結(jié)果。clamp(min1e-6)是防止mask全0的時候除零。4.4 訓練循環(huán)每一步都打印關(guān)鍵指標訓練循環(huán)我寫得比較啰嗦但每一步都記錄了關(guān)鍵信息def train(model, dataloader, optimizer, scheduler, num_epochs): for epoch in range(num_epochs): model.train() total_loss 0.0 for step, (input_ids, mask, labels) in enumerate(dataloader): input_ids torch.tensor(input_ids, dtypetorch.long) mask torch.tensor(mask, dtypetorch.float) labels torch.tensor(labels, dtypetorch.long) logits model(input_ids, mask) loss nn.functional.cross_entropy(logits, labels) optimizer.zero_grad() loss.backward() # 梯度裁剪 grad_norm clip_gradients(model.parameters(), max_norm1.0) optimizer.step() scheduler.step() total_loss loss.item() if step % 50 0: lr scheduler.get_last_lr()[0] print(fEpoch {epoch} Step {step} Loss {loss.item():.4f} fGradNorm {grad_norm:.4f} LR {lr:.6f}) avg_loss total_loss / len(dataloader) print(fEpoch {epoch} Avg Loss {avg_loss:.4f})這里有幾個實操心得每50步打印一次太頻繁會刷屏太稀疏會漏掉異常。梯度范數(shù)和學習率一起打印這樣能看出它們之間的關(guān)聯(lián)。驗證集評估放在每個epoch結(jié)束不要放在訓練循環(huán)里面不然會拖慢訓練速度。4.5 服務層用Flask封裝一個帶批處理的推理接口訓練完之后模型要能對外提供服務。我用Flask寫了一個最簡單的接口支持單條和批量推理from flask import Flask, request, jsonify import torch app Flask(__name__) model load_model(model.pt) model.eval() app.route(/predict, methods[POST]) def predict(): data request.get_json() texts data[texts] # 支持列表 if isinstance(texts, str): texts [texts] input_ids, mask tokenize_and_pad(texts) with torch.no_grad(): logits model(input_ids, mask) probs torch.softmax(logits, dim-1) preds torch.argmax(probs, dim-1) results [] for text, pred, prob in zip(texts, preds.tolist(), probs.tolist()): results.append({ text: text, label: pred, confidence: max(prob) }) return jsonify({results: results}) if __name__ __main__: app.run(host0.0.0.0, port5000)注意torch.no_grad()一定要加不然推理會構(gòu)建計算圖顯存會爆。批處理接口比單條接口吞吐量高很多因為GPU的并行能力只有在大batch下才能發(fā)揮出來。我實測下來batch_size32的時候QPS是單條的8倍左右。5. 常見問題與排查技巧實錄5.1 訓練loss不下降從數(shù)據(jù)到梯度的排查順序loss不降是最常見的問題我一般按這個順序排查先看數(shù)據(jù)隨機抽幾條樣本打印它們的token ids和label確認沒有錯位。我遇到過label和text反了的情況查了半天才發(fā)現(xiàn)是數(shù)據(jù)加載的時候字段名寫錯了。再看梯度打印每一層的梯度范數(shù)如果某一層梯度全是0說明那一層沒參與計算。常見原因是mask寫錯了或者某一層的輸入被detach了。然后看學習率如果學習率太大loss會震蕩太小則下降很慢。我一般會跑一個學習率掃描從1e-5到1e-2每個跑100步看哪個loss降得最快。最后看模型結(jié)構(gòu)如果以上都沒問題那可能是模型容量不夠或者結(jié)構(gòu)有bug。我會先用一個極小的數(shù)據(jù)集比如100條過擬合一下如果連100條都過擬合不了那肯定是代碼有問題。5.2 顯存不夠用幾個立竿見影的優(yōu)化手段顯存不夠的時候按這個優(yōu)先級來優(yōu)化減小batch_size最直接但會影響訓練穩(wěn)定性。我一般會配合梯度累積比如batch_size8累積4次等效batch_size32。用混合精度torch.cuda.amp能省一半顯存速度還快。但要注意有些操作在fp16下會溢出需要用GradScaler。檢查有沒有不必要的張量保留比如在訓練循環(huán)里把loss存到一個列表里如果loss是tensor那整個計算圖都會被保留。正確做法是存loss.item()。用梯度檢查點這個比較高級適合大模型。原理是不保存中間激活值反向傳播時重新計算。代價是訓練速度慢20%左右。5.3 推理延遲高從模型到服務的全鏈路優(yōu)化推理延遲高先定位瓶頸在哪排查點可能原因優(yōu)化手段模型前向?qū)訑?shù)太多、注意力計算量大剪枝、量化、換更小的模型數(shù)據(jù)預處理分詞慢、padding太多緩存分詞結(jié)果、動態(tài)padding服務框架單條推理、沒有批處理加批處理、用異步框架硬件CPU推理、顯存帶寬不夠換GPU、用TensorRT我實測下來動態(tài)padding對延遲的改善最明顯。因為大部分句子的長度都遠小于最大長度固定padding會浪費大量計算。動態(tài)padding就是每個batch按當前最長句子來padding能省30%到50%的計算量。5.4 常見問題速查表問題現(xiàn)象可能原因快速驗證方法解決方案loss變成NaN學習率太大、梯度爆炸打印梯度范數(shù)降低學習率、加梯度裁剪驗證loss上升過擬合對比訓練和驗證loss曲線加dropout、早停、數(shù)據(jù)增強預測結(jié)果全是同一類數(shù)據(jù)失衡、模型沒學到打印預測分布重采樣、換損失函數(shù)服務響應超時批處理太大、模型太慢打印每個請求的處理時間減小batch、加超時限制模型加載失敗參數(shù)名不匹配、形狀不對打印state_dict的key手動映射參數(shù)名6. 我踩過的坑和最后再分享幾個小技巧第一個坑是隨機種子沒固定全。Python的random、NumPy的np.random、PyTorch的torch.manual_seed都要設(shè)而且DataLoader的worker_init_fn也要設(shè)不然多進程加載數(shù)據(jù)的時候順序還是會變。我現(xiàn)在的做法是在訓練腳本開頭寫一個set_seed(42)函數(shù)把所有能設(shè)的種子都設(shè)一遍。第二個坑是學習率調(diào)度器的step位置。PyTorch的CosineAnnealingLR是按epoch調(diào)的但OneCycleLR是按step調(diào)的。如果搞混了學習率曲線會完全不對。我現(xiàn)在的習慣是每個step都打印學習率這樣一眼就能看出調(diào)度器有沒有正常工作。第三個坑是模型保存和加載的不一致。訓練的時候用了nn.DataParallel保存的state_dict的key會多一個module.前綴加載的時候如果不用DataParallel就會報錯。解決辦法是保存的時候用model.module.state_dict()或者加載的時候用OrderedDict把前綴去掉。最后分享一個小技巧在訓練循環(huán)里加一個異常捕獲把出錯的batch的數(shù)據(jù)和中間結(jié)果保存下來。這樣即使訓練崩了你也能復現(xiàn)問題。我一般會在try塊里跑訓練except塊里把當前batch的input_ids、mask、labels和loss都存成npy文件然后重新拋出異常。這個習慣幫我省了很多調(diào)試時間。還有一個技巧是用小數(shù)據(jù)集做快速迭代。每次改完代碼先用100條數(shù)據(jù)跑10個step確認沒有形狀錯誤和NaN再上全量數(shù)據(jù)。這樣能把調(diào)試周期從幾小時縮短到幾分鐘。我現(xiàn)在的流程是改代碼 - 小數(shù)據(jù)跑通 - 全量訓練 - 驗證集評估 - 服務部署每一步都有明確的檢查點不會等到最后才發(fā)現(xiàn)問題。