下的MoE路由穩(wěn)定與Rehearsal-free訓(xùn)練)
在大型語言模型的分布式指令微調(diào)場景里DistMoE 把三個(gè)原本分開的問題綁定到了同一個(gè)系統(tǒng)中多個(gè)數(shù)據(jù)方不能共享私有數(shù)據(jù)卻要協(xié)作微調(diào)一個(gè) Mixture-of-ExpertsMoE模型MoE 內(nèi)部的路由模塊要決定每個(gè) token 進(jìn)入哪些專家而這些路由決策在引入新任務(wù)后不能出現(xiàn)明顯漂移并且不能依賴回放舊數(shù)據(jù)來維持穩(wěn)定。單獨(dú)解決其中任何一個(gè)問題都有成熟方案但組合在一起之后難點(diǎn)就變成了路由穩(wěn)定性通???rehearsal 來維持而 rehearsal 又依賴舊數(shù)據(jù)與私有數(shù)據(jù)約束直接沖突。DistMoE 這個(gè)研究方向的核心就是在不接觸私有數(shù)據(jù)的前提下讓路由模塊在分布式指令微調(diào)過程中保持穩(wěn)定。下面會(huì)從 MoE 路由的基本原理講起搭建一個(gè)最小可復(fù)現(xiàn)的分布式指令微調(diào)實(shí)驗(yàn)框架然后說明如何用路由錨點(diǎn)分布實(shí)現(xiàn) rehearsal-free 的穩(wěn)定訓(xùn)練最后給出驗(yàn)證指標(biāo)、排查路徑和生產(chǎn)落地建議。1. 先理解 DistMoE 的三層技術(shù)背景分布式指令微調(diào)、MoE 路由、Rehearsal-free1.1 指令微調(diào)為什么需要分布式協(xié)作指令微調(diào)Instruction Tuning指的是用“指令 期望回答”這樣的監(jiān)督數(shù)據(jù)對基座大模型做有監(jiān)督微調(diào)讓模型學(xué)會(huì)按照用戶指令輸出有用回答。它與預(yù)訓(xùn)練不同預(yù)訓(xùn)練階段模型看到的是大規(guī)模無標(biāo)注文本學(xué)習(xí)的是語言統(tǒng)計(jì)規(guī)律指令微調(diào)階段模型看到的是少量、高質(zhì)量、任務(wù)導(dǎo)向的數(shù)據(jù)學(xué)習(xí)的是“知道何時(shí)該執(zhí)行什么動(dòng)作”。實(shí)際企業(yè)場景中指令數(shù)據(jù)往往分散在不同部門、不同組織甚至不同地區(qū)??头块T有客服對話數(shù)據(jù)風(fēng)控部門有風(fēng)控問答數(shù)據(jù)業(yè)務(wù)部門還有內(nèi)部規(guī)章制度數(shù)據(jù)。把數(shù)據(jù)集中到一個(gè)機(jī)房訓(xùn)練最簡單但隱私和合規(guī)成本很高。分布式指令微調(diào)要解決的問題就是“數(shù)據(jù)不動(dòng)模型動(dòng)”原始數(shù)據(jù)留在各自的本地?cái)?shù)據(jù)方只有模型參數(shù)更新或必要的統(tǒng)計(jì)信息參與協(xié)作。這種方式也被稱為數(shù)據(jù)隔離下的協(xié)作訓(xùn)練DistMoE 討論的分布式正是這種多個(gè)數(shù)據(jù)方共同參與訓(xùn)練的拓?fù)涠皇菃渭冎浮岸鄼C(jī)多卡做數(shù)據(jù)并行”。1.2 MoE 路由門控網(wǎng)絡(luò)如何決定 token 去哪個(gè)專家Mixture-of-ExpertsMoE的核心設(shè)計(jì)是把傳統(tǒng) Transformer 中的前饋網(wǎng)絡(luò)FFN替換成一組專家網(wǎng)絡(luò)并由一個(gè)路由模塊Router/Gating決定每個(gè) token 激活哪些專家。設(shè)輸入為x路由模塊輸出一個(gè)在E個(gè)專家上的概率分布然后取 top-k 個(gè)專家做加權(quán)求和。這樣做的好處是模型參數(shù)量可以很大但每個(gè) token 只計(jì)算其中一小部分專家推理和訓(xùn)練的計(jì)算成本都低于等參數(shù)量的稠密模型。一個(gè)容易誤解的地方是路由并不是簡單意義上的“語義分類器”。路由的輸入通常是當(dāng)前 token 的隱藏狀態(tài)輸出是對專家的 softmax 分布。在實(shí)際訓(xùn)練中如果只是簡單按路由概率加權(quán)會(huì)出現(xiàn)路由坍縮router collapse所有 token 都傾向去少數(shù)幾個(gè)專家其余專家閑置。因此 MoE 訓(xùn)練幾乎都會(huì)加輔助負(fù)載均衡損失讓專家利用率保持均勻。分布式場景下路由還有一種新的含義每個(gè) token 被路由到哪個(gè)專家可以看作模型如何處理這個(gè) token 的一種“行為指紋”。專家、路由分布、token 選擇三者結(jié)合在一起就構(gòu)成了路由行為的可觀察特征。DistMoE 之所以要針對 routing 單獨(dú)設(shè)計(jì)正是因?yàn)槁酚煞植技扔绊懩P托Ч謺?huì)在增量訓(xùn)練時(shí)發(fā)生漂移而且這種漂移和私有數(shù)據(jù)緊密相關(guān)。1.3 私有數(shù)據(jù)約束讓 rehearsal 不再可行Rehearsal回放/復(fù)習(xí)是持續(xù)學(xué)習(xí)中對抗災(zāi)難性遺忘的常見手段在模型學(xué)習(xí)新任務(wù)時(shí)混入一部分舊任務(wù)的樣本一起訓(xùn)練讓模型不忘記舊能力。這個(gè)方案在大模型微調(diào)中也很有效但它有一個(gè)硬前提舊樣本可以被訪問。分布式私有數(shù)據(jù)場景恰恰不滿足這個(gè)前提。原始數(shù)據(jù)不能離開本地甚至中間層的特征在某些合規(guī)要求下也不能直接外傳。于是出現(xiàn)了矛盾要穩(wěn)住路由就要復(fù)習(xí)舊任務(wù)要復(fù)習(xí)舊任務(wù)就要訪問舊樣本但私有數(shù)據(jù)約束不允許舊樣本離開本地。DistMoE 的路線是繞過“復(fù)習(xí)”這個(gè)動(dòng)作不要求模型重新看到舊樣本而是把舊任務(wù)的路由行為本身保存下來作為訓(xùn)練新任務(wù)時(shí)的約束。用一個(gè)概率向量描述舊路由模式再把這個(gè)約束嵌入新任務(wù)的損失函數(shù)。這樣就不需要回放數(shù)據(jù)也能約束路由不劇烈漂移。也就是說distributed 解決的是數(shù)據(jù)如何協(xié)作routing 解決的是 token 如何分配而 rehearsal-free 解決的則是“沒有舊數(shù)據(jù)時(shí)如何守住舊行為”。2. 實(shí)驗(yàn)環(huán)境與依賴準(zhǔn)備先搭一個(gè)可復(fù)現(xiàn)的分布式 MoE 最小工程2.1 學(xué)習(xí)環(huán)境的依賴版本基線由于原始論文沒有提供一份公開的精確依賴清單下面示例基于常見 PyTorch 生態(tài)組織。落地前要先確認(rèn)自己的 CUDA 驅(qū)動(dòng)、Python 版本和 PyTorch 版本是否匹配避免把時(shí)間花在環(huán)境問題而不是路由機(jī)制上。conda create -n distmoe python3.10 -y conda activate distmoe pip install torch2.0 transformers4.30 datasets scikit-learn pyyaml學(xué)習(xí)階段建議先用 CPU 跑通邏輯再切到 GPU。CPU 環(huán)境下把模型維度調(diào)小也可以完整驗(yàn)證路由分布和錨點(diǎn)約束的行為。生產(chǎn)環(huán)境才需要考慮多機(jī)多卡、通信壓縮、梯度聚合和審計(jì)。2.2 目錄結(jié)構(gòu)和訓(xùn)練配置文件項(xiàng)目目錄可以按這個(gè)方式組織核心是把“模型實(shí)現(xiàn)”“數(shù)據(jù)方”“訓(xùn)練邏輯”“評估邏輯”分開。distmoe-lab/ ├── configs/ │ └── tiny_moe.yaml ├── data/ │ ├── party_a/ │ │ └── train.jsonl │ └── party_b/ │ └── train.jsonl ├── moe/ │ ├── __init__.py │ ├── experts.py │ ├── router.py │ └── moe_layer.py ├── train_distributed.py └── eval_routing.py配置文件里需要同時(shí)描述模型規(guī)模、訓(xùn)練超參、數(shù)據(jù)方數(shù)量和隱私策略。下面是一個(gè)最小 YAML 示例。model: d_model: 128 d_ff: 512 num_experts: 4 top_k: 2 num_layers: 6 training: batch_size: 16 local_epochs: 1 lr: 5e-4 weight_decay: 0.01 aux_loss_alpha: 0.01 anchor_kl_alpha: 0.1 data: num_parties: 2 max_seq_len: 64 task_ratio: [0.7, 0.3] privacy: share_router_stats: true share_raw_gradients: falseshare_router_stats表示本地只允許把路由統(tǒng)計(jì)量發(fā)送出去share_raw_gradients表示不允許共享樣本級梯度。這里要特別說明共享梯度在聯(lián)邦學(xué)習(xí)里很常見但它并不安全攻擊者可以從梯度反推訓(xùn)練樣本。如果目標(biāo)是驗(yàn)證 DistMoE 的 rehearsal-free 路由機(jī)制更穩(wěn)妥的做法是只交換路由分布這樣的聚合信息。2.3 模擬數(shù)據(jù)方邊界與隱私策略在真實(shí)系統(tǒng)中每個(gè)數(shù)據(jù)方擁有自己的指令數(shù)據(jù)集。為了在實(shí)驗(yàn)里模擬最簡單的做法是把一個(gè)公開指令數(shù)據(jù)集按任務(wù)類別切分成兩份分別放到party_a和party_b。例如 A 方放“問答生成”類任務(wù)B 方放“摘要改寫”類任務(wù)。需要注意這種模擬只是“數(shù)據(jù)不跨方”并不等同于真實(shí)的隱私保護(hù)。實(shí)驗(yàn)里可以定義一條顯式規(guī)則任何離開數(shù)據(jù)方的對象只能是經(jīng)過聚合的路由分布統(tǒng)計(jì)量或者是模型參數(shù)更新。原始文本、逐樣本隱藏狀態(tài)、逐樣本梯度都不能外發(fā)。如果只是為了理解路由機(jī)制建議一開始不要使用完整的 7B 或 13B 模型先用一個(gè) 6 層小模型驗(yàn)證邏輯。小模型同樣能復(fù)現(xiàn)路由漂移現(xiàn)象訓(xùn)練速度快調(diào)試方便。等機(jī)制驗(yàn)證通過再替換成目標(biāo)規(guī)模的基座模型。3. 最小實(shí)現(xiàn)一個(gè)可訓(xùn)練的小型 MoE 指令微調(diào)循環(huán)3.1 實(shí)現(xiàn)專家網(wǎng)絡(luò)和路由模塊先用 PyTorch 實(shí)現(xiàn)一個(gè)最簡的 MoE 層包含專家網(wǎng)絡(luò)和路由模塊。這里只展示核心邏輯實(shí)際項(xiàng)目里還要加入 dropout、殘差和歸一化。import torch import torch.nn as nn import torch.nn.functional as F class Expert(nn.Module): def __init__(self, d_model, d_ff, dropout0.1): super().__init__() self.net nn.Sequential( nn.Linear(d_model, d_ff), nn.GELU(), nn.Dropout(dropout), nn.Linear(d_ff, d_model), ) def forward(self, x): return self.net(x) class Router(nn.Module): def __init__(self, d_model, num_experts): super().__init__() self.gate nn.Linear(d_model, num_experts) self.num_experts num_experts def forward(self, x): logits self.gate(x) # [batch, num_experts] return F.softmax(logits, dim-1), logits class MoELayer(nn.Module): def __init__(self, d_model, d_ff, num_experts, top_k2): super().__init__() self.router Router(d_model, num_experts) self.experts nn.ModuleList( [Expert(d_model, d_ff) for _ in range(num_experts)] ) self.top_k top_k def forward(self, x): probs, logits self.router(x) top_probs, top_idx torch.topk(probs, self.top_k, dim-1) out torch.zeros_like(x) for k in range(self.top_k): idx top_idx[:, k] weight top_probs[:, k].unsqueeze(-1) expert_outputs [] for b in range(x.size(0)): expert_outputs.append(self.experts[idx[b]](x[b])) expert_outputs torch.stack(expert_outputs) out out weight * expert_outputs return out, logits這個(gè)實(shí)現(xiàn)里每個(gè) token 會(huì)選top_k個(gè)專家并按路由概率加權(quán)求和。代碼中按 batch 維度做了循環(huán)方便理解真實(shí)訓(xùn)練里通常會(huì)改成一次計(jì)算所有專家輸出再按索引聚合或者使用torch.where、scatter等技術(shù)減少循環(huán)。要注意的是訓(xùn)練早期路由分布很不穩(wěn)定top_k索引經(jīng)常變化是正?,F(xiàn)象。3.2 加入輔助負(fù)載均衡損失路由模塊需要額外加一個(gè)負(fù)載均衡損失否則很容易出現(xiàn)路由坍縮。常用做法是計(jì)算每個(gè)專家的“被選擇比例”和“被路由概率均值”的乘積再乘上專家數(shù)量。def load_balance_loss(logits, num_experts): probs F.softmax(logits, dim-1) fraction probs.mean(dim0) load F.one_hot(probs.argmax(dim-1), num_experts).float().mean(dim0) return num_experts * (fraction * load).sum()這個(gè)損失的直觀含義是如果路由分布均勻每個(gè)專家被選擇的概率都接近1 / num_experts損失會(huì)趨近一個(gè)較小的值如果某些專家占用過高乘積就會(huì)變大梯度會(huì)推動(dòng) router 把 token 分散到其他專家。實(shí)際項(xiàng)目里這個(gè)損失通常乘以一個(gè)很小的系數(shù)alpha比如 0.01避免干擾主任務(wù)損失。3.3 用多數(shù)據(jù)方訓(xùn)練循環(huán)模擬分布式微調(diào)下面的訓(xùn)練循環(huán)是一個(gè)簡化版的多輪協(xié)作流程每一輪每個(gè)數(shù)據(jù)方先在本地?cái)?shù)據(jù)上訓(xùn)練若干 epoch然后計(jì)算路由統(tǒng)計(jì)量并參與聚合。為了演示 rehearsal-free 的效果還加入了可選的錨點(diǎn)約束。def run_local_epochs(model, loader, optimizer, cfg, anchorNone): model.train() for _ in range(cfg[local_epochs]): for batch in loader: optimizer.zero_grad() loss, aux_loss, router_logits model(batch, labelsbatch[labels]) anchor_loss torch.tensor(0.0, devicerouter_logits.device) if anchor is not None: anchor_loss kl_anchor(router_logits, anchor) total ( loss cfg[aux_loss_alpha] * aux_loss cfg[anchor_kl_alpha] * anchor_loss ) total.backward() optimizer.step() def train_round(model, party_loaders, cfg, anchorNone): optimizer torch.optim.AdamW( model.parameters(), lrcfg[lr], weight_decaycfg[weight_decay], ) for party_id, loader in party_loaders.items(): run_local_epochs(model, loader, optimizer, cfg, anchor)這里的anchor就是舊路由行為的概率向量。沒有anchor時(shí)模型只學(xué)習(xí)新任務(wù)路由分布會(huì)隨訓(xùn)練漂移有anchor時(shí)模型需要在擬合新任務(wù)和保持舊路由習(xí)慣之間取平衡。真正的分布式環(huán)境里train_round會(huì)用torch.distributed或參數(shù)服務(wù)器框架替代這個(gè)單進(jìn)程循環(huán)數(shù)據(jù)方之間只交換允許外發(fā)的統(tǒng)計(jì)量。4. 路由穩(wěn)定與 Rehearsal-free 的關(guān)鍵機(jī)制4.1 路由漂移是如何發(fā)生的路由漂移的本質(zhì)是增量訓(xùn)練改變了 router 的參數(shù)使得同樣一批舊 token 在新模型里被分配到不同的專家。指令數(shù)據(jù)進(jìn)入模型后通過反向傳播影響到所有層其中也包括 router。新任務(wù)的數(shù)據(jù)分布如果與舊任務(wù)差異較大router 會(huì)為了擬合新任務(wù)而調(diào)整決策邊界。漂移并不是絕對壞事。如果新任務(wù)確實(shí)需要新的專家組合那么適度調(diào)整路由是合理的。問題在于極端情況當(dāng)新任務(wù)的數(shù)據(jù)量很大、舊任務(wù)數(shù)據(jù)不可見時(shí)router 可能完全偏向新任務(wù)舊任務(wù)的路由模式被覆蓋最終導(dǎo)致舊任務(wù)能力明顯下降。這個(gè)過程與模型其他參數(shù)的災(zāi)難性遺忘類似但由于 router 是一個(gè)高維 softmax 分類器它的遺忘速度往往更快。4.2 用路由錨點(diǎn)分布做無復(fù)習(xí)正則Rehearsal-free 的關(guān)鍵是把舊的“行為”而不是舊的“數(shù)據(jù)”保留下來。最直接的做法是計(jì)算每個(gè) token 在舊模型上的路由概率分布聚合后形成一個(gè)錨點(diǎn)向量。這個(gè)向量可以看作模型對“歷史任務(wù)應(yīng)如何分配專家”的統(tǒng)計(jì)記憶。訓(xùn)練新任務(wù)時(shí)加入一個(gè) KL 散度約束讓當(dāng)前 router 的輸出不要偏離錨點(diǎn)太遠(yuǎn)。def kl_anchor(router_logits, anchor_probs): log_probs F.log_softmax(router_logits, dim-1) return F.kl_div( log_probs, anchor_probs.expand_as(log_probs), reductionbatchmean, )錨點(diǎn)向量是從所有舊任務(wù) token 上聚合出來的因此它不指向任何一條具體樣本。這個(gè)特點(diǎn)讓它可以用于私有數(shù)據(jù)場景只要聚合協(xié)議不泄露單個(gè) token 的隱藏狀態(tài)錨點(diǎn)向量本身對隱私的威脅遠(yuǎn)小于原始樣本。錨點(diǎn)計(jì)算方式如下在開始新任務(wù)訓(xùn)練之前用舊模型在本地驗(yàn)證集或測試集上跑一次前向記錄每個(gè) token 的 router softmax 輸出然后求均值。def compute_router_anchor(model, loader): model.eval() probs_sum 0.0 total_tokens 0 with torch.no_grad(): for batch in loader: hidden model.extract_hidden(batch) probs F.softmax(model.router(hidden), dim-1) probs_sum probs_sum probs.sum(dim0) total_tokens probs.size(0) return probs_sum / total_tokens需要強(qiáng)調(diào)的是anchor_kl_alpha這個(gè)系數(shù)要經(jīng)過實(shí)驗(yàn)調(diào)優(yōu)。系數(shù)過大模型會(huì)過度保持舊路由學(xué)習(xí)新任務(wù)的能力變差系數(shù)過小錨點(diǎn)約束形同虛設(shè)。常見做法是在驗(yàn)證集上同時(shí)觀察舊任務(wù)保留率和新任務(wù)準(zhǔn)確率選擇一個(gè)相對平衡點(diǎn)。4.3 只交換聚合統(tǒng)計(jì)量保留私有數(shù)據(jù)邊界在多個(gè)數(shù)據(jù)方協(xié)作時(shí)錨點(diǎn)向量不能由單方獨(dú)立計(jì)算后直接廣播給所有人因?yàn)閱蝹€(gè)數(shù)據(jù)方計(jì)算出的錨點(diǎn)只代表它自己的數(shù)據(jù)分布容易暴露該方的任務(wù)特征。更穩(wěn)妥的做法是每個(gè)數(shù)據(jù)方在本地計(jì)算局部路由統(tǒng)計(jì)量再通過安全聚合或聯(lián)邦平均得到全局錨點(diǎn)。全局錨點(diǎn)可以作為訓(xùn)練約束分發(fā)給所有數(shù)據(jù)方。這樣整個(gè)訓(xùn)練過程中跨數(shù)據(jù)方交換的對象只有兩類模型參數(shù)或梯度更新以及路由聚合統(tǒng)計(jì)量。哪些對象允許外發(fā)應(yīng)該在配置文件里顯式聲明而不是寫死在代碼里。對于合規(guī)要求嚴(yán)格的場景還需要考慮對統(tǒng)計(jì)量加入噪聲或做差分隱私處理因?yàn)榧词故蔷酆辖y(tǒng)計(jì)量在攻擊者擁有大量先驗(yàn)知識(shí)時(shí)也可能造成信息泄露。5. 運(yùn)行驗(yàn)證如何判斷路由是否穩(wěn)定、隱私邊界是否守住5.1 三個(gè)可量化的指標(biāo)路由 KL、專家利用率、任務(wù)保留率運(yùn)行階段至少要觀察三個(gè)指標(biāo)。第一是路由 KL 散度用于度量當(dāng)前路由分布與舊錨點(diǎn)之間的差異。KL 越大說明路由漂移越嚴(yán)重。第二是專家利用率變異系數(shù)用于判斷是否出現(xiàn)路由坍縮。變異系數(shù) 專家負(fù)載標(biāo)準(zhǔn)差 / 專家負(fù)載均值值越低說明專家負(fù)載越均衡一般低于 0.2 算比較健康。第三是舊任務(wù)保留率需要保留一小份舊任務(wù)的評估集在訓(xùn)練前后分別計(jì)算模型在舊任務(wù)上的指標(biāo)。這里要注意評估集可以放在受信任的評測方不一定回放給訓(xùn)練過程。def routing_metrics(router, loader, anchorNone, old_top1None): probs_list [] top1_list [] with torch.no_grad(): for batch in loader: hidden extract_hidden(batch) probs F.softmax(router(hidden), dim-1) probs_list.append(probs) top1_list.append(probs.argmax(dim-1)) probs torch.cat(probs_list, dim0) top1 torch.cat(top1_list, dim0) load torch.bincount(top1, minlengthrouter.num_experts).float() load_cv (load.std() / load.mean()).item() kl float(inf) if anchor is not None: kl F.kl_div( probs.log(), anchor.unsqueeze(0).expand_as(probs), reductionbatchmean, ).item() consistency None if old_top1 is not None: consistency (top1 old_top1).float().mean().item() return {routing_kl: kl, load_cv: load_cv, top1_consistency: consistency}這里的old_top1是訓(xùn)練前記錄下來的 top-1 專家索引用它計(jì)算一致性率可以更直觀地看到“舊 token 是否還去舊專家”。5.2 訓(xùn)練曲線中應(yīng)該看到的現(xiàn)象如果 rehearsal-free 機(jī)制有效訓(xùn)練曲線應(yīng)該呈現(xiàn)以下特征加入錨點(diǎn)約束后routing_kl在多個(gè)訓(xùn)練輪次中保持平穩(wěn)而不是在第一輪新任務(wù)訓(xùn)練后陡增。load_cv始終低于預(yù)設(shè)閾值說明沒有出現(xiàn)路由坍縮。舊任務(wù)評估指標(biāo)不會(huì)出現(xiàn)斷崖式下跌。新任務(wù)訓(xùn)練損失能正常下降說明錨點(diǎn)約束沒有過度壓制學(xué)習(xí)能力。如果看到routing_kl繼續(xù)上升但舊任務(wù)保留率沒有明顯惡化說明模型對新任務(wù)的適應(yīng)性更重要可以適當(dāng)調(diào)小anchor_kl_alpha。5.3 隱私保護(hù)檢查清單隱私邊界是否守住不能只靠代碼注釋要形成可檢查的清單。檢查項(xiàng)檢查內(nèi)容通過標(biāo)準(zhǔn)外發(fā)對象訓(xùn)練代碼里允許被發(fā)送出去的變量類型只有模型參數(shù)、路由聚合統(tǒng)計(jì)量原始文本日志、指標(biāo)、斷言里是否出現(xiàn)訓(xùn)練文本片段一律不打印、不落盤樣本級梯度是否共享了逐樣本梯度不允許只能共享聚合后梯度統(tǒng)計(jì)量粒度路由統(tǒng)計(jì)量是否做了跨方聚合單側(cè)統(tǒng)計(jì)量不直接廣播訪問控制參與方是否能讀取其他方的本地目錄目錄權(quán)限按數(shù)據(jù)方隔離生產(chǎn)環(huán)境里最好把隱私檢查做成自動(dòng)化腳本在訓(xùn)練任務(wù)啟動(dòng)前、訓(xùn)練完成后分別執(zhí)行。隱私保護(hù)是“沒有檢查就沒有保障”的領(lǐng)域不能只依賴開發(fā)者自覺。6. 常見問題與排查路徑6.1 路由坍縮導(dǎo)致一部分專家永遠(yuǎn)不被使用現(xiàn)象訓(xùn)練到中期部分專家對應(yīng)的負(fù)載接近 0模型有效參數(shù)量下降效果反而變差。問題現(xiàn)象常見原因檢查方式處理建議部分專家負(fù)載為 0輔助負(fù)載均衡損失系數(shù)過小或缺失打印load_cv和各專家被選次數(shù)調(diào)大aux_loss_alpha或改用基于 top-1 選擇的重采樣損失負(fù)載周期性抖動(dòng)token 數(shù)量太少統(tǒng)計(jì)波動(dòng)大查看每個(gè) batch 的專家負(fù)載直方圖增大 batch size 或梯度累積步數(shù)路由分布向某專家偏移該專家初始化或數(shù)據(jù)分布不均對比各專家輸出范數(shù)檢查專家初始化必要時(shí)增加專家 dropout最直接的修復(fù)方式是先把a(bǔ)ux_loss_alpha從 0.01 逐步調(diào)大觀察load_cv是否回落。注意不要一次調(diào)得太大否則主任務(wù)損失會(huì)被淹沒。6.2 路由漂移指標(biāo)不降反升現(xiàn)象加了錨點(diǎn)約束后routing_kl仍然很高而且舊任務(wù)指標(biāo)下降。問題現(xiàn)象常見原因檢查方式處理建議KL 不減反增anchor_kl_alpha過小打印 anchor loss 數(shù)值確認(rèn)它被計(jì)入總損失調(diào)大anchor_kl_alphaKL 正常但舊任務(wù)指標(biāo)下降錨點(diǎn)只約束了 router其他層仍然遺忘對比舊任務(wù)在新舊模型上的輸出 logits增加輸出分布蒸餾約束或凍結(jié)部分底層參數(shù)錨點(diǎn)本身計(jì)算錯(cuò)誤計(jì)算錨點(diǎn)時(shí)用了訓(xùn)練模式而不是 eval 模式檢查compute_router_anchor是否在torch.no_grad()下執(zhí)行統(tǒng)一用 eval 模式、固定 seed這里要區(qū)分路由穩(wěn)定和模型整體穩(wěn)定。路由穩(wěn)定只是必要條件如果下游層仍然遺忘舊任務(wù)光約束 router 效果有限。所以指標(biāo)設(shè)計(jì)時(shí)舊任務(wù)保留率比routing_kl更重要。6.3 多數(shù)據(jù)方訓(xùn)練中同步通信耗時(shí)過高現(xiàn)象模型訓(xùn)練本身不慢但每輪同步參數(shù)和統(tǒng)計(jì)量占用大量時(shí)間。問題現(xiàn)象常見原因檢查方式處理建議單輪耗時(shí)隨參與方增加快速上升同步次數(shù)過多每次傳輸全量參數(shù)打印同步耗時(shí)和后端類型降低同步頻率使用累積多步后一次同步通信量過大傳輸了中間層隱藏狀態(tài)或逐樣本統(tǒng)計(jì)量檢查外發(fā)對象類型和維度只傳輸 router 概率分布均值等低維統(tǒng)計(jì)量小 batch 下通信占比高計(jì)算時(shí)間短通信成為瓶頸觀察 GPU 利用率增大本地 batch 或梯度累積步數(shù)學(xué)習(xí)環(huán)境不需要過度優(yōu)化通信先把功能跑通。生產(chǎn)環(huán)境則要評估是網(wǎng)絡(luò)帶寬受限還是同步頻率過高再?zèng)Q定采用異步更新還是周期性同步。6.4 隱私統(tǒng)計(jì)量被誤當(dāng)作普通訓(xùn)練數(shù)據(jù)使用現(xiàn)象某數(shù)據(jù)方把錨點(diǎn)向量直接當(dāng)作監(jiān)督信號(hào)參與所有 loss 計(jì)算導(dǎo)致路由被錨點(diǎn)完全鎖死。問題現(xiàn)象常見原因檢查方式處理建議模型無法學(xué)習(xí)新任務(wù)錨點(diǎn)約束權(quán)重設(shè)置過高查看anchor_loss與主損失的數(shù)量級將anchor_kl_alpha降到 0.01 以下代碼里出現(xiàn)原始文本外發(fā)邏輯復(fù)用了集中式訓(xùn)練的 dataloader檢查數(shù)據(jù)加載器是否跨方訪問每個(gè)數(shù)據(jù)方獨(dú)立加載本地?cái)?shù)據(jù)禁止跨方目錄訪問審計(jì)日志缺失沒有記錄外發(fā)對象檢查日志中是否包含發(fā)送函數(shù)調(diào)用點(diǎn)為外發(fā)函數(shù)增加審計(jì) hook這類問題很難通過模型指標(biāo)發(fā)現(xiàn)需要靠代碼審查和日志審計(jì)。建議在代碼中把“外發(fā)對象”封裝成獨(dú)立的函數(shù)或接口而不是到處直接調(diào)send這樣審查時(shí)能明確知道哪些數(shù)據(jù)會(huì)離開本地。7. 從實(shí)驗(yàn)到生產(chǎn)的實(shí)踐建議7.1 發(fā)布前檢查清單在把實(shí)驗(yàn)代碼推廣到更大規(guī)模之前先用下面這份清單做一次體檢。配置文件里是否顯式聲明了允許外發(fā)的對象類型。路由錨點(diǎn)是否來自聚合后的統(tǒng)計(jì)量而不是單側(cè)數(shù)據(jù)。是否同時(shí)記錄了routing_kl、load_cv和舊任務(wù)保留率三個(gè)指標(biāo)。是否保留了一份與訓(xùn)練數(shù)據(jù)隔離的舊任務(wù)評測集。有沒有對不同數(shù)據(jù)方的目錄做權(quán)限隔離。訓(xùn)練日志里是否可能打印訓(xùn)練文本片段。是否備份了初始模型參數(shù)和每一輪的錨點(diǎn)向量方便回溯。是否準(zhǔn)備好回滾方案例如新任務(wù)效果異常時(shí)重新加載舊模型。這些項(xiàng)目看起來瑣碎但每一項(xiàng)都可能在生產(chǎn)環(huán)境里變成事故源。7.2 學(xué)習(xí)環(huán)境、實(shí)驗(yàn)環(huán)境與生產(chǎn)環(huán)境的差異維度學(xué)習(xí)環(huán)境實(shí)驗(yàn)環(huán)境生產(chǎn)環(huán)境模型規(guī)模6 層小模型CPU 可跑單卡到單機(jī)多卡多機(jī)多卡可能需要模型并行數(shù)據(jù)規(guī)模幾百條模擬數(shù)據(jù)萬級公開指令數(shù)據(jù)多數(shù)據(jù)方真實(shí)業(yè)務(wù)數(shù)據(jù)通信單進(jìn)程即可單機(jī)多卡 DDP安全聚合、異步通信、斷點(diǎn)續(xù)訓(xùn)隱私不涉及真實(shí)隱私模擬隱私邊界合規(guī)審核、審計(jì)日志、差分隱私指標(biāo)看訓(xùn)練收斂看 KL、負(fù)載、保留率看線上效果、延遲、資源消耗學(xué)習(xí)環(huán)境追求快速理解機(jī)制所以模型越小越好實(shí)驗(yàn)環(huán)境要驗(yàn)證機(jī)制有效性所以要保留完整的指標(biāo)和可復(fù)現(xiàn)腳本生產(chǎn)環(huán)境則需要考慮隱私合規(guī)、監(jiān)控、告警和回滾已經(jīng)不是單純改進(jìn)算法能解決的問題。7.3 可擴(kuò)展方向如果一個(gè)標(biāo)準(zhǔn)的錨點(diǎn)正則已經(jīng)能緩解路由漂移下一步可以沿著三個(gè)方向擴(kuò)展。第一把錨點(diǎn)向量升級成更細(xì)粒度的錨點(diǎn)結(jié)構(gòu)。例如按任務(wù)類型分別保存路由分布在訓(xùn)練新任務(wù)時(shí)只約束與舊任務(wù)重疊的 token 類別而不是強(qiáng)制所有 token 都保持舊分布。第二引入動(dòng)態(tài)專家擴(kuò)展。當(dāng)新任務(wù)確實(shí)需要新能力時(shí)與其強(qiáng)制復(fù)用舊專家的路由模式不如動(dòng)態(tài)增加專家并讓新任務(wù)主要分配新專家從而減少對舊路由的擾動(dòng)。這一點(diǎn)與 MoE 的容量設(shè)計(jì)和稀疏激活天然契合。第三結(jié)合差分隱私與安全聚合實(shí)現(xiàn)強(qiáng)隱私保證。路由統(tǒng)計(jì)量雖然在實(shí)踐上比原始樣本安全但并不是絕對無泄露。如果參與方數(shù)量少、任務(wù)分布可辨識(shí)聚合統(tǒng)計(jì)量也可能泄露信息。生產(chǎn)環(huán)境里要做好隱私風(fēng)險(xiǎn)評估再?zèng)Q定是否需要加噪聲。DistMoE 的核心價(jià)值不在于某一個(gè)符號(hào)或公式而在于它把“分布式數(shù)據(jù)協(xié)作”“MoE 路由穩(wěn)定性”“持續(xù)學(xué)習(xí)的災(zāi)難性遺忘”三個(gè)問題納入了同一個(gè)設(shè)計(jì)框架。對開發(fā)者的啟發(fā)是當(dāng)數(shù)據(jù)不能移動(dòng)時(shí)與其保存舊數(shù)據(jù)不如保存舊行為當(dāng)路由可能漂移時(shí)與其強(qiáng)制凍結(jié)不如用分布約束讓新任務(wù)和舊習(xí)慣共存。順著這條思路可以先在小模型上復(fù)現(xiàn)路由漂移現(xiàn)象再用錨點(diǎn)約束驗(yàn)證效果最后才考慮擴(kuò)到真實(shí)分布式系統(tǒng)。這個(gè)從小到大、從行為到機(jī)制的驗(yàn)證路徑比直接在大模型上跑實(shí)驗(yàn)要有效得多。