🤖 AI/ML

Continual Learning

새로운 지식을 지속적으로 학습하면서 기존 지식을 유지하는 학습

📖 상세 설명

Continual Learning(지속 학습)은 AI 모델이 새로운 데이터와 태스크를 순차적으로 학습하면서도 이전에 학습한 지식을 잊지 않도록 하는 기계학습 패러다임입니다. Lifelong Learning, Incremental Learning이라고도 불리며, 실제 세계의 지속적인 변화에 적응할 수 있는 AI 시스템 구축의 핵심 기술입니다.

전통적인 딥러닝에서 새로운 데이터로 학습하면 기존 지식이 급격히 손실되는 "파국적 망각(Catastrophic Forgetting)" 현상이 발생합니다. 1989년 McCloskey와 Cohen이 이 문제를 처음 발견한 이후, 신경망의 근본적 한계로 여겨져 왔으며, 이를 해결하기 위한 다양한 연구가 진행 중입니다.

주요 접근법은 세 가지입니다. Replay 방식은 과거 데이터를 저장하거나 생성하여 새 데이터와 함께 학습합니다. Regularization 방식(EWC, SI)은 중요한 파라미터의 변화를 제한합니다. Architecture 방식은 새 태스크마다 네트워크를 확장하거나 분리된 서브넷을 할당합니다.

실무에서 Continual Learning은 추천 시스템의 실시간 업데이트, 사기 탐지 모델의 새로운 패턴 적응, 자율주행의 새로운 환경 학습, 의료 AI의 새 질병 데이터 통합 등에 적용됩니다. 특히 데이터 프라이버시로 과거 데이터 저장이 어려운 환경에서 중요하며, LLM의 지식 업데이트 문제와도 연결됩니다.

💻 코드 예제

import torch
import torch.nn as nn
from copy import deepcopy

class EWC:
    """Elastic Weight Consolidation - Continual Learning 기법"""

    def __init__(self, model: nn.Module, dataloader, device, lambda_ewc: float = 1000):
        self.model = model
        self.lambda_ewc = lambda_ewc
        self.device = device

        # 이전 태스크 학습 후 파라미터 저장
        self.params_old = {n: p.clone().detach()
                          for n, p in model.named_parameters() if p.requires_grad}

        # Fisher Information Matrix 계산 (파라미터 중요도)
        self.fisher = self._compute_fisher(dataloader)

    def _compute_fisher(self, dataloader) -> dict:
        """Fisher Information 근사 계산"""
        fisher = {n: torch.zeros_like(p)
                  for n, p in self.model.named_parameters() if p.requires_grad}

        self.model.eval()
        for data, target in dataloader:
            data, target = data.to(self.device), target.to(self.device)
            self.model.zero_grad()
            output = self.model(data)
            loss = nn.functional.cross_entropy(output, target)
            loss.backward()

            for n, p in self.model.named_parameters():
                if p.requires_grad and p.grad is not None:
                    fisher[n] += p.grad.data ** 2

        # 평균화
        for n in fisher:
            fisher[n] /= len(dataloader)

        return fisher

    def penalty(self) -> torch.Tensor:
        """EWC 패널티 항 계산 - 기존 손실에 추가"""
        loss = 0
        for n, p in self.model.named_parameters():
            if p.requires_grad:
                # 중요한 파라미터일수록(Fisher 높음) 변화에 큰 페널티
                loss += (self.fisher[n] * (p - self.params_old[n]) ** 2).sum()

        return self.lambda_ewc * loss

# 사용 예시
def train_with_ewc(model, new_task_loader, ewc=None, epochs=10):
    optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

    for epoch in range(epochs):
        for data, target in new_task_loader:
            optimizer.zero_grad()
            output = model(data)
            loss = nn.functional.cross_entropy(output, target)

            # EWC 패널티 추가 (이전 태스크 지식 보존)
            if ewc is not None:
                loss += ewc.penalty()

            loss.backward()
            optimizer.step()

    return model

# Task 1 학습 후 EWC 초기화
# ewc = EWC(model, task1_loader, device)
# Task 2 학습 (이전 지식 보존하며)
# model = train_with_ewc(model, task2_loader, ewc)

🗣️ 실무 대화 예시

ML 시스템 설계 미팅에서

"추천 모델을 매일 전체 재학습하면 비용이 너무 커요. Continual Learning으로 새 데이터만 학습하되, EWC나 Replay 버퍼로 catastrophic forgetting을 막으면 학습 시간을 80% 줄일 수 있습니다."

기술 면접에서

"Continual Learning의 핵심 도전은 stability-plasticity 딜레마입니다. 새 지식을 잘 배우려면(plasticity) 파라미터가 많이 바뀌어야 하는데, 기존 지식을 유지하려면(stability) 파라미터를 덜 바꿔야 하죠. EWC는 중요도 기반으로 이 균형을 조절합니다."

데이터 규정 논의에서

"GDPR 때문에 고객 데이터를 오래 보관하기 어렵습니다. Generative Replay 방식으로 합성 데이터를 만들어 학습하면, 원본 없이도 과거 지식을 유지할 수 있어요. 개인정보 이슈도 피할 수 있고요."

⚠️ 주의사항

1
태스크 간 간섭 관리

새 태스크와 기존 태스크가 충돌하면 양쪽 성능이 모두 저하될 수 있습니다. 태스크 유사도를 분석하고, 심하게 다른 태스크는 별도 모델을 고려하세요.

2
성능 모니터링 필수

이전 태스크 성능을 지속적으로 측정해야 망각을 조기 탐지할 수 있습니다. 벤치마크 테스트셋을 유지하고 정기적으로 평가하세요.

3
메모리 vs 성능 트레이드오프

Replay 방식은 과거 데이터 저장이 필요하고, Architecture 방식은 모델 크기가 증가합니다. 시스템 제약에 맞는 기법을 선택하세요.

🔗 관련 용어

📚 더 배우기