PromptTuningConfig API

PromptTuningConfig API

PEFT의 프롬프트 튜닝은 PromptTuningConfig 로 구성하고, get_peft_model() 로 모델을 감싸는 방식으로 실행돼요. 가상 프롬프트 파라미터만 학습되고 사전학습 가중치는 얼려져요.

핵심 개념

프롬프트 튜닝은 입력에 태스크별 가상 프롬프트(임베딩 공간의 학습 가능한 벡터)를 추가해요. 가상 토큰 파라미터는 사전학습 모델과 무관하게 갱신돼요. 논문에 따르면 수십억 파라미터를 넘어서면 모델 튜닝(full fine-tuning)과 격차를 좁혀 성능이 맞춰져요.

구성 만들기

from peft import PromptTuningConfig, PromptTuningInit, get_peft_model

config = PromptTuningConfig(
    task_type="CAUSAL_LM",
    num_virtual_tokens=8,
    prompt_tuning_init=PromptTuningInit.TEXT,
    prompt_tuning_init_text="Classify if the tweet is a complaint or not:",
    tokenizer_name_or_path="bigscience/bloomz-560m",
)
model = get_peft_model(base_model, config)
model.print_trainable_parameters()   # trainable params: 8192  (매우 작음!)

초기화 방식

prompt_tuning_init 으로 입력 임베딩 초기화를 고를 수 있어요.

  • TEXT: 초기 텍스트로 시작.
  • SAMPLE_VOCAB: 어휘에서 무작위 샘플링.
  • RANDOM: 연속 소프트 토큰 무작위(임베딩 다양체 밖일 수 있음).

저장 용량

프롬프트 튜닝 결과는 수십 KB 수준이에요. 어댑터만 저장해 모델별로 하나씩 바꿔 쓰는 게 가능해요.

더 알아보기