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)
"추천 모델을 매일 전체 재학습하면 비용이 너무 커요. Continual Learning으로 새 데이터만 학습하되, EWC나 Replay 버퍼로 catastrophic forgetting을 막으면 학습 시간을 80% 줄일 수 있습니다."
"Continual Learning의 핵심 도전은 stability-plasticity 딜레마입니다. 새 지식을 잘 배우려면(plasticity) 파라미터가 많이 바뀌어야 하는데, 기존 지식을 유지하려면(stability) 파라미터를 덜 바꿔야 하죠. EWC는 중요도 기반으로 이 균형을 조절합니다."
"GDPR 때문에 고객 데이터를 오래 보관하기 어렵습니다. Generative Replay 방식으로 합성 데이터를 만들어 학습하면, 원본 없이도 과거 지식을 유지할 수 있어요. 개인정보 이슈도 피할 수 있고요."
새 태스크와 기존 태스크가 충돌하면 양쪽 성능이 모두 저하될 수 있습니다. 태스크 유사도를 분석하고, 심하게 다른 태스크는 별도 모델을 고려하세요.
이전 태스크 성능을 지속적으로 측정해야 망각을 조기 탐지할 수 있습니다. 벤치마크 테스트셋을 유지하고 정기적으로 평가하세요.
Replay 방식은 과거 데이터 저장이 필요하고, Architecture 방식은 모델 크기가 증가합니다. 시스템 제약에 맞는 기법을 선택하세요.