計(jì)算:從數(shù)組運(yùn)算到向量化性能優(yōu)化)
先回答一個(gè)我隔三差五就能刷到的問題Python 這么慢為什么搞科學(xué)計(jì)算和 AI 的人還在天天用它嚴(yán)格來說這句話只說對(duì)了一半。Python 慢的是你手寫的那些循環(huán)它底層真正干重活的 C 庫(kù)一點(diǎn)都不慢。而 NumPy恰恰就是那個(gè)把“話事權(quán)”交還給 C 庫(kù)的關(guān)鍵角色。你在 Python 里寫個(gè) for 循環(huán)逐個(gè)加數(shù)值相當(dāng)于用解釋器一格一格地看視頻改成 NumPy 的數(shù)組運(yùn)算等于直接把整份數(shù)據(jù)丟給一個(gè)高度優(yōu)化的批處理流水線。這篇文章我就從“為什么你需要 NumPy”開始一路聊到怎么正確安裝、ndarray 的內(nèi)存布局和軸概念、行列式計(jì)算與線性代數(shù)操作、索引切片和廣播里的隱藏坑最后再送上一套我自己常用的提速實(shí)戰(zhàn)思路。無論你是剛接觸 Python 的數(shù)據(jù)新手還是被循環(huán)慢到懷疑人生的老開發(fā)這篇都能幫你把 NumPy 用得更明白。1. 為什么需要 NumPy先跑一次百倍性能差的對(duì)比實(shí)驗(yàn)1.1 一個(gè)數(shù)值計(jì)算任務(wù)的兩種寫法為了把 NumPy 的價(jià)值講清楚我們先做一個(gè)非常樸素的實(shí)驗(yàn)生成 1000 萬個(gè)隨機(jī)浮點(diǎn)數(shù)然后計(jì)算它們的平方和。如果完全用純 Python代碼大概是這樣的import random from time import perf_counter n 10_000_000 data [random.random() for _ in range(n)] start perf_counter() total 0.0 for x in data: total x * x print(f純Python循環(huán)耗時(shí): {perf_counter() - start:.3f} 秒)同樣的事情換用 NumPyimport numpy as np from time import perf_counter arr np.random.random(n) start perf_counter() total np.sum(arr * arr) print(fNumPy向量化耗時(shí): {perf_counter() - start:.3f} 秒)我在自己的筆記本上跑過很多次純 Python 版本通常在 1.5 到 2.5 秒之間浮動(dòng)NumPy 版本基本穩(wěn)定在 10 到 20 毫秒。一百倍的差距而且數(shù)據(jù)量越大差距越離譜。注意這兩段代碼的邏輯完全一樣區(qū)別只在于一個(gè)用解釋器逐條執(zhí)行 Python 字節(jié)碼一個(gè)把運(yùn)算下沉到了 C 語(yǔ)言實(shí)現(xiàn)的底層函數(shù)里。有朋友可能會(huì)說10 毫秒和 2 秒對(duì)我的小項(xiàng)目來說好像都挺快。這話沒錯(cuò)但當(dāng)你處理的是百萬級(jí)矩陣、千萬級(jí)時(shí)間序列、或者深度學(xué)習(xí)里動(dòng)輒幾十 GB 的預(yù)處理數(shù)據(jù)時(shí)循環(huán)版本就不是慢一點(diǎn)的問題了而是根本等不起。NumPy 一開始就是為了解決這類數(shù)值計(jì)算而生的它不是一個(gè)“可用可不用”的優(yōu)化選項(xiàng)而是 Python 科學(xué)計(jì)算生態(tài)的地基。1.2 為什么 NumPy 可以快這么多連續(xù)內(nèi)存與內(nèi)存局部性這里面的道理值得稍微展開講講。Python 內(nèi)置的 list 是一個(gè)“對(duì)象數(shù)組”每個(gè)元素其實(shí)是一個(gè)指向 PyObject 的指針而這些 PyObject 散落在堆內(nèi)存的各個(gè)角落。你在遍歷 list 的時(shí)候解釋器每碰一個(gè)元素都要做類型檢查、引用計(jì)數(shù)增減、取真實(shí)數(shù)值然后再運(yùn)算最后還要?jiǎng)?chuàng)建一個(gè)新的臨時(shí)對(duì)象。每一步都有開銷積少成多就成了肉眼可見的慢。NumPy 的 ndarray 是完全不同的設(shè)計(jì)它是一整塊連續(xù)的內(nèi)存所有元素按同樣的 dtype數(shù)據(jù)類型緊密排列。就好比讀一沓裝訂好的 A4 紙你可以順著頁(yè)碼一目十行而 Python list 像是一本貼滿了便簽的書每讀一條都要翻去另一個(gè)章節(jié)。計(jì)算機(jī) CPU 讀連續(xù)內(nèi)存的時(shí)候cache 命中率極高現(xiàn)代編譯器還能自動(dòng)生成 SIMD單指令流多數(shù)據(jù)流指令一次處理多個(gè)數(shù)值。這還不算完NumPy 很多線性代數(shù)運(yùn)算背后直接接的是 BLAS、LAPACK 這類被優(yōu)化了幾十年的數(shù)值庫(kù)???Python 自己不可能有這種性能。1.3 什么時(shí)候不用急著上 NumPy當(dāng)然我也不是讓你所有代碼都強(qiáng)行 NumPy 化。如果你只是處理幾十個(gè)商品價(jià)格、幾個(gè)學(xué)生的成績(jī)、或者做一些字符串操作那直接用 Python list 就行引入 NumPy 反而增加依賴和心智負(fù)擔(dān)。真正的分界線在于一旦出現(xiàn)“批量數(shù)值運(yùn)算、多維數(shù)組、矩陣變換、統(tǒng)計(jì)分析、數(shù)據(jù)處理”NumPy 就應(yīng)該成為默認(rèn)選項(xiàng)。另一個(gè)判斷標(biāo)準(zhǔn)是性能——如果同一個(gè)操作你發(fā)現(xiàn)自己在寫雙層甚至三層 for 循環(huán)而且每層都要做浮點(diǎn)運(yùn)算那大概率寫錯(cuò)了換成 NumPy 一行就能解決。2. 裝對(duì) NumPy安裝方法、版本不匹配與多環(huán)境排查2.1 三種安裝方式與統(tǒng)一驗(yàn)證手法很多人在 NumPy 安裝上卡殼其實(shí)不是不會(huì)裝而是裝錯(cuò)了地方。我最推薦的安裝方式是用 Python 自己的模塊工具而不是直接敲 pippython -m pip install --upgrade pip python -m pip install numpy為什么強(qiáng)調(diào)python -m pip因?yàn)橹苯忧胮ip install時(shí)你沒法保證這個(gè) pip 屬于當(dāng)前正在使用的那個(gè) Python。尤其是 macOS 和 Linux 這類自帶多個(gè) Python 的機(jī)器上系統(tǒng)里可能同時(shí)存在 3.8、3.9、3.11 好幾個(gè)解釋器pip指到的那個(gè)未必是你的項(xiàng)目用的那個(gè)。用python -m pip相當(dāng)于明確告訴系統(tǒng)請(qǐng)?zhí)嫖疫@個(gè) Python 安裝。如果你是 Anaconda 用戶也可以走 condaconda install -c conda-forge numpyconda 的好處是會(huì)自動(dòng)解析依賴尤其在后續(xù)裝了 pandas、scipy、opencv 這一大家子的時(shí)候能在很大程度上避免依賴沖突。安裝完成后建議統(tǒng)一用下面這段驗(yàn)證別只看安裝日志python -c import numpy; print(numpy.__version__)如果這行能正常輸出版本號(hào)說明當(dāng)前 Python 環(huán)境下已經(jīng)有可用的 NumPy 了。2.2 版本不匹配的兩個(gè)典型癥狀NumPy 的版本坑最近兩年尤其值得注意。自 NumPy 2.0 發(fā)布以來因?yàn)樗鼘?duì) C API 做了一些不兼容調(diào)整如果你環(huán)境里還有舊版 pandas、opencv、scipy 或者 TensorFlow 跑在 1.x 時(shí)代很容易撞出問題。最常見的兩個(gè)癥狀I(lǐng)mportError: numpy.core.multiarray failed to importAttributeError: module numpy has no attribute float第二個(gè)尤其常見。早年很多代碼里會(huì)寫np.float、np.int、np.bool在新版本 NumPy 里這些名字已經(jīng)被正式移除了正確做法是直接用 Python 內(nèi)置的float、int、bool。如果一運(yùn)行就報(bào)這種錯(cuò)誤優(yōu)先檢查代碼里有沒有用廢棄別名而不是急著降級(jí) NumPy。排查版本問題有個(gè)通用思路先把依賴樹看清楚。在項(xiàng)目環(huán)境里執(zhí)行python -m pip list | grep -i numpy python -c import numpy; print(np.__version__)如果用的是 conda可以conda list | grep numpy配合conda update --all嘗試解決依賴關(guān)系。一般原則是優(yōu)先升級(jí)那些依賴 NumPy 的庫(kù)而不是偷偷降級(jí) NumPy因?yàn)樾马?xiàng)目可能已經(jīng)依賴 2.x 的某些特性。2.3 多 Python 環(huán)境下的“裝錯(cuò)地方”問題還有一個(gè)特別常見的坑系統(tǒng)里存在多個(gè) Python你在終端里運(yùn)行 Python 顯示 3.9結(jié)果 pip 裝完卻跑到了 3.8 的環(huán)境。排查思路很簡(jiǎn)單先確認(rèn)當(dāng)前 Python 的真實(shí)路徑which python python -c import sys; print(sys.executable)然后再執(zhí)行python -m pip install numpy確保裝進(jìn)sys.executable指向的那個(gè)環(huán)境。我遇到過一次很典型的場(chǎng)景代碼在 VS Code 里調(diào)試一切正常但轉(zhuǎn)到 Jupyter Notebook 之后import numpy直接報(bào)錯(cuò)。原因就是 Jupyter 內(nèi)核選錯(cuò)了它啟動(dòng)的是另一個(gè) Python 解釋器而 numpy 裝在了 VS Code 用的那個(gè)解釋器上。解決方式不是反復(fù)pip install而是把 Jupyter 的內(nèi)核切換到正確環(huán)境或者在虛擬環(huán)境里重新安裝 ipykernel。記住一句話遇到 import 失敗先查sys.executable再查版本。3. ndarray 到底快在哪dtype、軸與 NCHW 布局3.1 ndarray 的核心組成數(shù)據(jù)塊、dtype、shape 與 stridesNumPy 的核心抽象是一個(gè)叫 ndarrayN 維數(shù)組對(duì)象的東西。它并不是一個(gè)簡(jiǎn)單的“列表套列表”而是由幾個(gè)關(guān)鍵字段組成的內(nèi)存結(jié)構(gòu)字段作用data指向一塊連續(xù)內(nèi)存的指針dtype每個(gè)元素的類型決定元素占多少字節(jié)shape每個(gè)維度的大小例如 (3, 4) 表示 3 行 4 列strides沿每個(gè)維度移動(dòng)一步需要跳過的字節(jié)數(shù)可以用一小段代碼觀察這些信息import numpy as np arr np.zeros((3, 4), dtypenp.float32) print(arr.dtype) # float32 print(arr.shape) # (3, 4) print(arr.strides) # (16, 4)這里的 strides 有點(diǎn)意思第一維步長(zhǎng)是 16 字節(jié)說明從第 0 行跳到第 1 行要跳過 4 個(gè) float3216 字節(jié)第二維步長(zhǎng)是 4 字節(jié)說明在同一行內(nèi)移動(dòng)一個(gè)元素要跨過 4 字節(jié)。NumPy 很多操作本質(zhì)上只是改 strides 而不動(dòng)數(shù)據(jù)比如轉(zhuǎn)置、reshape這也是它們能那么快的原因之一。對(duì)比一下 Python list 和 ndarray 的差異會(huì)更清晰維度Python listNumPy ndarray存儲(chǔ)方式對(duì)象指針數(shù)組元素散落在堆上連續(xù)內(nèi)存塊同類型數(shù)據(jù)緊鄰排列元素類型可以混著 int/str/對(duì)象必須統(tǒng)一由 dtype 決定逐元素操作解釋器循環(huán)慢C 層循環(huán)快適用場(chǎng)景異構(gòu)數(shù)據(jù)、小數(shù)據(jù)量、邏輯拼接同構(gòu)數(shù)值、批量運(yùn)算、多維矩陣3.2 dtype 選型一份隱藏的性能賬本dtype 是你最容易忽略、但影響極大的一個(gè)維度。同樣一個(gè) 1024×1024×3 的 RGB 圖像如果以u(píng)int80~255 無符號(hào)整數(shù)存儲(chǔ)占 3MB轉(zhuǎn)成float32變成 12MB如果哪一步不小心轉(zhuǎn)成了默認(rèn)的float64直接翻到 24MB。當(dāng)你手里有幾千張圖片、幾十個(gè)特征矩陣時(shí)這個(gè)差距就是幾十 GB 與幾百 GB 的區(qū)別。我的實(shí)操建議是沒有小數(shù)精度需求的數(shù)據(jù)能用uint8或int64就別用浮點(diǎn)深度學(xué)習(xí)預(yù)處理階段用float32就夠了沒必要扛著float64讀取 CSV 時(shí) pandas 經(jīng)常給數(shù)值列默認(rèn)float64如果只是為了算均值、喂模型可以主動(dòng)astype(float32)省一半內(nèi)存對(duì)于大矩陣乘法float32還能順便享受更寬的 SIMD 吞吐部分機(jī)器上速度也會(huì)提升。3.3 axis 與 NCHW 布局從圖像到張量很多做深度學(xué)習(xí)的人第一次接觸 NumPy 的多維數(shù)組會(huì)卡在“軸axis”這個(gè)概念上。一張 CHW 格式的圖片在 NumPy 里就是一個(gè)三維數(shù)組三個(gè)軸分別表示通道、高度、寬度。如果再疊一個(gè) batch 維度就變成了四維的 NCHW 布局形狀是 (Batch, Channel, Height, Width)。NCHW 這個(gè)詞在 PyTorch、TensorFlow 的底層層層出現(xiàn)但剝開看它就是一個(gè)四維 ndarray 的 shape 約定。你只需要記住每個(gè) axis 代表什么就能順暢操作import numpy as np # 假裝是一張 4x4 的 RGB 圖通道數(shù) 3按 CHW 排 img np.random.randint(0, 255, (3, 4, 4), dtypenp.uint8) red_channel img[0] # 取 R 通道形狀 (4, 4) pixel img[:, 2, 3] # 取第 3 行第 4 列的三個(gè)通道值 hwc_img img.transpose(1, 2, 0) # 變成 HWC 布局形狀 (4, 4, 3) batch np.stack([img, img], axis0) # 變成 NCHW形狀 (2, 3, 4, 4)處理視頻時(shí)還會(huì)多一個(gè)時(shí)間軸 T變成五維的 (N, C, T, H, W)。但只要理解了 axis 的順序這些無非是 shape 里多了一個(gè)數(shù)字而已。我之前給新人講 axis 的時(shí)候總用一句話axis0 是“最外層”axis-1 是“最內(nèi)層”從外往里看數(shù)據(jù)準(zhǔn)沒錯(cuò)。4. 從行列式到線性代數(shù)NumPy 的一行方案 vs 純 Python 手寫4.1 行列式一句 linalg.det vs 一段手寫展開有個(gè)話題在社區(qū)里經(jīng)常被討論行列式計(jì)算能不能不用 NumPy能但沒必要。先看看手寫版本有多啰嗦。2×2 的行列式還算簡(jiǎn)單def det2(a): return a[0][0] * a[1][1] - a[0][1] * a[1][0]3×3 就已經(jīng)需要按行展開一次了def det3(a): return ( a[0][0] * (a[1][1]*a[2][2] - a[1][2]*a[2][1]) - a[0][1] * (a[1][0]*a[2][2] - a[1][2]*a[2][0]) a[0][2] * (a[1][0]*a[2][1] - a[1][1]*a[2][0]) )再往上按代數(shù)余子式展開的復(fù)雜度是 O(n!)10×10 矩陣就已經(jīng)慢到?jīng)]法用了。就算你改進(jìn)成高斯消元法也要自己處理部分主元選擇、浮點(diǎn)誤差、零值判斷等一系列問題。而 NumPy 只需要一行import numpy as np A np.array([[1., 2.], [3., 4.]]) det_A np.linalg.det(A) # 輸出 -2.0np.linalg.det背后調(diào)用的 LAPACK 庫(kù)里的 LU 分解實(shí)現(xiàn)行列式等于對(duì)角元素的乘積乘以符號(hào)修正既快又穩(wěn)定。這不是“用高級(jí)工具偷懶”而是把成熟的數(shù)值算法直接拿來用。4.2 一個(gè)能直接用的線性代數(shù)工具箱NumPy 的linalg模塊基本覆蓋了你會(huì)用到的所有線性代數(shù)需求需求推薦用法矩陣乘法a b或np.matmul(a, b)轉(zhuǎn)置a.T或np.transpose(a)行列式np.linalg.det(a)逆矩陣np.linalg.inv(a)解線性方程組np.linalg.solve(A, b)特征值/特征向量np.linalg.eig(a)或np.linalg.eigh(a)奇異值分解np.linalg.svd(a)最小二乘解np.linalg.lstsq(a, b)這里我想特別強(qiáng)調(diào)一個(gè)新手容易踩的坑解線性方程組時(shí)不要寫成x np.linalg.inv(A) b。雖然數(shù)學(xué)上等價(jià)但在數(shù)值上直接求逆再乘會(huì)把誤差放大而且多花一倍以上的時(shí)間。正確做法是x np.linalg.solve(A, b)它內(nèi)部走 LU 分解穩(wěn)定性和效率都好得多。A np.array([[3., 1.], [1., 2.]]) b np.array([9., 8.]) x np.linalg.solve(A, b) print(x) # [2. 3.]4.3 數(shù)值穩(wěn)定性一個(gè)容易被忽略的“隱形坑”講一個(gè)比較實(shí)際的例子希爾伯特矩陣。這類矩陣每個(gè)元素是1 / (i j 1)看著很干凈但條件數(shù)大得離譜稍微大一點(diǎn)的行列式用樸素方法算出來幾乎不可信。你完全可以用np.linalg.cond檢查一個(gè)矩陣的病態(tài)程度比如H np.array([[1 / (i j 1) for j in range(10)] for i in range(10)]) print(np.linalg.cond(H)) # 輸出會(huì)是一個(gè)巨大的數(shù)字條件數(shù)越大說明矩陣對(duì)數(shù)值誤差越敏感。遇到這種矩陣再牛的庫(kù)也會(huì)算得勉強(qiáng)。這也提醒我們用 NumPy 不等于“永遠(yuǎn)準(zhǔn)確”理解背后的數(shù)值原理才能解釋為什么有些結(jié)果看起來不對(duì)勁。5. 索引、切片與廣播從“怎么寫代碼”到“怎么省時(shí)間”5.1 切片到底復(fù)制了嗎視圖與副本的經(jīng)典陷阱如果你是從 Python 轉(zhuǎn)過來的這里有個(gè)特別容易踩的坑Python 的 list 切片list[:]會(huì)生成一個(gè)全新的列表而 NumPy 的切片通常返回的是原數(shù)組的“視圖”也就是說它不復(fù)制數(shù)據(jù)只是給你一個(gè)新的“窗口”去看同一塊內(nèi)存。a np.arange(12).reshape(3, 4) b a[1:, :] # 第二行開始的所有列 b[0, 0] 999 print(a[1, 0]) # 999原數(shù)組 a 也被改了我當(dāng)時(shí)第一次遇到時(shí)排查了很久差點(diǎn)以為是內(nèi)存被外部破壞了。要判斷兩個(gè)數(shù)組是否共享底層數(shù)據(jù)可以這樣print(np.shares_memory(a, b)) # True如果確實(shí)希望得到獨(dú)立副本記得顯式調(diào)用.copy()b a[1:, :].copy()現(xiàn)在我可以直接給一條經(jīng)驗(yàn)凡是對(duì)切片結(jié)果有“我接下來要修改它”的打算先想清楚是想要視圖還是副本。視圖省內(nèi)存、速度快適合讀副本適合寫但會(huì)占用額外空間。這個(gè)思維一旦建立很多詭異的 bug 都能避免。5.2 廣播機(jī)制NumPy 的“隱性復(fù)制”廣播是 NumPy 里最優(yōu)雅、也最讓人困惑的東西。簡(jiǎn)單來說當(dāng)兩個(gè)數(shù)組形狀不一致時(shí)NumPy 會(huì)嘗試把較小的數(shù)組“擴(kuò)展”到較大的形狀再做運(yùn)算。這個(gè)過程并不會(huì)真的復(fù)制內(nèi)存而是在計(jì)算層面虛擬地“鋪開”。規(guī)則其實(shí)就一句話從最后一個(gè)維度往前看兩個(gè)維度要么相等要么其中一個(gè)是 1要么其中一個(gè)沒有這個(gè)維度。比如一個(gè)形狀 (2, 3) 的矩陣減去一個(gè)形狀 (3,) 的向量NumPy 會(huì)自動(dòng)把向量沿第一維復(fù)制一份完成逐行操作scores np.array([[80, 90, 100], [70, 85, 95]]) # 兩個(gè)學(xué)生的三科成績(jī) mean_score np.mean(scores, axis1).reshape(-1, 1) # 變成 (2, 1) centered scores - mean_score # (2, 3) - (2, 1) - 廣播如果忘了reshape(-1, 1)直接用形狀 (2,) 的均值去減 (2, 3) 的矩陣NumPy 會(huì)直接報(bào)錯(cuò)提示這兩個(gè)形狀無法廣播。這時(shí)候先在紙上畫出每個(gè)數(shù)組的 shape再?zèng)Q定要不要加維度基本不會(huì)錯(cuò)。5.3 用向量化思維改寫三層 for 循環(huán)很多人一開始寫數(shù)值代碼腦子還是 C 語(yǔ)言那套循環(huán)思維。舉個(gè)最常見的歸一化例子# 反例逐元素循環(huán) out np.empty_like(x) for i in range(n): out[i] (x[i] - mu) / sigma用向量化一行就搞定out (x - mu) / sigma這不是魔法原因是 NumPy 的-和/運(yùn)算符底層用 C 幫你遍歷了每一個(gè)元素。更復(fù)雜一點(diǎn)的場(chǎng)景比如按行減均值并除以標(biāo)準(zhǔn)差也是同樣的思路# 高維矩陣按列標(biāo)準(zhǔn)化 mean x.mean(axis0) std x.std(axis0) y (x - mean) / std # 整段操作幾毫秒完成如果你發(fā)現(xiàn)自己還在堆for i in range還嫌慢第一反應(yīng)應(yīng)該是能不能把這個(gè)循環(huán)變成數(shù)組運(yùn)算很多時(shí)候答案都是能。6. 提速三板斧與實(shí)戰(zhàn)心得向量化、dtype 與可復(fù)現(xiàn)性6.1 三板斧向量化、選 dtype、確認(rèn)隱藏的拷貝我?guī)蛣e人調(diào) NumPy 性能問題的時(shí)候基本就是三板斧。第一板斧是向量化把能寫成數(shù)組表達(dá)式的運(yùn)算全部寫成數(shù)組表達(dá)式原理上文已說過。第二板斧是選擇合適 dtype別讓全工程默默用著 float64。第三板斧是留意那些隱藏的拷貝。這里要重點(diǎn)提一下np.array和np.asarray的區(qū)別y np.asarray(x) # 如果 x 本來就是 ndarray返回同一個(gè)對(duì)象不復(fù)制 y np.array(x) # 無論 x 是什么總是復(fù)制一份在寫函數(shù)的時(shí)候我傾向于用np.asarray做輸入轉(zhuǎn)換因?yàn)樗茉诓恍枰獜?fù)制的時(shí)候替我省掉一大塊內(nèi)存和時(shí)間。反過來如果確定要獨(dú)立修改輸入再用np.array避免不小心改了調(diào)用方的原始數(shù)據(jù)。比較隱蔽的另一個(gè)拷貝點(diǎn)是切片之后做排序、翻轉(zhuǎn)這類操作時(shí)。比如arr[::-1]從語(yǔ)義上看是倒序遍歷但如果你對(duì)它再調(diào)用.sort()或者賦值操作就可能產(chǎn)生中間副本。建議遇到性能異常時(shí)用np.shares_memory和arr.flags先看看數(shù)據(jù)是不是真的連續(xù)了。6.2 可復(fù)現(xiàn)性seed 與新版 RNG聊性能之外還有一個(gè)每個(gè)做實(shí)驗(yàn)的人都會(huì)遇到的痛點(diǎn)隨機(jī)數(shù)的可復(fù)現(xiàn)性。老式寫法是這樣np.random.seed(42) samples np.random.normal(size1000)問題在于np.random.seed設(shè)置的是全局隨機(jī)狀態(tài)只要中間任何第三方庫(kù)偷偷調(diào)用了np.random你的“復(fù)現(xiàn)”就失效了。更穩(wěn)妥的做法是用較新的default_rng接口它創(chuàng)建的是一個(gè)獨(dú)立、可控的隨機(jī)數(shù)生成器rng np.random.default_rng(42) samples rng.normal(size1000)我實(shí)際遇到過不止一次明明在開頭設(shè)置了np.random.seed(42)后面每次跑出來的結(jié)果還是不一致最后發(fā)現(xiàn)是某個(gè)數(shù)據(jù)處理庫(kù)在你不知情的時(shí)候用了全局隨機(jī)狀態(tài)。切換到default_rng之后這類問題從此絕跡。這也是我在帶新項(xiàng)目時(shí)強(qiáng)烈推薦的做法隨機(jī)狀態(tài)自己管理不讓全局狀態(tài)背鍋。6.3 一個(gè)真實(shí)項(xiàng)目的優(yōu)化記錄三層循環(huán)改成向量化最后分享一個(gè)我印象很深的優(yōu)化案例。之前接手一套圖像預(yù)處理的舊代碼邏輯本身不復(fù)雜讀一批圖片對(duì)每張圖做歸一化、裁剪、翻轉(zhuǎn)增強(qiáng)生成訓(xùn)練數(shù)據(jù)。原始代碼用三層 for 循環(huán)寫外層遍歷圖片中層遍歷通道內(nèi)層逐像素計(jì)算。本身邏輯完全正確但跑 3 萬張圖要將近四十分鐘嚴(yán)重影響試驗(yàn)迭代速度。我把整個(gè)流程改成 NumPy 的 batch 化處理先把所有圖片讀進(jìn)一個(gè)四維數(shù)組 (N, C, H, W)歸一化直接用(x - mean) / std裁剪和翻轉(zhuǎn)用切片和np.flip水平翻轉(zhuǎn)直接x[:, :, :, ::-1]。改完以后同樣的 3 萬張圖預(yù)處理耗時(shí)降到幾十秒提速非常明顯。而且代碼還更短了可讀性反而更好。那次之后我養(yǎng)成了一個(gè)習(xí)慣寫任何數(shù)值計(jì)算代碼第一版就盡量用數(shù)組表達(dá)式而不是循環(huán)。如果實(shí)在無法避免循環(huán)至少把最內(nèi)層的運(yùn)算向量化。很多時(shí)候真輪不到上并行、上 GPU先把 NumPy 的向量化吃透性能就已經(jīng)夠用了。