
在大型語言模型的分布式指令微調場景里DistMoE 把三個原本分開的問題綁定到了同一個系統中多個數據方不能共享私有數據卻要協作微調一個 Mixture-of-ExpertsMoE模型MoE 內部的路由模塊要決定每個 token 進入哪些專家而這些路由決策在引入新任務后不能出現明顯漂移并且不能依賴回放舊數據來維持穩定。單獨解決其中任何一個問題都有成熟方案但組合在一起之后難點就變成了路由穩定性通常靠 rehearsal 來維持而 rehearsal 又依賴舊數據與私有數據約束直接沖突。DistMoE 這個研究方向的核心就是在不接觸私有數據的前提下讓路由模塊在分布式指令微調過程中保持穩定。下面會從 MoE 路由的基本原理講起搭建一個最小可復現的分布式指令微調實驗框架然后說明如何用路由錨點分布實現 rehearsal-free 的穩定訓練最后給出驗證指標、排查路徑和生產落地建議。1. 先理解 DistMoE 的三層技術背景分布式指令微調、MoE 路由、Rehearsal-free1.1 指令微調為什么需要分布式協作指令微調Instruction Tuning指的是用“指令 期望回答”這樣的監督數據對基座大模型做有監督微調讓模型學會按照用戶指令輸出有用回答。它與預訓練不同預訓練階段模型看到的是大規模無標注文本學習的是語言統計規律指令微調階段模型看到的是少量、高質量、任務導向的數據學習的是“知道何時該執行什么動作”。實際企業場景中指令數據往往分散在不同部門、不同組織甚至不同地區。客服部門有客服對話數據風控部門有風控問答數據業務部門還有內部規章制度數據。把數據集中到一個機房訓練最簡單但隱私和合規成本很高。分布式指令微調要解決的問題就是“數據不動模型動”原始數據留在各自的本地數據方只有模型參數更新或必要的統計信息參與協作。這種方式也被稱為數據隔離下的協作訓練DistMoE 討論的分布式正是這種多個數據方共同參與訓練的拓撲而不是單純指“多機多卡做數據并行”。1.2 MoE 路由門控網絡如何決定 token 去哪個專家Mixture-of-ExpertsMoE的核心設計是把傳統 Transformer 中的前饋網絡FFN替換成一組專家網絡并由一個路由模塊Router/Gating決定每個 token 激活哪些專家。設輸入為x路由模塊輸出一個在E個專家上的概率分布然后取 top-k 個專家做加權求和。這樣做的好處是模型參數量可以很大但每個 token 只計算其中一小部分專家推理和訓練的計算成本都低于等參數量的稠密模型。一個容易誤解的地方是路由并不是簡單意義上的“語義分類器”。路由的輸入通常是當前 token 的隱藏狀態輸出是對專家的 softmax 分布。在實際訓練中如果只是簡單按路由概率加權會出現路由坍縮router collapse所有 token 都傾向去少數幾個專家其余專家閑置。因此 MoE 訓練幾乎都會加輔助負載均衡損失讓專家利用率保持均勻。分布式場景下路由還有一種新的含義每個 token 被路由到哪個專家可以看作模型如何處理這個 token 的一種“行為指紋”。專家、路由分布、token 選擇三者結合在一起就構成了路由行為的可觀察特征。DistMoE 之所以要針對 routing 單獨設計正是因為路由分布既影響模型效果又會在增量訓練時發生漂移而且這種漂移和私有數據緊密相關。1.3 私有數據約束讓 rehearsal 不再可行Rehearsal回放/復習是持續學習中對抗災難性遺忘的常見手段在模型學習新任務時混入一部分舊任務的樣本一起訓練讓模型不忘記舊能力。這個方案在大模型微調中也很有效但它有一個硬前提舊樣本可以被訪問。分布式私有數據場景恰恰不滿足這個前提。原始數據不能離開本地甚至中間層的特征在某些合規要求下也不能直接外傳。于是出現了矛盾要穩住路由就要復習舊任務要復習舊任務就要訪問舊樣本但私有數據約束不允許舊樣本離開本地。DistMoE 的路線是繞過“復習”這個動作不要求模型重新看到舊樣本而是把舊任務的路由行為本身保存下來作為訓練新任務時的約束。用一個概率向量描述舊路由模式再把這個約束嵌入新任務的損失函數。這樣就不需要回放數據也能約束路由不劇烈漂移。也就是說distributed 解決的是數據如何協作routing 解決的是 token 如何分配而 rehearsal-free 解決的則是“沒有舊數據時如何守住舊行為”。2. 實驗環境與依賴準備先搭一個可復現的分布式 MoE 最小工程2.1 學習環境的依賴版本基線由于原始論文沒有提供一份公開的精確依賴清單下面示例基于常見 PyTorch 生態組織。落地前要先確認自己的 CUDA 驅動、Python 版本和 PyTorch 版本是否匹配避免把時間花在環境問題而不是路由機制上。conda create -n distmoe python3.10 -y conda activate distmoe pip install torch2.0 transformers4.30 datasets scikit-learn pyyaml學習階段建議先用 CPU 跑通邏輯再切到 GPU。CPU 環境下把模型維度調小也可以完整驗證路由分布和錨點約束的行為。生產環境才需要考慮多機多卡、通信壓縮、梯度聚合和審計。2.2 目錄結構和訓練配置文件項目目錄可以按這個方式組織核心是把“模型實現”“數據方”“訓練邏輯”“評估邏輯”分開。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配置文件里需要同時描述模型規模、訓練超參、數據方數量和隱私策略。下面是一個最小 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表示本地只允許把路由統計量發送出去share_raw_gradients表示不允許共享樣本級梯度。這里要特別說明共享梯度在聯邦學習里很常見但它并不安全攻擊者可以從梯度反推訓練樣本。如果目標是驗證 DistMoE 的 rehearsal-free 路由機制更穩妥的做法是只交換路由分布這樣的聚合信息。2.3 模擬數據方邊界與隱私策略在真實系統中每個數據方擁有自己的指令數據集。為了在實驗里模擬最簡單的做法是把一個公開指令數據集按任務類別切分成兩份分別放到party_a和party_b。例如 A 方放“問答生成”類任務B 方放“摘要改寫”類任務。需要注意這種模擬只是“數據不跨方”并不等同于真實的隱私保護。實驗里可以定義一條顯式規則任何離開數據方的對象只能是經過聚合的路由分布統計量或者是模型參數更新。原始文本、逐樣本隱藏狀態、逐樣本梯度都不能外發。如果只是為了理解路由機制建議一開始不要使用完整的 7B 或 13B 模型先用一個 6 層小模型驗證邏輯。小模型同樣能復現路由漂移現象訓練速度快調試方便。等機制驗證通過再替換成目標規模的基座模型。3. 最小實現一個可訓練的小型 MoE 指令微調循環3.1 實現專家網絡和路由模塊先用 PyTorch 實現一個最簡的 MoE 層包含專家網絡和路由模塊。這里只展示核心邏輯實際項目里還要加入 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這個實現里每個 token 會選top_k個專家并按路由概率加權求和。代碼中按 batch 維度做了循環方便理解真實訓練里通常會改成一次計算所有專家輸出再按索引聚合或者使用torch.where、scatter等技術減少循環。要注意的是訓練早期路由分布很不穩定top_k索引經常變化是正常現象。3.2 加入輔助負載均衡損失路由模塊需要額外加一個負載均衡損失否則很容易出現路由坍縮。常用做法是計算每個專家的“被選擇比例”和“被路由概率均值”的乘積再乘上專家數量。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()這個損失的直觀含義是如果路由分布均勻每個專家被選擇的概率都接近1 / num_experts損失會趨近一個較小的值如果某些專家占用過高乘積就會變大梯度會推動 router 把 token 分散到其他專家。實際項目里這個損失通常乘以一個很小的系數alpha比如 0.01避免干擾主任務損失。3.3 用多數據方訓練循環模擬分布式微調下面的訓練循環是一個簡化版的多輪協作流程每一輪每個數據方先在本地數據上訓練若干 epoch然后計算路由統計量并參與聚合。為了演示 rehearsal-free 的效果還加入了可選的錨點約束。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時模型只學習新任務路由分布會隨訓練漂移有anchor時模型需要在擬合新任務和保持舊路由習慣之間取平衡。真正的分布式環境里train_round會用torch.distributed或參數服務器框架替代這個單進程循環數據方之間只交換允許外發的統計量。4. 路由穩定與 Rehearsal-free 的關鍵機制4.1 路由漂移是如何發生的路由漂移的本質是增量訓練改變了 router 的參數使得同樣一批舊 token 在新模型里被分配到不同的專家。指令數據進入模型后通過反向傳播影響到所有層其中也包括 router。新任務的數據分布如果與舊任務差異較大router 會為了擬合新任務而調整決策邊界。漂移并不是絕對壞事。如果新任務確實需要新的專家組合那么適度調整路由是合理的。問題在于極端情況當新任務的數據量很大、舊任務數據不可見時router 可能完全偏向新任務舊任務的路由模式被覆蓋最終導致舊任務能力明顯下降。這個過程與模型其他參數的災難性遺忘類似但由于 router 是一個高維 softmax 分類器它的遺忘速度往往更快。4.2 用路由錨點分布做無復習正則Rehearsal-free 的關鍵是把舊的“行為”而不是舊的“數據”保留下來。最直接的做法是計算每個 token 在舊模型上的路由概率分布聚合后形成一個錨點向量。這個向量可以看作模型對“歷史任務應如何分配專家”的統計記憶。訓練新任務時加入一個 KL 散度約束讓當前 router 的輸出不要偏離錨點太遠。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, )錨點向量是從所有舊任務 token 上聚合出來的因此它不指向任何一條具體樣本。這個特點讓它可以用于私有數據場景只要聚合協議不泄露單個 token 的隱藏狀態錨點向量本身對隱私的威脅遠小于原始樣本。錨點計算方式如下在開始新任務訓練之前用舊模型在本地驗證集或測試集上跑一次前向記錄每個 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需要強調的是anchor_kl_alpha這個系數要經過實驗調優。系數過大模型會過度保持舊路由學習新任務的能力變差系數過小錨點約束形同虛設。常見做法是在驗證集上同時觀察舊任務保留率和新任務準確率選擇一個相對平衡點。4.3 只交換聚合統計量保留私有數據邊界在多個數據方協作時錨點向量不能由單方獨立計算后直接廣播給所有人因為單個數據方計算出的錨點只代表它自己的數據分布容易暴露該方的任務特征。更穩妥的做法是每個數據方在本地計算局部路由統計量再通過安全聚合或聯邦平均得到全局錨點。全局錨點可以作為訓練約束分發給所有數據方。這樣整個訓練過程中跨數據方交換的對象只有兩類模型參數或梯度更新以及路由聚合統計量。哪些對象允許外發應該在配置文件里顯式聲明而不是寫死在代碼里。對于合規要求嚴格的場景還需要考慮對統計量加入噪聲或做差分隱私處理因為即使是聚合統計量在攻擊者擁有大量先驗知識時也可能造成信息泄露。5. 運行驗證如何判斷路由是否穩定、隱私邊界是否守住5.1 三個可量化的指標路由 KL、專家利用率、任務保留率運行階段至少要觀察三個指標。第一是路由 KL 散度用于度量當前路由分布與舊錨點之間的差異。KL 越大說明路由漂移越嚴重。第二是專家利用率變異系數用于判斷是否出現路由坍縮。變異系數 專家負載標準差 / 專家負載均值值越低說明專家負載越均衡一般低于 0.2 算比較健康。第三是舊任務保留率需要保留一小份舊任務的評估集在訓練前后分別計算模型在舊任務上的指標。這里要注意評估集可以放在受信任的評測方不一定回放給訓練過程。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是訓練前記錄下來的 top-1 專家索引用它計算一致性率可以更直觀地看到“舊 token 是否還去舊專家”。5.2 訓練曲線中應該看到的現象如果 rehearsal-free 機制有效訓練曲線應該呈現以下特征加入錨點約束后routing_kl在多個訓練輪次中保持平穩而不是在第一輪新任務訓練后陡增。load_cv始終低于預設閾值說明沒有出現路由坍縮。舊任務評估指標不會出現斷崖式下跌。新任務訓練損失能正常下降說明錨點約束沒有過度壓制學習能力。如果看到routing_kl繼續上升但舊任務保留率沒有明顯惡化說明模型對新任務的適應性更重要可以適當調小anchor_kl_alpha。5.3 隱私保護檢查清單隱私邊界是否守住不能只靠代碼注釋要形成可檢查的清單。檢查項檢查內容通過標準外發對象訓練代碼里允許被發送出去的變量類型只有模型參數、路由聚合統計量原始文本日志、指標、斷言里是否出現訓練文本片段一律不打印、不落盤樣本級梯度是否共享了逐樣本梯度不允許只能共享聚合后梯度統計量粒度路由統計量是否做了跨方聚合單側統計量不直接廣播訪問控制參與方是否能讀取其他方的本地目錄目錄權限按數據方隔離生產環境里最好把隱私檢查做成自動化腳本在訓練任務啟動前、訓練完成后分別執行。隱私保護是“沒有檢查就沒有保障”的領域不能只依賴開發者自覺。6. 常見問題與排查路徑6.1 路由坍縮導致一部分專家永遠不被使用現象訓練到中期部分專家對應的負載接近 0模型有效參數量下降效果反而變差。問題現象常見原因檢查方式處理建議部分專家負載為 0輔助負載均衡損失系數過小或缺失打印load_cv和各專家被選次數調大aux_loss_alpha或改用基于 top-1 選擇的重采樣損失負載周期性抖動token 數量太少統計波動大查看每個 batch 的專家負載直方圖增大 batch size 或梯度累積步數路由分布向某專家偏移該專家初始化或數據分布不均對比各專家輸出范數檢查專家初始化必要時增加專家 dropout最直接的修復方式是先把aux_loss_alpha從 0.01 逐步調大觀察load_cv是否回落。注意不要一次調得太大否則主任務損失會被淹沒。6.2 路由漂移指標不降反升現象加了錨點約束后routing_kl仍然很高而且舊任務指標下降。問題現象常見原因檢查方式處理建議KL 不減反增anchor_kl_alpha過小打印 anchor loss 數值確認它被計入總損失調大anchor_kl_alphaKL 正常但舊任務指標下降錨點只約束了 router其他層仍然遺忘對比舊任務在新舊模型上的輸出 logits增加輸出分布蒸餾約束或凍結部分底層參數錨點本身計算錯誤計算錨點時用了訓練模式而不是 eval 模式檢查compute_router_anchor是否在torch.no_grad()下執行統一用 eval 模式、固定 seed這里要區分路由穩定和模型整體穩定。路由穩定只是必要條件如果下游層仍然遺忘舊任務光約束 router 效果有限。所以指標設計時舊任務保留率比routing_kl更重要。6.3 多數據方訓練中同步通信耗時過高現象模型訓練本身不慢但每輪同步參數和統計量占用大量時間。問題現象常見原因檢查方式處理建議單輪耗時隨參與方增加快速上升同步次數過多每次傳輸全量參數打印同步耗時和后端類型降低同步頻率使用累積多步后一次同步通信量過大傳輸了中間層隱藏狀態或逐樣本統計量檢查外發對象類型和維度只傳輸 router 概率分布均值等低維統計量小 batch 下通信占比高計算時間短通信成為瓶頸觀察 GPU 利用率增大本地 batch 或梯度累積步數學習環境不需要過度優化通信先把功能跑通。生產環境則要評估是網絡帶寬受限還是同步頻率過高再決定采用異步更新還是周期性同步。6.4 隱私統計量被誤當作普通訓練數據使用現象某數據方把錨點向量直接當作監督信號參與所有 loss 計算導致路由被錨點完全鎖死。問題現象常見原因檢查方式處理建議模型無法學習新任務錨點約束權重設置過高查看anchor_loss與主損失的數量級將anchor_kl_alpha降到 0.01 以下代碼里出現原始文本外發邏輯復用了集中式訓練的 dataloader檢查數據加載器是否跨方訪問每個數據方獨立加載本地數據禁止跨方目錄訪問審計日志缺失沒有記錄外發對象檢查日志中是否包含發送函數調用點為外發函數增加審計 hook這類問題很難通過模型指標發現需要靠代碼審查和日志審計。建議在代碼中把“外發對象”封裝成獨立的函數或接口而不是到處直接調send這樣審查時能明確知道哪些數據會離開本地。7. 從實驗到生產的實踐建議7.1 發布前檢查清單在把實驗代碼推廣到更大規模之前先用下面這份清單做一次體檢。配置文件里是否顯式聲明了允許外發的對象類型。路由錨點是否來自聚合后的統計量而不是單側數據。是否同時記錄了routing_kl、load_cv和舊任務保留率三個指標。是否保留了一份與訓練數據隔離的舊任務評測集。有沒有對不同數據方的目錄做權限隔離。訓練日志里是否可能打印訓練文本片段。是否備份了初始模型參數和每一輪的錨點向量方便回溯。是否準備好回滾方案例如新任務效果異常時重新加載舊模型。這些項目看起來瑣碎但每一項都可能在生產環境里變成事故源。7.2 學習環境、實驗環境與生產環境的差異維度學習環境實驗環境生產環境模型規模6 層小模型CPU 可跑單卡到單機多卡多機多卡可能需要模型并行數據規模幾百條模擬數據萬級公開指令數據多數據方真實業務數據通信單進程即可單機多卡 DDP安全聚合、異步通信、斷點續訓隱私不涉及真實隱私模擬隱私邊界合規審核、審計日志、差分隱私指標看訓練收斂看 KL、負載、保留率看線上效果、延遲、資源消耗學習環境追求快速理解機制所以模型越小越好實驗環境要驗證機制有效性所以要保留完整的指標和可復現腳本生產環境則需要考慮隱私合規、監控、告警和回滾已經不是單純改進算法能解決的問題。7.3 可擴展方向如果一個標準的錨點正則已經能緩解路由漂移下一步可以沿著三個方向擴展。第一把錨點向量升級成更細粒度的錨點結構。例如按任務類型分別保存路由分布在訓練新任務時只約束與舊任務重疊的 token 類別而不是強制所有 token 都保持舊分布。第二引入動態專家擴展。當新任務確實需要新能力時與其強制復用舊專家的路由模式不如動態增加專家并讓新任務主要分配新專家從而減少對舊路由的擾動。這一點與 MoE 的容量設計和稀疏激活天然契合。第三結合差分隱私與安全聚合實現強隱私保證。路由統計量雖然在實踐上比原始樣本安全但并不是絕對無泄露。如果參與方數量少、任務分布可辨識聚合統計量也可能泄露信息。生產環境里要做好隱私風險評估再決定是否需要加噪聲。DistMoE 的核心價值不在于某一個符號或公式而在于它把“分布式數據協作”“MoE 路由穩定性”“持續學習的災難性遺忘”三個問題納入了同一個設計框架。對開發者的啟發是當數據不能移動時與其保存舊數據不如保存舊行為當路由可能漂移時與其強制凍結不如用分布約束讓新任務和舊習慣共存。順著這條思路可以先在小模型上復現路由漂移現象再用錨點約束驗證效果最后才考慮擴到真實分布式系統。這個從小到大、從行為到機制的驗證路徑比直接在大模型上跑實驗要有效得多。