Federated Learning実装
🔐 Federated Learning実装:ジム戦で学ぶ分散学習
Section titled “🔐 Federated Learning実装:ジム戦で学ぶ分散学習”📝 この章について
Section titled “📝 この章について”**Federated Learning(連合学習)**を「ポケモンのジムリーダーが各地で独自に修行し、その成果だけを共有する」イメージで解説します。データを中央に集めずに、分散したデバイスでモデルを訓練できます。
ポケモンでいえば
Section titled “ポケモンでいえば”- 中央集権型学習 = 全てのポケモンを1箇所に集めて集中訓練
- Federated Learning = 各ジムで独自に訓練し、戦術だけを共有
- プライバシー保護 = ポケモンの個体値は秘密(技構成だけ共有)
- モデル集約 = 各ジムの知見を統合して最強の戦術を編み出す
🎯 この章で学べること
Section titled “🎯 この章で学べること”| 項目 | 内容 |
|---|---|
| FLの基礎 | 中央集権学習との違い、プライバシー保護の仕組み |
| 実装方法 | TensorFlow Federated、PySyft等のフレームワーク活用 |
| ユースケース | 医療、金融、スマホキーボード予測等 |
| 課題と対策 | 通信コスト、デバイス異質性、悪意ある参加者への対処 |
🎯 Federated Learningとは
Section titled “🎯 Federated Learningとは”教科書的な説明
Section titled “教科書的な説明”Federated Learningは、データを中央サーバーに集めずに、各デバイス上でモデルを訓練し、更新されたモデルパラメータ(重み)だけを中央サーバーに送信して集約する分散機械学習手法です。
ポケモン版・マサカリ的実務視点
Section titled “ポケモン版・マサカリ的実務視点”| 比較項目 | 中央集権型学習 | Federated Learning |
|---|---|---|
| データの場所 | 全データを中央サーバーに集約 | 各デバイスに分散(移動しない) |
| プライバシー | データが外部に送信される | データは外部に出ない |
| 通信コスト | データ転送で膨大(GB単位) | モデル更新のみ(MB単位) |
| ポケモンの例え | 全ポケモンを育て屋に預ける | 各トレーナーが自分で育成 |
| 実例 | ImageNetで画像を集めて訓練 | スマホのキーボード予測(Gboard) |
🏗️ Federated Learningの仕組み
Section titled “🏗️ Federated Learningの仕組み”基本アルゴリズム: FedAvg
Section titled “基本アルゴリズム: FedAvg”# Federated Averaging (FedAvg)の疑似コード
# サーバー側def federated_training(global_model, clients, rounds=100): for round in range(rounds): # 1. クライアントをランダムに選択(全体の10%など) selected_clients = random.sample(clients, k=len(clients) // 10)
# 2. グローバルモデルを各クライアントに配布 client_updates = [] for client in selected_clients: updated_model = client.train(global_model) client_updates.append(updated_model)
# 3. クライアントの更新を集約(平均化) global_model = aggregate(client_updates)
return global_model
# クライアント側class FederatedClient: def __init__(self, local_data): self.data = local_data
def train(self, global_model, epochs=5): # ローカルデータでモデルを訓練 model = clone(global_model) for epoch in range(epochs): model.fit(self.data)
return model.get_weights() # データではなく重みだけを返す
# 集約関数def aggregate(client_updates): # 各クライアントの重みを平均化 avg_weights = [] for layer_weights in zip(*client_updates): avg_weights.append(np.mean(layer_weights, axis=0))
return avg_weights💻 実装1: TensorFlow Federatedを使った実装
Section titled “💻 実装1: TensorFlow Federatedを使った実装”import tensorflow as tfimport tensorflow_federated as tff
# 1. モデル定義def create_keras_model(): return tf.keras.models.Sequential([ tf.keras.layers.Dense(128, activation='relu', input_shape=(784,)), tf.keras.layers.Dense(10, activation='softmax') ])
def model_fn(): keras_model = create_keras_model() return tff.learning.models.from_keras_model( keras_model, input_spec=federated_train_data[0].element_spec, loss=tf.keras.losses.SparseCategoricalCrossentropy(), metrics=[tf.keras.metrics.SparseCategoricalAccuracy()] )
# 2. クライアントデータの準備(各病院のデータを想定)def create_federated_data(num_clients=10): # MNISTデータを10の病院に分割 (x_train, y_train), _ = tf.keras.datasets.mnist.load_data() x_train = x_train.reshape(-1, 784).astype('float32') / 255.0
client_data = [] samples_per_client = len(x_train) // num_clients
for i in range(num_clients): start = i * samples_per_client end = start + samples_per_client
client_dataset = tf.data.Dataset.from_tensor_slices( (x_train[start:end], y_train[start:end]) ).batch(32)
client_data.append(client_dataset)
return client_data
federated_train_data = create_federated_data(num_clients=10)
# 3. Federatedトレーニングプロセスを構築iterative_process = tff.learning.algorithms.build_weighted_fed_avg( model_fn, client_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=0.02), server_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=1.0))
# 4. 訓練実行state = iterative_process.initialize()
for round_num in range(50): # 各ラウンドでクライアントの一部を選択 sampled_clients = random.sample(federated_train_data, k=5)
# 訓練を実行 result = iterative_process.next(state, sampled_clients) state = result.state metrics = result.metrics
print(f'Round {round_num + 1}:') print(f' Train loss: {metrics["client_work"]["train"]["loss"]:.4f}') print(f' Train accuracy: {metrics["client_work"]["train"]["sparse_categorical_accuracy"]:.4f}')
# 5. 最終モデルの取得final_model = create_keras_model()final_model.set_weights(state.model.trainable)🔒 実装2: PySyftでプライバシー強化
Section titled “🔒 実装2: PySyftでプライバシー強化”import syft as syimport torchimport torch.nn as nnimport torch.optim as optim
# 1. Syft環境のセットアップhook = sy.TorchHook(torch)
# 仮想クライアント(病院A、病院B、病院C)hospital_a = sy.VirtualWorker(hook, id="hospital_a")hospital_b = sy.VirtualWorker(hook, id="hospital_b")hospital_c = sy.VirtualWorker(hook, id="hospital_c")
# 2. データを各病院に分散x_train_a = torch.tensor([[1.0, 2.0], [3.0, 4.0]]).send(hospital_a)y_train_a = torch.tensor([0, 1]).send(hospital_a)
x_train_b = torch.tensor([[5.0, 6.0], [7.0, 8.0]]).send(hospital_b)y_train_b = torch.tensor([1, 0]).send(hospital_b)
x_train_c = torch.tensor([[9.0, 10.0], [11.0, 12.0]]).send(hospital_c)y_train_c = torch.tensor([0, 1]).send(hospital_c)
# 3. モデル定義class SimpleModel(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(2, 2)
def forward(self, x): return self.fc(x)
model = SimpleModel()
# 4. Federated Learningトレーニングdef train_federated(model, workers, datasets, epochs=10): optimizer = optim.SGD(model.parameters(), lr=0.01) criterion = nn.CrossEntropyLoss()
for epoch in range(epochs): epoch_loss = 0.0
# 各ワーカー(病院)で訓練 for worker, (x, y) in zip(workers, datasets): # モデルをワーカーに送信 model.send(worker)
# ローカルで訓練 optimizer.zero_grad() output = model(x) loss = criterion(output, y) loss.backward() optimizer.step()
# モデルを取り戻す model.get()
epoch_loss += loss.get().item()
print(f'Epoch {epoch + 1}, Loss: {epoch_loss / len(workers):.4f}')
# 実行workers = [hospital_a, hospital_b, hospital_c]datasets = [(x_train_a, y_train_a), (x_train_b, y_train_b), (x_train_c, y_train_c)]
train_federated(model, workers, datasets, epochs=10)🏥 ユースケース1: 医療画像診断
Section titled “🏥 ユースケース1: 医療画像診断”import flwr as fl # Flower: シンプルなFLフレームワークfrom typing import Dict, List, Tuple
class MedicalImageClient(fl.client.NumPyClient): """各病院のクライアント"""
def __init__(self, hospital_id: str, local_data): self.hospital_id = hospital_id self.model = self.create_model() self.x_train, self.y_train = local_data
def create_model(self): model = tf.keras.Sequential([ tf.keras.layers.Conv2D(32, 3, activation='relu', input_shape=(224, 224, 3)), tf.keras.layers.MaxPooling2D(), tf.keras.layers.Conv2D(64, 3, activation='relu'), tf.keras.layers.MaxPooling2D(), tf.keras.layers.Flatten(), tf.keras.layers.Dense(128, activation='relu'), tf.keras.layers.Dense(3, activation='softmax') # 正常/肺炎/COVID-19 ]) model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy']) return model
def get_parameters(self, config): """モデルパラメータを取得""" return self.model.get_weights()
def fit(self, parameters, config): """ローカルデータで訓練""" self.model.set_weights(parameters)
# HIPAA準拠: データは外部に送信しない self.model.fit(self.x_train, self.y_train, epochs=5, batch_size=32, verbose=0)
return self.model.get_weights(), len(self.x_train), {}
def evaluate(self, parameters, config): """ローカルデータで評価""" self.model.set_weights(parameters) loss, accuracy = self.model.evaluate(self.x_train, self.y_train, verbose=0) return loss, len(self.x_train), {"accuracy": accuracy}
# サーバー起動def start_federated_server(): fl.server.start_server( server_address="0.0.0.0:8080", config=fl.server.ServerConfig(num_rounds=50), strategy=fl.server.strategy.FedAvg( fraction_fit=0.3, # 各ラウンドで30%の病院を選択 min_available_clients=3 # 最低3病院が必要 ) )
# クライアント起動(各病院で実行)def start_hospital_client(hospital_id: str): local_data = load_hospital_data(hospital_id) # 各病院のローカルデータ client = MedicalImageClient(hospital_id, local_data) fl.client.start_numpy_client(server_address="localhost:8080", client=client)
# 使用例# 病院A: python federated_client.py --hospital=A# 病院B: python federated_client.py --hospital=B# 病院C: python federated_client.py --hospital=C📱 ユースケース2: スマホキーボード予測(Gboard方式)
Section titled “📱 ユースケース2: スマホキーボード予測(Gboard方式)”class KeyboardPredictionClient: """スマホ上でのローカル学習"""
def __init__(self, user_id: str): self.user_id = user_id self.model = self.load_model() self.typing_data = [] # ユーザーのタイピングデータ
def collect_typing_data(self, text: str): """ユーザーのタイピングを記録(デバイス上で完結)""" self.typing_data.append(text)
def train_locally(self): """デバイス上でモデルを訓練(夜間充電中など)""" if len(self.typing_data) < 100: return # データ不足
# ローカルデータでファインチューニング tokenized_data = self.tokenize(self.typing_data) self.model.fit(tokenized_data, epochs=3, batch_size=16)
return self.model.get_weights() # 重みだけを返す(データは送信しない)
def participate_in_federated_round(self, global_weights): """Federatedラウンドに参加""" # グローバルモデルを受信 self.model.set_weights(global_weights)
# ローカルで訓練 local_weights = self.train_locally()
# 差分だけを送信(通信量削減) weight_diff = [local - global for local, global in zip(local_weights, global_weights)]
# さらにスパース化(重要な差分だけ送信) sparse_diff = self.sparsify(weight_diff, threshold=0.01)
return sparse_diff
def sparsify(self, weights, threshold): """小さな変更は送信しない(通信量を90%削減)""" return [np.where(np.abs(w) > threshold, w, 0) for w in weights]
# サーバー側の集約class FederatedKeyboardServer: def aggregate_updates(self, client_updates: List[np.ndarray], num_clients: int): """クライアントの更新を集約""" # Secure Aggregation: 個々のクライアント更新は見えない aggregated = [np.sum(updates, axis=0) / num_clients for updates in zip(*client_updates)]
return aggregated⚠️ Federated Learningの課題と対策
Section titled “⚠️ Federated Learningの課題と対策”課題1: 非IIDデータ(データ分布の偏り)
Section titled “課題1: 非IIDデータ(データ分布の偏り)”# 問題: 各病院で扱う症例が異なる(病院Aは高齢者、病院Bは小児が多いなど)
# 対策1: Personalized Federated Learningclass PersonalizedFLClient: def __init__(self): self.global_model = create_model() self.local_model = create_model() # 個別最適化モデル
def train(self, global_weights): # グローバルモデルを基にローカルモデルをファインチューニング self.global_model.set_weights(global_weights)
# ローカルデータでファインチューニング for layer in self.global_model.layers[:-2]: layer.trainable = False # 下層は固定
self.local_model.fit(local_data, epochs=5)
return self.global_model.get_weights() # グローバル層のみ共有課題2: 悪意ある参加者(Byzantine攻撃)
Section titled “課題2: 悪意ある参加者(Byzantine攻撃)”# 問題: 一部のクライアントが意図的に間違ったモデル更新を送信
# 対策: Robust Aggregation(中央値ベース)def robust_aggregate(client_updates: List[np.ndarray]): """外れ値を除外して集約""" # 各パラメータの中央値を計算 median_weights = [np.median(updates, axis=0) for updates in zip(*client_updates)]
# 中央値から大きく外れたクライアントを除外 filtered_updates = [] for update in client_updates: deviation = np.mean([np.abs(u - m).mean() for u, m in zip(update, median_weights)])
if deviation < threshold: # 正常なクライアント filtered_updates.append(update)
# フィルタ後に平均化 return [np.mean(updates, axis=0) for updates in zip(*filtered_updates)]課題3: 通信コスト
Section titled “課題3: 通信コスト”# 対策: モデル圧縮とスパース化class CommunicationEfficientClient: def compress_update(self, model_update: np.ndarray, compression_ratio: float = 0.1): """上位10%の重要な更新だけを送信""" # Top-K sparsification flat_update = model_update.flatten() k = int(len(flat_update) * compression_ratio)
# 絶対値が大きい上位K個だけを保持 top_k_indices = np.argsort(np.abs(flat_update))[-k:]
sparse_update = np.zeros_like(flat_update) sparse_update[top_k_indices] = flat_update[top_k_indices]
return sparse_update.reshape(model_update.shape)
# 通信量削減:# Before: 100MB/round# After: 10MB/round(90%削減)📊 Federated Learning vs 中央集権学習
Section titled “📊 Federated Learning vs 中央集権学習”| 項目 | 中央集権学習 | Federated Learning |
|---|---|---|
| プライバシー | ❌ データが外部流出 | ✅ データは各デバイスに留まる |
| 通信コスト | 🔴 高い(GB単位) | 🟢 低い(MB単位) |
| 訓練速度 | 🟢 高速 | 🔴 通信待ちで遅い |
| モデル精度 | 🟢 高精度 | 🟡 やや劣る(非IIDデータ) |
| 規制対応 | ❌ GDPR/HIPAA違反リスク | ✅ 準拠しやすい |
重要ポイント
Section titled “重要ポイント”-
プライバシー保護が最優先
医療・金融等のセンシティブデータはFederated Learningで活用 -
通信コストに注意
モデル圧縮、スパース化で通信量を90%削減 -
データ分布の偏りに対処
Personalized FLやRobust Aggregationを活用 -
セキュリティ対策
悪意ある参加者からの攻撃に備える
次のステップ
Section titled “次のステップ”- 26. AIエージェントと自律システム - 自律的なAIシステムの構築
- 22. AIセキュリティと敵対的攻撃 - FLのセキュリティ強化
- 18. AIオプトアウトとデータガバナンス - プライバシー管理の基礎
ポケモンマスターの知恵
「最強のジムリーダーを育てるには、各地のジムが独自に修行し、その成果を共有すればいい。ポケモンを1箇所に集める必要はない。」