Skip to content
Greg's Space
Go back

什麼是 SIGReg?用手電筒影子檢查 AI 有沒有偷懶

現在很多 AI 模型是「自監督學習」:沒有人告訴它標準答案,它要自己從一大堆圖片、文字裡摸索出規律。聽起來很酷,但有個很麻煩的問題——模型很會偷懶。

SIGReg 就是專門用來抓「偷懶」的檢查員。這篇文章會用最白話的方式,帶你看懂它在做什麼,再一行一行拆解它的 PyTorch 程式碼。

先認識 embedding:把東西變成一串數字

電腦看不懂照片或句子,所以模型會先把每一筆資料變成一串數字,這串數字叫 embedding。

你可以把它想成每個人的「特徵清單」:

同學身高體重愛吃甜食程度…
小明0.8-0.31.2…
小華-1.10.5-0.4…

在我們的程式裡,一次會處理 64 筆資料(batch_size = 64),每一筆有 256 個數字(dim = 256)。所以 embedding z 就是一張 64 × 256 的大表格:64 個人,每人 256 項特徵。

問題:模型偷懶了(collapse)

好的特徵清單,應該要能讓大家看得出誰跟誰不一樣。但模型有時候會走捷徑:

第二種叫做維度塌縮——原本有 256 個維度可以用,實際上只用了 16 個,其他 240 個都浪費掉了。

我們範例程式一開始就故意做了這種「壞掉的 embedding」:

z = torch.zeros(batch_size, dim)
z[:, 0:16] = torch.randn(batch_size, 16)          # 前 16 項:認真填
z[:, 16:256] = torch.randn(batch_size, 240) * 0.01  # 後 240 項:幾乎是 0

SIGReg 的工作,就是看出這個 embedding 有問題,並且給一個很大的「懲罰分數」,逼模型改進。

SIGReg 的想法:要求 embedding 長得「圓滾滾」

那什麼樣的 embedding 才叫健康?SIGReg 的答案是:像一顆圓滾滾的標準常態分布雲團。

這種圓滾滾的雲團有個專有名詞叫 isotropic Gaussian(isotropic 就是「各方向一樣」的意思)。整個 SIGReg 的全名 Sketched Isotropic Gaussian Regularization,就是「用抽樣的方式,逼 embedding 變成各向同性的高斯分布」。

怎麼檢查 256 維的雲團?用手電筒照影子!

問題來了:256 維的東西,我們根本想像不出來,要怎麼檢查它圓不圓?

想像你手上有一個看不清楚的立體物體,你可以拿手電筒從不同角度照它,看牆上的影子:

這就是 SIGReg 的策略:不直接檢查 256 維,而是隨機挑一些方向,把資料「投影」成一維的影子,再檢查影子像不像鐘形曲線。

因為我們只是「隨機抽幾個方向來看」而不是把所有方向都看一遍,所以叫 Sketched(素描,抓大概的輪廓)。

怎麼判斷影子像不像鐘形曲線?

判斷一堆數字像不像鐘形曲線,這裡用的是 Epps-Pulley 檢定。它的原理有點抽象,我用一個比喻:

把數線像捲尺一樣「捲」成一個圓圈,每個數字都會落在圈上的某個位置。然後算出所有落點的重心。

實際重心 跟 理想位置 差多遠,就是這個影子的「不像程度」。

那個 t 決定「捲的鬆緊程度」,也就是用多粗、多細的眼光去看。只用一種 t 可能被騙過去,所以我們會用好幾種不同的 t(頻率)都檢查一遍,再加總起來。

一步一步看程式

完整程式有 7 個步驟,我們一個一個來。

先看參數:

batch_size = 64        # 一次處理 64 筆資料
dim = 256              # 每筆資料 256 個數字
num_directions = 100   # 隨機挑 100 個方向來照影子
t_max = 3.0            # 「捲尺鬆緊」最大取到 3
n_points = 9           # 在 0 到 3 之間,取 9 種鬆緊程度來檢查

第 1 步:準備 embedding

就是前面那個故意做壞的 z。真實訓練的時候,z 會是 encoder(把資料變成數字的模型)的輸出,而且要加上 requires_grad_(True),這樣才算得出梯度、才能回頭調整模型。

第 2 步:隨機挑手電筒的方向

directions = torch.randn(dim, num_directions)     # (256, 100)
directions = directions / directions.norm(dim=0)   # 每一欄變成長度 1

我們隨機生出 100 個方向,每個方向是一個 256 維的箭頭。除以 norm(長度)是把每個箭頭的長度都調成 1,這樣每支手電筒才「公平」,不會有的亮有的暗。

第 3 步:照影子(投影)

projected = z @ directions   # (64, 256) @ (256, 100) -> (64, 100)

一個矩陣相乘,就同時把 64 個人、投影到 100 個方向上。結果是一張 64 × 100 的表格:每一格是「第 i 個人,在第 j 個方向上的影子位置」。

原本用 for 迴圈要寫 64 × 100 = 6400 次,現在一行搞定。

第 4 步:準備固定的常數

test_frequencies = torch.linspace(0, t_max, n_points)   # 0, 0.375, 0.75, ..., 3.0

dt = t_max / (n_points - 1)
trapezoid_weight = torch.full((n_points,), 2 * dt)
trapezoid_weight[0] = dt
trapezoid_weight[-1] = dt

phi_at_t = torch.exp(-0.5 * test_frequencies**2)         # 理想重心的位置
final_weight = trapezoid_weight * phi_at_t

這一步做了三件事:

  1. test_frequencies:要檢查的 9 種「鬆緊程度」。
  2. phi_at_t:對每一種鬆緊程度,鐘形曲線的重心「理想上」該在哪裡(就是 exp(-t²/2))。
  3. final_weight:加總 9 種鬆緊程度時,每一種要乘多少權重。

權重是怎麼來的?有兩個部分:

這些數字從頭到尾都不會變,所以正式使用時,通常在模型初始化的時候就先算好存起來,不用每次算 loss 都重算一遍。

第 5 步:一次算完所有「方向 × 頻率」(關鍵!)

signal = projected.unsqueeze(-1) * test_frequencies   # (64, 100, 1) * (9,) -> (64, 100, 9)

cos_mean = torch.cos(signal).mean(dim=0)   # (100, 9)
sin_mean = torch.sin(signal).mean(dim=0)   # (100, 9)

difference = (cos_mean - phi_at_t) ** 2 + sin_mean ** 2   # (100, 9)

這是整份程式最精彩的地方,要拆開看:

(1) unsqueeze(-1) 與廣播(broadcasting)

projected 的形狀是 (64, 100),unsqueeze(-1) 在最後面多加一個維度,變成 (64, 100, 1)。再跟形狀 (9,) 的 test_frequencies 相乘時,PyTorch 會自動把它們「補齊、複製」成 (64, 100, 9)。

也就是說,一行程式就算出了64 個人 × 100 個方向 × 9 種鬆緊程度,共 57,600 個組合。

(2) cos 與 sin:把數字捲成圓圈

還記得前面「把數線捲成圓圈」的比喻嗎?每個數字落在圓圈上的位置,可以用兩個座標表示:cos(左右)跟 sin(上下)。

(3) .mean(dim=0):算重心

對 64 個人取平均,就得到「重心」的位置,形狀變成 (100, 9)——每個方向、每種鬆緊程度,各有一個重心。

(4) difference:重心跟理想位置差多遠

理想的重心在 (phi_at_t, 0)。實際重心是 (cos_mean, sin_mean)。兩點距離的平方就是:

(cos_mean - phi_at_t)² + (sin_mean - 0)²

也就是程式裡那一行。差越多,數字越大。

第 6 步:加權加總

per_direction_score = (difference @ final_weight) * batch_size   # (100, 9) @ (9,) -> (100,)

把 9 種鬆緊程度的差距,依照權重加起來,每個方向得到一個分數。

至於最後為什麼要乘 batch_size?這是 Epps-Pulley 檢定公式本身就有的設計:樣本數越多,越有把握判斷,所以差距要被放大。

具體用小例子算一次(不用想像,直接看數字)

前面講的都是形狀、原理,這裡我們縮小規模,真的把數字全部算出來,讓你看到「一個方向的分數」是怎麼從頭跑到尾的。

假設方向 j 上,投影值只有 4 筆(不是真實的 64 筆,純粹方便手算):

P_{:,j} = [-1, -0.3, 0.3, 1]

t_max=3、n_points=9 跟真實程式一樣,所以還是那 9 個 t:0, 0.375, 0.75, 1.125, 1.5, 1.875, 2.25, 2.625, 3.0。

對每一個 t,都重複同樣四個動作:① t · p_i → ② 代入 cos/sin → ③ 對 4 筆資料取平均(cos_mean、sin_mean) → ④ 跟 phi_at_t 比較,算出 difference。算完 9 個 t 之後,再各自乘上 final_weight:

表格每一欄,是怎麼算出來的:

欄位公式說明
cos_mean(cos(t·p_1) + cos(t·p_2) + cos(t·p_3) + cos(t·p_4)) / 44 筆資料的投影值,各自乘上 t、代入 cos,再取平均
sin_mean(sin(t·p_1) + sin(t·p_2) + sin(t·p_3) + sin(t·p_4)) / 4同上,只是換成 sin
phi_at_texp(-0.5 × t²)直接代公式算,標準常態分布的 ground truth,不需要資料
difference(cos_mean - phi_at_t)² + sin_mean²「用資料估計出來的值」跟「ground truth」差多遠,平方是因為只在乎差距大小、不在乎正負
trapezoid_weight頭尾是 dt,中間是 2·dt(其中 dt = t_max/(n_points-1) = 3/8 = 0.375)只跟 t 在 9 個點裡的位置有關,跟資料、跟 phi_at_t 都無關
final_weighttrapezoid_weight × phi_at_t把「這個點該佔多少範圍」跟「這個點可不可信」兩件事合併成一個權重
difference × final_weight就是字面上的相乘這一欄 9 個數字加起來,就是這個方向的分數(還沒乘 batch_size 之前)

實際數字:

tcos_meansin_meanphi_at_tdifferencetrapezoid_weightfinal_weightdifference × final_weight
0.0001.00000.00001.00000.0000000.375(頭)0.37500.000000
0.3750.96210.00000.93210.0008990.7500.69910.000629
0.7500.85320.00000.75480.0096830.7500.56610.005482
1.1250.68740.00000.53110.0244250.7500.39830.009729
1.5000.48560.00000.32470.0259020.7500.24350.006307
1.8750.27320.00000.17240.0101550.7500.12930.001313
2.2500.07630.00000.07960.0000110.7500.05970.000001
2.625-0.08190.00000.03190.0129590.7500.02390.000310
3.000-0.18420.00000.01110.0381420.375(尾)0.00420.000159

最右邊那欄,9 個數字加起來:0.023929。

再乘上 batch_size(這裡的小例子只有 4 筆資料,batch_size=4):

方向 j 的分數 = 0.023929 × 4 ≈ 0.0957

這個 0.0957,就是方向 j 這一個方向的最終分數——對應到真實程式裡 per_direction_score 這個長度 100 的向量,其中一格的值。真實情況下,64 筆資料、100 個方向,也是重複一模一樣的四個動作,只是規模大很多,而且全部用矩陣運算一次做完,不會真的寫迴圈。

幾個從這張表可以直接看到的重點:

第 7 步:所有方向取平均

sigreg_loss = per_direction_score.mean()

100 個方向的分數平均,就是最終的 SIGReg loss——一個單一的數字。

這個分數到底代表什麼?

我把幾種不同的 embedding 丟進去算,結果大概是這樣(實際數字會因為亂數而略有不同):

embedding 的樣子SIGReg loss(大約)
健康的鐘形雲團(標準常態)約 1
前面故意做壞的(只有 16 維有用)約 20
全部人都交一模一樣的答案約 40

有兩點值得注意:

健康的也不是 0。 因為只有 64 個樣本,隨機抽出來的東西本來就不可能完美,所以會有一點點抽樣誤差。

為什麼壞掉的那個分數會這麼高? 因為手電筒隨便照一個方向,影子的寬度其實取決於「有多少能量在這個方向上」。壞掉的 embedding 裡,256 個維度只有 16 個有東西,所以隨便一個方向的影子,寬度只剩正常的 16/256 = 1/16 左右——影子被壓得又瘦又扁,跟寬寬的鐘形曲線差非常遠,分數自然就爆高。

最重要的一件事:它可以拿來「訓練」

分數高低只是診斷,真正的重點是它能微分。整個計算過程只用了矩陣相乘、cos、sin、平均這些 PyTorch 認得的運算,所以:

sigreg_loss.backward()
print(z.grad.shape)   # (64, 256),跟 z 一模一樣

backward() 之後,z.grad 會告訴 z 裡面每一個數字:「往哪個方向改,可以讓 loss 變小」。

實際訓練時,這個梯度會一路傳回 encoder,逼它不要只用 16 個維度偷懶,要把 256 個維度都用起來。SIGReg 通常會跟模型本來的訓練目標加在一起,變成:

總 loss = 原本的學習目標 + λ × SIGReg loss

其中 λ 是用來調整「檢查員有多嚴格」的係數。

跟官方套件對答案

程式最後,拿自己寫的版本跟官方 lejepa 套件比較:

import lejepa

official_test = lejepa.univariate.epps_pulley.EppsPulley(t_max=t_max, n_points=n_points)
official_score = official_test(projected.detach()).mean()

print(torch.allclose(sigreg_loss.detach(), official_score, atol=1e-4))   # True

兩個數字一致,代表我們這份向量化版本沒有算錯。

為什麼說這是「向量化版本」?

原本最直覺的寫法會用三層 for 迴圈(每個方向、每個頻率、每個樣本)。這份程式把「一個一個算」換成「一次全部算」:

步驟用到的技巧
產生所有方向torch.randn(dim, num_directions) 一次生出來
投影矩陣相乘 @
方向 × 頻率的所有組合unsqueeze + 廣播
對樣本取平均.mean(dim=0)
頻率加權加總矩陣乘向量 @

邏輯完全一樣,但在 GPU 上快非常多,這才是真的能拿來訓練的寫法。

一句話總結

SIGReg 就像一個很有耐心的檢查員:拿手電筒從很多隨機角度照你的 embedding,看影子是不是都是漂亮的鐘形曲線。如果有任何一個角度的影子太瘦、太歪,就扣分;而且扣分的方式可以微分,能直接告訴模型「哪裡該改」,逼它別偷懶,把每一個維度都好好用起來。



Previous Post
Agent Skill 的 SKILL.md 怎麼寫?資料結構與完整範例