Split Learning
スプリットラーニングが複数のデバイス間でニューラルネットワークをどのように分割して協調型AIをサポートするのかを学びながら、プライバシーリスク、トレーニングワークフロー、アプリケーション、および設計上の選択肢について解説します。
スプリットラーニングは、ニューラルネットワークを2つ以上のコンピューティングロケーションに分割する分散型機械学習のアプローチです。クライアントはモデルの初期レイヤーを通じてプライベートな入力データを処理し、中間活性化のみをサーバーに送信し、ローカルレイヤーのトレーニングを継続するために必要な勾配を受信します。これにより、組織やデバイスは、生のトレーニングデータを直接転送することなく共同作業を行うことができます。
このアプローチは、機械学習がプライバシー、所有権、帯域幅、またはハードウェアの境界を越えて動作する必要がある場合に特に重要です。たとえば、病院は医療画像をローカルに保持し、より強力なサーバーがモデルの計算負荷の高い部分を実行することができます。ただし、中間表現が機密情報を依然として漏洩する可能性があるため、生データをローカルに維持しても、自動的にデータプライバシーが保証されるわけではありません。
スプリットラーニングの仕組み#
ニューラルネットワークは、選択されたカットレイヤーで分割されます。カット前のレイヤーはクライアントで実行され、カット後のレイヤーはサーバーで実行されます。クライアント側のネットワークの出力は、アクティベーション、中間表現、またはスマッシュデータと呼ばれることがよくあります。
トレーニングのステップは次のシーケンスに従います。
- クライアントは、生入力からカットレイヤーまでフォワードパスを実行します。
- 結果のアクティベーションをサーバーに送信します。
- サーバーはフォワードパスを完了し、損失を計算します。
- 誤差逆伝播法の実行中、サーバーはカットレイヤーのアクティベーションの勾配を計算して返します。
- クライアントはその勾配を使用してローカルレイヤーを更新します。
このプロセスは、PyTorchの自動微分などのシステムによって実装されているのと同じ連鎖律に依存しています。違いは、トレーニング中にアクティベーションと勾配がネットワークの境界を越える点です。
次の単一プロセス例は、その境界をシミュレートしています。
import torch
from torch import nn
client_model = nn.Sequential(nn.Linear(8, 16), nn.ReLU())
server_model = nn.Sequential(nn.Linear(16, 2))
client_optimizer = torch.optim.SGD(client_model.parameters(), lr=0.01)
server_optimizer = torch.optim.SGD(server_model.parameters(), lr=0.01)
inputs = torch.randn(4, 8)
targets = torch.tensor([0, 1, 0, 1])
client_optimizer.zero_grad()
server_optimizer.zero_grad()
client_activations = client_model(inputs)
sent_activations = client_activations.detach().requires_grad_()
predictions = server_model(sent_activations)
loss = nn.CrossEntropyLoss()(predictions, targets)
loss.backward()
client_activations.backward(sent_activations.grad)
server_optimizer.step()
client_optimizer.step()
print(loss.item())client_activationsをデタッチすることは、別のシステムへの送信を表します。返されたアクティベーション勾配は、最適化のために2つの半分を再接続します。本番環境の実装では、ネットワーク、認証、暗号化、障害処理、およびプライバシー制御を追加する必要があります。
スプリットラーニングと関連アプローチの比較#
スプリットラーニングは分散学習のより広い分野に属しますが、計算のパーティション分割の方法が異なります。
- フェデレーテッドラーニング: 通常、各参加者は完全なローカルモデルをトレーニングし、集約のためにモデルの更新を送信します。スプリットラーニングでは、各参加者にモデルの一部のみを与え、中間アクティベーションと勾配を交換します。
- パイプライン並列処理: 両方のアプローチとも、異なるデバイスに異なるレイヤーを配置します。パイプライン並列処理は主に信頼できる環境でのスケールやハードウェアの利用率を向上させますが、スプリットラーニングは通常、データ所有者とコンピュートプロバイダーを分離します。
- データ並列トレーニング: PyTorch DistributedDataParallelやTensorFlow分散トレーニングなどのフレームワークは、モデルを複製し、更新を同期します。これらは通常、初期レイヤーを元のデータのそばだけに排他的に保持することはありません。
スプリットラーニングは、組織が一致するレコードに対して異なる特徴量を保持する、垂直方向にパーティション分割されたデータもサポートできます。SecretFlowのスプリットラーニングワークフローはこの配置を示しています。
実世界での利用例#
-
協調医療画像処理: 病院は、X線やスキャンを自身のインフラストラクチャ内に保持したまま、共有のコンピュータビジョンシステムをトレーニングできます。各病院は最初のレイヤーをローカルで実行し、中央サーバーが中間特徴量からトレーニングを完了します。MITのスプリットラーニング概要では、放射線科センターを使用してこのアーキテクチャを説明しています。
-
リソース制約のある産業用カメラ: 工場のカメラやゲートウェイは、コンパクトな特徴抽出器をローカルで実行し、サーバーがオブジェクト検出のための残りのレイヤーをトレーニングすることができます。これにより、生のビデオ転送とクライアントの計算を削減でき、複数施設で稼働するエッジAIシステムに関連するアプローチとなります。
利点、リスク、および設計上の選択#
カットレイヤーによって、クライアントのワークロード、サーバーのワークロード、通信量、および情報の公開のバランスが決まります。早い段階でのカットはクライアントの計算量を削減しますが、入力に似た大きめのアクティベーションを生成する可能性があります。後の段階でのカットはより抽象的な特徴を作成できますが、より強力なクライアントハードウェアが必要になります。
中間アクティベーションと勾配は、再構築、推論、または操作に対して脆弱なままになる可能性があります。したがって、チームはスプリットラーニングを完全なプライバシーソリューションとして扱うのではなく、アクセス制御、暗号化された転送、アクティベーションの保護、監査ログ、および参加者の信頼を評価する必要があります。NISTプライバシーフレームワークとNIST AIリスク管理フレームワークは、これらのリスクを評価するためのより広いプロセスを提供します。
すべてのトレーニングステップで双方向通信が必要になる場合があるため、帯域幅とレイテンシも重要です。遅い、または信頼性の低いクライアントはシステム全体を遅延させる可能性があり、不一貫なデータ分布は収束に影響を与える可能性があります。
Ultralytics YOLOは、ターンキーのスプリットラーニングオーケストレーションを提供しません。これを実装するには、YOLOアーキテクチャを慎重にパーティション分割し、リモートのフォワードおよびバックワードパスを調整し、文書化されたカスタムトレーナーワークフローを拡張する必要がある場合があります。データを持続的に所有するハードウェアに保持することのみを必要とするプロジェクトの場合、Ultralyticsプラットフォームモデルトレーニングはストリーミングメトリクスを使用したローカル訓練をサポートしますが、モデル自体が参加者間で分割されていないため、ローカル訓練はスプリットラーニングではありません。






