Skip to content

Federated Learning実装

🔐 Federated Learning実装:ジム戦で学ぶ分散学習

Section titled “🔐 Federated Learning実装:ジム戦で学ぶ分散学習”

**Federated Learning(連合学習)**を「ポケモンのジムリーダーが各地で独自に修行し、その成果だけを共有する」イメージで解説します。データを中央に集めずに、分散したデバイスでモデルを訓練できます。

  • 中央集権型学習 = 全てのポケモンを1箇所に集めて集中訓練
  • Federated Learning = 各ジムで独自に訓練し、戦術だけを共有
  • プライバシー保護 = ポケモンの個体値は秘密(技構成だけ共有)
  • モデル集約 = 各ジムの知見を統合して最強の戦術を編み出す
項目内容
FLの基礎中央集権学習との違い、プライバシー保護の仕組み
実装方法TensorFlow Federated、PySyft等のフレームワーク活用
ユースケース医療、金融、スマホキーボード予測等
課題と対策通信コスト、デバイス異質性、悪意ある参加者への対処

Federated Learningは、データを中央サーバーに集めずに、各デバイス上でモデルを訓練し、更新されたモデルパラメータ(重み)だけを中央サーバーに送信して集約する分散機械学習手法です。

ポケモン版・マサカリ的実務視点

Section titled “ポケモン版・マサカリ的実務視点”
比較項目中央集権型学習Federated Learning
データの場所全データを中央サーバーに集約各デバイスに分散(移動しない)
プライバシーデータが外部に送信されるデータは外部に出ない
通信コストデータ転送で膨大(GB単位)モデル更新のみ(MB単位)
ポケモンの例え全ポケモンを育て屋に預ける各トレーナーが自分で育成
実例ImageNetで画像を集めて訓練スマホのキーボード予測(Gboard)

# 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 tf
import 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 sy
import torch
import torch.nn as nn
import 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 Learning
class 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)]
# 対策: モデル圧縮とスパース化
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違反リスク✅ 準拠しやすい

  1. プライバシー保護が最優先
    医療・金融等のセンシティブデータはFederated Learningで活用

  2. 通信コストに注意
    モデル圧縮、スパース化で通信量を90%削減

  3. データ分布の偏りに対処
    Personalized FLやRobust Aggregationを活用

  4. セキュリティ対策
    悪意ある参加者からの攻撃に備える


ポケモンマスターの知恵
「最強のジムリーダーを育てるには、各地のジムが独自に修行し、その成果を共有すればいい。ポケモンを1箇所に集める必要はない。」