본문으로 건너뛰기
김신건의 로그

[FL] FedAvg (Federated Averaging)

· 수정 · 📖 약 4분 · 1,460자/단어 #ml #federated-learning #distributed #algorithm
FedAvg, Federated Averaging, FL FedAvg, FedSGD, 연합 평균, communication-efficient FL, McMahan FedAvg

정의

FedAvg (Federated Averaging) 는 McMahan et al. (2017) 이 제안한 연합 학습의 기본 알고리즘 입니다. 원본 데이터를 중앙 서버로 보내지 않고 각 클라이언트가 로컬 데이터로 모델을 학습한 뒤 weight 을 가중 평균 하여 글로벌 모델을 갱신합니다.

“local 여러 step, 서버 평균 한 번” 이라는 단순 구조로 FedSGD 대비 통신 라운드 수를 10-100 배 줄여, 실용적 연합 학습의 문을 열었습니다.

FedSGD 와의 차이

FedSGD (naive baseline)

  • 각 라운드에서 클라이언트가 1 step gradient 만 계산해 서버로 전송
  • 서버가 gradient 를 평균해 global weight 갱신
  • 각 SGD step 마다 통신 -> 매우 비효율

FedAvg (핵심 개선)

  • 각 클라이언트가 로컬에서 여러 epoch 을 돌리고
  • 최종 weight 를 서버에 전송
  • 서버는 weight 을 가중 평균
  • 한 라운드에 로컬 여러 step 이 들어가므로 통신 횟수 대폭 감소

핵심 통찰: 데이터가 IID 에 가깝다면 로컬 여러 step 이 낭비가 아니라 이득 (서버 라운드 = 통신 = 병목이므로).

알고리즘

Server:

Initialize w_0
For each round t = 0, 1, 2, ...:
    S_t = random subset of K clients (fraction C of N clients)
    For each client k in S_t (in parallel):
        w_k^{t+1} = ClientUpdate(k, w_t)
    w_{t+1} = sum over k in S_t of (n_k / n) * w_k^{t+1}

Client (ClientUpdate(k, w)):

Split D_k into batches of size B
For each local epoch e = 1..E:
    For each batch b:
        w <- w - eta * grad_l(w; b)
Return w

하이퍼파라미터

기호의미관용 값
전체 클라이언트 수100 ~ 10^6+ (cross-device)
라운드당 참여 fraction0.001 ~ 0.1 (cross-device), 1.0 (cross-silo)
라운드당 실제 참여 클라이언트 수
로컬 epoch 수1 ~ 20
로컬 minibatch 크기10 ~ 200
로컬 learning rate0.001 ~ 0.1
총 라운드 수수백 ~ 수만

서버 집계 공식

가중 평균 (weighted average):

  • : 클라이언트 의 데이터 개수
  • : 이번 라운드 참여 클라이언트 총 데이터 개수

왜 개수 비율로 가중: 데이터가 많은 클라이언트의 로컬 optimum 이 실제 loss 표면에 대해 더 많은 정보를 담고 있다고 가정 (empirical risk minimization 관점의 자연스러운 가중).

주의: 이 가중은 개수 편향 을 유발할 수 있습니다. 대형 클라이언트가 지배적이면 소수 클라이언트의 분포가 무시됩니다. 문제가 있으면 uniform 또는 클래스 균형 가중으로 대체.

왜 로컬 여러 epoch 이 통하는가

IID 데이터 (모든 클라이언트가 같은 분포에서 샘플링) 라면 각 클라이언트의 로컬 optimum 이 글로벌 optimum 과 근접합니다. 여러 step 이 낭비가 아니라 서버 라운드 사이의 진전.

Non-IID 라면 각 클라이언트가 자기 로컬 optimum 으로 수렴하려 하고, 서로 다른 방향으로 이동합니다 (client drift). 평균이 잘못된 방향으로 이동해 수렴 속도 저하 또는 성능 하락. 이 문제는 FL Non-IID 위키에서 상세 다룹니다.

Convergence 이론

Li et al. (2019) “On the Convergence of FedAvg on Non-IID Data” 는 convex loss 가정 하에 다음을 보였습니다.

  • IID + full participation: FedSGD 와 유사한 rate 로 수렴
  • Non-IID + partial participation: decay learning rate 가 필요, 아니면 수렴하지 않을 수 있음
  • 통신 라운드 관점에서 (strongly convex), (convex)

Non-convex (딥러닝) 는 정확한 이론이 없고 empirical validation 에 의존합니다.

FedSGD 로부터의 도출

E = 1, B = 무한 (전체 배치) 이면 FedAvg = FedSGD. 즉 FedSGD 는 FedAvg 의 특수 케이스.

E 를 늘리면:

  • 통신 감소 (좋음)
  • Client drift 증가 (나쁨, non-IID 에서)
  • Batch normalization 통계 로컬 편향 심화

Sweet spot: E = 1-5, B = 클라이언트 데이터의 1/10 정도.

통신 비용 분석

한 라운드 통신량:

  • 서버 -> 클라이언트: global weight 다운로드
  • 클라이언트 -> 서버: 로컬 학습 후 weight 업로드

는 모델 파라미터 총 크기 (float32 기준 4 bytes/param).

최적화 방법:

  • Model compression: quantization (8-bit, 4-bit), sparsification (top-k)
  • Structured updates: low-rank, random mask
  • FetchSGD: Count sketch 로 gradient 압축
  • Federated Dropout: 부분 모델만 학습/전송

배치 정규화의 함정

Batch Normalization 은 배치 통계에 의존합니다. 로컬 데이터가 non-IID 이면 로컬 BN 통계가 글로벌 분포를 대표하지 않아 성능 저하.

해결:

  • GroupNorm: 배치 무관 (BN 대체 자주 사용)
  • LayerNorm
  • FedBN: BN 파라미터는 로컬 유지, 다른 파라미터만 aggregate

실전 구현 (Flower)

import flwr as fl
import torch
from torch.utils.data import DataLoader

class Client(fl.client.NumPyClient):
    def __init__(self, model, train_ds, val_ds):
        self.model = model
        self.train_dl = DataLoader(train_ds, batch_size=32, shuffle=True)
        self.val_dl = DataLoader(val_ds, batch_size=64)

    def get_parameters(self, config):
        return [p.detach().cpu().numpy() for p in self.model.parameters()]

    def set_parameters(self, parameters):
        for p, new in zip(self.model.parameters(), parameters):
            p.data = torch.tensor(new)

    def fit(self, parameters, config):
        self.set_parameters(parameters)
        # 로컬 E epoch 학습
        opt = torch.optim.SGD(self.model.parameters(), lr=config["lr"])
        for epoch in range(config["local_epochs"]):
            for x, y in self.train_dl:
                opt.zero_grad()
                loss = self.model.loss(x, y)
                loss.backward()
                opt.step()
        return (
            self.get_parameters({}),
            len(self.train_dl.dataset),  # n_k, 서버가 가중 평균에 사용
            {}
        )

    def evaluate(self, parameters, config):
        self.set_parameters(parameters)
        loss, acc = self.eval()
        return float(loss), len(self.val_dl.dataset), {"acc": acc}

# 서버
strategy = fl.server.strategy.FedAvg(
    fraction_fit=0.1,           # C = 0.1
    min_fit_clients=10,
    min_available_clients=100,
    on_fit_config_fn=lambda rnd: {"lr": 0.01, "local_epochs": 1},
)

fl.server.start_server(
    server_address="0.0.0.0:8080",
    config=fl.server.ServerConfig(num_rounds=200),
    strategy=strategy,
)

FedAvg 의 개선/변형

변형개선점
FedProxNon-IID client drift 완화 (proximal term)
SCAFFOLDControl variate 로 drift correction
FedYogi / FedAdam서버측 adaptive optimizer
FedNovaLocal step 수 불균형 정규화
FedBNBatch Norm 통계 로컬 유지
Personalized FedAvg클라이언트별 로컬 fine-tuning

각 변형은 특정 조건 (non-IID 극심, system heterogeneity 등) 에 특화. 자세한 알고리즘 유도는 원 논문 참조.

함정

WARNING

Non-IID 극심 하면 FedAvg 자체가 수렴하지 않을 수 있습니다. Learning rate decay + FedProx/SCAFFOLD 등 변형 사용 고려.

CAUTION

가중 평균의 편향. 대형 클라이언트 데이터가 노이즈/편향이면 글로벌 모델이 그 편향을 흡수. Robust aggregation (median, trimmed mean) 을 검토.

WARNING

Client dropout. 라운드 중 클라이언트가 이탈하면 그 클라이언트의 weight 만 누락 -> 편향. Secure Aggregation 은 dropout 을 threshold secret sharing 으로 처리.

IMPORTANT

BN 통계 문제. 딥러닝 모델에 BN 이 있으면 성능 저하 위험. GroupNorm 대체 또는 FedBN 사용.

관련 위키

이 글의 용어 (6개)
[FL] Frameworks (Flower, TFF, NVFlare, FATE, PySyft)ml
정의 연합 학습 프레임워크 는 서버-클라이언트 통신, 집계 알고리즘, 프라이버시 도구, 시뮬레이션 환경, 실 배포 인프라를 제공하는 라이브러리/플랫폼입니다. 알고리즘을 직접 구현…
[FL] Non-IID Data & Client Driftml
정의 Non-IID data 는 연합 학습에서 각 클라이언트의 데이터가 서로 다른 분포에서 추출되는 상황을 말합니다. FedAvg 를 비롯한 대부분 FL 알고리즘의 최대 난관이며…
[FL] Personalized Federated Learningml
정의 Personalized FL 은 모든 클라이언트가 동일한 글로벌 모델을 쓰지 않고, 각 클라이언트가 자신의 데이터 분포에 맞춰 조정된 모델 을 학습하는 연합 학습 계열입니다…
[FL] Secure Aggregationml
정의 Secure Aggregation 은 연합 학습 서버가 개별 클라이언트 업데이트 를 보지 못하고 오직 합 (또는 평균) 만 볼 수 있도록 보장하는 암호학 프로토콜입니다. 서…
Differential Privacy: (ε, δ) 로 정량화하는 프라이버시 보장ml
정의 Differential Privacy (DP, 차분 프라이버시) 는 데이터셋에 대한 질의 (query) 결과에 calibrated noise 를 추가함으로써, 한 개인의 데…
Federated Learning: 분산 학습 without central dataml
정의 Federated Learning (FL) 은 데이터를 중앙에 모으지 않고 각 클라이언트 (edge device, 병원, 은행 등) 가 로컬 데이터로 모델을 학습한 뒤 모델…

💬 댓글

사이트 검색 / 명령어

검색

스크롤 = 확대/축소 · 드래그 = 이동 · 0 = 원래 크기 · ESC = 닫기