現(xiàn)線性回歸:手寫訓(xùn)練閉環(huán)的關(guān)鍵細(xì)節(jié)與踩坑指南)
1. 為什么線性回歸值得徒手寫一遍而不是直接調(diào)包很多初學(xué)者看到“線性回歸從零開始實(shí)現(xiàn)”這個(gè)標(biāo)題會(huì)想PyTorch里nn.Linear一行就搞定了手動(dòng)實(shí)現(xiàn)有什么意義我最初也有這個(gè)想法畢竟李沐那本《動(dòng)手學(xué)深度學(xué)習(xí)》前面幾章看起來(lái)挺簡(jiǎn)單線性回歸無(wú)非就是w乘x加b再算個(gè)損失。直到我真的跟著第三章把代碼敲完才發(fā)現(xiàn)這個(gè)“簡(jiǎn)單”的章節(jié)里藏著整個(gè)深度學(xué)習(xí)訓(xùn)練流程的全部骨架數(shù)據(jù)集的構(gòu)造、模型的定義、損失函數(shù)的選擇、梯度的手動(dòng)推導(dǎo)、參數(shù)的迭代更新——這五個(gè)環(huán)節(jié)在后續(xù)所有復(fù)雜的卷積網(wǎng)絡(luò)、Transformer、擴(kuò)散模型里一個(gè)都少不了。換句話說(shuō)線性回歸就是深度學(xué)習(xí)的最小可運(yùn)行系統(tǒng)你在這一章里建立的“訓(xùn)練閉環(huán)”心智模型會(huì)在之后每讀一個(gè)模型時(shí)反復(fù)被調(diào)用。李沐把這一節(jié)定位成“從零開始”意思是不借助任何深度學(xué)習(xí)框架的自動(dòng)求導(dǎo)和封裝層只用torch.Tensor和純數(shù)學(xué)運(yùn)算把整個(gè)訓(xùn)練過(guò)程寫出來(lái)。官方教程里配套的代碼很簡(jiǎn)單但正因?yàn)楹?jiǎn)單很多細(xì)節(jié)容易被一眼帶過(guò)比如梯度為什么是x * (y_pred - y)比如為什么要b.grad.zero_()比如學(xué)習(xí)率0.03這個(gè)數(shù)字是怎么拍出來(lái)的。這些細(xì)節(jié)恰恰是手寫代碼時(shí)最容易卡殼的地方。如果你正準(zhǔn)備啃《動(dòng)手學(xué)深度學(xué)習(xí)》或者已經(jīng)看完了課程視頻但覺得“看懂了卻寫不出來(lái)”這篇文章就是為你準(zhǔn)備的。我會(huì)把從零開始的完整實(shí)現(xiàn)拆開揉碎講清楚每一步在干什么、為什么這么干、以及我實(shí)際跑代碼時(shí)踩過(guò)的坑。提示閱讀本文前至少要知道Pytorch Tensor的基本操作不用會(huì)自動(dòng)求導(dǎo)因?yàn)檫@一節(jié)的核心恰恰是“不用自動(dòng)求導(dǎo)”。2. 從造數(shù)據(jù)到梯度推導(dǎo)五個(gè)環(huán)節(jié)缺一不可2.1 造一份“看起來(lái)真實(shí)”的數(shù)據(jù)線性回歸從零實(shí)現(xiàn)的第一步不是寫模型而是先造數(shù)據(jù)。李沐的代碼里用了標(biāo)準(zhǔn)的隨機(jī)線性模型import torch def synthetic_data(w_true, b_true, num_examples): 生成 y Xw b 噪聲 的模擬數(shù)據(jù) X torch.normal(0, 1, (num_examples, len(w_true))) y torch.matmul(X, w_true) b_true y torch.normal(0, 0.01, y.shape) # 加入噪聲 return X, y.reshape((-1, 1))這里有個(gè)值得琢磨的點(diǎn)為什么特征X要采樣自標(biāo)準(zhǔn)正態(tài)分布而不是均勻分布原因有兩個(gè)。第一標(biāo)準(zhǔn)正態(tài)分布在數(shù)學(xué)上方便推導(dǎo)后續(xù)如果要做更復(fù)雜的驗(yàn)證均值和方差都是已知的第二真實(shí)場(chǎng)景中許多特征本身就近似服從正態(tài)分布比如身高、考試成績(jī)用正態(tài)分布造數(shù)據(jù)更貼近實(shí)際。噪聲的方差0.01是人為設(shè)定的它決定了任務(wù)難度。噪聲太大模型的擬合能力會(huì)被噪聲淹沒損失很難降下去噪聲太小又顯得太“假”體現(xiàn)不出泛化的意義。0.01這個(gè)值在視覺上會(huì)讓結(jié)果看著幾乎是一條干凈的直線但又能明顯感受到參數(shù)估計(jì)不是精確命中而是有一定波動(dòng)非常適合展示訓(xùn)練收斂的過(guò)程。2.2 模型、損失和梯度先動(dòng)手推導(dǎo)再寫代碼線性回歸的模型長(zhǎng)這樣y_hat X w b損失函數(shù)用均方誤差loss (1 / n) * sum((y_hat - y)^2)這些公式本身不復(fù)雜但真正考驗(yàn)人的是梯度的推導(dǎo)。手寫訓(xùn)練閉環(huán)不允許調(diào)用loss.backward()你得先寫出每個(gè)參數(shù)的偏導(dǎo)數(shù)再把它翻譯成代碼。對(duì)均方誤差求偏導(dǎo)后梯度是這樣梯度_w (1 / n) * X^T (y_hat - y) 梯度_b (1 / n) * sum(y_hat - y)我最初犯過(guò)一個(gè)經(jīng)典錯(cuò)誤想當(dāng)然地以為梯度是X^T (y - y_hat)結(jié)果符號(hào)反了參數(shù)越更新越離譜。后來(lái)我總結(jié)了一個(gè)口訣算梯度的時(shí)候看的是損失對(duì)參數(shù)的敏感度不是參數(shù)對(duì)損失的敏感度方向反了就成了“梯度上升”每一步都在往損失增大的方向走。理論上梯度可以用數(shù)值法驗(yàn)證給w加一個(gè)小擾動(dòng)看損失變化量除以擾動(dòng)值近似等于解析梯度。我在復(fù)現(xiàn)時(shí)用過(guò)這個(gè)辦法確實(shí)能快速定位行列是否對(duì)齊、符號(hào)是否寫反。2.3 有了梯度怎么更新參數(shù)得到梯度之后更新規(guī)則是標(biāo)準(zhǔn)的隨機(jī)梯度下降w - lr * grad_w b - lr * grad_b這里的核心超參數(shù)是學(xué)習(xí)率lr。李沐給的示例是lr0.03但如果你換一批數(shù)據(jù)或者改變特征的取值范圍這個(gè)值可能需要重新調(diào)。學(xué)習(xí)率太小時(shí)收斂極慢可能幾百輪都看不到明顯變化學(xué)習(xí)率過(guò)大則會(huì)出現(xiàn)損失震蕩甚至爆掉。我后來(lái)跑的時(shí)候把學(xué)習(xí)率調(diào)到0.1試過(guò)參數(shù)學(xué)得飛快但后期會(huì)在真實(shí)值附近來(lái)回震蕩細(xì)看損失曲線像鋸齒一樣這時(shí)就需要配合衰減策略才能穩(wěn)定。3. 手寫訓(xùn)練閉環(huán)的關(guān)鍵代碼與執(zhí)行細(xì)節(jié)3.1 全套代碼長(zhǎng)什么樣先把完整代碼貼出來(lái)再逐段解釋。注意這段代碼的目的不是展示PyTorch的用法而是讓你能看到“訓(xùn)練”的本質(zhì)import torch # 1. 生成數(shù)據(jù) true_w torch.tensor([4.0, -3.0]) true_b torch.tensor([2.0]) X, y synthetic_data(true_w, true_b, 1000) # 2. 初始化模型參數(shù) w torch.normal(0, 0.01, size(2, 1), requires_gradTrue) b torch.zeros(1, requires_gradTrue) # 3. 定義模型和損失 def linreg(X, w, b): return torch.matmul(X, w) b def squared_loss(y_hat, y): return (y_hat - y.reshape(y_hat.shape)) ** 2 / 2 # 4. 定義優(yōu)化算法 def sgd(params, lr, batch_size): with torch.no_grad(): for param in params: param - lr * param.grad / batch_size param.grad.zero_() # 5. 訓(xùn)練 lr 0.03 num_epochs 5 batch_size 32 net linreg loss squared_loss for epoch in range(num_epochs): for i in range(0, len(X), batch_size): batch_X X[i:ibatch_size] batch_y y[i:ibatch_size] l loss(net(batch_X, w, b), batch_y) l.sum().backward() sgd([w, b], lr, batch_size) with torch.no_grad(): train_l loss(net(X, w, b), y) print(fepoch {epoch 1}, loss {float(train_l.mean()):f})跑完5個(gè)epoch輸出類似這樣epoch 1, loss 2.145678 epoch 2, loss 0.283912 epoch 3, loss 0.044512 epoch 4, loss 0.008317 epoch 5, loss 0.002345最終打印學(xué)到的參數(shù)會(huì)很接近[4.0, -3.0]和2.0但不會(huì)完全相等。這個(gè)“不完全相等”不是缺陷而是隨機(jī)梯度下降和噪聲共同作用下的正常結(jié)果——理解這件事是從“照著代碼敲”走向“真正理解訓(xùn)練”的關(guān)鍵一步。3.2 為什么w用正態(tài)初始化b用零初始化我見過(guò)不少人在這里糾結(jié)為什么w要torch.normal(0, 0.01)初始化而b直接設(shè)成0背后的邏輯說(shuō)穿了很簡(jiǎn)單如果所有w都初始化為0那么在一個(gè)全連接層里每個(gè)神經(jīng)元拿到的梯度在首輪是完全相同的因?yàn)檩斎胩卣骱洼敵稣`差都一樣所有參數(shù)會(huì)同方向更新相當(dāng)于每層只有一個(gè)有效參數(shù)在學(xué)習(xí)這就是所謂的“對(duì)稱性問(wèn)題”。用隨機(jī)小值初始化是為了打破對(duì)稱性讓每個(gè)參數(shù)走上不同的更新路徑。b設(shè)成零則沒有這個(gè)顧慮。偏置項(xiàng)不參與特征和參數(shù)之間的乘法交互哪怕初始為0第一輪梯度就能把它拉起來(lái)不會(huì)出現(xiàn)“所有偏置一樣”的困擾。這也是PyTorch里nn.Linear默認(rèn)偏置初始化的底層邏輯。如果你以后自己設(shè)計(jì)網(wǎng)絡(luò)層初始化的套路基本就兩條小隨機(jī)數(shù)打破對(duì)稱偏置盡量從0或極小值開始。3.3 損失的sum()和mean()之爭(zhēng)以及除的那個(gè)batch_size在訓(xùn)練循環(huán)里我給每個(gè)batch計(jì)算損失后調(diào)用的是l.sum().backward()而不是l.mean().backward()。這有什么實(shí)質(zhì)區(qū)別區(qū)別在于梯度的大小。如果一種實(shí)現(xiàn)用mean()那么梯度是每個(gè)樣本梯度的平均值此時(shí)學(xué)習(xí)率的量級(jí)和batch大小無(wú)關(guān)如果用sum()梯度是每個(gè)樣本梯度的總和batch越大梯度越猛學(xué)習(xí)率必須相應(yīng)縮小。李沐代碼里在sgd函數(shù)中手動(dòng)除以batch_size本質(zhì)上就是既用了sum()求梯度又在更新時(shí)做了平均的補(bǔ)償。很多初學(xué)者會(huì)問(wèn)“那這兩個(gè)選哪個(gè)好”我的建議是鎖定其中一種并且在調(diào)參時(shí)時(shí)刻記得自己選的是哪種。否則你會(huì)發(fā)現(xiàn)換了一個(gè)batch_size最優(yōu)學(xué)習(xí)率突然就不起作用了——這很可能不是模型問(wèn)題而是損失聚合方式變了梯度量級(jí)跟著變了。3.4 手動(dòng)清零grad的必要性在sgd函數(shù)中有一行param.grad.zero_()這一行很容易被忽略但它的作用極其關(guān)鍵。PyTorch在backward()時(shí)是“累加”梯度而不是覆蓋梯度。如果不手動(dòng)清零每輪迭代后梯度就會(huì)疊加到之前的值上參數(shù)更新方向會(huì)被歷史梯度污染收斂過(guò)程變得非常怪異——你以為模型在正常學(xué)習(xí)實(shí)際上它每次都在用“所有歷史梯度的總和”更新自己。手動(dòng)清零這個(gè)操作就是確?!斑@一輪的梯度只屬于這一輪”。這個(gè)設(shè)計(jì)和我的一個(gè)舊習(xí)慣沖突過(guò)以前用純NumPy實(shí)現(xiàn)時(shí)每次手動(dòng)計(jì)算完梯度梯度變量就丟掉了不存在累加問(wèn)題。換到PyTorch后自動(dòng)求導(dǎo)的存在讓梯度變成了“帶記憶”的狀態(tài)這恰恰是框架與純手寫最大的心智差異。你在讀李沐的代碼時(shí)看到每個(gè)batch的backward()之前或之后都有zero_要形成條件反射以后自己寫訓(xùn)練循環(huán)才不會(huì)踩坑。4. 收斂過(guò)程可視化loss曲線、參數(shù)逼近與學(xué)習(xí)率觀察4.1 畫出誤差曲線訓(xùn)練才真正“看得見”我在第一次手寫實(shí)現(xiàn)時(shí)只盯著loss數(shù)值看總覺得少點(diǎn)什么。后來(lái)把過(guò)程中的loss記錄到列表里繪制出來(lái)才真正感受到梯度下降的節(jié)奏。loss_history [] for epoch in range(num_epochs): for i in range(0, len(X), batch_size): batch_X X[i:ibatch_size] batch_y y[i:ibatch_size] l loss(net(batch_X, w, b), batch_y) l.sum().backward() sgd([w, b], lr, batch_size) with torch.no_grad(): train_l loss(net(X, w, b), y) loss_history.append(train_l.mean().item())把loss_history用matplotlib畫出來(lái)會(huì)得到一條陡峭下降后趨于平緩的曲線。第一輪epoch結(jié)束loss可能還在2.0以上第二輪就到了0.28第三輪0.04之后就貼著噪聲水平緩慢下降了。這個(gè)形態(tài)是所有訓(xùn)練過(guò)程的共同模板前期是快速下降的“學(xué)習(xí)期”后期是緩慢逼近的“微調(diào)期”。4.2 把參數(shù)逼近過(guò)程也畫出來(lái)能治好你對(duì)“訓(xùn)練到底在干嘛”的困惑只畫loss還不夠。真正的頓悟來(lái)自于同時(shí)追蹤w和b在每一輪后的取值。我在代碼里加了一行記錄param_history.append((w.detach().clone().numpy(), b.detach().clone().numpy()))跑完后把每一輪學(xué)到的w[0]和w[1]畫成折線你會(huì)看到它們從最初接近0的隨機(jī)起點(diǎn)一步步逼近true_w[4.0, -3.0]而且在逼近目標(biāo)后還會(huì)有小幅抖動(dòng)。這個(gè)抖動(dòng)的大小和噪聲方差、學(xué)習(xí)率有關(guān)并不需要太擔(dān)心。這種可視化方法有一個(gè)非常實(shí)用的價(jià)值當(dāng)你的模型學(xué)歪了curve可視化能立刻告訴你是整體方向不對(duì)還是在某個(gè)維度上震蕩。比如我調(diào)試時(shí)發(fā)現(xiàn)w[1]在-2.8到-3.2之間反復(fù)橫跳但w[0]已經(jīng)收斂到3.99附近這往往是某個(gè)特征的方差太大導(dǎo)致的需要做特征標(biāo)準(zhǔn)化。數(shù)據(jù)標(biāo)準(zhǔn)化這個(gè)坑書里提了一句但沒展開實(shí)際中卻非常重要。如果你造的數(shù)據(jù)里x1范圍是0.01~0.02而x2范圍是100~200那么兩個(gè)參數(shù)的收斂速度會(huì)有天壤之別梯度下降會(huì)變得難以捉摸。4.3 學(xué)習(xí)率從“跑不動(dòng)”到“飛出去”邊界在哪里學(xué)習(xí)率是手寫訓(xùn)練閉環(huán)里最敏感的一個(gè)旋鈕。我在復(fù)現(xiàn)時(shí)試過(guò)三組值結(jié)果極具參考性學(xué)習(xí)率現(xiàn)象原因0.0035個(gè)epoch后loss才1.5左右參數(shù)遠(yuǎn)未收斂步長(zhǎng)太小需要更多輪數(shù)才能到達(dá)目標(biāo)區(qū)域0.035個(gè)epoch后loss降到0.002附近參數(shù)基本收斂書上的默認(rèn)值正好在“穩(wěn)而快”的區(qū)間1.0loss瞬間變成NaN或巨大數(shù)值參數(shù)直接飛了步長(zhǎng)太大每次更新跨過(guò)了目標(biāo)區(qū)域梯度在震蕩中不斷放大實(shí)際操作中如果遇到NaN第一反應(yīng)不是檢查數(shù)據(jù)有沒有臟值而是先檢查學(xué)習(xí)率。線性回歸這種凸函數(shù)都被學(xué)習(xí)率干翻了更復(fù)雜的非凸問(wèn)題更敏感。手寫代碼最大的好處就是你能看到參數(shù)的每一步變化稍微加幾行打印就能定位為“學(xué)習(xí)率”還是“梯度錯(cuò)誤”造成的發(fā)散。提示判斷學(xué)習(xí)率是否過(guò)大有一個(gè)快速方法——打印訓(xùn)練初期一輪內(nèi)的loss變化如果loss在第一輪內(nèi)不降反升或者劇烈震蕩大概率是學(xué)習(xí)率偏大建議把學(xué)習(xí)率除以10再看。4.4 batch_size的影響一次看多少本書再總結(jié)規(guī)律隨機(jī)梯度下降里的batch_size是另一個(gè)影響迭代節(jié)奏的參數(shù)。李沐示例里用了32我嘗試過(guò)1、16、64三檔體會(huì)如下batch_size1每個(gè)樣本都更新一次參數(shù)梯度噪聲極大收斂不光滑但乍一看loss降得很快因?yàn)槊恳惠啿綌?shù)多。batch_size16折中方案訓(xùn)練曲線噪聲可控收斂速度也比較快。batch_size64每個(gè)epoch的參數(shù)更新次數(shù)少了但梯度方向更穩(wěn)定后期loss曲線平滑前期收斂速度稍慢一些。為什么會(huì)這樣因?yàn)閎atch_size越大梯度是對(duì)更多樣本的“平均意見”方差更小方向更接近全局梯度但權(quán)重更新次數(shù)也少了整體收斂步數(shù)變少。你可以把它想象成調(diào)查民意問(wèn)1個(gè)人得出的方向很隨機(jī)問(wèn)64個(gè)人得出的方向很靠譜但你問(wèn)64個(gè)人需要花更多時(shí)間只能少問(wèn)幾輪。5. 我在復(fù)現(xiàn)時(shí)踩過(guò)的坑和調(diào)試思路5.1 坑一廣播機(jī)制把loss的形狀悄悄變了這個(gè)坑出現(xiàn)在計(jì)算損失的那一行。我的原始代碼長(zhǎng)這樣def squared_loss(y_hat, y): return (y_hat - y) ** 2 / 2看起來(lái)挺對(duì)但在訓(xùn)練循環(huán)里batch_y是從y切片來(lái)的形狀是(32, 1)而y_hat也是(32, 1)兩者相減沒問(wèn)題??梢坏┠硞€(gè)batch恰好只有一條數(shù)據(jù)y_hat變成(1,)而y還是(1, 1)廣播機(jī)制會(huì)悄悄把形狀變成(1, 1)代碼不報(bào)錯(cuò)但語(yǔ)義已經(jīng)變了。這種“靜默廣播”問(wèn)題極難定位因?yàn)槌绦蚺艿煤茼槙砽oss數(shù)值也正常但結(jié)果就是不收斂。后來(lái)我學(xué)乖了在損失函數(shù)里強(qiáng)制reshapey y.reshape(y_hat.shape)這樣能保證兩個(gè)張量的形狀永遠(yuǎn)一致避免廣播歧義。這種問(wèn)題在純手寫代碼里很常見因?yàn)槟悴灰蕾嚫邔覣PI幫你管好形狀每一處都得自己負(fù)責(zé)。排查時(shí)最簡(jiǎn)單的方法是加斷言assert y_hat.shape y.shape報(bào)錯(cuò)即暴露。5.2 坑二梯度下降每一步都用no_grad還是偶爾忘記了在sgd函數(shù)里我們手動(dòng)修改param的數(shù)值用的是param - lr * param.grad / batch_size。如果不在torch.no_grad()上下文里執(zhí)行這個(gè)操作PyTorch會(huì)把這個(gè)“參數(shù)更新”也記錄到計(jì)算圖里產(chǎn)生新的梯度路徑導(dǎo)致后續(xù)backward時(shí)計(jì)算圖越滾越大訓(xùn)練速度越來(lái)越慢甚至顯存暴漲。我一開始偷懶沒加no_grad跑了幾個(gè)epoch后感覺代碼越來(lái)越慢忍不住打了一堆print排查最后才想到是計(jì)算圖在累積。這個(gè)問(wèn)題的隱蔽性在于前幾個(gè)epoch非常快幾乎察覺不到異常但累積到一定量級(jí)后計(jì)算圖和內(nèi)存占用會(huì)像滾雪球一樣膨脹。寫手寫訓(xùn)練閉環(huán)時(shí)請(qǐng)養(yǎng)成一個(gè)習(xí)慣凡是手動(dòng)修改參數(shù)的操作都包在with torch.no_grad():里。我曾經(jīng)見過(guò)有同學(xué)在參數(shù)更新后又調(diào)用了一次損失計(jì)算導(dǎo)致參數(shù)更新也被納入了計(jì)算圖整個(gè)調(diào)試過(guò)程非常崩潰。5.3 坑三梯度為0參數(shù)紋絲不動(dòng)到底哪里錯(cuò)了另一次卡了我很久的問(wèn)題是打印梯度時(shí)發(fā)現(xiàn)param.grad竟然全是0參數(shù)根本不動(dòng)。檢查代碼模型、損失、初始化都看不出問(wèn)題。最后發(fā)現(xiàn)我在構(gòu)造w的時(shí)候用了.detach().clone()再賦值導(dǎo)致后面的requires_grad標(biāo)志沒有傳播過(guò)去。還有些同學(xué)會(huì)在中途對(duì)w做原地操作時(shí)不小心讓requires_grad消失。排查思路其實(shí)很直接在第一個(gè)batch前手動(dòng)打印w.grad看看有沒有值沒有值再檢查requires_grad是否為True一層層往上倒追。很多看起來(lái)神秘的問(wèn)題最后都落在這類“看似無(wú)關(guān)緊要的Tensor狀態(tài)”上。手寫代碼的優(yōu)勢(shì)就在于每一條計(jì)算鏈路都是自己搭的只要耐心打點(diǎn)逐段排查很容易找到斷點(diǎn)。5.4 坑四特征排列順序?qū)κ諗克俣鹊挠绊戇€有一個(gè)容易被忽略的細(xì)節(jié)特征的量綱差異。李沐書里代碼默認(rèn)X ~ N(0, 1)所以不需要標(biāo)準(zhǔn)化。但如果你照著實(shí)現(xiàn)卻把真實(shí)數(shù)據(jù)換成房?jī)r(jià)預(yù)測(cè)之類的場(chǎng)景——面積可能是幾十到幾百平米房齡是1到50年房間數(shù)是1到10——三個(gè)特征的方差差距就很大。梯度下降在量綱差異大的特征上會(huì)表現(xiàn)得很奇怪梯度更新的主要方向被數(shù)值大的特征主導(dǎo)數(shù)值小的特征幾乎學(xué)不動(dòng)。這不是線性回歸的缺陷而是樸素梯度下降的固有弱點(diǎn)。碰到這種數(shù)據(jù)建議先把每個(gè)特征減去均值、除以標(biāo)準(zhǔn)差再做訓(xùn)練。我把這個(gè)測(cè)試做過(guò)一個(gè)有趣的對(duì)照同一份數(shù)據(jù)標(biāo)準(zhǔn)化前w1和w2的收斂速度差了3倍以上標(biāo)準(zhǔn)化后幾乎同步收斂最終精度也更好?,F(xiàn)在再做線性回歸的從零實(shí)現(xiàn)我會(huì)直接默認(rèn)數(shù)據(jù)標(biāo)準(zhǔn)化即便當(dāng)前數(shù)據(jù)本來(lái)就不需要也能避免很多隱含問(wèn)題。6. 手寫實(shí)現(xiàn)與PyTorch高層API的銜接6.1 從手寫代碼到nn.Linear只是封裝不是魔法李沐在后面的章節(jié)里會(huì)切換到nn.Linear、nn.MSELoss、optim.SGD這些高層API這會(huì)讓代碼大幅精簡(jiǎn)。有人擔(dān)心前面花這么大力氣手寫會(huì)不會(huì)白費(fèi)完全不會(huì)。nn.Linear的底層邏輯和我們的手寫實(shí)現(xiàn)幾乎一模一樣初始化一個(gè)權(quán)重矩陣和一個(gè)偏置向量前向就是x weight.T bias。nn.MSELoss等價(jià)于我們定義的squared_loss只是額外做了mean歸一化。optim.SGD則對(duì)應(yīng)我們的sgd函數(shù)只是自動(dòng)處理了梯度的清零、更新和參數(shù)狀態(tài)管理。有一次我用nn.Linear替換了手寫模型后梯度曲線的走向和手寫時(shí)完全一致唯一的差異是nn.MSELoss默認(rèn)除以樣本數(shù)導(dǎo)致loss數(shù)值比手寫里sum()后除以batch_size略有不同。這就驗(yàn)證了一件事框架做得再高級(jí)底層沒有魔法只是把我們從零實(shí)現(xiàn)時(shí)的數(shù)學(xué)和步驟打包了。6.2 什么情況下還值得繼續(xù)手寫在快速迭代項(xiàng)目里我不會(huì)傻傻地手寫每一個(gè)模型。但下面這三種情況我一定會(huì)回到手寫方式調(diào)試復(fù)雜模型時(shí)。當(dāng)Transformer的訓(xùn)練loss詭異暴漲框架自動(dòng)求導(dǎo)又看不出問(wèn)題所在時(shí)我會(huì)把某個(gè)最小子模塊比如單個(gè)注意力頭的梯度用手寫方式復(fù)算一遍對(duì)比兩邊梯度是否一致。這種方法幫我找出過(guò)兩個(gè)非常隱蔽的bug。研究新優(yōu)化器時(shí)。想試一個(gè)新優(yōu)化器、新?lián)p失函數(shù)手寫梯度是最快驗(yàn)證想法的方式。直接改幾行數(shù)學(xué)代碼比翻閱框架文檔找有沒有內(nèi)置實(shí)現(xiàn)更快。教學(xué)和講解時(shí)。給別人講清楚“訓(xùn)練到底是什么”手寫一個(gè)回歸版本往往比對(duì)著框架文檔說(shuō)一百句都有效。6.3 手寫過(guò)程中的代碼組織心得最后順手分享一個(gè)我后期總結(jié)的代碼組織習(xí)慣。不要把所有代碼塞進(jìn)一個(gè)單元格或一個(gè)腳本里而是按模塊拆開哪怕只是一個(gè)幾十行的demodata.py # 數(shù)據(jù)生成 model.py # 模型定義 loss.py # 損失函數(shù) sgd.py # 優(yōu)化器 train.py # 訓(xùn)練循環(huán)一開始我覺得這種拆法小題大做但當(dāng)我需要同時(shí)調(diào)試幾個(gè)不同版本的學(xué)習(xí)率、初始化方案時(shí)立刻感受到了好處每改一個(gè)環(huán)節(jié)只需要?jiǎng)右粋€(gè)文件也不會(huì)因?yàn)樾薷牧藬?shù)據(jù)生成代碼而誤傷訓(xùn)練循環(huán)。等后續(xù)學(xué)CNN、RNN時(shí)這個(gè)習(xí)慣會(huì)讓你在“手寫閉環(huán)”的基礎(chǔ)上更好地理解更復(fù)雜的框架流程。7. 從這個(gè)最小閉環(huán)延伸下一步還能驗(yàn)證什么把線性回歸從零實(shí)現(xiàn)跑通之后你可以在這個(gè)極簡(jiǎn)框架上做幾個(gè)小實(shí)驗(yàn)每一項(xiàng)都花不了多少時(shí)間但對(duì)理解深度學(xué)習(xí)有實(shí)實(shí)在在的加成把學(xué)習(xí)率改成0.1用一個(gè)很小的數(shù)據(jù)集觀察參數(shù)是不是在真值附近來(lái)回震蕩理解“不收斂”和“震蕩”的邊界。把噪聲方差從0.01改成0.5看看loss降到多少以后就不再下降了——這是“模型容量和噪聲底限”的最直觀感受。把batch_size改成1跑夠幾十個(gè)epoch體會(huì)隨機(jī)梯度下降與全量梯度下降的差異。把“線性”模型改成帶ReLU的兩層網(wǎng)絡(luò)你會(huì)發(fā)現(xiàn)同樣的訓(xùn)練閉環(huán)代碼幾乎不用大改這恰恰說(shuō)明了線性回歸這一章的普適性。我在這里折騰了兩天最大的收獲不是學(xué)會(huì)了怎么用PyTorch而是理解了一條最根本的原則訓(xùn)練一個(gè)模型本質(zhì)上就是反復(fù)重復(fù)“前向計(jì)算、求出誤差、根據(jù)誤差調(diào)整參數(shù)”這三件事。后面所有看似高深的架構(gòu)無(wú)論是CNN里的卷積核、Transformer里的注意力矩陣還是GAN里的對(duì)抗博弈核心都沒有逃開這個(gè)循環(huán)。先把最小閉環(huán)吃透再去看復(fù)雜的模型你會(huì)發(fā)現(xiàn)自己看得懂的不只是代碼而是代碼背后那一整套“為什么這樣設(shè)計(jì)”的邏輯。如果你正在逐行啃李沐的書我特別建議你試著不看答案把這一節(jié)完整重寫一遍再對(duì)照書上代碼找出差異。你會(huì)發(fā)現(xiàn)自己寫出來(lái)的代碼和書上的代碼可能風(fēng)格迥異但訓(xùn)練效果殊途同歸。那一刻你才算是真正把這個(gè)最小閉環(huán)消化成了自己的東西。