분산 환경에서의 연합 학습: FedProx와 SCAFFOLD 알고리즘 비교

AG News 데이터셋을 활용하여 연합 학습(Federated Learning) 환경에서 FedProx와 SCAFFOLD 알고리즘의 성능을 비교 분석하였다. IID와 Non-IID 데이터 분포 환경에서의 실험을 통해 각 알고리즘의 특성과 한계를 파악하고, Bayesian Optimization을 활용한 하이퍼파라미터 튜닝을 수행하여 최적의 설정을 도출하였다.

1. 서론

1.1 연합 학습이란?

연합 학습(Federated Learning, FL)은 분산된 데이터를 중앙 서버로 수집하지 않고도 여러 클라이언트(디바이스)에서 협력적으로 머신러닝 모델을 학습하는 패러다임이다. Google이 2016년 처음 제안한 이 방식은 데이터 프라이버시를 보장하면서도 대규모 분산 데이터를 활용할 수 있다는 점에서 학계와 산업계 모두에서 큰 주목을 받고 있다.

전통적인 머신러닝은 모든 데이터를 중앙 서버에 모아서 학습하는 방식이다. 하지만 이 방식은 몇 가지 심각한 문제를 안고 있다:

  1. 프라이버시 문제: 개인의 민감한 데이터(의료 기록, 금융 정보, 위치 데이터, 개인 메시지 등)를 중앙에 모으는 것은 심각한 프라이버시 침해 위험이 있다. 데이터 유출 사고가 발생할 경우 막대한 피해가 발생할 수 있으며, 사용자들의 신뢰를 잃을 수 있다.

  2. 통신 비용: 대용량 데이터를 중앙 서버로 전송하는 것은 막대한 네트워크 비용을 발생시킨다. 특히 모바일 환경에서는 제한된 대역폭과 데이터 요금 문제가 있어 대규모 데이터 전송이 현실적으로 어려울 수 있다.

  3. 법적 규제: GDPR(General Data Protection Regulation) 등 데이터 보호 규정으로 인해 데이터를 특정 지역 외부로 이동시키기 어려운 경우가 많다. 특히 의료, 금융 분야에서는 데이터 이동에 대한 엄격한 규제가 존재한다.

  4. 실시간성: 중앙 집중식 학습은 데이터 수집과 처리에 시간이 걸려 실시간 서비스에 적합하지 않을 수 있다.

연합 학습은 이러한 문제들을 해결하기 위해 제안되었다. 기본 아이디어는 다음과 같다:

  1. 중앙 서버가 초기화된 글로벌 모델을 각 클라이언트에 배포한다.
  2. 각 클라이언트는 자신의 로컬 데이터로 모델을 독립적으로 학습한다.
  3. 클라이언트는 학습된 모델의 파라미터(또는 그래디언트)만 서버에 전송한다. 원본 데이터는 절대 전송되지 않는다.
  4. 서버는 수집된 파라미터를 집계(aggregation)하여 글로벌 모델을 업데이트한다.
  5. 이 과정을 모델이 수렴할 때까지 반복한다.

이렇게 하면 원본 데이터는 각 클라이언트에 그대로 남아있고, 모델 파라미터만 공유되므로 프라이버시가 보호된다. 또한 데이터 전송량이 크게 줄어들어 통신 비용도 절감된다.

1.2 연합 학습의 실제 활용 사례

연합 학습은 이미 다양한 실제 서비스에서 활용되고 있다:

Google Gboard 키보드: 가장 대표적인 연합 학습 활용 사례이다. 수백만 대의 안드로이드 기기에서 사용자의 타이핑 패턴을 학습하여 다음 단어 예측 모델을 개선한다. 사용자의 실제 입력 데이터는 기기를 떠나지 않으며, 학습된 모델 업데이트만 서버로 전송된다.

Apple Siri: 음성 인식 모델을 개선하기 위해 연합 학습을 활용한다. 사용자의 음성 데이터를 기기에서 직접 처리하여 프라이버시를 보호하면서도 모델 성능을 지속적으로 향상시킨다.

의료 분야: 여러 병원이 환자 데이터를 공유하지 않으면서도 협력하여 질병 진단 모델을 학습할 수 있다. 이는 의료 데이터의 민감성과 규제를 고려할 때 매우 중요한 적용 사례이다.

금융 분야: 여러 금융 기관이 고객 데이터를 공유하지 않으면서도 사기 탐지 모델을 공동으로 학습할 수 있다.

1.3 연합 학습의 핵심 도전 과제

연합 학습은 이론적으로는 매력적이지만, 실제 구현에서는 여러 도전 과제에 직면한다:

Non-IID 데이터 분포

가장 큰 도전 과제 중 하나는 Non-IID(Non-Independent and Identically Distributed) 데이터 분포 문제이다. 실제 환경에서 각 클라이언트의 데이터는 서로 매우 다른 분포를 가질 수 있다.

예를 들어, 스마트폰 키보드 예측 모델을 학습한다고 가정해보자:

  • 한 사용자는 주로 영어로 비즈니스 이메일을 작성하고
  • 다른 사용자는 한국어로 일상적인 메시지를 주로 보내며
  • 또 다른 사용자는 기술 관련 용어와 코드를 많이 입력하고
  • 어떤 사용자는 이모티콘과 줄임말을 자주 사용한다

이처럼 각 클라이언트의 데이터 분포가 다르면, 로컬 학습 과정에서 모델이 각자의 데이터에 과적합되어 글로벌 모델과 크게 벗어나는 클라이언트 드리프트(Client Drift) 현상이 발생한다.

클라이언트 드리프트 문제의 메커니즘

클라이언트 드리프트는 Non-IID 환경에서 발생하는 핵심 문제이다. 그 메커니즘을 상세히 살펴보면:

  1. 로컬 학습 수행: 각 클라이언트가 글로벌 모델을 받아 로컬 데이터로 여러 에폭의 학습을 수행한다.

  2. 편향된 업데이트 생성: Non-IID 데이터로 인해 각 로컬 모델은 자신의 데이터 분포에 최적화된다. 이는 글로벌 최적점과 다른 방향으로의 업데이트를 의미한다.

  3. 누적되는 편향: 로컬 학습 에폭이 많을수록 각 클라이언트 모델은 글로벌 모델에서 더 멀어진다. 이것이 "드리프트"이다.

  4. 비효율적인 집계: 서로 다른 방향으로 드리프트된 업데이트들이 집계되면, 상충되는 그래디언트로 인해 글로벌 모델의 업데이트가 비효율적이 된다.

  5. 수렴 문제: 심한 경우 모델이 전혀 수렴하지 않거나, 수렴하더라도 성능이 중앙 집중식 학습보다 크게 떨어질 수 있다.

이 문제를 해결하기 위해 FedProx, SCAFFOLD, FedNova, FedOpt 등 다양한 알고리즘이 제안되었다. 이 프로젝트에서는 그 중 FedProx와 SCAFFOLD를 비교 분석한다.

기타 도전 과제들

시스템 이질성(System Heterogeneity): 클라이언트마다 컴퓨팅 파워, 네트워크 속도, 배터리 상태가 다르다. 일부 클라이언트는 학습을 완료하지 못하거나 탈락(dropout)할 수 있다.

통신 효율성: 모델 파라미터를 주고받는 통신 비용을 최소화해야 한다. 특히 대규모 모델의 경우 이 문제가 심각해진다.

프라이버시 공격: 모델 파라미터만으로도 원본 데이터에 대한 정보가 유출될 수 있다(gradient leakage attack). 추가적인 프라이버시 보호 기법(차분 프라이버시 등)이 필요할 수 있다.

악의적 클라이언트: 일부 클라이언트가 의도적으로 잘못된 업데이트를 보내는 공격(Byzantine attack)에 대응해야 한다.

1.4 프로젝트 목표

이 프로젝트의 구체적인 목표는 다음과 같다:

  1. 텍스트 분류 데이터셋에서의 연합 학습 벤치마크 구축: 기존 연합 학습 연구가 주로 이미지 분류(MNIST, CIFAR-10)에 집중된 것과 달리, AG News 텍스트 분류 데이터셋을 활용하여 NLP 도메인에서의 연합 학습을 검증한다.

  2. IID/Non-IID 데이터 구성 방법 구현: Dirichlet 분포를 활용하여 다양한 수준의 데이터 이질성을 시뮬레이션할 수 있는 데이터 분할 방법을 구현하고, 그 특성을 분석한다.

  3. FedProx와 SCAFFOLD 알고리즘 비교: 두 가지 대표적인 클라이언트 드리프트 완화 알고리즘을 구현하고, 다양한 환경에서 성능을 비교 분석한다.

  4. Bayesian Optimization을 활용한 체계적인 하이퍼파라미터 튜닝: wandb의 sweep 기능을 활용하여 과학적이고 체계적인 하이퍼파라미터 최적화를 수행한다.

  5. 실용적 가이드라인 제시: 실험 결과를 바탕으로 언제 어떤 알고리즘을 선택해야 하는지에 대한 실용적인 가이드라인을 제시한다.

2. 데이터셋 및 전처리

2.1 AG News 데이터셋

AG News 데이터셋은 텍스트 분류 작업을 위한 대표적인 벤치마크 데이터셋이다. 원래 AG's corpus of news articles에서 추출되었으며, 학술 연구 목적으로 가공된 버전이 널리 사용된다.

이 데이터셋은 뉴스 헤드라인과 설명으로 구성되어 있으며, 각 샘플은 4개의 카테고리 중 하나로 레이블링되어 있다:

레이블 카테고리 설명 예시
0 World 세계 뉴스, 국제 정치, 외교 "UN Security Council meets to discuss..."
1 Sports 스포츠 뉴스, 경기 결과, 선수 소식 "Lakers defeat Celtics in overtime..."
2 Business 비즈니스, 경제, 금융 뉴스 "Stock market reaches record high as..."
3 Sci/Tech 과학, 기술, IT 뉴스 "Apple announces new iPhone with..."

데이터셋의 상세 특성:

  • 훈련 세트: 120,000개 샘플 (각 카테고리당 정확히 30,000개)
  • 테스트 세트: 7,600개 샘플 (각 카테고리당 정확히 1,900개)
  • 균형 잡힌 분포: 모든 클래스가 동일한 수의 샘플을 가짐
  • 텍스트 길이: 평균 약 40~50 단어, 최대 수백 단어

AG News를 선택한 이유:

  1. 텍스트 분류의 표준 벤치마크: 많은 NLP 연구에서 사용되어 비교가 용이
  2. 적절한 난이도: 4개 클래스로 너무 쉽지도 어렵지도 않음
  3. 균형 잡힌 데이터: 클래스 불균형 문제 없이 순수하게 알고리즘 성능을 평가 가능
  4. 적절한 규모: 연합 학습 실험에 충분한 크기

2.2 텍스트 전처리 파이프라인

텍스트 데이터를 모델에 입력하기 위해 체계적인 전처리 파이프라인을 구축하였다.

토크나이저 선택

BERT 토크나이저(bert-base-uncased)를 사용하여 토큰화를 수행하였다:

from transformers import BertTokenizer

tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')

BERT 토크나이저를 선택한 이유:

  1. WordPiece 알고리즘: 미등록어(OOV) 문제를 효과적으로 처리
  2. 표준화된 어휘: 30,522개의 풍부한 어휘 제공
  3. 서브워드 분할: 희귀 단어도 의미있는 서브워드로 분할
  4. 널리 사용됨: 재현성과 비교 용이성
토큰화 설정

토큰화 과정에서 다음 설정을 적용하였다:

encoding = tokenizer(
    text,
    truncation=True,        # 긴 시퀀스 절단
    padding='max_length',   # 짧은 시퀀스 패딩
    max_length=128,         # 최대 시퀀스 길이
    return_tensors='pt'     # PyTorch 텐서 반환
)
  • 최대 시퀀스 길이 128 토큰: AG News의 대부분의 샘플이 128 토큰 이내에 포함됨. 더 긴 길이는 계산 비용 대비 성능 향상이 미미함.
  • 패딩: 짧은 시퀀스는 [PAD] 토큰으로 최대 길이까지 채움
  • 절단: 긴 시퀀스는 최대 길이에서 뒷부분을 절단
  • 특수 토큰: [CLS], [SEP] 토큰이 자동으로 추가됨

2.3 커스텀 PyTorch Dataset 클래스

데이터셋을 PyTorch의 Dataset 형식으로 래핑하여 효율적인 데이터 로딩을 구현하였다:

class AGNewsDataset(Dataset):
    def __init__(self, texts, labels, tokenizer, max_length=128):
        self.texts = texts
        self.labels = labels
        self.tokenizer = tokenizer
        self.max_length = max_length

    def __len__(self):
        return len(self.texts)

    def __getitem__(self, idx):
        text = self.texts[idx]
        label = self.labels[idx]

        encoding = self.tokenizer(
            text,
            truncation=True,
            padding='max_length',
            max_length=self.max_length,
            return_tensors='pt'
        )

        return {
            'input_ids': encoding['input_ids'].squeeze(),
            'labels': torch.tensor(label, dtype=torch.long)
        }

이 클래스의 설계 특징:

  1. 지연 토큰화(Lazy Tokenization): 메모리 효율을 위해 __getitem__에서 토큰화 수행
  2. 배치 호환성: squeeze()를 통해 배치 처리에 적합한 텐서 형태 반환
  3. 유연성: max_length 파라미터로 시퀀스 길이 조절 가능

3. IID 및 Non-IID 데이터 분할

연합 학습 실험에서 가장 중요한 부분 중 하나는 클라이언트 간 데이터 분포를 어떻게 설정하느냐이다. 이 프로젝트에서는 IID와 Non-IID 두 가지 시나리오를 모두 구현하여 알고리즘의 성능을 종합적으로 평가하였다.

3.1 IID 데이터 분할

IID(Independent and Identically Distributed) 설정에서는 모든 클라이언트가 전체 데이터셋의 분포를 동일하게 반영하는 로컬 데이터셋을 가진다. 이는 가장 이상적인 연합 학습 환경을 나타내며, 알고리즘의 기본 성능을 평가하는 베이스라인으로 사용된다.

IID 데이터 분할 구현:

def split_data_iid(dataset: Dataset, num_users: int) -> List[Dataset]:
    # 사용자당 샘플 수 계산
    num_items = len(dataset) // num_users

    # 데이터셋 인덱스를 무작위로 섞기
    indices = np.random.permutation(len(dataset))

    # 섞인 인덱스를 동일한 크기의 서브셋으로 분할
    return [
        torch.utils.data.Subset(dataset, indices[i * num_items:(i + 1) * num_items])
        for i in range(num_users)
    ]

이 함수의 동작 방식:

  1. 샘플 수 계산: 전체 샘플 수를 사용자 수로 나누어 각 사용자가 받을 샘플 수를 결정한다. 예: 120,000개 / 10명 = 12,000개/명
  2. 무작위 셔플링: NumPy의 random.permutation을 사용하여 데이터셋의 인덱스를 무작위로 섞는다. 이를 통해 원래 데이터의 순서에 의한 편향을 제거한다.
  3. 균등 분할: 섞인 인덱스를 동일한 크기의 서브셋으로 나눈다.

IID 분할의 수학적 특성:

  • 각 사용자가 동일한 수의 샘플을 받음: |D_k| = |D| / K
  • 각 사용자의 레이블 분포가 전체 데이터셋의 분포와 유사함: P_k(y) ≈ P(y)
  • 사용자 간 데이터 편향이 없음

예시 출력 (10명의 사용자, 120,000개 샘플):

<mark class="highlight"> IID Data Distribution </mark>
User 0: 12000 samples, Label distribution: [2971 3037 2981 3011]
User 1: 12000 samples, Label distribution: [2985 3021 2930 3064]
User 2: 12000 samples, Label distribution: [2976 3038 3032 2954]
User 3: 12000 samples, Label distribution: [2884 2965 3087 3064]
User 4: 12000 samples, Label distribution: [3079 3000 3000 2921]
User 5: 12000 samples, Label distribution: [2996 2994 2986 3024]
User 6: 12000 samples, Label distribution: [3020 2978 3007 2995]
User 7: 12000 samples, Label distribution: [3046 3040 2904 3010]
User 8: 12000 samples, Label distribution: [3046 2943 2999 3012]
User 9: 12000 samples, Label distribution: [2997 2984 3074 2945]

모든 사용자가 4개 클래스에서 거의 균등한 분포(약 3,000개씩, 이상적으로는 정확히 3,000개)를 가지고 있음을 확인할 수 있다. 약간의 편차는 무작위 셔플링의 자연스러운 결과이다.

3.2 Non-IID 데이터 분할 (Dirichlet 분포 기반)

실제 연합 학습 환경에서는 각 클라이언트가 서로 다른 데이터 분포를 가지는 Non-IID 상황이 일반적이다. 이를 시뮬레이션하기 위해 Dirichlet 분포를 활용한 데이터 분할 방법을 구현하였다.

Dirichlet 분포의 수학적 배경

Dirichlet 분포는 확률 벡터(합이 1인 양수 벡터)를 생성하는 확률 분포이다. K차원 Dirichlet 분포는 다음과 같이 정의된다:

$$Dir(\alpha_1, \alpha_2, ..., \alpha_K)$$

여기서 α_i는 집중 파라미터(concentration parameter)이다.

모든 α_i = α인 대칭 Dirichlet 분포의 경우:

  • α가 작을 때 (예: 0.1): 확률이 특정 요소에 집중되는 경향이 있음 (높은 이질성). 생성된 벡터는 희소(sparse)한 경향을 보임.
  • α = 1: 균등 분포(uniform distribution). 모든 확률 벡터가 동일한 확률로 생성됨.
  • α가 클 때 (예: 10.0): 확률이 균등하게 분산됨 (낮은 이질성, IID에 가까움). 생성된 벡터의 각 요소가 1/K에 가까움.

Non-IID 데이터 분할 구현:

def split_data_non_iid(dataset: Dataset, num_users: int, alpha: float) -> List[Dataset]:
    # 텍스트 길이 계산 (추가적인 구조적 편향)
    text_lengths = [len(dataset.texts[idx]) for idx in range(len(dataset))]

    # 텍스트 길이 기준으로 데이터셋 인덱스 정렬
    sorted_indices = np.argsort(text_lengths)

    # Dirichlet 분포 확률 생성
    dirichlet_dist = np.random.dirichlet([alpha] * num_users, size=len(sorted_indices))

    # Dirichlet 확률에 따라 샘플을 사용자에게 할당
    user_indices = [[] for _ in range(num_users)]
    for idx, probabilities in enumerate(dirichlet_dist):
        for user_id, prob in enumerate(probabilities):
            if np.random.rand() < prob:
                user_indices[user_id].append(sorted_indices[idx])
                break

    # 각 사용자의 서브셋 반환
    return [torch.utils.data.Subset(dataset, indices) for indices in user_indices]

이 구현의 핵심 요소 상세 설명:

  1. 텍스트 길이 기반 정렬: 먼저 텍스트 길이를 기준으로 데이터를 정렬한다. 이는 실제 환경에서 특정 클라이언트가 더 긴 또는 짧은 텍스트에 접근할 수 있는 상황을 시뮬레이션한다. 예를 들어, 트위터 사용자는 짧은 텍스트를, 블로그 작성자는 긴 텍스트를 가질 수 있다.

  2. Dirichlet 분포 기반 확률 생성: 각 샘플에 대해 num_users 차원의 확률 벡터를 Dirichlet 분포에서 샘플링한다. 이 확률 벡터는 해당 샘플이 각 사용자에게 할당될 확률을 나타낸다.

  3. 확률적 할당: 생성된 확률에 따라 각 샘플이 특정 사용자에게 할당된다. 이 과정에서 추가적인 무작위성이 도입되어 더 다양한 분포를 생성한다.

α 파라미터의 영향 상세 분석

α = 0.1 (높은 이질성): - 일부 사용자가 특정 레이블의 대부분을 차지 - 일부 사용자는 특정 레이블의 샘플이 거의 없거나 전혀 없을 수 있음 - 극단적인 Non-IID 시나리오로, 연합 학습의 스트레스 테스트에 적합 - 실제 환경의 예: 특정 지역 사용자는 특정 언어의 뉴스만 접함

α = 0.5 (중간 이질성): - 사용자 간 레이블 분포에 부분적 편향 존재 - 모든 사용자가 어느 정도의 다양성을 유지하면서도 편향이 있음 - 가장 현실적인 시나리오 중 하나 - 실제 환경의 예: 사용자마다 관심사가 다르지만 다양한 뉴스도 접함

α = 1.0 (균등 분포): - 분포가 IID에 가까워짐 - 약간의 무작위 편차만 존재 - 연합 학습이 거의 중앙 집중식 학습과 유사하게 동작

α = 10.0 이상 (매우 낮은 이질성): - 사실상 IID와 구분하기 어려움 - 모든 사용자가 거의 동일한 분포를 가짐

예시 출력 (α = 0.5, 10명의 사용자):

<mark class="highlight"> Non-IID Data Distribution </mark>
User 0: 11821 samples, Label distribution: [2923 2974 3000 2924]
User 1: 11018 samples, Label distribution: [2704 2768 2776 2770]
User 2: 10132 samples, Label distribution: [2563 2541 2500 2528]
User 3: 9228 samples, Label distribution: [2356 2286 2244 2342]
User 4: 8434 samples, Label distribution: [2103 2085 2196 2050]
User 5: 7756 samples, Label distribution: [1989 1918 1905 1944]
User 6: 7046 samples, Label distribution: [1753 1781 1742 1770]
User 7: 6349 samples, Label distribution: [1583 1589 1608 1569]
User 8: 5556 samples, Label distribution: [1348 1375 1406 1427]
User 9: 5290 samples, Label distribution: [1342 1331 1304 1313]

사용자마다 샘플 수가 크게 다르고(5,290~11,821개), 이는 실제 환경에서 활발한 사용자와 비활발한 사용자가 존재하는 상황을 반영한다.

4. 모델 아키텍처

4.1 GRU 모델 선택 이유

이 프로젝트에서는 텍스트 분류를 위해 GRU(Gated Recurrent Unit) 기반 모델을 사용하였다. GRU를 선택한 구체적인 이유:

  1. 순차 데이터 처리 능력: GRU는 순환 신경망의 일종으로, 텍스트와 같은 순차 데이터를 자연스럽게 처리할 수 있다. 단어의 순서와 문맥을 고려한 학습이 가능하다.

  2. 장기 의존성 학습: GRU의 게이팅 메커니즘(업데이트 게이트, 리셋 게이트)은 긴 시퀀스에서도 중요한 정보를 유지할 수 있게 해준다. 이는 뉴스 기사처럼 긴 텍스트를 처리할 때 중요하다.

  3. 경량 아키텍처: LSTM(Long Short-Term Memory)보다 게이트가 하나 적어(LSTM: 3개, GRU: 2개) 파라미터 수가 약 25% 적다. 이는 연합 학습 환경에서의 통신 비용과 클라이언트 메모리 사용량을 줄일 수 있다.

  4. 빠른 학습 속도: 파라미터가 적으므로 학습이 빠르다. 리소스가 제한된 클라이언트 디바이스(스마트폰 등)에서도 효율적으로 학습할 수 있다.

  5. Transformer 대비 단순성: Transformer 모델(BERT 등)은 성능은 좋지만 파라미터 수가 매우 많아(BERT-base: 110M) 연합 학습에는 비효율적이다. GRU는 적절한 성능과 효율성의 균형을 제공한다.

4.2 GRU의 동작 원리 상세 설명

GRU는 기본 RNN의 기울기 소실 문제를 해결하기 위해 Cho et al.(2014)이 제안한 아키텍처이다. 두 가지 게이트를 사용하여 정보의 흐름을 제어한다:

1. 업데이트 게이트 (Update Gate, z): $$z_t = \sigma(W_z \cdot [h_{t-1}, x_t] + b_z)$$

  • 역할: 이전 은닉 상태에서 얼마나 많은 정보를 유지할지 결정
  • 값이 1에 가까우면: 이전 정보를 많이 유지 (장기 기억에 적합)
  • 값이 0에 가까우면: 새로운 정보로 대체 (새로운 정보에 집중)
  • 직관적 이해: "이 시점에서 새로운 정보를 얼마나 반영할 것인가?"

2. 리셋 게이트 (Reset Gate, r): $$r_t = \sigma(W_r \cdot [h_{t-1}, x_t] + b_r)$$

  • 역할: 이전 은닉 상태를 얼마나 무시할지 결정
  • 값이 0에 가까우면: 이전 상태를 무시하고 새로운 입력에 집중
  • 값이 1에 가까우면: 이전 상태를 완전히 고려
  • 직관적 이해: "과거 정보를 얼마나 참고할 것인가?"

3. 후보 은닉 상태 계산: $$\tilde{h}_t = \tanh(W \cdot [r_t \odot h_{t-1}, x_t] + b)$$

  • 리셋 게이트로 필터링된 이전 상태와 현재 입력을 결합
  • tanh 활성화로 -1 ~ 1 범위로 정규화

4. 최종 은닉 상태 계산: $$h_t = (1 - z_t) \odot h_{t-1} + z_t \odot \tilde{h}_t$$

  • 업데이트 게이트를 사용하여 이전 상태와 후보 상태를 선형 보간
  • 이를 통해 필요에 따라 정보를 유지하거나 업데이트

이러한 게이팅 메커니즘 덕분에 GRU는 기존 RNN의 기울기 소실 문제를 완화하고, 긴 시퀀스에서도 효과적으로 학습할 수 있다.

4.3 모델 구현

GRU 기반 텍스트 분류 모델의 PyTorch 구현:

class GRUModel(nn.Module):
    def __init__(self, vocab_size, embedding_dim=128, hidden_dim=256, num_layers=1):
        super(GRUModel, self).__init__()
        self.embedding = nn.Embedding(vocab_size, embedding_dim)
        self.gru = nn.GRU(embedding_dim, hidden_dim, num_layers, batch_first=True)
        self.fc = nn.Linear(hidden_dim, vocab_size)

    def forward(self, input_ids):
        embedded = self.embedding(input_ids)
        gru_out, _ = self.gru(embedded)
        logits = self.fc(gru_out)
        return logits

모델의 구성 요소 상세 설명:

1. 임베딩 레이어 (Embedding Layer): - 입력: 토큰 ID (정수, 0 ~ vocab_size-1) - 출력: 밀집 벡터 표현 (128차원) - 역할: 이산적인 토큰을 연속적인 벡터 공간으로 매핑 - 파라미터 수: vocab_size × embedding_dim (약 30,522 × 128 = 3.9M)

2. GRU 레이어: - 입력: 임베딩 벡터 시퀀스 [batch_size, seq_length, embedding_dim] - 출력: 은닉 상태 시퀀스 [batch_size, seq_length, hidden_dim] - 설정: hidden_dim=256, num_layers=1, batch_first=True - 역할: 순차적 정보 처리 및 문맥 인코딩 - 파라미터 수: 3 × (embedding_dim × hidden_dim + hidden_dim × hidden_dim + 2 × hidden_dim)

3. 완전 연결 레이어 (Fully Connected Layer): - 입력: GRU 출력 (256차원) - 출력: 어휘 크기만큼의 로짓 (vocab_size차원) - 역할: 최종 분류 점수 계산 - 파라미터 수: hidden_dim × vocab_size

4.4 손실 함수: Cross-Entropy Loss

분류 작업을 위해 Cross-Entropy Loss를 사용하였다. 이 손실 함수는 예측 확률 분포와 실제 레이블 분포 사이의 차이를 측정한다.

수학적 정의: $$\mathcal{L} = -\frac{1}{N} \sum_{i=1}^{N} \sum_{c=1}^{C} y_{i,c} \log(\hat{y}_{i,c})$$

여기서:

  • N: 배치 내 샘플 수
  • C: 클래스 수
  • y_{i,c}: 샘플 i의 실제 레이블이 c인지 나타내는 원-핫 인코딩
  • ŷ_{i,c}: 모델이 예측한 샘플 i가 클래스 c일 확률

Cross-Entropy Loss의 핵심 특성:

  1. 로그 패널티의 직관적 이해: - 정답에 0.9 확률 부여: -log(0.9) ≈ 0.105 (작은 패널티) - 정답에 0.5 확률 부여: -log(0.5) ≈ 0.693 (중간 패널티) - 정답에 0.1 확률 부여: -log(0.1) ≈ 2.302 (큰 패널티) - 정답에 0.01 확률 부여: -log(0.01) ≈ 4.605 (매우 큰 패널티)

  2. 확률적 해석: KL 발산(Kullback-Leibler divergence)과 밀접한 관계. 예측 분포가 실제 분포에서 멀어질수록 손실이 증가.

  3. 그래디언트 특성: softmax와 결합하면 그래디언트가 (예측 - 실제)의 형태로 단순해져 효율적인 학습 가능.

PyTorch 구현:

criterion = nn.CrossEntropyLoss()

# Forward pass
outputs = model(input_ids)  # [batch_size, seq_length, vocab_size]
outputs = outputs.view(-1, outputs.size(-1))  # Flatten to [N, C]
labels = labels.view(-1)  # Flatten to [N]

# Compute loss
loss = criterion(outputs, labels)

5. 연합 학습 알고리즘

5.1 FedAvg: 기본 알고리즘과 그 한계

연합 학습의 가장 기본적인 알고리즘인 FedAvg(Federated Averaging)를 먼저 이해하는 것이 중요하다. McMahan et al.(2017)이 제안한 이 알고리즘은 다음과 같이 동작한다:

FedAvg 알고리즘: 1. 서버가 글로벌 모델 θ^(t)를 K개의 클라이언트 중 일부(m개)에 배포 2. 각 선택된 클라이언트 k가 로컬 데이터로 E 에폭 학습: θ_k^(t+1) = θ^(t) - η∇L_k(θ^(t)) 3. 서버가 가중 평균 계산: θ^(t+1) = Σ(n_k/n) × θ_k^(t+1)

여기서 n_k는 클라이언트 k의 샘플 수, n은 총 샘플 수이다.

FedAvg의 장점: - 구현이 단순함 - 통신 효율적 (여러 에폭 후 한 번만 통신) - IID 데이터에서 효과적

FedAvg의 한계 - 클라이언트 드리프트:

Non-IID 환경에서 FedAvg는 심각한 문제에 직면한다. 클라이언트 드리프트의 수학적 분석:

각 클라이언트의 로컬 최적점을 θ_k*라 하면, Non-IID 데이터에서:

  • 각 θ_k는 글로벌 최적점 θ와 다름
  • 로컬 학습이 θ_k* 방향으로 진행됨
  • E 에폭이 많을수록 θ_k^(t+1)은 θ_k에 가까워지고 θ에서 멀어짐
  • 이런 편향된 업데이트들의 평균은 θ*를 향하지 않을 수 있음

이 문제를 해결하기 위해 FedProx와 SCAFFOLD가 제안되었다.

5.2 FedProx: 근접 정규화를 통한 드리프트 완화

FedProx(Federated Proximal)는 Li et al.(2020)이 제안한 알고리즘으로, 로컬 손실 함수에 근접 항(proximal term)을 추가하여 클라이언트 드리프트를 완화한다.

핵심 아이디어

로컬 모델이 글로벌 모델에서 너무 멀어지지 않도록 명시적인 제약을 건다. 이는 L2 정규화와 유사하지만, 글로벌 모델을 기준점으로 사용한다는 점이 다르다.

수학적 정의

FedProx의 로컬 목적 함수:

$$\min_\theta \mathcal{L}_k(\theta) + \frac{\mu}{2} \|\theta - \theta_{global}\|^2$$

여기서:

  • L_k(θ): 클라이언트 k의 로컬 손실 함수
  • μ: 근접 항의 가중치 (핵심 하이퍼파라미터)
  • θ: 로컬 모델 파라미터
  • θ_global: 글로벌 모델 파라미터
근접 항의 효과 분석

근접 항 (μ/2)||θ - θ_global||²의 그래디언트는 μ(θ - θ_global)이다.

즉, 로컬 업데이트에 글로벌 모델 방향으로의 "인력"이 추가된다:

  • θ가 θ_global에서 멀어지면 → 강한 복원력 발생
  • μ가 크면 → 복원력이 강해져 드리프트 억제
  • μ가 작으면 → 로컬 데이터에 더 적응
구현
if self.method == 'FedProx':
    prox_term = 0.0
    # 근접 항 계산
    for name, param in self.model.named_parameters():
        prox_term += torch.norm(param - global_params[name]) ** 2
    # 표준 손실에 근접 항 추가
    loss += (self.mu / 2) * prox_term
FedProx의 이론적 수렴 보장

Li et al.의 분석에 따르면, FedProx는 다음 조건에서 수렴이 보장된다:

  • 손실 함수가 Lipschitz 연속 그래디언트를 가짐
  • μ가 충분히 큼
  • 학습률이 적절히 작음
FedProx의 장단점

장점: 1. 구현이 단순함 (FedAvg에 몇 줄만 추가) 2. 추가 통신 오버헤드 없음 3. 중간 수준의 Non-IID에서 효과적 4. 이론적 수렴 보장 존재

단점: 1. 극단적 Non-IID에서 한계 2. μ 하이퍼파라미터 튜닝 필요 3. μ가 너무 크면 로컬 적응 능력 저하

5.3 SCAFFOLD: 제어 변량을 통한 명시적 드리프트 보정

SCAFFOLD(Stochastic Controlled Averaging for Federated Learning)는 Karimireddy et al.(2020)이 제안한 알고리즘으로, 제어 변량(control variates)을 사용하여 클라이언트 드리프트를 명시적으로 보정한다.

핵심 아이디어

각 클라이언트와 서버에서 "제어 변량"을 유지하고, 이를 사용하여 로컬 그래디언트를 조정한다. 제어 변량은 로컬 그래디언트와 글로벌 그래디언트의 차이를 추정하며, 이 차이를 보정함으로써 드리프트를 제거한다.

제어 변량의 수학적 정의
  • c: 서버의 글로벌 제어 변량 (글로벌 그래디언트의 추정치)
  • c_k: 클라이언트 k의 로컬 제어 변량 (로컬 그래디언트의 추정치)

제어 변량은 다음을 만족하도록 업데이트된다:

  • c ≈ (1/K) Σ ∇L_k(θ) (글로벌 평균 그래디언트)
  • c_k ≈ ∇L_k(θ) (로컬 그래디언트)
그래디언트 보정

SCAFFOLD에서 로컬 업데이트는 다음과 같이 보정된다:

$$\theta \leftarrow \theta - \eta (\nabla L_k(\theta) - c_k + c)$$

이 보정의 효과:

  • ∇L_k(θ) - c_k: 로컬 그래디언트에서 로컬 편향 제거
    • c: 글로벌 방향 추가
  • 결과적으로 로컬 업데이트가 글로벌 최적화 방향과 일치
제어 변량 업데이트

로컬 학습 후, 제어 변량은 다음과 같이 업데이트된다:

클라이언트 측: $$c_k^{new} = c_k - c + \frac{1}{\eta K} (\theta_{global} - \theta_k)$$

서버 측: $$c^{new} = c + \frac{1}{S} \sum_{k \in S} (c_k^{new} - c_k)$$

여기서 S는 현재 라운드에 참여한 클라이언트 집합이다.

구현

클라이언트 측 (로컬 학습):

if self.method == 'SCAFFOLD' and global_control is not None:
    for name, param in self.model.named_parameters():
        if param.grad is not None:
            # 그래디언트 보정: g - c_local + c_global
            param.grad += self.c_local[name] - global_control[name]

서버 측 (글로벌 제어 변량 업데이트):

if self.method == 'SCAFFOLD' and delta_cs is not None:
    for name in self.c_global.keys():
        for delta_c in delta_cs:
            self.c_global[name] += delta_c[name] / len(delta_cs)
SCAFFOLD의 이론적 분석

Karimireddy et al.의 분석에 따르면:

  • SCAFFOLD는 분산 감소(variance reduction) 기법의 일종
  • 데이터 이질성에 관계없이 선형 수렴 속도 달성
  • FedAvg보다 최대 O(K)배 빠른 수렴 (K: 클라이언트 수)
SCAFFOLD의 장단점

장점: 1. 극단적 Non-IID에서도 효과적 2. 이론적으로 더 강한 수렴 보장 3. 하이퍼파라미터 튜닝이 덜 필요 4. 빠른 수렴 속도

단점: 1. 제어 변량 전송으로 인한 통신 오버헤드 (약 2배) 2. 구현이 더 복잡 3. 클라이언트 메모리 사용량 증가

5.4 FedProx vs SCAFFOLD 이론적 비교

측면 FedProx SCAFFOLD
드리프트 대응 방식 간접적 (정규화) 직접적 (그래디언트 보정)
수학적 원리 근접 연산자 분산 감소
통신 복잡도 O(d) O(2d)
공간 복잡도 O(d) O(2d)
수렴 속도 (이론) O(1/√T) O(1/T)
Non-IID 강건성 중간 높음
구현 복잡성 낮음 중간
추가 하이퍼파라미터 μ 없음

6. 실험 설계

6.1 Bayesian Optimization을 활용한 하이퍼파라미터 튜닝

하이퍼파라미터 최적화를 위해 wandb(Weights & Biases)Bayesian Optimization 기능을 활용하였다. 총 20회의 sweep을 수행하여 최적의 하이퍼파라미터 조합을 탐색하였다.

Bayesian Optimization의 원리

Bayesian Optimization은 목적 함수 f(x)의 최적값을 찾기 위한 확률적 최적화 기법이다:

  1. 대리 모델(Surrogate Model): 가우시안 프로세스(GP)를 사용하여 f(x)의 사후 분포 근사 $$f(x) \sim GP(\mu(x), k(x, x'))$$

  2. 획득 함수(Acquisition Function): 다음에 평가할 점 결정 - Expected Improvement (EI): E[max(f(x) - f(x*), 0)] - Upper Confidence Bound (UCB): μ(x) + κσ(x)

  3. 반복 과정: a. 현재까지의 관측으로 GP 업데이트 b. 획득 함수 최대화하는 x_next 선택 c. f(x_next) 평가 d. 수렴할 때까지 반복

Bayesian Optimization의 장점:

  • 샘플 효율성: Grid Search(O(n^d)) 대비 훨씬 적은 평가로 좋은 결과
  • 불확실성 고려: 탐색과 활용의 균형
  • 비볼록 함수에 적합: 복잡한 하이퍼파라미터 공간에서도 효과적

6.2 실험 환경 설정

wandb Sweep 설정
sweep_config = {
    'method': 'bayes',
    'metric': {
        'name': 'test_accuracy',
        'goal': 'maximize'
    },
    'parameters': {
        'alpha': {'values': [0.1, 0.5, 1.0]},
        'epochs': {'values': [5, 10, 15]},
        'fraction': {'values': [0.2, 0.5, 0.8]},
        'iid': {'values': [True, False]},
        'learning_rate': {'values': [0.001, 0.01, 0.1]},
        'mu': {'values': [0.0, 0.1, 0.5]},
        'num_users': {'values': [5, 10, 20]}
    }
}
하이퍼파라미터 탐색 공간
파라미터 설명
alpha (α) 0.1, 0.5, 1.0 Non-IID 정도 (Dirichlet 파라미터)
epochs 5, 10, 15 로컬 학습 에폭 수
fraction 0.2, 0.5, 0.8 라운드당 참여 클라이언트 비율
iid True, False IID/Non-IID 데이터 분포
learning_rate 0.001, 0.01, 0.1 학습률
mu (μ) 0.0, 0.1, 0.5 FedProx 근접 항 가중치
num_users 5, 10, 20 총 클라이언트 수

총 조합 수: 3 × 3 × 3 × 2 × 3 × 3 × 3 = 1,458개

Bayesian Optimization을 통해 이 중 가장 유망한 20개 조합만 평가하여 효율적으로 최적값을 탐색하였다.

평가 지표
  • Test Accuracy: 테스트 세트에서의 분류 정확도 (주 지표)
  • F1 Score: 정밀도와 재현율의 조화 평균 (클래스별 성능 평가)
  • Train Loss: 학습 손실 (수렴 모니터링)

7. 실험 결과

7.1 FedProx 결과 분석

FedProx 하이퍼파라미터 스윕 정확도 곡선

위 그래프는 FedProx 알고리즘의 다양한 하이퍼파라미터 조합에 따른 테스트 정확도 곡선을 보여준다. 각 선은 서로 다른 하이퍼파라미터 설정을 나타내며, wandb의 실험 ID로 구분된다.

주요 관찰:

  • 대부분의 설정에서 5 라운드 이내에 80% 이상의 정확도에 도달
  • 일부 설정(약 25% 수준에서 정체)은 학습 실패를 나타냄
  • 최적 설정들은 약 87%의 정확도를 달성하며 안정적으로 수렴

FedProx 하이퍼파라미터 스윕 병렬 좌표

위 병렬 좌표 그래프는 각 하이퍼파라미터와 테스트 정확도 간의 관계를 시각화한다.

병렬 좌표 그래프 해석:

  • 각 세로축은 하나의 하이퍼파라미터
  • 각 선은 하나의 실험 (하이퍼파라미터 조합)
  • 선의 색상은 최종 정확도 (노란색: 높음, 보라색: 낮음)
  • 특정 값에 밝은 색 선이 집중되면 그 값이 좋은 성능과 연관됨

주요 발견:

  1. IID vs Non-IID: IID 설정(true)에서 일관되게 높은 정확도 (밝은 색 선 집중)
  2. 학습률: 0.1에서 가장 좋은 성능 (0.001은 수렴이 느림)
  3. μ 값: 0.0에서 가장 좋은 결과 (이 데이터셋에서는 강한 정규화가 불필요)
  4. 클라이언트 수: 10~20명에서 안정적인 성능

7.2 SCAFFOLD 결과 분석

SCAFFOLD 하이퍼파라미터 스윕 정확도 곡선

SCAFFOLD의 정확도 곡선은 FedProx보다 더 일관된 수렴 양상을 보인다:

  • 대부분의 설정이 비슷한 속도와 궤적으로 수렴
  • 실패하는 설정(25% 정체)이 거의 없음
  • 최종 정확도가 85~87% 범위에 밀집되어 분산이 작음

SCAFFOLD 하이퍼파라미터 스윕 병렬 좌표

SCAFFOLD의 병렬 좌표 그래프 분석:

  1. α 파라미터: 0.1(높은 이질성)에서도 밝은 색 선이 많음 → Non-IID에서 강건
  2. 에폭 수: 15 에폭에서 가장 좋은 결과
  3. 클라이언트 참여율: 0.5~0.8에서 안정적
  4. 학습률: 0.1에서 최적 성능

7.3 상세 결과표

FedProx 상위 성능 설정
Alpha Epochs Fraction IID Learning Rate Mu Users F1 Accuracy
1.0 15 0.5 True 0.1 0.0 10 0.871 87.13%
0.5 10 0.5 True 0.1 0.0 20 0.870 87.00%
0.1 10 0.8 True 0.1 0.0 20 0.868 86.77%
0.5 10 0.2 True 0.1 0.0 20 0.868 86.77%
SCAFFOLD 상위 성능 설정
Alpha Epochs Fraction IID Learning Rate Users F1 Accuracy
0.1 15 0.8 True 0.1 10 0.866 86.64%
0.1 15 0.5 True 0.1 10 0.866 86.60%
0.1 15 0.8 False 0.1 10 0.865 86.56%
0.1 15 0.8 True 0.1 10 0.866 86.64%

7.4 알고리즘 비교 종합

IID 환경 비교
알고리즘 최고 정확도 평균 정확도 표준편차 수렴 라운드
FedProx 87.13% 85.2% 2.1% 5-7
SCAFFOLD 86.64% 85.8% 0.9% 4-6

IID 환경에서는 FedProx가 약간 더 높은 최고 정확도를 달성했다. 이는 SCAFFOLD의 제어 변량 오버헤드가 IID 환경에서는 불필요하기 때문으로 해석된다.

Non-IID 환경 비교
알고리즘 α=1.0 α=0.5 α=0.1 Non-IID 평균
FedProx 86.4% 85.2% 84.8% 85.5%
SCAFFOLD 85.6% 85.9% 86.6% 86.0%

Non-IID 환경(특히 α=0.1)에서 SCAFFOLD가 FedProx보다 약 1.8%p 높은 정확도를 달성했다. 이는 SCAFFOLD의 제어 변량이 클라이언트 드리프트를 효과적으로 보정함을 보여준다.

8. 결론

이 프로젝트에서는 AG News 텍스트 분류 데이터셋을 사용하여 연합 학습 환경에서 FedProx와 SCAFFOLD 알고리즘의 성능을 비교 분석하였다.

주요 결과 요약

  1. IID 환경: 두 알고리즘 모두 약 87%의 정확도를 달성하며 비슷한 성능을 보였다. FedProx가 약간 더 높은 최고 정확도(87.13%)를 기록했다.

  2. Non-IID 환경: SCAFFOLD가 FedProx보다 더 안정적이고 일관된 성능을 보였다. 특히 극단적 Non-IID(α=0.1)에서 SCAFFOLD(86.56%)가 FedProx(84.8%)보다 약 1.7%p 높은 정확도를 달성했다.

  3. 수렴 속도: SCAFFOLD가 제어 변량을 통한 그래디언트 조정으로 더 빠르게 수렴하는 경향을 보였다.

  4. 안정성: SCAFFOLD가 하이퍼파라미터 선택에 덜 민감하고 더 안정적인 학습을 제공했다.

실용적 권장사항

  • 일반적인 경우: SCAFFOLD를 권장 (안정성과 Non-IID 강건성)
  • 리소스 제약 환경: FedProx를 권장 (낮은 통신 오버헤드)
  • 극단적 Non-IID: SCAFFOLD 필수 (클라이언트 드리프트 보정)

연합 학습은 프라이버시 보존 머신러닝의 핵심 패러다임으로, 앞으로 더욱 중요해질 것이다. 이 연구가 연합 학습 알고리즘 선택과 하이퍼파라미터 튜닝에 실질적인 지침을 제공하기를 바란다.

참고문헌

  1. Kairouz, P., McMahan, H. B., et al., "Advances and open problems in federated learning," Foundations and Trends® in Machine Learning, vol. 14, no. 1–2, pp. 1–210, 2021.

  2. McMahan, H. B., et al., "Communication-efficient learning of deep networks from decentralized data," AISTATS, pp. 1273–1282, 2017.

  3. Li, T., et al., "Federated learning: Challenges, methods, and future directions," IEEE Signal Processing Magazine, vol. 37, no. 3, pp. 50–60, 2020.

  4. Karimireddy, S. P., et al., "SCAFFOLD: Stochastic controlled averaging for federated learning," ICML, pp. 5132–5143, 2020.

  5. Li, T., et al., "Fair resource allocation in federated learning," ICLR, 2020.

  6. Zhao, Y., et al., "Federated learning with non-IID data," arXiv preprint arXiv:1806.00582, 2018.

9. 심화 분석

9.1 하이퍼파라미터별 상세 분석

학습률(Learning Rate)의 영향

실험 결과, 학습률은 연합 학습의 성능에 매우 큰 영향을 미쳤다:

학습률 0.1: - FedProx: 평균 정확도 86.2%, 최고 87.13% - SCAFFOLD: 평균 정확도 85.9%, 최고 86.64% - 대부분의 실험에서 안정적으로 수렴 - 5~7 라운드 내에 최종 성능 도달

학습률 0.01: - FedProx: 평균 정확도 83.5%, 최고 86.39% - SCAFFOLD: 평균 정확도 85.2%, 최고 85.61% - 수렴이 느리지만 안정적 - 10~15 라운드 필요

학습률 0.001: - FedProx: 평균 정확도 45.2%, 최고 65.69% - SCAFFOLD: 평균 정확도 42.1%, 최고 41.65% - 수렴이 매우 느려 15 라운드 내에 수렴 실패 - 많은 실험에서 학습 실패

분석: AG News 데이터셋과 GRU 모델 조합에서는 상대적으로 높은 학습률(0.1)이 효과적이었다. 이는 모델이 비교적 단순하고 데이터셋이 명확한 분류 경계를 가지기 때문으로 보인다. 더 복잡한 모델이나 어려운 데이터셋에서는 더 낮은 학습률이 필요할 수 있다.

μ(FedProx 근접 항)의 영향

FedProx의 핵심 하이퍼파라미터인 μ에 대한 상세 분석:

μ = 0.0 (사실상 FedAvg): - IID: 평균 86.5%, 최고 87.13% - Non-IID: 평균 84.8%, 최고 86.77% - 근접 항이 없어도 좋은 성능

μ = 0.1: - IID: 평균 75.2%, 최고 81.96% - Non-IID: 평균 70.5%, 최고 74.4% - 성능 저하 관찰

μ = 0.5: - IID: 평균 27.8%, 최고 28.68% - Non-IID: 평균 26.2%, 최고 27.97% - 학습 거의 실패 (4클래스에서 랜덤 = 25%)

분석: 예상과 달리 μ > 0에서 성능이 저하되었다. 이는 다음과 같은 이유로 해석된다:

  1. 데이터 특성: AG News는 클래스 간 분리가 명확하여 드리프트가 심하지 않음
  2. 모델 용량: GRU 모델이 상대적으로 단순하여 과도한 정규화가 해로움
  3. 학습 동역학: 높은 μ가 로컬 적응을 과도하게 억제하여 학습 방해

실제 적용 시에는 μ 값을 0.001~0.1 범위에서 세밀하게 탐색할 필요가 있다.

클라이언트 수(num_users)의 영향

5명의 클라이언트: - FedProx: 평균 83.2%, 표준편차 3.1% - SCAFFOLD: 평균 85.1%, 표준편차 1.8% - 각 클라이언트가 많은 데이터를 보유 (24,000개/클라이언트) - 로컬 학습이 효과적이나 다양성 부족

10명의 클라이언트: - FedProx: 평균 85.8%, 표준편차 2.4% - SCAFFOLD: 평균 86.1%, 표준편차 0.9% - 적절한 균형

20명의 클라이언트: - FedProx: 평균 84.5%, 표준편차 3.2% - SCAFFOLD: 평균 85.5%, 표준편차 1.2% - 다양성 증가하나 각 클라이언트 데이터 감소 (6,000개/클라이언트)

분석: 10명의 클라이언트에서 가장 좋은 균형을 보였다. 클라이언트가 너무 적으면 집계의 이점이 줄어들고, 너무 많으면 각 클라이언트의 데이터가 부족해져 로컬 학습이 불안정해진다.

클라이언트 참여율(fraction)의 영향

20% 참여 (fraction = 0.2): - FedProx: 평균 84.1%, 변동성 높음 - SCAFFOLD: 평균 85.2%, 상대적으로 안정적 - 라운드당 소수 클라이언트만 참여 - 통신 비용 절감, 수렴 불안정

50% 참여 (fraction = 0.5): - FedProx: 평균 85.5%, 적절한 변동성 - SCAFFOLD: 평균 85.8%, 안정적 - 가장 균형 잡힌 설정

80% 참여 (fraction = 0.8): - FedProx: 평균 85.2%, 안정적 - SCAFFOLD: 평균 86.0%, 매우 안정적 - 높은 통신 비용, 안정적인 수렴

분석: SCAFFOLD는 낮은 참여율에서도 안정적인 성능을 유지했다. 이는 제어 변량이 참여하지 않는 클라이언트의 정보도 간접적으로 반영하기 때문이다. FedProx는 참여율에 더 민감하게 반응했다.

9.2 수렴 행동 분석

학습 곡선 패턴

실험에서 관찰된 주요 학습 곡선 패턴:

패턴 1: 빠른 수렴 (성공적 학습) - 특징: 1~3 라운드에서 급격한 성능 향상, 이후 안정화 - 조건: 높은 학습률(0.1), 적절한 로컬 에폭(10-15) - FedProx와 SCAFFOLD 모두에서 관찰

패턴 2: 점진적 수렴 (느린 학습) - 특징: 10+ 라운드에 걸쳐 서서히 성능 향상 - 조건: 낮은 학습률(0.01), 적은 로컬 에폭(5) - 최종 성능은 양호하나 통신 비용 증가

패턴 3: 정체 (학습 실패) - 특징: 25% 근처(랜덤 성능)에서 정체 - 조건: 매우 낮은 학습률(0.001), 높은 μ(0.5) - 그래디언트가 너무 작거나 과도한 정규화

패턴 4: 진동 - 특징: 성능이 오르내림 반복 - 조건: Non-IID + 낮은 참여율 - FedProx에서 더 자주 관찰, SCAFFOLD는 안정적

조기 종료 지점 분석

실험 데이터를 기반으로 한 조기 종료 권장:

설정 권장 라운드 수 이유
IID, lr=0.1 7 5라운드에 90% 성능 도달
Non-IID, lr=0.1 10 추가 수렴 필요
SCAFFOLD 6 FedProx보다 빠른 수렴
FedProx (μ=0) 8 표준 속도

9.3 통계적 유의성 검정

동일 설정에서의 반복 실험 결과를 바탕으로 t-검정 수행:

FedProx vs SCAFFOLD (Non-IID, α=0.1): - FedProx 평균: 84.8%, 표준편차: 1.2% - SCAFFOLD 평균: 86.5%, 표준편차: 0.8% - t-통계량: 3.21, p-값: 0.012 - 결론: SCAFFOLD가 통계적으로 유의하게 우수 (p < 0.05)

FedProx vs SCAFFOLD (IID): - FedProx 평균: 86.5%, 표준편차: 1.5% - SCAFFOLD 평균: 86.2%, 표준편차: 0.9% - t-통계량: 0.52, p-값: 0.61 - 결론: 유의한 차이 없음 (p > 0.05)

9.4 실패 사례 분석

FedProx 학습 실패 사례

사례 1: μ=0.5, lr=0.01, Non-IID (α=0.1) - 증상: 정확도 25%에서 정체 - 원인 분석: - 높은 μ로 인해 로컬 학습이 과도하게 제한됨 - 낮은 학습률로 그래디언트가 작음 - Non-IID 데이터로 인한 추가 어려움 - 해결책: μ를 0.01 이하로 낮추거나 학습률 증가

사례 2: lr=0.001, epochs=5, 20 users - 증상: 정확도 40~50%에서 느린 개선 - 원인 분석: - 학습률이 너무 낮아 수렴이 느림 - 짧은 로컬 에폭으로 각 라운드 업데이트가 작음 - 많은 사용자로 인해 각 클라이언트 데이터 부족 - 해결책: 학습률 증가 또는 로컬 에폭 증가

SCAFFOLD 학습 실패 사례

SCAFFOLD는 학습 실패 사례가 거의 없었으나, 하나의 예외적 케이스:

사례: lr=0.001, fraction=0.2, α=0.1 - 증상: 정확도 41.65%에서 느린 개선 - 원인 분석: - 극도로 낮은 학습률 - 극도로 낮은 참여율 - 제어 변량 업데이트가 불충분 - 해결책: 학습률 증가 필요

9.5 계산 비용 분석

시간 복잡도

한 라운드당 계산 비용:

FedProx: - 클라이언트: O(E × n_k × d) - E 에폭, n_k 샘플, d 파라미터 - 서버: O(m × d) - m 참여 클라이언트 - 근접 항 계산: O(d) 추가

SCAFFOLD: - 클라이언트: O(E × n_k × d) + O(d) 제어 변량 보정 - 서버: O(m × d) + O(d) 제어 변량 업데이트 - 총 오버헤드: 약 5~10% 증가

공간 복잡도

FedProx: - 클라이언트: O(d) - 모델 파라미터 - 서버: O(d) - 글로벌 모델

SCAFFOLD: - 클라이언트: O(2d) - 모델 + 로컬 제어 변량 - 서버: O(2d) - 글로벌 모델 + 글로벌 제어 변량

통신 비용

FedProx: 라운드당 O(m × d) 파라미터 전송 SCAFFOLD: 라운드당 O(m × 2d) 파라미터 + 제어 변량 전송

SCAFFOLD의 통신 비용이 약 2배이나, 더 빠른 수렴으로 총 통신량은 유사하거나 적을 수 있다.

10. 향후 연구 방향

10.1 알고리즘 확장

  1. FedProx + SCAFFOLD 하이브리드: 두 알고리즘의 장점 결합 - 근접 항으로 안정성 확보 - 제어 변량으로 드리프트 보정 - 통신 오버헤드와 성능의 균형

  2. 적응형 μ 스케줄링: FedProx의 μ를 학습 과정에서 동적 조정 - 초기: 작은 μ로 빠른 적응 - 후기: 큰 μ로 안정적 수렴

  3. 압축 기법 적용: 통신 효율성 향상 - 그래디언트 양자화 - 희소화(Sparsification) - Top-k 선택

10.2 다른 도메인으로의 확장

  1. 이미지 분류: CIFAR-10, ImageNet 데이터셋
  2. 자연어 이해: 감성 분석, 질의응답
  3. 시계열 예측: 금융 데이터, 센서 데이터
  4. 추천 시스템: 개인화된 추천

10.3 실제 배포 고려사항

  1. 디바이스 이질성: 다양한 하드웨어 성능 대응
  2. 네트워크 불안정성: 연결 끊김, 지연 시간 처리
  3. 프라이버시 강화: 차분 프라이버시, 동형 암호화 적용
  4. 악의적 클라이언트 방어: Byzantine-robust 집계 기법

11. 최종 결론

이 프로젝트를 통해 연합 학습의 두 가지 대표적인 알고리즘인 FedProx와 SCAFFOLD를 체계적으로 비교 분석하였다. 주요 결론은 다음과 같다:

알고리즘 선택 가이드

  1. 데이터 분포가 IID에 가깝다면: FedProx (μ=0, 사실상 FedAvg) 추천 - 단순하고 효율적 - 추가 오버헤드 없음

  2. 중간 수준의 Non-IID라면: FedProx (μ=0.001~0.01) 추천 - 약간의 정규화로 안정성 확보 - SCAFFOLD보다 통신 효율적

  3. 극단적 Non-IID라면: SCAFFOLD 필수 - 제어 변량으로 드리프트 명시적 보정 - 약간의 통신 오버헤드 감수

하이퍼파라미터 권장값

파라미터 권장값 이유
학습률 0.05~0.1 빠른 수렴, 안정성
로컬 에폭 10~15 충분한 로컬 학습
참여율 0.3~0.5 통신-성능 균형
클라이언트 수 10~20 다양성-데이터량 균형
μ (FedProx) 0~0.01 과도한 정규화 방지

연구 기여

  1. 텍스트 분류 도메인에서의 연합 학습 벤치마크 제공
  2. Dirichlet 분포 기반 Non-IID 데이터 생성 방법 구현
  3. Bayesian Optimization을 통한 체계적인 하이퍼파라미터 탐색 수행
  4. FedProx와 SCAFFOLD의 실증적 비교 분석 제시
  5. 실용적인 알고리즘 선택 가이드라인 도출

연합 학습은 프라이버시 보존 머신러닝의 핵심 기술로서, 의료, 금융, 모바일 서비스 등 다양한 분야에서 점점 더 중요해지고 있다. 이 연구가 연합 학습을 실제로 적용하고자 하는 연구자와 엔지니어들에게 유용한 지침을 제공하기를 기대한다.

12. 구현 세부사항

12.1 전체 시스템 아키텍처

연합 학습 시스템은 다음과 같은 구성 요소로 이루어진다:

서버 컴포넌트
class FederatedServer:
    def __init__(self, model, method='FedAvg'):
        self.global_model = model
        self.method = method
        if method == 'SCAFFOLD':
            self.c_global = {name: torch.zeros_like(param) 
                           for name, param in model.named_parameters()}

    def aggregate(self, client_updates, sample_counts):
        """클라이언트 업데이트를 가중 평균으로 집계"""
        total_samples = sum(sample_counts)

        # 가중 평균 계산
        new_params = {}
        for name in self.global_model.state_dict().keys():
            new_params[name] = sum(
                client_updates[i][name] * (sample_counts[i] / total_samples)
                for i in range(len(client_updates))
            )

        self.global_model.load_state_dict(new_params)
        return self.global_model.state_dict()

    def update_control_variates(self, delta_cs):
        """SCAFFOLD: 글로벌 제어 변량 업데이트"""
        if self.method != 'SCAFFOLD':
            return

        for name in self.c_global.keys():
            for delta_c in delta_cs:
                self.c_global[name] += delta_c[name] / len(delta_cs)
클라이언트 컴포넌트
class FederatedClient:
    def __init__(self, client_id, dataset, model, method='FedAvg', mu=0.0):
        self.client_id = client_id
        self.dataset = dataset
        self.model = copy.deepcopy(model)
        self.method = method
        self.mu = mu

        if method == 'SCAFFOLD':
            self.c_local = {name: torch.zeros_like(param) 
                          for name, param in model.named_parameters()}

    def local_train(self, global_params, global_control=None, epochs=1, lr=0.01):
        """로컬 학습 수행"""
        self.model.load_state_dict(global_params)
        optimizer = torch.optim.SGD(self.model.parameters(), lr=lr)
        criterion = nn.CrossEntropyLoss()

        dataloader = DataLoader(self.dataset, batch_size=16, shuffle=True)

        for epoch in range(epochs):
            for batch in dataloader:
                input_ids = batch['input_ids']
                labels = batch['labels']

                optimizer.zero_grad()
                outputs = self.model(input_ids)
                loss = criterion(outputs.view(-1, outputs.size(-1)), labels.view(-1))

                # FedProx: 근접 항 추가
                if self.method == 'FedProx' and self.mu > 0:
                    prox_term = 0.0
                    for name, param in self.model.named_parameters():
                        prox_term += torch.norm(param - global_params[name]) ** 2
                    loss += (self.mu / 2) * prox_term

                loss.backward()

                # SCAFFOLD: 그래디언트 보정
                if self.method == 'SCAFFOLD' and global_control is not None:
                    for name, param in self.model.named_parameters():
                        if param.grad is not None:
                            param.grad += self.c_local[name] - global_control[name]

                optimizer.step()

        return self.model.state_dict()

12.2 wandb 통합

실험 추적 및 하이퍼파라미터 최적화를 위한 wandb 통합:

import wandb

def run_experiment(config=None):
    with wandb.init(config=config):
        config = wandb.config

        # 데이터 분할
        if config.iid:
            client_datasets = split_data_iid(train_dataset, config.num_users)
        else:
            client_datasets = split_data_non_iid(train_dataset, config.num_users, config.alpha)

        # 서버 및 클라이언트 초기화
        server = FederatedServer(model, method=config.method)
        clients = [
            FederatedClient(i, client_datasets[i], model, 
                          method=config.method, mu=config.mu)
            for i in range(config.num_users)
        ]

        # 학습 루프
        for round in range(config.num_rounds):
            # 클라이언트 선택
            selected = np.random.choice(
                config.num_users, 
                int(config.num_users * config.fraction), 
                replace=False
            )

            # 로컬 학습
            client_updates = []
            sample_counts = []
            for i in selected:
                update = clients[i].local_train(
                    server.global_model.state_dict(),
                    global_control=server.c_global if config.method == 'SCAFFOLD' else None,
                    epochs=config.epochs,
                    lr=config.learning_rate
                )
                client_updates.append(update)
                sample_counts.append(len(clients[i].dataset))

            # 집계
            server.aggregate(client_updates, sample_counts)

            # 평가
            accuracy = evaluate(server.global_model, test_dataset)
            wandb.log({'test_accuracy': accuracy, 'round': round})

        return accuracy

# Sweep 실행
sweep_id <mark class="highlight"><strong><u> wandb.sweep(sweep_config, project</u></strong></mark>'federated-learning')
wandb.agent(sweep_id, run_experiment, count=20)

12.3 재현성 보장

실험의 재현성을 위해 다음 조치를 취하였다:

def set_seed(seed=42):
    """모든 무작위성 소스에 시드 설정"""
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    torch.backends.cudnn.deterministic = True
    torch.backends.cudnn.benchmark = False

또한 각 실험의 설정과 결과를 wandb에 기록하여 언제든 재현할 수 있도록 하였다.

12.4 메모리 최적화

연합 학습에서 메모리 효율성을 위한 기법들:

  1. 모델 복사 최소화: 필요할 때만 깊은 복사 수행
  2. 그래디언트 누적: 작은 배치로 큰 효과적 배치 크기 달성
  3. 혼합 정밀도 학습: FP16 사용으로 메모리 절감 (NVIDIA GPU)
# 혼합 정밀도 학습 예시
from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler()

with autocast():
    outputs = model(input_ids)
    loss = criterion(outputs, labels)

scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()

13. 실험 환경

13.1 하드웨어

  • CPU: AMD EPYC 7742 64-Core Processor
  • RAM: 256GB DDR4
  • GPU: NVIDIA A100 40GB × 4
  • Storage: NVMe SSD 2TB

13.2 소프트웨어

소프트웨어 버전
Python 3.9.16
PyTorch 2.0.1
Transformers 4.30.0
NumPy 1.24.3
wandb 0.15.4
CUDA 11.8

13.3 실험 시간

설정 1 라운드 시간 총 실험 시간
5 클라이언트, 5 에폭 ~30초 ~8분
10 클라이언트, 10 에폭 ~1분 ~15분
20 클라이언트, 15 에폭 ~3분 ~45분

전체 sweep (20회 실험): 약 6시간

14. 코드 및 데이터 가용성

14.1 프로젝트 구조

federated-learning-benchmark/
├── data/
│   ├── ag_news/
│   │   ├── train.csv
│   │   └── test.csv
│   └── preprocess.py
├── models/
│   ├── gru.py
│   └── __init__.py
├── fl/
│   ├── server.py
│   ├── client.py
│   ├── fedprox.py
│   └── scaffold.py
├── utils/
│   ├── data_split.py
│   ├── metrics.py
│   └── seed.py
├── configs/
│   └── sweep_config.yaml
├── experiments/
│   └── run_sweep.py
├── notebooks/
│   └── analysis.ipynb
├── requirements.txt
└── README.md

14.2 핵심 코드 스니펫

IID 데이터 분할:

def split_data_iid(dataset, num_users):
    num_items = len(dataset) // num_users
    indices = np.random.permutation(len(dataset))
    return [Subset(dataset, indices[i*num_items:(i+1)*num_items]) 
            for i in range(num_users)]

Non-IID 데이터 분할 (Dirichlet):

def split_data_non_iid(dataset, num_users, alpha):
    labels = np.array([dataset[i][1] for i in range(len(dataset))])
    num_classes = len(np.unique(labels))

    # 각 클래스별로 Dirichlet 분포에 따라 분할
    client_indices = [[] for _ in range(num_users)]
    for c in range(num_classes):
        class_indices = np.where(labels == c)[0]
        np.random.shuffle(class_indices)

        proportions = np.random.dirichlet([alpha] * num_users)
        proportions = (np.cumsum(proportions) * len(class_indices)).astype(int)[:-1]

        for i, chunk in enumerate(np.split(class_indices, proportions)):
            client_indices[i].extend(chunk.tolist())

    return [Subset(dataset, indices) for indices in client_indices]

이 프로젝트의 전체 코드와 실험 결과는 학술 연구 목적으로 활용될 수 있다.

비슷한 글 추천

Comments (0)

No comments yet. Be the first to comment!