연속 학습에서 망각을 제어하는 모델 운영 설계
Continual Learning과 Lifelong Learning에서 Catastrophic Forgetting을 줄이기 위한 정규화, 리플레이, 격리 전략과 운영 설계 기준
2026-08-14 · 최초 발행 2024-04-29
새 지식이 기존 모델을 덮어쓰지 않게 하려면
Continual Learning(CL)과 Lifelong Learning(LLL)은 시간에 따라 달라지는 태스크, 도메인, 클래스 분포를 따라 모델을 순차적으로 학습하는 패러다임이다. 새 데이터를 반영하면서도 이미 배운 지식을 유지해야 하므로, 지식 보존과 신규 지식 습득 사이의 균형이 중심 과제가 된다.
이때 신규 데이터 학습 뒤 이전 태스크의 성능이 크게 떨어지는 현상을 Catastrophic Forgetting이라 한다. 파라미터를 공유하는 모델에서 업데이트가 충돌하거나, 분포가 변하고, 메모리 예산이 제한될 때 발생한다.
문제의 형태는 태스크 식별 정보와 클래스 변화 여부에 따라 나뉜다.
- Task-Incremental: 태스크 ID가 제공되며 태스크별 헤드를 분리할 수 있다.
- Domain-Incremental: 태스크 ID는 제공되지 않고 클래스는 같지만 분포가 달라진다.
- Class-Incremental: 태스크 ID 없이 클래스 집합이 확장되며 난이도가 높다.
평가에서는 최종 평균 정확도인 ACC, 과거 성능 변화를 나타내는 BWT(Backward Transfer), 사전 전이 이득을 보는 FWT(Forward Transfer)를 사용한다. 메모리 예산(MB/샘플 수), 업데이트 지연(latency), 재학습 비용(연산량)도 함께 모니터링해야 한다.
망각을 줄이는 전략과 운영 구성
가중치 변화를 제한하는 방식은 중요한 파라미터가 급격히 이동하지 않도록 한다. EWC, SI, MAS는 중요 파라미터의 변화를 억제하고, LwF는 지식 증류를 이용한다. 기존 태스크의 Fisher/importance를 바탕으로 가중치 이동에 페널티를 부여하는 접근이다.
리플레이는 샘플 버퍼, 요약 샘플, GAN/VAE 같은 생성 모델로 과거 분포를 다시 학습에 포함한다. 이 경우 클래스·태스크 균형을 지키는 일과 프라이버시 제약을 함께 고려해야 한다.
파라미터 자체를 분리하는 방법도 있다. PackNet의 태스크별 마스크·프루닝, Prefix/Prompt Tuning의 어댑터·프로브·프롬프트 삽입은 태스크 사이의 간섭을 줄이고 빠른 전개를 돕는다. Progressive Networks나 모듈 추가처럼 용량을 늘리는 방식은 점진적 확장이 가능하지만 리소스 상한과 지연 요구를 기준으로 동적으로 판단해야 한다.
태스크 전환과 분포 변화도 학습 전략에 직접 영향을 준다. ADWIN, CUSUM, PSI 등으로 분포의 변곡점을 감지하고, 태스크·클래스 증가를 추정하는 로직을 둔다. 감지 결과에 따라 학습률을 단계화하고, 정규화 강도를 조절하며, 버퍼를 리밸런싱하거나 파라미터 동결·해제 규칙을 적용한다.
버퍼는 reservoir/knapsack 샘플링과 클래스별 균형 유지로 관리한다. coreset, herding처럼 정보량을 높이는 요약도 사용할 수 있다. 데이터 주권이나 프라이버시 때문에 원본 데이터를 보관할 수 없다면 생성 리플레이나 피처 리플레이가 대안이 된다.
운영 측면에서는 모델 레지스트리와 버전 관리, 온라인 A/B, 롤백, 게이트드 릴리즈가 필요하다. 성능 지표와 망각 지표를 동시에 보며, 체크포인트와 Fisher/importance 스냅샷, 버퍼 메타데이터까지 감사 추적이 가능해야 한다.
데이터 스트림에서 릴리스까지의 흐름
입력에는 데이터 스트림의 샘플과 라벨 또는 약라벨, 태스크·드리프트 시그널, 메모리 예산과 정규화 강도 같은 정책 파라미터가 포함된다.
처리 단계에서는 먼저 분포 변화를 감지하고 태스크 상태를 추정한다. 이어 EWC/SI 같은 정규화, 버퍼·생성 기반 리플레이, 격리·어댑터 중 전략을 선택한다. 증분 학습에서는 신규 데이터와 버퍼 데이터를 미니배치로 섞고, 중요도에 따라 가중치 락을 적용하거나 완화하며 기준셋으로 검증한다. 이후 성능과 망각 지표를 평가해 하이퍼파라미터를 다시 조정한다.
출력은 업데이트된 모델 아티팩트, Fisher/importance 스냅샷, 리플레이 버퍼 상태, 지표 리포트다.
급격한 성능 하락이 발생하면 정규화 강도를 높이고 학습률을 낮추며 버퍼를 리밸런싱하거나 직전 버전으로 롤백한다. 메모리가 초과되면 reservoir 확률을 높이거나 요약 샘플링으로 전환하고 피처 리플레이를 사용한다. 태스크 경계가 분명하지 않다면 비지도 드리프트 감지를 강화하고 임시 어댑터 계층을 격리한다.
중요 파라미터에는 soft lock(EWC/SI penalty)을 적용하고 비중요 파라미터에 업데이트를 집중한다. 옵티마이저 상태인 모멘텀도 버전과 동기화하며, 레지스트리 트랜잭션이 커밋된 뒤 트래픽을 전환한다.
제약 조건에 따른 방법 선택
| 방법 계열 | 성능(망각 억제) | 확장성 | 일관성(지식 보존) | 안정성(학습 변동) | 운영 편의 |
|---|---|---|---|---|---|
| 정규화(EWC/SI/LwF) | 중~상 | 상 | 중~상 | 상 | 상 |
| 리플레이(버퍼/생성) | 상 | 중 | 상 | 중 | 중 |
| 파라미터 격리(마스크/프루닝) | 중 | 중 | 상 | 상 | 중 |
| 동적 확장(Progressive/모듈) | 상 | 중~하 | 상 | 상 | 하 |
| 프롬프트/어댑터(PLM) | 중 | 상 | 중 | 중 | 상 |
데이터 프라이버시 제약이 있으면 생성·피처 리플레이가 유리하다. 온라인 지연 제약이 강하면 정규화와 어댑터 방식을 우선할 수 있다.
운영 환경에서 만나는 적용 장면
추천 시스템의 일간 업데이트에서는 신상품과 계절 변화를 반영하면서 버퍼 기반 리플레이로 과거 선호를 유지한다. ACC를 유지하고 업데이트 시간을 줄이는 것이 목표다.
제조 결함 탐지에서는 신규 라인이 추가될 때 도메인 증분 학습을 적용할 수 있다. SI와 소량 버퍼를 함께 사용해 기존 라인 성능을 유지한다.
보안 침입 탐지에서는 새로운 공격 패턴이 나타날 때 빠른 증분 업데이트가 필요하다. Drift 감지 뒤 정규화를 강화하고 버퍼를 리샘플링하는 파이프라인으로 연결할 수 있다.
고객 서비스 챗봇은 어댑터나 프롬프트로 신규 의도를 확장하면서 원 모델을 고정할 수 있다. 이 방식은 롤백과 규정 준수에 유리하다.
의료 영상처럼 프라이버시 제약으로 원본 리플레이가 금지된 환경에서는 생성 리플레이와 특징 리플레이를 적용하고 모델과 증거의 추적을 강화한다.
목표 지표와 설계 기준
전량 재학습과 비교하면 학습 비용을 30~60% 절감하고 업데이트 지연을 분 단위 수준으로 달성할 수 있다. 성능 유지 목표는 평균 BWT -3pp 이내, 최종 ACC 기존 대비 ±2pp 내 안정화로 둘 수 있다.
버퍼 예산은 전체 데이터의 0.5~2% 수준으로 유지하면서 스토리지·네트워크 비용을 줄인다. 배포와 롤백은 블루/그린 방식으로 무중단 전환을 구성하고, 위험 구간에는 자동 경보를 연결한다.
언어·비전 대규모 사전학습 모델에는 어댑터나 프롬프트를 우선 적용하고, 필요할 때 소량 리플레이를 병행한다. 경량 엣지 모델은 EWC/SI 같은 정규화 기반 방식을 우선하며 메모리 예산 안에서 소형 버퍼를 둔다. 규정 또는 프라이버시 제한이 있다면 생성·피처 리플레이와 LwF 기반 지식 증류를 강화한다.
정규화 강도(lambda)는 101000 범위에서 로그스윕하고 태스크 난이도와 버퍼 크기를 교차 튜닝한다. 버퍼는 클래스당 20200 샘플을 두고 reservoir sampling을 사용할 수 있다. 신규 데이터와 버퍼 데이터의 혼합 비율은 3:1~5:1로 잡되 분포 이동 크기에 따라 동적으로 조정한다.
온라인에서는 최근 윈도우 ACC, BWT 추정치, drift score, 업데이트 지연을 관찰한다. ACC 하락이 >3pp이거나 drift score 상위 1% 상태가 지속되면 재학습 모드를 강화하는 경보 기준을 둘 수 있다.
PyTorch로 구현한 리플레이와 EWC
환경: Python 3.10+, PyTorch 2.x, CUDA 선택.
import torch, random
from torch import nn, optim
from collections import deque, defaultdict
class Net(nn.Module):
def __init__(self, d_in=784, n_cls=10):
super().__init__()
self.net = nn.Sequential(nn.Linear(d_in, 256), nn.ReLU(), nn.Linear(256, n_cls))
def forward(self, x): return self.net(x)
# Reservoir buffer
class ReplayBuffer:
def __init__(self, capacity=2000):
self.capacity, self.n_seen = capacity, 0
self.data = []
def add_batch(self, x, y):
for xi, yi in zip(x, y):
self.n_seen += 1
if len(self.data) < self.capacity: self.data.append((xi.detach().cpu(), yi.detach().cpu()))
else:
j = random.randint(0, self.n_seen - 1)
if j < self.capacity: self.data[j] = (xi.detach().cpu(), yi.detach().cpu())
def sample(self, k):
k = min(k, len(self.data))
idx = random.sample(range(len(self.data)), k) if k > 0 else []
if not idx: return None, None
xs, ys = zip(*[self.data[i] for i in idx])
return torch.stack(xs).to(device), torch.stack(ys).to(device)
def estimate_fisher(model, data_loader, n_samples=1024):
model.eval()
fisher = {n: torch.zeros_like(p, device=device) for n, p in model.named_parameters() if p.requires_grad}
cnt = 0
for x, y in data_loader:
x, y = x.to(device), y.to(device)
logits = model(x)
logp = torch.log_softmax(logits, dim=-1)
idx = torch.arange(x.size(0), device=device)
loss = -logp[idx, y].mean()
model.zero_grad(set_to_none=True)
loss.backward()
for (n, p) in model.named_parameters():
if p.grad is not None and p.requires_grad:
fisher[n] += p.grad.detach()**2
cnt += x.size(0)
if cnt >= n_samples: break
for n in fisher: fisher[n] /= max(1, cnt)
return fisher
def ewc_penalty(model, fisher, prev_params, lam=100.0):
loss = 0.0
for (n, p) in model.named_parameters():
if p.requires_grad and n in fisher:
loss = loss + (fisher[n] * (p - prev_params[n]).pow(2)).sum()
return lam * loss
# Continual loop (single-head classification)
device = "cuda" if torch.cuda.is_available() else "cpu"
model = Net().to(device)
opt = optim.AdamW(model.parameters(), lr=3e-4)
buf = ReplayBuffer(capacity=2000)
prev_fisher, prev_params = None, None
def snapshot_params(model):
return {n: p.detach().clone() for n, p in model.named_parameters() if p.requires_grad}
def train_stream(stream_loaders, replay_k=64, lam=200.0, epochs=1):
global prev_fisher, prev_params
for t_id, loader in enumerate(stream_loaders):
model.train()
for _ in range(epochs):
for x, y in loader:
x, y = x.to(device), y.to(device)
xr, yr = buf.sample(replay_k)
if xr is not None:
x = torch.cat([x, xr], 0); y = torch.cat([y, yr], 0)
logits = model(x)
ce = nn.CrossEntropyLoss()(logits, y)
reg = ewc_penalty(model, prev_fisher, prev_params, lam) if prev_fisher is not None else 0.0
loss = ce + reg
opt.zero_grad(set_to_none=True)
loss.backward()
opt.step()
buf.add_batch(x.detach(), y.detach())
# After task: snapshot importance
prev_fisher = estimate_fisher(model, loader)
prev_params = snapshot_params(model)
신규 데이터와 버퍼 데이터를 섞은 미니배치는 망각을 억제한다. 태스크 경계 뒤 Fisher와 파라미터 스냅샷을 저장하면 다음 태스크에서 EWC 항이 복원력을 제공한다. 서비스 환경에서는 드리프트 감지, 클래스 균형 버퍼, 검증 루프, 레지스트리 연동까지 포함해야 한다.
연속 학습은 정규화·리플레이·격리 또는 확장 전략을 조합해 Catastrophic Forgetting을 제어하는 문제다. 분포 변화 감지, 버퍼 관리, 중요도 기반 락, 운영형 SLO를 하나의 파이프라인으로 다루고 규정·리소스·지연 제약에 맞춰 전략을 선택한다.