機器學習和科學領域的許多問題都歸結為相同的任務:你擁有一組資料點,並希望還原它們所來自的分佈,即哪些值常見,哪些值稀有。確定該分佈意味著估計兩個量:分佈的密度,以及隨著維度增加而變得更有用的分數。密度是直方圖的平滑版本,在點聚集的地方高,在點稀疏的地方低。分數是對數密度的梯度,指向密度上升最快的方向:沿著分數移動一個點,它就會朝著更可能發生的區域前進。
基於擴散的生成模型(如 Stable Diffusion 和 DALL-E 等 AI 圖像生成器背後的技術)從隨機雜訊開始,並重複遵循分數,將雜訊轉化為逼真的圖像。同樣的分數也驅動著貝氏取樣以及用於模擬電漿等系統的粒子模擬。
從有限樣本中提取密度和分數極具挑戰性,而現今的工具需要在泛化能力和準確性之間做出權衡。一種經典方法是核密度估計 (KDE),它根據周圍的資料點計算任何位置的密度:點越近越多,密度就越高。KDE 無需訓練,適用於任何分佈,但其準確性會隨著維度增加而急劇下降。
另一種方法是神經分數匹配模型,它們經過訓練來預測分數,即使在高維度下也能保持準確,但每個模型都需要學習特定的分佈,並且必須為不同的分佈從頭開始重新訓練。
我們引入了一種名為 DiScoFormer(密度與分數 Transformer)的新解決方案,這是一個單一模型,給定一組資料點後,它能在單次前向傳播中估計分佈的密度和分數,而無需重新訓練。
DiScoFormer 利用堆疊的 Transformer 區塊層,將整個樣本映射到其背後分佈的密度和分數。該模型採用交叉注意力機制,使其能夠在任何點(而不僅僅是資料點所在位置)評估密度和分數。分數和密度之間存在數學關係:分數是對數密度的梯度。我們利用這一點,設計了一個共享骨幹網路,並配備兩個輸出頭,一個用於密度,一個用於分數。
這種耦合不僅節省了參數。分數頭必須在每個查詢點與對數密度頭的梯度相匹配,因此它們之間的任何差異都是一種無需標籤的一致性損失。我們在推論時利用這一點:固定上下文,對該一致性損失執行幾個梯度步驟,DiScoFormer 就能即時適應分佈外輸入,無需真實值的密度或分數。
Transformer 架構之所以適合這項任務,有其數學上的原因。核密度估計 (KDE) 具有單一頻寬,即每個點的影響範圍,預先固定並在各處相同應用。注意力機制是其嚴格的推廣:我們分析證明,單一注意力頭的權重幾乎是資料上的高斯核,因此一個交叉注意力區塊就能重現 KDE 的密度和分數。
在此基礎上,該模型更進一步,同時學習多個此類尺度並將其適應於資料。DiScoFormer 並非拋棄經典方法而採用黑盒子,而是將 KDE 作為一個特例包含在內並加以改進。
我們使用什麼資料來訓練 DiScoFormer?我們主要基於兩個原因依賴高斯混合模型 (GMM)。首先,GMM 是通用的密度近似器,只要有足夠的組件,它們就能以任意小的誤差匹配幾乎任何平滑分佈。其次,GMM 具有封閉形式的密度和分數,因此我們總是有一個精確的目標可以進行監督。
我們利用這兩個特性,為每個批次繪製一個新的 GMM,為模型提供幾乎無限的目標分佈範例,並針對給定 GMM 的精確密度和分數進行監督。
總體而言,DiScoFormer 在密度和分數估計方面都超越了 KDE,而且差距在 KDE 表現不佳的地方尤其顯著。在 100 維度下,DiScoFormer 的表現遙遙領先,相較於最佳手動調整的 KDE,它將分數誤差降低了約 6.5 倍,密度誤差降低了超過 37 倍,並且隨著樣本的增加而持續改進,而 KDE 則會耗盡記憶體。
它也能遠離其訓練資料,在具有比訓練期間見過更多模式的混合分佈以及像拉普拉斯和學生 t 分佈等非高斯形狀上保持準確。KDE 的主要優勢仍然是速度,尤其是在資料集較小的時候。
我們認為 DiScoFormer 最有前景的部分在於,分數估計是生成模型、貝氏推論和科學計算等許多領域的共同依賴項。一個預訓練的、可插拔的估計器,在高維度下仍能保持準確,並消除了針對每個問題重新訓練的需求,這可以同時降低所有這些領域的成本,一個模型,在任何出現分數和密度的地方重複使用。
我們鼓勵您閱讀我們的技術報告以獲取更多詳細資訊。



