본문으로 건너뛰기

[SGLang] Sampling Parameters: 전체 샘플링 파라미터 정리

들어가며

LLM 서빙에서 샘플링 파라미터는 출력의 품질과 다양성을 결정한다. SGLang의 SamplingParams 클래스는 temperature, top-k, top-p, min-p, 반복 페널티, 문법 제약 등 모든 파라미터를 한 곳에서 관리하며, 검증과 정규화 로직까지 포함한다.

이 글에서는 python/sglang/srt/sampling/sampling_params.py를 중심으로 각 파라미터의 역할과 상호작용을 분석한다.

파라미터 분류 구조도

SamplingParams
├── 토큰 선택 파라미터
│   ├── temperature     (확률 분포 평탄화)
│   ├── top_k           (상위 K개 후보)
│   ├── top_p           (누적 확률 P 이내)
│   └── min_p           (최대 확률 대비 비율)
│
├── 반복 제어 파라미터
│   ├── frequency_penalty   (빈도 비례 감점)
│   ├── presence_penalty    (등장 여부 감점)
│   └── repetition_penalty  (스케일링 페널티)
│
├── 길이 제어 파라미터
│   ├── max_new_tokens     (최대 생성 길이)
│   ├── min_new_tokens     (최소 생성 길이)
│   └── ignore_eos         (EOS 무시 여부)
│
├── 정지 조건 파라미터
│   ├── stop              (정지 문자열)
│   ├── stop_token_ids    (정지 토큰 ID)
│   └── stop_regex        (정지 정규식)
│
├── 문법 제약 (상호 배타적)
│   ├── json_schema
│   ├── regex
│   └── ebnf
│
└── 기타
    ├── n                 (병렬 생성 수)
    ├── custom_params     (사용자 정의 파라미터)
    ├── logit_bias        (특정 토큰 가중치)
    └── sampling_seed     (결정론적 시드)

핵심 코드 분석

생성자: 특수 케이스 처리

SamplingParams.__init__에서 주목할 부분은 temperature 0 처리와 top_k 기본값 변환이다.

_SAMPLING_EPS = 1e-6
TOP_K_ALL = 1 << 30  # 약 10억

def __init__(self, ..., temperature=1.0, top_k=-1, ...):
    self.temperature = temperature
    self.top_k = top_k
    # ...

    if 0 <= self.temperature < _SAMPLING_EPS:
        self.temperature = 1.0
        self.top_k = 1
    if self.top_k == -1:
        self.top_k = TOP_K_ALL

temperature가 0이면 greedy로 전환된다. 내부적으로 temperature=1.0, top_k=1로 설정하여 softmax를 정상 수행하되 상위 1개만 선택하는 방식이다. top_k=-1은 전체 어휘를 의미하며, 1 << 30이라는 충분히 큰 값으로 변환된다.

파라미터 검증

verify() 메서드는 유효 범위를 엄격하게 검증한다.

def verify(self, vocab_size):
    if self.temperature < 0.0:
        raise ValueError(f"temperature must be non-negative, got {self.temperature}.")
    if not 0.0 < self.top_p <= 1.0:
        raise ValueError(f"top_p must be in (0, 1], got {self.top_p}.")
    if not 0.0 <= self.min_p <= 1.0:
        raise ValueError(f"min_p must be in [0, 1], got {self.min_p}.")
    if not -2.0 <= self.frequency_penalty <= 2.0:
        raise ValueError(...)
    if not 0.0 <= self.repetition_penalty <= 2.0:
        raise ValueError(...)

주요 유효 범위를 정리하면 다음과 같다.

파라미터 유효 범위 기본값
temperature >= 0.0 1.0
top_p (0, 1] 1.0
top_k >= 1 또는 -1 -1 (전체)
min_p [0, 1] 0.0
frequency_penalty [-2, 2] 0.0
presence_penalty [-2, 2] 0.0
repetition_penalty [0, 2] 1.0

문법 제약은 상호 배타적으로 검증된다.

grammars = [self.json_schema, self.regex, self.ebnf]
if sum(x is not None for x in grammars) > 1:
    raise ValueError("Only one of regex, json_schema, or ebnf can be set.")

정지 조건 정규화

normalize() 메서드는 stop 문자열과 stop regex의 최대 길이를 사전 계산한다.

def normalize(self, tokenizer):
    if self.stop_strs is None:
        self.stop_strs = []
        self.stop_str_max_len = 0
    else:
        if isinstance(self.stop_strs, str):
            self.stop_strs = [self.stop_strs]
        stop_str_max_len = 0
        for stop_str in self.stop_strs:
            stop_str_ids = tokenizer.encode(stop_str, add_special_tokens=False)
            stop_str_max_len = max(stop_str_max_len, len(stop_str_ids))
        self.stop_str_max_len = stop_str_max_len

stop_str_max_len은 디코딩 시 버퍼링할 토큰 수를 결정한다. 정지 문자열이 여러 토큰에 걸칠 수 있으므로 그 최대 길이만큼 출력을 유보해야 한다.

정규식 최대 길이 계산

정지 정규식의 최대 매칭 길이를 파싱하여 계산한다.

def get_max_seq_length(regex_str: str):
    return _max_length_from_subpattern(sre_parse.parse(regex_str))

def _max_length_from_subpattern(subpattern):
    total = 0
    for token, value in subpattern:
        if token in {sre_parse.LITERAL, sre_parse.IN, sre_parse.ANY}:
            total += 1
        elif token == sre_parse.SUBPATTERN:
            _, _, _, inner_subpattern = value
            total += _max_length_from_subpattern(inner_subpattern)
        elif token == sre_parse.BRANCH:
            _, branches = value
            total += max(_max_length_from_subpattern(b) for b in branches)
        elif token in {sre_parse.MAX_REPEAT, sre_parse.MIN_REPEAT}:
            _, max_num_repeat, inner = value
            if max_num_repeat == sre_parse.MAXREPEAT:
                total += MAX_LEN  # 2^30, 사실상 무한
            else:
                total += max_num_repeat * _max_length_from_subpattern(inner)
    return total

*+ 같은 무한 반복은 2^30으로 설정한다. 이 값은 버퍼링 상한을 결정하므로 실제로 무한 버퍼링이 발생하지는 않는다.

파라미터 상호작용 흐름

사용자 요청
    │
    ▼
SamplingParams.__init__()
    │  ┌── temperature=0 → top_k=1 (greedy)
    │  └── top_k=-1 → TOP_K_ALL
    ▼
verify(vocab_size)
    │  ┌── 범위 검증
    │  └── 문법 상호 배타 검증
    ▼
normalize(tokenizer)
    │  ├── stop_strs 정규화
    │  └── stop_regex 최대 길이 계산
    ▼
SamplingBatchInfo 구성
    │  ├── temperatures 텐서
    │  ├── top_ks 텐서
    │  ├── top_ps 텐서
    │  └── min_ps 텐서
    ▼
Sampler.forward() 에서 실제 적용

설계 근거

설계 선택 이유
temperature=0을 top_k=1로 변환 softmax(0으로 나누기) 대신 정상 softmax + top-1 선택으로 수치적 안정성 확보
TOP_K_ALL = 1 << 30 top_k 비활성화 시 조건 분기 없이 동일 코드 경로 사용
문법 상호 배타 검증 json_schema, regex, ebnf가 동시에 적용되면 충돌 가능
stop_str_max_len 사전 계산 스트리밍 시 정지 문자열 매칭을 위한 버퍼 크기 결정

관련 포스트

참고

댓글

관련 포스트

SGLang 의 다른글