연합학습
Federated Learning
데이터를 중앙 집중하지 않고 분산 학습. 개인정보 보호에 유리.
Federated Learning
데이터를 중앙 집중하지 않고 분산 학습. 개인정보 보호에 유리.
연합학습(Federated Learning)은 데이터를 중앙 서버로 수집하지 않고, 각 디바이스(클라이언트)에서 로컬 데이터로 모델을 학습한 뒤 모델 업데이트(가중치 또는 그래디언트)만 서버로 전송하여 글로벌 모델을 개선하는 분산 학습 패러다임입니다. Google이 2017년 제안하였으며, 개인정보 보호와 데이터 주권 문제를 해결하는 핵심 기술입니다.
연합학습의 기본 프로세스는 다음과 같습니다: (1) 서버가 글로벌 모델을 클라이언트에 배포, (2) 각 클라이언트가 로컬 데이터로 학습, (3) 클라이언트가 모델 업데이트를 서버로 전송, (4) 서버가 FedAvg 등의 알고리즘으로 업데이트를 집계하여 글로벌 모델 갱신. 이 과정을 여러 라운드 반복합니다.
연합학습은 Non-IID(비독립동일분포) 데이터, 시스템 이질성(디바이스 성능 차이), 통신 효율성, 보안 공격(Poisoning, Inference Attack) 등 고유한 도전 과제가 있습니다. 차분 프라이버시(Differential Privacy), Secure Aggregation, 압축 기법 등이 이를 해결하기 위해 연구되고 있습니다.
실제 활용 사례로는 스마트폰 키보드의 다음 단어 예측(Google Gboard), 의료 기관 간 협력 학습, 금융 사기 탐지, 자율주행 차량의 경험 공유 등이 있습니다. GDPR, 개인정보보호법 등 규제 환경에서 데이터를 직접 공유하지 않으면서도 협력 학습이 가능하다는 점에서 주목받고 있습니다.
# Flower 프레임워크를 이용한 연합학습 구현
import flwr as fl
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from collections import OrderedDict
# 간단한 모델 정의
class SimpleNet(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(784, 128)
self.fc2 = nn.Linear(128, 10)
def forward(self, x):
x = torch.relu(self.fc1(x))
return self.fc2(x)
# Flower 클라이언트 정의
class FlowerClient(fl.client.NumPyClient):
def __init__(self, model, train_loader, test_loader):
self.model = model
self.train_loader = train_loader
self.test_loader = test_loader
self.criterion = nn.CrossEntropyLoss()
self.optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
def get_parameters(self, config):
"""현재 모델 파라미터를 서버로 전송"""
return [val.cpu().numpy() for val in self.model.state_dict().values()]
def set_parameters(self, parameters):
"""서버로부터 받은 글로벌 모델로 업데이트"""
params_dict = zip(self.model.state_dict().keys(), parameters)
state_dict = OrderedDict({k: torch.tensor(v) for k, v in params_dict})
self.model.load_state_dict(state_dict, strict=True)
def fit(self, parameters, config):
"""로컬 데이터로 학습 (서버 라운드마다 호출)"""
self.set_parameters(parameters)
# 로컬 에포크 수 (보통 1-5)
local_epochs = config.get("local_epochs", 1)
self.model.train()
for _ in range(local_epochs):
for data, target in self.train_loader:
self.optimizer.zero_grad()
output = self.model(data)
loss = self.criterion(output, target)
loss.backward()
self.optimizer.step()
# 업데이트된 파라미터와 데이터 개수 반환
return self.get_parameters(config), len(self.train_loader.dataset), {}
def evaluate(self, parameters, config):
"""글로벌 모델 평가"""
self.set_parameters(parameters)
self.model.eval()
loss, correct = 0.0, 0
with torch.no_grad():
for data, target in self.test_loader:
output = self.model(data)
loss += self.criterion(output, target).item()
correct += (output.argmax(1) == target).sum().item()
accuracy = correct / len(self.test_loader.dataset)
return loss, len(self.test_loader.dataset), {"accuracy": accuracy}
# 클라이언트 실행 (각 디바이스에서)
def start_client(client_id):
# 각 클라이언트별 로컬 데이터 로드
train_loader = load_partition(client_id, "train")
test_loader = load_partition(client_id, "test")
model = SimpleNet()
client = FlowerClient(model, train_loader, test_loader)
# 서버에 연결
fl.client.start_numpy_client(
server_address="localhost:8080",
client=client
)
# 서버 실행 (중앙 서버에서)
def start_server():
strategy = fl.server.strategy.FedAvg(
fraction_fit=0.5, # 각 라운드에 참여할 클라이언트 비율
min_fit_clients=2, # 최소 참여 클라이언트 수
min_available_clients=3,
)
fl.server.start_server(
server_address="0.0.0.0:8080",
config=fl.server.ServerConfig(num_rounds=10), # 10라운드 학습
strategy=strategy
)
# 실행: python server.py (서버) / python client.py (각 클라이언트)
| 항목 | 일반적인 값 | 비고 |
|---|---|---|
| 글로벌 라운드 수 | 50-500 라운드 | 수렴까지 필요한 통신 횟수 |
| 로컬 에포크 | 1-5 에포크 | 라운드당 클라이언트 학습량 |
| 통신 오버헤드 | 모델 크기 x 2 / 라운드 | 다운로드 + 업로드 |
| 정확도 저하 (vs 중앙집중) | 1-5% | Non-IID 정도에 따라 상이 |
| 참여 클라이언트 수 | 10-1000+ | 더 많을수록 일반화 성능 향상 |
"FedAvg로 모델 집계하고 있는데, Non-IID라 FedProx로 바꿔볼까요?"
"클라이언트 드롭아웃이 심해서 robust aggregation 적용해야 해요."
"DP 적용하면 프라이버시는 보장되는데 정확도가 좀 떨어져요."
"통신 비용 줄이려고 그래디언트 압축 적용했어요."
1. Non-IID 과소평가: 클라이언트별 데이터 분포가 다르면 수렴이 어렵습니다. FedProx, SCAFFOLD 등 Non-IID에 강건한 알고리즘을 사용하세요.
2. 보안 공격 무시: 악의적 클라이언트가 잘못된 업데이트를 보내면 모델이 오염됩니다. Byzantine-robust aggregation이나 이상치 탐지를 적용하세요.
3. 프라이버시 과신: 그래디언트만 전송해도 역추론 공격으로 원본 데이터 복원이 가능할 수 있습니다. 차분 프라이버시나 Secure Aggregation을 함께 적용하세요.
4. 시스템 이질성: 느린 클라이언트가 병목이 됩니다. 비동기 학습이나 클라이언트 선택 전략을 고려하세요.
머신러닝 (Machine Learning) - 연합학습의 기반 기술
개인정보 보호책임자 (DPO) - 연합학습 도입 시 협력 필요
데이터 거버넌스 (Data Governance) - 분산 데이터 관리
딥러닝 (Deep Learning) - 연합학습에 주로 사용되는 모델
에포크 (Epoch) - 로컬 학습의 반복 단위