[논문 리뷰] Pre-trained Vision and Language Transformers Are Few-Shot Incremental Learners

jiwoong·2026년 3월 29일

논문 리뷰

목록 보기
5/5

https://arxiv.org/pdf/2404.02117

Abstract

소수 샘플 기반 점진적 학습(Few-Shot Class Incremental Learning, FSCIL)은 각 클래스에 대해 소수의 샘플만 주어졌을 때, 망각 없이 새로운 클래스를 점진적으로 학습해야 하는 작업

FSCIL의 두 가지 중요 문제

  1. catastrophic forgetting
  2. overfitting

기존 연구에서는 ResNet-18과 같은 얕은 모델에 의존

적은 파라미터의 모델이 망각과 과적합 문제를 완화할 수 있기 때문.

하지만, 수 샘플 점진적 학습 세션 동안 부적절한 지식 전달로 이어짐

본 논문에서는 대규모 데이터 셋으로 pre-training된 비전 및 언어 모델이 훌륭한 few-shot incremental learner가 될 수 있다고 주장함

이를 위해 프롬프팅 함수와 지식 증류를 활용한 사전 학습 비전 및 언어 transformer인 PriViLege라는 새로운 FSCIL 프레임워크를 제안

새로운 사전 학습 지식 튜닝(pre-trained knowledge tuning, PKT)

엔트로피 기반 발산 손실(entropy-based divergence loss) 및 의미론적 지식 증류 손실(semantic knowledge distillation loss)을 통해 대형 모델의 치명적 망각 및 과적합 문제를 효과적으로 해결

PriViLege가 CUB200에서 +9.38%, CIFAR-100에서 +20.58%, miniImageNet에서 +13.36%의 큰 성능 향상을 보이며 기존 최첨단 방법들을 크게 능가함을 보임

1. Introduction

FSCIL

  • base session과 incremental session으로 구성
  • base session에서는 충분한 데이터로 많은 클래스를 학습
  • 이후 incremental session에서는 각 클래스당 few-shot 데이터만으로 새로운 클래스를 학습해야 함.
  • 이 과정에서 기존 클래스를 잊지 않는 것이 핵심

FSCIL의 핵심 문제

  • catastrophic forgetting
    • 새로운 클래스를 순차적으로 배우면서, 이전에 배운 지식을 심하게 잊는 문제
  • overfitting
    • incremental session의 데이터가 너무 적어서 소수 샘플에 과도하게 적합되는 문제

기존 방법의 한계

  • 기존 FSCIL 방법들은 주로 얕은 모델(shallow model) 에 의존 (ex. ResNet-18)
  • 장점
    • 파라미터 수가 적어서 forgetting 완화에 유리
    • overfitting도 어느 정도 줄일 수 있음
  • 한계
    • 모델 용량이 작아서 base session의 충분한 지식을 incremental session으로 잘 전달하지 못함
    • 즉, transferability 부족

large pre-trained model에 대한 문제의식

  • 최근에는 ViT와 CLIP 같은 large pre-trained model이 비전 분야에서 강한 성능을 보임
  • 이런 모델은 base session의 지식을 더 풍부하게 학습하고 전달할 가능성이 있음.
  • 하지만 FSCIL에 그대로 쓰기에는 trade-off가 존재
    • fine-tuning을 많이 하면 useful pre-trained knowledge를 잊기 쉽고,
    • freeze를 많이 하면 domain-specific knowledge를 충분히 학습하지 못함

본 논문은 CIFAR-100 5-way 5-shot 설정에서 pre-trained ViT-B를 기존 FSCIL 방법들에 적용

  • WaRP: selective freezing 방식, incremental session에서 severe forgetting 발생
  • CEC: 전체 network를 freeze해서 forgetting은 줄지만 전체 세션에서 useful knowledge학습 부족
  • L2P / DualPrompt: prompting 기반 방법도 FSCIL에서는 기대만큼 강하지 않음
  • 결론
    • 큰 모델이 문제인 게 아니라, 큰 모델을 FSCIL에 맞게 쓰는 방법이 부족

PriViLege 제안

  • Pre-trained Vision and Language transformers with prompting functions and knowledge distillation
  • 핵심 구성
    1. PKT
    2. Entropy-based Divergence Loss
    3. Semantic Knowledge Distillation Loss

2. Related Work

Few-Shot Class Incremental Learning

FSCIL 은 class incremental learning의 일종이지만, novel class를 few-shot으로 배워야 한다는 점에서 더욱 어려운 문제

기존 FSCIL의 방법론

  • dynamic network structure-based methods
    • 학습 중 네트워크 구조를 조정하며 forgetting을 줄이려는 방식
  • feature / feature-space based methods
    • feature extractor의 일반화 능력을 높여 새로운 클래스 적응을 돕는 방식
  • prototype-based methods
    • prototype과 classifier weight를 정렬하거나, prototype을 이용해 분류 성능을 높이는 방식
  • 하지만 기존 방법들 대부분 shallow model에서 forgetting과 overfitting을 줄이는 게 초점. 성능은 제한적

Prompt Engineering for Vision Transformer

VIT의 prompt engineering의 대표적 방식

  • prompt tuning
    • 입력 시퀀스에 learnable prompt를 추가
  • prefix tuning
    • attention의 key, value 앞에 prompt를 붙여 attention pattern에 직접적인 영향

Semantic Guidance from Language Models

본 논문에서는 pre-trained language model의 semantic knowledge를 visual space로 직접 distillation하는 방식을 제안

Method Overview

  • backbone
    • pre-trained ViT
  • 학습 요소
    • B-Prompt
    • VL-Prompt
    • selected ViT layers
  • 세 핵심 구성의 역할
    • PKT: pre-trained knowledge 보존 + domain knowledge 학습
    • LED: vision token이 독립적인 discriminative feature를 갖도록 유도
    • LSKD: language model의 semantic knowledge를 visual learning에 주입

3. Method

FSCIL 문제 설정

전체 training dataset

D={D0,D1,,DT}D = \{D^0, D^1, \dots, D^T\}

  • D0D^0: base session 데이터셋
  • DtD^t: t번째 incremental session 데이터셋
  • TT: incremental session의 총 개수

incremental session

Dt={(xi,yi)}i=1DtD^t = \{(x_i, y_i)\}_{i=1}^{|D^t|}

  • xix_i: 입력 이미지
  • yiy_i: 정답 레이블

Dt|D^t|: 해당 세션의 샘플 수

  • incremental session에서는 각 클래스당 샘플 수가 적음

Dt=kCt|D^t| = k \cdot |C^t|

  • Ct|C^t|: t번째 session의 novel class 수
  • kk: novel class당 샘플 수

목표

  • 적은 샘플만으로 새로운 클래스를 순차적으로 학습
  • 동시에 이전까지 등장한 모든 클래스의 분류 성능 유지

3.1 Pre-trained Knowledge Tuning

PKT의 핵심 아이디어

  • large pre-trained model 전체를 다 바꾸는 것이 아닌, 필요한 일부 layer와 prompt만 학습해서
    • useful pre-trained knowledge는 보존
    • domain-specific knowledge는 새로 학습하도록 설계

사용되는 prompt

(1) Base Prompt

PBRL×2×DP_B \in \mathbb{R}^{L \times 2 \times D}

  • LL: 학습할 ViT layer 수
  • 22: key / value prompt
  • DD: embedding dimension

의미

  • 선택된 LL개 layer의 key/value 앞에 붙는 prefix prompt
  • base session의 domain knowledge를 학습하고 incremental session으로 전달하는 역할

(2) Vision-Language Prompt

PVLR2×DP_{VL} \in \mathbb{R}^{2 \times D}

  • 2의 의미
    • vision token 1개
    • language token 1개를 의미

의미

  • vision token은 이후 LED의 대상
  • language token은 이후 LSKD의 대상

modulation prompt 도입 이유

  • 단순 prefix tuning만으로는 B-Prompt가 충분히 잘 학습되지 않을 수 있었음.
    • adaptation 속도가 느릴 수 있음
    • fine-tuned layer가 prompt를 무시할 수도 있음

⇒ prompt modulation을 도입

modulation 관련 수식

Self-attention output

hMSA=MSA(hQ,hK,hV)h^{MSA} = \mathrm{MSA}(h_Q, h_K, h_V)

  • hQ,hK,hVh_Q, h_K, h_V: query, key, value 입력
  • hMSAh^{MSA}: multi-head self-attention의 출력

head-specific modulation prompt

PMS=[g1S(h1MSA);;gHS(hHMSA)]P_M^S = [g_1^S(h_1^{MSA}); \dots ; g_H^S(h_H^{MSA})]

  • HH: attention head 수
  • hMSAh^{MSA}: h번째 head의 출력
  • ghSg_h^S: point-wise convolution

의미

  • 각 head마다 따로 modulation signal을 만듦

generic modulation prompt

hMLP=MLP(hMSA)h^{MLP} = \mathrm{MLP}(h^{MSA})

PMG=gG(hMLP)P_M^G = g^G(h^{MLP})

  • hMLPh^{MLP}: MLP 출력
  • gGg^G: point-wise convolution

의미

  • 더 전역적인 modulation 역할

modulation을 B-Prompt에 적용

PˉK=PMSPBK\bar{P}'_K = P_M^S \odot P_B^K

PˉV=PMGPBV\bar{P}'_V = P_M^G \odot P_B^V

  • ⊙: element-wise multiplication
  • PK,PˉV{P}'_K, \bar{P}'_V: modulation이 적용된 새로운 prompt

최종 attention

${h}^{out}

\mathrm{MSA}
(
[P_{VL}^Q; h_Q],
[\bar{P}'_K; h_K],
[\bar{P}'_V; h_V]
)$

의미:

  • query에는 VL-Prompt
  • key/value에는 modulation된 B-Prompt를 넣어 attention 수행

PKT 정리

  • 일부 layer만 selective tuning
  • B-Prompt로 domain knowledge 학습
  • VL-Prompt로 vision / language token 도입
  • modulation prompt로 prompt의 representation learning 강화

⇒ PKT는 사전학습 지식을 최대한 보존하면서, 필요한 부분만 최소한으로 조정해 domain knowledge를 얻는 장치

3.2 Entropy-based Divergence Loss

제안 배경

  • PKT 학습 과정에서 [CLS] token과 vision token을 함께 사용
  • 그런데 둘이 같은 분류 목표만 공유하면
    • 두 token의 feature가 점점 비슷해지고
    • vision token이 [CLS]의 복사본처럼 될 위험이 있음.

prototype classifier 구성

각 base class cjc_j의 prototype

$proto_{c_j}

\frac{1}{N{c_j}}
\sum
{k=1}^{N_{c_j}} f_k^{cls}$

  • NcjN_{c_j}: 클래스 cj의 샘플 수
  • fkclsf_k^{cls}: 해당 샘플의 [CLS] feature

prototype classifier

ψ=[protoc1;protoc2;;protocC0]\psi = [proto_{c_1}; proto_{c_2}; \dots ; proto_{c_{|C^0|}}]

  • 각 base class prototype을 모은 classifier
  • 논문에서는 이를 stable basis로 사용

[CLS]와 vision token의 logits

y^icls=ψ(ficls),\hat{y}_i^{cls} = \psi(f_i^{cls}),
y^ivis=ψ(fivis)\hat{y}_i^{vis} = \psi(f_i^{vis})

  • ficlsf_i^{cls}: [CLS] feature
  • fivisf_i^{vis}: vision token feature

LED 수식

$\mathcal{L}_{ED}

\log\left(
\frac{
\mathcal{L}_{CE}(\hat{y}_i^{vis}, y_i)

  • \mathcal{L}{CE}(\hat{y}_i^{cls}, y_i)
    }{
    \mathcal{L}
    {KL}(\delta(\hat{y}_i^{vis}), \delta(\hat{y}_i^{cls}))
    }
    +1
    \right)$
  • LCE\mathcal{L}_{CE}: cross-entropy loss
  • LKL\mathcal{L}_{KL}: KL divergence
  • δ()\delta(\cdot): softmax

식의 의미

  • 분자의 CE loss를 줄이려면
    • [CLS]도 정답을 맞추고
    • vision token도 정답을 맞춰야 함
  • 분모의 KL divergence를 크게 하려면
    • 두 token의 softmax 분포가 너무 비슷해지면 안 됨

즉 LED는

둘 다 정확해야 하지만, 완전히 같은 방식으로 예측해서는 안 된다는 제약

⇒ vision token이 독립적인 discriminative knowledge를 가지도록 유도.

3.3 Semantic Knowledge Distillation Loss

제안 배경

  • few-shot novel class를 학습할 때는 데이터가 너무 적어서
    visual signal만으로는 representation이 부족할 수 있음.
  • 그래서 논문은 class name이 담고 있는 semantic information을 활용했음.

class name embedding 추출

클래스 이름을 텍스트 프롬프트로 바꾸고

wordcni=“a photo of [cni]”word_{c_{n_i}} = \text{“a photo of [}c_{n_i}\text{]”}

PLM에 넣어 embedding을 얻음.

wcni=fϕ(wordcni)w_{c_{n_i}} = f_\phi(word_{c_{n_i}})

  • fϕf_\phi: pre-trained language model (예: BERT)
  • wcniw_{c_{n_i}}: class name의 semantic embedding

backbone 내부의 language token feature

filang=fθ(xi)[2]f_i^{lang} = f_\theta(x_i)[2]

  • fθf_\theta: ViT backbone
  • filangf_i^{lang}: VL-Prompt의 language token feature

문제점

  • wcniw_{c_{n_i}}: language space의 embedding
  • filangf_i^{lang}: visual backbone 내부 feature

즉 두 벡터는 서로 다른 embedding space에 있음.

그래서 단순 distillation만으로는 부족

prototype classifier 활용

y^ilang=ψ(filang)\hat{y}_i^{lang} = \psi(f_i^{lang})

  • language token feature도 prototype classifier에 통과시켜
  • 실제 visual classification 기준으로도 올바른 방향을 갖게함

LSKD 수식

$\mathcal{L}_{SKD}

\mathcal{L}{KD}(f_i^{lang}, w{c_{n_i}})

  • \gamma \cdot \mathcal{L}_{CE}(\hat{y}_i^{lang}, y_i)$
  • 첫 번째 항: language token이 class-name semantic embedding을 닮게 함
  • 두 번째 항: language token이 classification에도 실제로 유효한 방향을 갖게 함
  • γ=0.1\gamma = 0.1 사용

⇒ LSKD는 class name의 의미 정보를 language model로부터 받아 visual feature learning에 주입하는 장치

4. Experiments

학습 loss

base session

$\mathcal{L}{CE}(\hat{y}_i, y_i)+
\alpha \cdot \mathcal{L}
{ED}

  • \beta \cdot \mathcal{L}_{SKD}$

incremental session

$\mathcal{L}_{inc}

\mathcal{L}_{CE}(\hat{y}_i, y_i)

  • \beta \cdot \mathcal{L}_{SKD}$
  • incremental session에서는 LED를 사용하지 않음
  • 이유: few-shot 데이터만으로는 discriminative feature를 안정적으로 학습하기 어렵기 때문

4.1 Experimental Settings

Datasets and Metrics

  • CUB200
  • CIFAR-100
  • miniImageNet

Metrics

  • ABaseA_{Base}
    • base session accuracy
  • ALastA_{Last}
    • 마지막 session accuracy
  • AAvgA_{Avg}
    • 전체 session accuracy 평균

ALastA_{Last}는 forgetting을 가장 잘 보여주는 지표

Baselines

  • CEC
  • WaRP
  • NC-FSCIL
  • L2P
  • DualPrompt

4.2 Main Expeimental Results

  • PriViLege는 세 데이터셋 모두에서 기존 SOTA를 큰 폭으로 앞섰음.
  • 특히 마지막 session accuracy까지 높았다는 점에서, 단순 평균 상승이 아니라 forgetting 완화에 성공

4.3 Ablation Study

PriViLege는 ViT 전용 기법이 아니라, pre-trained vision-language model 전반에 적용 가능

4.4 Analysis

LED는 단순 loss 추가가 아니라, 실제 feature geometry를 더 discriminative하게 만드는 역할

LSKD를 넣었을 때 성능이 많이 오른 클래스 이름에는 다음 단어들이 자주 등장

  • red

  • green

  • yellow

  • white

  • headed

  • tailed

  • billed

  • fish

  • 클래스 이름 자체가 시각적으로 informative할수록 LSKD의 효과가 더

  • 특히 fine-grained dataset인 CUB200에서 의미가 큼

Conclusion

기존 FSCIL 연구는 forgetting과 overfitting을 줄이기 위해 주로 얕은 모델에 의존

하지만 얕은 모델은 base session의 풍부한 지식을 incremental session으로 충분히 전달하지 못함

  • 본 논문은 large pre-trained vision / language transformer도 FSCIL에서 매우 강력한 learner가 될 수 있다고 주장
  • PriViLege를 제안
    • PKT
      • pre-trained knowledge를 보존하면서 domain knowledge를 학습하고
    • LED
      • vision token의 구분력을 강화
    • LSKD
      • class name의 semantic knowledge를 visual learning에 주입했음.
  • 실험 결과 PriViLege는 CUB200, CIFAR-100, miniImageNet 모두에서 기존 SOTA를 크게 넘어섰고, CLIP에도 적용 가능함을 보임
profile
학생

0개의 댓글