소량 데이터 학습을 위한 메타러닝: MAML과 Prototypical Networks

MAML과 Prototypical Networks의 에피소드 학습 구조, 비용·확장성 차이, 소량 데이터 환경의 운영 전략을 정리한다.

2026-08-15 · 최초 발행 2024-04-29

적은 데이터로 새 태스크에 적응시키는 학습 방식

메타러닝(meta-learning)은 새로운 태스크에 대한 적응 속도를 높이는 학습 프레임워크다. 여러 태스크를 에피소드(episode) 단위로 학습하면서, 초기 파라미터나 임베딩 자체에 일반화 능력을 담는다.

이 글에서 다루는 두 방식은 접근점이 다르다. MAML(Model-Agnostic Meta-Learning)은 적은 경사하강 스텝만으로도 성능을 빠르게 올릴 수 있는 초기 파라미터 세트 θ⋆를 학습한다. 반면 Prototypical Networks는 각 클래스의 임베딩 평균을 프로토타입으로 두고, 쿼리 데이터와 프로토타입의 거리를 비교해 분류한다.

에피소드 학습은 보통 N-way K-shot 설정으로 서포트 셋과 쿼리 셋을 나눈다. 이때 태스크 분포를 얼마나 잘 근사하느냐가 샘플러 설계에 달려 있다. 클래스 불균형이나 누락을 다룰 재샘플링 정책, 에피소드 난이도 스케줄링도 함께 고려한다.

대규모 데이터로 ResNet, ViT 같은 백본을 사전학습한 뒤 메타학습 파인튜닝을 적용할 수 있다. Weight decay, Dropout과 Mixup, RandAugment는 과적합을 억제하는 데 사용한다.

같은 에피소드를 다르게 처리하는 MAML과 ProtoNet

MAML의 내측 루프는 현재 태스크의 서포트 데이터로 몇 차례 경사하강을 수행해 적응된 파라미터 θ′를 만든다. 외측 루프에서는 θ′로 쿼리 손실을 계산하고, θ에 대해 2차 미분 또는 1차 근사를 사용해 메타 업데이트를 수행한다.

Prototypical Networks는 임베딩 네트워크 f(·)로 입력을 임베딩 공간으로 옮긴다. 클래스별 평균 벡터가 프로토타입이 되며, 쿼리와 프로토타입 간 거리는 보통 Euclidean 또는 Cosine 함수로 계산한다. 그 결과를 소프트맥스 분류에 사용한다.

MAMLProtoNet클래스 < N클래스 충분입력: 태스크 분포 D, N-wayK-shot 에피소드에피소드 샘플러:Support/Query 분할분기내측 루프: Support로 θ - θ'외측 루프: Query 손실로 θ업데이트출력: 메타 파라미터 θ*임베딩 f(x) 계산클래스별 프로토타입 μ_c =mean(f(S_c))쿼리-프로토타입 거리 기반분류출력: 임베딩 f, 프로토타입기반 결정유효성 검사재샘플링/스킵 처리학습 진행

MAML에서는 태스크별 support K샷, query Q개, 초기 파라미터 θ가 에피소드에 들어온다. support 손실로 θ를 θ′로 옮긴 뒤, θ′가 만든 query 손실을 원래 초기화 θ의 메타 업데이트에 사용한다. 학습 중에는 gradient clipping을 적용하고 FOMAML이나 Reptile을 선택해 안정성을 조정할 수 있으며, 결과는 새로운 태스크에 빠르게 적응하기 위한 초기화 θ*다.

Prototypical Networks에서는 support와 query를 임베딩 함수 fθ에 통과시킨다. support 임베딩을 클래스별로 평균내 프로토타입을 만들고, query와 각 프로토타입 사이의 거리에서 로짓을 계산한다. cross-entropy는 프로토타입이 잘 분리되는 임베딩 함수 θ*를 갱신하며, 추론에서는 이 거리 계산 경로를 사용한다. 배치 내 클래스 불균형과 프로토타입 분산, temperature와 metric 선택은 결정 경계와 성능에 영향을 준다.

학습 중 클래스 수가 부족하거나 샷 구성이 불균형하면 해당 에피소드를 재샘플링하거나 건너뛴다. 에피소드 단위의 분포를 일관되게 유지하고 샘플러 시드를 관리해야 재현성을 확보할 수 있다.

선택 기준은 적응 비용과 임베딩 운영 방식에 있다

항목 MAML Prototypical Networks
성능(1-shot 일반화) 다양한 도메인에서 강력하나 최적화 민감. 최신 SOTA와 상대 비교는 최신 정보 확인 필요 임베딩 품질이 우수할 때 일관된 성능. 거리 함수 선택 영향 큼
계산 비용 내·외측 루프와 2차 그라디언트로 비용 높음 단순 전방향+평균/거리 계산으로 비용 낮음
확장성(way/shot 증가) 에피소드 크기 증가 시 메모리/시간 비용 급증 선형 확장 경향, 대규모 에피소드 처리 용이
일관성/안정성 학습율·스텝 수 등 하이퍼파라미터 민감 비교적 안정, 하이퍼파라미터 수 적음
운영 편의 온라인 도입 시 태스크별 빠른 적응 가능 추론 시 프로토타입 업데이트만으로 클래스 증감 대응 용이

태스크 간 분포 편차가 크고 빠른 태스크별 적응이 필요하면 MAML을 검토할 수 있다. 반대로 클래스가 늘고 줄어드는 운영 환경에서 추론 효율과 안정성을 우선하면 Prototypical Networks가 더 단순한 선택지가 된다.

데이터가 부족한 현장에서의 적용

제조·검사에서는 신규 불량 유형의 샘플이 부족한 경우 15장으로 탐지 모델을 신속히 부트스트랩할 수 있다. Prototypical Networks는 클래스 추가·삭제 시 프로토타입을 갱신하는 방식으로 운영을 단순화한다. 산업 비전 불량 검출에서 신제품이나 신공정 전환으로 라벨이 희소해질 때는 임베딩 백본을 고정한 뒤 신규 클래스를 프로토타입 등록으로 다룰 수 있으며, 현장 미세 조정에는 MAML 초기화로 510 스텝 적응을 적용해 30분 내 현장 재적용을 목표로 할 수 있다.

의료 영상과 이상 탐지에서는 레이블 비용이 높은 희귀 질환 분류에 MAML 초기화를 적용해 적은 추가 데이터로 고성능을 달성할 수 있다. 임베딩-거리 기반 판별은 프로토타입 근접도를 통해 설명 가능성도 제공한다. 의료 영상 희귀 질환 분류처럼 domain shift가 강한 경우에는 MAML 초기화가 센터 간 이질성을 완화하는 데 쓰이며, 데이터 보안 제약이 있다면 클라이언트별 에피소드 연합 학습(Federated meta-learning)을 적용할 수 있다.

NLP 인텐트·슬롯의 신규 도메인 온보딩에서는 클라이언트별 K-shot 데이터로 빠르게 적응할 수 있다. MAML은 파라미터 효율적 파인튜닝에, 멀티링구얼 임베딩과 ProtoNet의 결합은 저자원 언어 대응에 활용된다. 개인화 텍스트 분류나 음성 명령에서는 사용자별 소량 라벨로 개인화를 수행하며, 온라인 운영에서는 에피소드 재샘플링과 주기적 메타 업데이트로 개체 편향을 완화한다.

보안·사기 탐지처럼 신종 패턴이 나타나는 영역에서는 에피소드 재학습으로 대응 지연을 줄일 수 있다. 데이터 민감성 때문에 라벨링이 제한되는 환경에서도 데이터 효율을 높이는 접근이다.

라벨링 비용은 3070% 절감 가능성이 있으며, 이는 도메인과 사전학습 품질에 의존하고 최신 정보 확인이 필요하다. 신규 태스크 온보딩 시간은 50% 이상 단축되고, 파라미터 업데이트는 510 스텝 내 수렴한 사례가 다수 있다. ProtoNet의 추론 비용은 O(C) 프로토타입 비교이며, C가 커질 경우 ANN 탐색으로 O(log C) 근사가 가능하다.

벤치마크 기준으로 데이터 효율은 520배 적은 라벨로 동일 수준 정확도 달성이 가능하다. 신규 태스크의 fine-tune은 520 스텝 내 수렴할 수 있고, 온디바이스 환경에서는 100~500ms 수준의 적응 시간이 가능하다. 클래스 추가 시 ProtoNet은 재학습 없이 프로토타입을 갱신하므로 GPU 시간을 50%+ 절약할 수 있다. 배포 민첩성과 콜드스타트 리스크를 개선하고, 도메인 변동에 대한 내성을 높이며 운영 복잡도를 줄이는 효과도 기대할 수 있다.

구현 전제와 MAML 스켈레톤

구현 환경은 Python 3.10+, PyTorch 2.2+이며 CUDA는 선택 사항이다. 에피소드 샘플러(episode_loader)는 사용자가 구현해야 한다. torch.manual_seed와 데이터셋 시드를 고정해 재현성을 관리한다.

# Python 3.10, PyTorch 2.2
import torch
from torch import nn, optim
from typing import Iterable, Tuple

class SimpleNet(nn.Module):
    def __init__(self, in_dim=784, hidden=64, out_dim=5):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(in_dim, hidden), nn.ReLU(),
            nn.Linear(hidden, out_dim)
        )
    def forward(self, x): return self.net(x)

def clone_params(model: nn.Module):
    return {n: p.clone().detach().requires_grad_(True) for n, p in model.named_parameters()}

def forward_with_params(model, x, params):
    idx = 0
    out = x
    # Works with Sequential of Linear/ReLU. For generality, use functorch or higher library.
    for m in model.net:
        if isinstance(m, nn.Linear):
            W = params[f'net.{idx}.weight']; b = params[f'net.{idx}.bias']
            out = out @ W.T + b
            idx += 1
        elif isinstance(m, nn.ReLU):
            out = torch.relu(out)
    return out

def maml_step(model: nn.Module, episode, inner_lr=0.01, inner_steps=5):
    support_x, support_y, query_x, query_y = episode
    params = clone_params(model)
    # inner loop
    for _ in range(inner_steps):
        logits = forward_with_params(model, support_x, params)
        loss = nn.functional.cross_entropy(logits, support_y)
        grads = torch.autograd.grad(loss, params.values(), create_graph=True)
        params = {k: p - inner_lr * g for (k, p), g in zip(params.items(), grads)}
    # outer loss on query
    q_logits = forward_with_params(model, query_x, params)
    q_loss = nn.functional.cross_entropy(q_logits, query_y)
    return q_loss

def train_maml(model: nn.Module, episodes: Iterable, meta_lr=1e-3, iters=1000):
    opt = optim.Adam(model.parameters(), lr=meta_lr)
    model.train()
    for it, episode in enumerate(episodes):
        opt.zero_grad()
        loss = maml_step(model, episode)
        loss.backward()
        nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        opt.step()
        if it % 100 == 0:
            print(f'iter={it} loss={loss.item():.4f}')
        if it >= iters: break

# episodes = episode_loader(...)  # 사용자 구현
# model = SimpleNet(in_dim=784, out_dim=N_way)
# train_maml(model, episodes)

MAML은 2차 미분 비용이 크므로 1차 근사인 FOMAML이나 Reptile을 고려할 수 있다. 합성곱 또는 트랜스포머 백본을 사용한다면 functorch, higher 등으로 파라미터와 함수를 분리하는 방법을 권장한다.

FOMAML 구현에서는 에피소드 샘플러가 (support_x, support_y, query_x, query_y)를 반환한다고 가정한다.

# pip install torch torchvision
import torch
import torch.nn as nn
import torch.nn.functional as F
from copy import deepcopy

# 간단 CNN 임베딩
class EmbeddingNet(nn.Module):
    def __init__(self, out_dim=64):
        super().__init__()
        self.net = nn.Sequential(
            nn.Conv2d(1, 32, 3, padding=1), nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Conv2d(32, 64, 3, padding=1), nn.ReLU(),
            nn.AdaptiveAvgPool2d(1)
        )
        self.fc = nn.Linear(64, out_dim)

    def forward(self, x):
        h = self.net(x).view(x.size(0), -1)
        return self.fc(h)

# -------- Prototypical Networks --------
def prototypical_step(emb, sx, sy, qx, qy, temp=1.0):
    # 임베딩
    z_s = emb(sx)  # [S, d]
    z_q = emb(qx)  # [Q, d]
    classes = torch.unique(sy)
    # 프로토타입
    protos = torch.stack([z_s[sy == c].mean(dim=0) for c in classes])  # [N, d]
    # 거리 -> 로짓
    # Euclidean 거리의 음수 사용
    dists = torch.cdist(z_q, protos)  # [Q, N]
    logits = -dists / temp
    loss = F.cross_entropy(logits, torch.bucketize(qy, classes) - 1)
    acc = (logits.argmax(1) == (torch.bucketize(qy, classes) - 1)).float().mean().item()
    return loss, acc

# -------- MAML (First-Order: FOMAML) --------
def maml_fomaml_step(model, sx, sy, qx, qy, inner_lr=0.01, inner_steps=1):
    criterion = nn.CrossEntropyLoss()
    fast = deepcopy(model)  # 태스크별 파라미터 복사
    opt_inner = torch.optim.SGD(fast.parameters(), lr=inner_lr)

    # Inner loop
    for _ in range(inner_steps):
        logits_s = fast(sx)
        loss_s = criterion(logits_s, sy)
        opt_inner.zero_grad()
        loss_s.backward()
        torch.nn.utils.clip_grad_norm_(fast.parameters(), 5.0)
        opt_inner.step()

    # Query 손실로 외부 그래디언트 생성(1차 근사)
    logits_q = fast(qx)
    loss_q = criterion(logits_q, qy)

    # Accuracy 계산
    acc = (logits_q.argmax(1) == qy).float().mean().item()

    # 메타 업데이트는 호출자에서 model 파라미터에 대해 loss_q.backward() 호출
    # 여기서는 graph 끊기 방지 위해 fast에서 model로의 연계 생략(FOMAML 근사)
    return loss_q, acc

# 사용 예시 (에피소드 샘플러는 사용자 구현)
if __name__ == "__main__":
    N, K, Q = 5, 1, 15  # 5-way 1-shot
    emb = EmbeddingNet(out_dim=64)
    clf = nn.Sequential(emb, nn.Linear(64, N))  # MAML용 단순 분류기

    # dummy episode
    sx = torch.randn(N*K, 1, 28, 28); sy = torch.repeat_interleave(torch.arange(N), K)
    qx = torch.randn(N*Q, 1, 28, 28); qy = torch.repeat_interleave(torch.arange(N), Q)

    # ProtoNet
    loss_p, acc_p = prototypical_step(emb, sx, sy, qx, qy)
    opt_p = torch.optim.Adam(emb.parameters(), lr=1e-3)
    opt_p.zero_grad(); loss_p.backward(); opt_p.step()

    # MAML (FOMAML 외부 업데이트)
    opt_m = torch.optim.Adam(clf.parameters(), lr=1e-3)
    loss_m, acc_m = maml_fomaml_step(clf, sx, sy, qx, qy, inner_lr=0.01, inner_steps=1)
    opt_m.zero_grad(); loss_m.backward(); opt_m.step()

    print(f"ProtoNet acc={acc_p:.2f}, MAML acc={acc_m:.2f}")

2계 MAML이 필요하면 higher, functorch를 이용한다. 메모리 초과 상황에서는 gradient checkpointing과 mixed precision을 적용할 수 있다. ProtoNet은 temperature, metric, episodic batch 설계가 핵심이며, 클래스 불균형에는 class-balanced sampling을 권장한다. 두 방법 모두 실제 운영 데이터에서 에피소드를 샘플링해 태스크 분포 미스매치를 줄이고 데이터 누수를 막아야 한다.

임베딩과 거리로 학습하는 ProtoNet 스켈레톤

# Python 3.10, PyTorch 2.2
import torch
from torch import nn
from typing import Tuple

class EmbedNet(nn.Module):
    def __init__(self, in_dim=784, emb_dim=128):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(in_dim, 256), nn.ReLU(),
            nn.Linear(256, emb_dim)
        )
    def forward(self, x): return nn.functional.normalize(self.net(x), dim=-1)

def proto_loss(emb: nn.Module, support: Tuple[torch.Tensor, torch.Tensor],
               query: Tuple[torch.Tensor, torch.Tensor], metric='euclidean'):
    sx, sy = support; qx, qy = query
    s_emb = emb(sx)  # [S, D]
    q_emb = emb(qx)  # [Q, D]
    classes = torch.unique(sy)
    prototypes = torch.stack([s_emb[sy==c].mean(0) for c in classes])  # [C, D]

    if metric == 'euclidean':
        # dist^2 = ||q - p||^2 = q^2 + p^2 - 2qp
        q2 = (q_emb**2).sum(-1, keepdim=True)
        p2 = (prototypes**2).sum(-1).unsqueeze(0)
        logits = - (q2 + p2 - 2 * (q_emb @ prototypes.T))  # [Q, C]
    elif metric == 'cosine':
        logits = (q_emb @ prototypes.T)
    else:
        raise ValueError('unknown metric')

    # Map labels to 0..C-1
    label_map = {c.item(): i for i, c in enumerate(classes)}
    qy_mapped = torch.tensor([label_map[int(y)] for y in qy], device=qx.device)
    loss = nn.functional.cross_entropy(logits, qy_mapped)
    acc = (logits.argmax(-1) == qy_mapped).float().mean()
    return loss, acc

def train_protonet(emb: nn.Module, episodes, lr=1e-3, iters=1000):
    opt = torch.optim.Adam(emb.parameters(), lr=lr)
    emb.train()
    for it, (support, query) in enumerate(episodes):
        opt.zero_grad()
        loss, acc = proto_loss(emb, support, query, metric='euclidean')
        loss.backward()
        opt.step()
        if it % 100 == 0:
            print(f'iter={it} loss={loss.item():.4f} acc={acc.item():.3f}')
        if it >= iters: break

# episodes = episode_loader(...)  # 사용자 구현
# emb = EmbedNet(in_dim=784, emb_dim=128)
# train_protonet(emb, episodes)

운영에서는 데이터 분포와 자원 제약을 함께 본다

사전학습 백본과 ProtoNet의 조합은 추론 효율과 안정성 측면에서 유리하다. 태스크 간 분포 편차가 크다면 MAML 도입을 검토한다.

MAML은 inner steps 3~10, inner_lr 1e-2~1e-1, meta_lr 1e-4~1e-3 범위를 사용한다. ProtoNet의 emb_dim128~1024이며, 거리 함수는 Euclidean 기본이고 Cosine은 정규화가 필수다.

데이터는 클래스 균형 에피소드 샘플링, Hard episode mining, soft labels를 통한 라벨 노이즈 완화 전략으로 구성할 수 있다. 대규모 클래스에서는 ProtoNet에 ANN(FAISS, ScaNN)을 도입하고, MAML은 FOMAML 또는 Reptile 근사와 메모리 체크포인트를 사용해 확장한다.

도메인 변동이 큰 환경에서는 MAML과 FOMAML, Reptile을 우선 검토할 수 있다. 신규 클래스가 자주 생기는 환경이라면 임베딩과 프로토타입을 중심으로 운영하는 Prototypical Networks가 맞다. 어느 쪽이든 에피소드 설계, 샘플링 일관성, 데이터 누수 방지, 리소스 제약 아래의 AMP와 체크포인트 최적화가 적용의 기반이 된다.

메타러닝은 데이터 희소 환경에서 도입 가치가 높은 전략이다. MAML은 빠른 태스크 적응에, ProtoNet은 단순하고 견고한 임베딩-거리 기반 분류에 강점이 있다. 사전학습 백본과 에피소드 학습을 결합한 뒤, 추론 지연·클래스 변동성·학습 자원 제약에 따라 두 방식을 병행하거나 단계적으로 도입할 수 있다.

메타러닝소수샷 학습MAMLPrototypical Networks머신러닝