커스텀 로짓 프로세서

커스텀 로짓 프로세서 (Custom Logits Processors)

언어 모델이 다음 토큰을 고를 때, 때로는 그 확률 분포를 우리가 원하는 방향으로 조정하고 싶을 때가 있어요. 그 역할을 하는 게 로짓 프로세서(logits processor) 인데, vLLM은 사용자가 직접 작성한 커스텀 로짓 프로세서를 vLLM 소스 코드를 수정하거나 컴파일하지 않고도 로드해서 쓸 수 있게 해줍니다. 이 페이지에서는 커스텀 로짓 프로세서를 작성하고, 로드하고, 사용하는 방법을 설명할게요.

출처: vLLM 공식 문서 — custom_logitsprocs

⚠️ 중요: 일부 로짓 프로세서 설계 변경이 아직 진행 중이며, API가 가까운 미래에 바뀔 수 있어요. 이 API 부분은 조만간 안정화되기를 기대하고 있습니다.

로짓 프로세서 배경 (Logits Processors Background)

로짓 프로세서는 다음 토큰의 확률 분포를 조정해서, 모델을 원하는 유형의 행동으로 이끄는 데 쓰여요.

vLLM에서 로짓 프로세서는 배치 단위(batch granularity) 로 동작합니다. 주어진 엔진 스텝 동안 로짓 프로세서는 모델이 출력한 (num_requests) x (vocab_size) 형태의 원시 로짓 텐서를 소비해요. 로짓 프로세서를 활성화한 모든 요청에 대해, 로짓 프로세서는 로짓 텐서의 해당 행에 변환을 적용하고 다른 행은 그대로 둡니다. 변환된 로짓 텐서는 이후 softmax로 전달되죠.

커스텀 로짓 프로세서 만들기 (Creating a Custom Logits Processor)

커스텀 로짓 프로세서는 vllm.v1.sample.logits_processor.LogitsProcessor를 상속하고, 최소한 다음 메서드들을 정의해야 해요.

  • validate_params(cls, sampling_params: SamplingParams):

    • SamplingParams에 로짓 프로세서가 쓰는 잘못된 인자(특히 커스텀 인자)가 있으면 ValueError를 발생시켜요.
    • 요청이 엔트리포인트로 보내지면 validate_params()SamplingParams를 검증하고 잘못된 인자를 가진 요청을 거부합니다.
    • 참고: 잘못된 파라미터가 커스텀 로짓 프로세서에서 예상치 못한 동작을 일으킬 수 있으므로 validate_params()를 구현하는 게 중요해요.
  • __init__(self, vllm_config: VllmConfig, device: torch.device, is_pin_memory: bool):

    • vllm_config: 엔진 설정 데이터 구조
    • device: 하드웨어 가속기 디바이스 정보
    • is_pin_memory: 로짓 프로세서 구현을 지원하는 pin memory 사용 가능 여부 플래그
  • apply(self, logits: torch.Tensor) -> torch.Tensor:

    • (num_requests) x (vocab_size) 로짓 텐서(logits)를 소비합니다.
    • 배치 단위로 로짓 프로세서 변환을 적용합니다.
    • 변환된 (num_requests) x (vocab_size) 로짓 텐서를 반환합니다.
    • 입력 로짓 프로세서를 in-place 또는 out-of-place로 수정할 수 있는데, in-place가 더 메모리 효율적이에요.
  • is_argmax_invariant(self) -> bool:

    • 로짓 프로세서가 argmax 불변(주어진 요청의 가장 높은 로짓 값을 가진 토큰 ID를 절대 바꾸지 않음)이면 True, argmax를 수정할 수 있으면 False를 반환해요.
    • is_argmax_invariant()는 시작 시 한 번 평가되며, True면 모든 요청이 greedy 샘플링을 쓸 때 vLLM은 해당 스텝에서 이 로짓 프로세서를 건너뜁니다.
  • update_state(self, batch_update: Optional["BatchUpdate"]) -> None:

    • 현재 엔진 스텝 시작 시의 영구적인 배치 상태 변화를 나타내는 BatchUpdate 데이터 구조를 소비합니다.
    • BatchUpdate 멤버로 로짓 프로세서의 내부 상태를 업데이트합니다.
    • 참고: batch update 데이터 구조가 None일 수도 있는데, 이는 배치 구성의 변화가 없다는 뜻이에요. 그 경우에도 LogitsProcessor는 추가될 때 보관해 둔 업데이트된 output_token_ids 목록을 바탕으로 상태를 업데이트하고 싶을 수 있어요.

vLLM 엔진이 BatchUpdate 구조를 만드는 방법

⚠️ 중요: 일부 로짓 프로세서 설계 변경이 아직 진행 중이에요. 미래에는 로짓 프로세서를 구현할 때 배치 상태 변화를 고려할 필요가 없어질 것이고, 이 섹션의 정보는 무의미해질 것으로 예상합니다.

로짓 프로세서의 update_state() 구현은 모델 러너가 영구 배치 상태를 업데이트하는 다음 모델(BatchUpdate 추상화로 표현)을 가정해야 해요.

  1. 현재 엔진 스텝에서 완료된 요청의 인덱스를 식별
  2. 현재 스텝에서 새로 들어온 요청을 식별
  3. Add 연산으로 가능한 한 많은 완료된 요청을 새 요청으로 교체 (가장 낮은 인덱스부터 순서대로)
  4. 새 요청과 완료된 요청의 상대적 개수에 따라:
    • 개수가 같으면 다음 스텝 진행
    • 새 요청이 더 많으면: 완료된 요청을 교체하지 않은 나머지 새 요청으로 배치를 확장. current_max_batch_index + 1부터 시작해 이 새 요청에 연속 인덱스 할당
    • 새 요청이 더 적으면:
      • 새 요청으로 교체되지 않은 완료된 요청에 Remove 연산 적용. 제거된 요청 인덱스는 이전 단계에서 교체된 완료 요청의 가장 큰 인덱스보다 큼. Remove는 배치를 비연속 상태로 만들 수 있음
      • 배치를 연속으로 "응축(condense)": 가장 낮은 인덱스의 빈 슬롯부터, 배치의 현재 가장 높은 비어있지 않은 슬롯에서 빈 슬롯으로 일방향 Move를 적용. 빈 슬롯 목적지 인덱스가 증가하고 비어있지 않은 슬롯 소스 인덱스가 감소하는 순서로 추가 일방향 Move를 배치가 연속이 될 때까지 진행
      • 배치 축소: 응축의 부수 효과로 Remove 연산에서 생긴 빈 슬롯이 배치 배열 끝에 연속 블록으로 모임. 따라서 응축 후 BatchUpdate.batch_size를 비어있지 않은 슬롯 수로 업데이트
  5. 효율성을 위해 배치 재정렬. 어텐션 백엔드 구현과 배치 특성에 따라 0개 이상의 Swap Move 연산이 배치 재정렬에 적용될 수 있음

참고사항:

  • 로짓 프로세서 update_state() 메서드는 배치 업데이트 연산을 removes, adds, moves 순서로 처리해야 해요.
  • Add 연산의 index 인자는 Add가 발생한 시점의 인덱스, 즉 Move 연산 이전의 인덱스를 가리켜요.
  • Move 연산은 BatchUpdate.moved에 나타나는 순서대로 적용된다고 가정할 수 있어요.
  • 새/완료 요청이 없고 배치 재정렬도 없다면, 로짓 프로세서의 배치 업데이트는 None이 됩니다.

커스텀 로짓 프로세서에 커스텀 인자 전달하기

내장 로짓 프로세서와 달리, 커스텀 로짓 프로세서는 SamplingParams나 vLLM 서버 REST API에 하드코딩되지 않은 구성 인자가 필요할 수 있어요. 이를 해결하기 위해 커스텀 로짓 프로세서는 vLLM의 custom arguments 지원을 활용해 사용자로부터 설정을 받을 수 있습니다 (SamplingParams의 기존 필드를 활용하도록 설계해도 자유로워요).

예시 커스텀 로짓 프로세서 구현

아래 예시는 (num_requests) × (vocab_size) 로짓 텐서를 소비해 target_token 하나를 제외한 모든 토큰을 float(-inf)로 마스킹하는 커스텀 로짓 프로세서예요. target_token을 지정하지 않은 요청에 대해서는 비활성화됩니다. 로짓 프로세서가 활성화됐는지, 어떤 토큰을 마스킹하지 않을지는 SamplingParams.extra_argstarget_token 커스텀 인자를 통해 판단합니다.

import torch
from vllm.config import VllmConfig
from vllm.sampling_params import SamplingParams
from vllm.v1.sample.logits_processor import (BatchUpdate,
                                            LogitsProcessor,
                                            MoveDirectionality)

class DummyLogitsProcessor(LogitsProcessor):
    """Fake logit processor to support unit testing and examples"""

    @classmethod
    def validate_params(cls, params: SamplingParams):
        target_token: int | None = params.extra_args and params.extra_args.get(
            "target_token"
        )
        if target_token is not None and not isinstance(target_token, int):
            raise ValueError(f"target_token value {target_token} is not int")

    def __init__(self, vllm_config: "VllmConfig", device: torch.device,
                is_pin_memory: bool):
        self.req_info: dict[int, int] = {}

    def is_argmax_invariant(self) -> bool:
        """Never impacts greedy sampling"""
        return False

    def update_state(self, batch_update: BatchUpdate | None):
        if not batch_update:
            return

        # Process added requests.
        for index, params, _, _ in batch_update.added:
            assert params is not None
            self.validate_params(params)
            if params.extra_args and (target_token :=
                                    params.extra_args.get("target_token")):
                self.req_info[index] = target_token
            else:
                self.req_info.pop(index, None)

        if self.req_info:
            # Process removed requests.
            for index in batch_update.removed:
                self.req_info.pop(index, None)

            # Process moved requests, unidirectional move (a->b) and swap
            # (a<->b)
            for adx, bdx, direct in batch_update.moved:
                a_val = self.req_info.pop(adx, None)
                b_val = self.req_info.pop(bdx, None)
                if a_val is not None:
                    self.req_info[bdx] = a_val
                if direct == MoveDirectionality.SWAP and b_val is not None:
                    self.req_info[adx] = b_val

    def apply(self, logits: torch.Tensor) -> torch.Tensor:
        if not self.req_info:
            return logits

        # Save target values before modification
        cols = torch.tensor(
            list(self.req_info.values()), dtype=torch.long, device=logits.device
        )
        rows = torch.tensor(
            list(self.req_info.keys()), dtype=torch.long, device=logits.device
        )
        values_to_keep = logits[rows, cols].clone()

        # Mask all but target tokens
        logits[rows] = float('-inf')
        logits[rows, cols] = values_to_keep

        return logits

DummyLogitsProcessor.update_state() 구현은 self.req_info 딕셔너리에 배치 요청의 "희소(sparse)" 표현을 유지해요. target_token 값을 지정한 요청만 딕셔너리에 키로 존재하죠. update_state()는 영구 배치에 대한 Add, Remove, Move 연산에 응답해 저장된 요청 인덱스와 target_token 값(각각 self.req_info의 키와 값)을 조정합니다.

기존 요청 레벨 로짓 프로세서 래핑하기 (Wrapping an Existing Request-Level Logits Processor)

vLLM 엔진은 로짓 프로세서를 배치 단위로 적용하지만, 어떤 사용자는 요청 레벨(request-level) 로짓 프로세서 구현, 즉 개별 요청에 대해 동작하는 구현을 함께 쓰고 싶을 수 있어요. 특히 vLLM 0 버전용으로 개발된 로짓 프로세서가 그런 경우가 많죠 (v0에서는 Callable이어야 했어요).

RequestLogitsProcessor = Union[
    # (output token ids, logits tensor) -> logits tensor
    Callable[[list[int], Tensor], Tensor],
    # (prompt token ids, output token ids, logits tensor) -> logits tensor
    Callable[[list[int], list[int], Tensor], Tensor],
]

요청 레벨 로짓 프로세서는 vLLM 엔진에서 명시적으로 지원되지 않지만, vLLM은 기존 Callable 요청 레벨 로짓 프로세서를 래핑해 vLLM과 호환되는 배치 레벨 로짓 프로세서를 만드는 편리한 과정을 제공해요. Callable은 위의 타입 애너테이션을 따라야 하며, 다른 인터페이스라면 래핑을 위해 수정하거나 추가 래퍼 레이어를 구현해야 할 수 있어요.

AdapterLogitsProcessor를 서브클래싱해서 요청 레벨 로짓 프로세서를 래핑할 수 있어요.

  • AdapterLogitsProcessor.validate_params(cls, params)를 오버라이드해 요청의 샘플링 파라미터를 검증
  • AdapterLogitsProcessor.is_argmax_invariant(self)를 오버라이드해 요청 레벨 로짓 프로세서가 가장 높은 값의 로짓을 가진 토큰에 영향을 줄 수 있는지 정확히 반영
  • AdapterLogitsProcessor.new_req_logits_processor(self, params)를 오버라이드해 SamplingParams 인스턴스에서 새 요청 레벨 로짓 프로세서 인스턴스 생성
from vllm.v1.sample.logits_processor import (
    AdapterLogitsProcessor, # Wrapper base-class
    RequestLogitsProcessor, # Request-level logitsproc type annotation
)

# Stand-in for your request-level logits processor:
class DummyPerReqLogitsProcessor:
    """The request-level logits processor masks out all logits except the
    token id identified by `target_token`"""

    def __init__(self, target_token: int) -> None:
        """Specify `target_token`"""
        self.target_token = target_token

    def __call__(
        self,
        output_ids: list[int],
        logits: torch.Tensor,
    ) -> torch.Tensor:
        val_to_keep = logits[self.target_token].item()
        logits[:] = float("-inf")
        logits[self.target_token] = val_to_keep
        return logits

# Example of wrapping the request-level logits processor:
class WrappedPerReqLogitsProcessor(AdapterLogitsProcessor):
    """Example of wrapping a fake request-level logit processor to create a
    batch-level logits processor"""

    @classmethod
    def validate_params(cls, params: SamplingParams):
        target_token: Any | None = params.extra_args and params.extra_args.get(
            "target_token"
        )
        if target_token is not None and not isinstance(target_token, int):
            raise ValueError(
                f"target_token value {target_token} is not int"
            )

    def is_argmax_invariant(self) -> bool:
        return False

    def new_req_logits_processor(
        self,
        params: SamplingParams,
    ) -> Optional[RequestLogitsProcessor]:
        """This method returns a new request-level logits processor, customized
        to the `target_token` value associated with a particular request.

        Returns None if the logits processor should not be applied to the
        particular request.
        """
        target_token: Any | None = params.extra_args and params.extra_args.get(
            "target_token"
        )
        if target_token is None:
            return None
        return DummyPerReqLogitsProcessor(target_token)

📌 참고: new_req_logits_processor() 오버라이드는 None을 반환해 래핑된 로짓 프로세서가 해당 요청에 적용되지 않아야 한다는 신호를 보낼 수 있어요.

vLLM에서 커스텀 로짓 프로세서 로드하는 방법

로짓 프로세서는 초기화 시 로드됩니다. 중요한 점은, 로드된 로짓 프로세서 집합은 vLLM 엔진 로딩이 끝난 뒤에는 수정할 수 없고, 개별 요청에 대한 주문형 로드는 불가능하다는 것이에요.

방법 1: 초기화 시 FQCN(정규화된 클래스 이름) 전달

이 방법은 vLLM 오프라인과 온라인 사용 시나리오 모두에서 지원돼요. 커스텀 로짓 프로세서의 FQCN(dotted.path.to.module:ClassName 형태)을 LLMAsyncLLM Python 생성자에 인자로, 또는 vllm serve의 CLI 인자로 전달할 수 있어요.

vllm serve ... --logits_processors <logits processor 1> <logits processor 2> ...

FQCN에 대한 요구사항은 (1) Python의 importlib.import_module()이 FQCN의 dotted path 부분을 모듈로 로드할 수 있어야 하고, (2) FQCN의 클래스 이름 부분을 로드된 모듈에서 import할 수 있어야 하며, (3) FQCN이 가리키는 객체가 LogitsProcessor의 서브클래스여야 한다는 것뿐이에요.

# Pass in FQCN
llm = LLM(
    model="facebook/opt-125m",
    logits_processors=["your.module.path:DummyLogitsProcessor"],
)
# Pass in FQCN
engine_args = AsyncEngineArgs(model="facebook/opt-125m",
                              logits_processors=["your.module.path:DummyLogitsProcessor"])
async_llm = AsyncLLM.from_engine_args(engine_args)
vllm serve facebook/opt-125m --logits_processors your.module.path:DummyLogitsProcessor

방법 2: 엔트리포인트로 자동 감지

setuptools는 설치된 패키지가 "엔트리 포인트(entry points)"라는 메타데이터를 통해 다른 Python 프로그램에 플러그인으로 자신을 노출할 수 있게 해줘요.

초기화 중에 vLLM은 vllm.logits_processors 엔트리 포인트 그룹을 자동으로 스캔해 발견한 설치된 로짓 프로세서를 로드합니다. 커스텀 로짓 프로세서를 담은 Python 패키지를 개발했다면, 각 로짓 프로세서에 고유한 엔트리포인트를 추가해 노출하면 돼요.

[project.entry-points."vllm.logits_processors"]
dummy_logits_processor = "your.module.path:DummyLogitsProcessor"

패키지가 설치되면 vLLM이 초기화될 때마다 커스텀 로짓 프로세서가 자동으로 로드됩니다. 엔트리 포인트로 노출했다면 생성자나 서버에 명시적으로 전달할 필요가 없어요.

📌 참고: vLLM은 vllm.logits_processors 그룹 아래 엔트리포인트로 노출된 모든 로짓 프로세서를 항상 로드합니다.

방법 3 (오프라인 전용): Python 클래스 객체 전달

하나 이상의 커스텀 로짓 프로세서 클래스 객체LLMAsyncLLM 생성자에 전달할 수 있어요. 이 옵션은 매우 유연한데, 클래스가 (1) LLM/AsyncLLM이 인스턴스화되는 같은 Python 소스 파일에 로컬로 정의되거나, (2) Python 패키지에서 import될 수 있기 때문이에요.

# Import custom logits processor
from some.module import DummyLogitsProcessor

# ...or...

# Define custom logits processor locally
from vllm.v1.sample.logits_processor import LogitsProcessor

class DummyLogitsProcessor(LogitsProcessor):
    # See DummyLogitsProcessor implementation above
    ...

# Pass class object to LLM constructor
llm = LLM(
    model="facebook/opt-125m",
    logits_processors=[DummyLogitsProcessor],
)

# Pass class object to AsyncLLM constructor
engine_args = AsyncEngineArgs(model="facebook/opt-125m",
                              logits_processors=[DummyLogitsProcessor])
async_llm = AsyncLLM.from_engine_args(engine_args)

요청에 커스텀 로짓 프로세서 호출하기 (Invoking a Custom Logits Processor)

커스텀 로짓 프로세서의 설계에 따라, 특정 요청에 대해 로짓 프로세서를 활성화/비활성화할지와 프로세서를 구성할 인자를 결정해야 해요.

아래 예시들은 DummyLogitsProcessor에 커스텀 인자(target_token)를 전달해서 (1) 해당 요청에 대해 로짓 프로세서를 활성화하고 (2) 동작을 제어하는 방법을 보여줘요.

vLLM REST API:

curl http://localhost:8000/v1/completions \
    -H "Content-Type: application/json" \
    -d '{
        "model": "Qwen/Qwen2.5-1.5B-Instruct",
        ...
        "vllm_xargs": {"target_token": 67}
    }'

OpenAI SDK:

batch = await client.completions.create(
    model="Qwen/Qwen2.5-1.5B-Instruct",
    ...,
    extra_body={
        "vllm_xargs": {
            "target_token": 67
        }
    }
)

오프라인 LLM 요청:

outputs_logitproc = llm.generate("your prompt",
                                 SamplingParams(...,
                                    extra_args={"target_token": 67}))

오프라인 AsyncLLM 요청:

async for out in engine.generate(request_id="your request id",
                                 prompt="your prompt",
                                 sampling_params=SamplingParams(...,
                                    extra_args={"target_token": 67})):

    # Process async request outputs
    ...

커스텀 로짓 프로세서 작성 모범 사례 (Best Practices)

vLLM이 초기화 중에 로짓 프로세서를 로드하면, 이후 모든 엔진 스텝에서 그 로짓 프로세서에 대해 update_state()apply()를 호출해요. 두 메서드 모두 현재 vLLM 영구 배치에 있는 모든 요청에 대해 동작하므로, 효율적으로 구현하는 게 중요합니다.

  • 로짓 프로세서가 배치 단위로 동작한다는 점을 고려해 효율적인 apply()update_state()를 작성하세요.
    • 예를 들어 apply()에 효율적인 벡터화 연산을 쓰거나 update_state()에서 내부 상태 벡터를 업데이트할 수 있어요.
    • 하지만 로짓 프로세서가 드물게 사용될 것 같다면 "희소(sparse)" 요청 상태 표현이 적절할 수 있어요. 클래스가 로짓 프로세서를 활성화한 요청에 대한 메타데이터만 저장하는 딕셔너리로 요청 구성을 표현하는 방식이죠.
    • 참고: 래핑된 요청 레벨 로짓 프로세서는 apply()update_state()를 구현할 필요가 없어요. 기본 AdapterLogitsProcessor.update_state() 구현은 요청 상태의 희소 표현을 유지하고, 기본 AdapterLogitsProcessor.apply() 구현은 요청 레벨 로짓 프로세서를 입력 로짓의 각 행에 순차 적용해 출력 로짓 텐서를 조립합니다. 이 기본 구현의 성능이 충분하지 않다면, 요청 레벨 로짓 프로세서를 래핑하지 말고 배치 단위로 동작하는 최적화된 apply()/update_state()를 가진 LogitsProcessor 서브클래스로 다시 구현하는 편이 좋아요.
  • 로짓 프로세서 작성자는 다음을 결정해야 해요: (1) 해당 요청에 대한 로짓 프로세서 동작을 구성하는 요청별 속성 (예: 위의 target_token), 그리고 (2) 그 구성 인자가 없는 요청에 대해 어떻게 처리할지. 이 섹션 및 문서 전반이 그 설계 패턴을 보여줍니다.

더 알아보기 (Learn more)