Text Generation Inference (TGI)

Text Generation Inference (TGI)

HuggingFace의 Text Generation Inference(TGI)를 Vast.ai 서버리스로 돌리면, 대규모 언어 모델 추론을 쉽게 서빙할 수 있어요. TGI 서버리스 템플릿은 필요한 환경 변수만 채우면 바로 쓰기 좋게 준비되어 있어요. 여기서는 환경 변수 설정과 /generate/, /generate_stream/ 두 엔드포인트를 다루는 방법을 정리해 볼게요.

출처: Vast.ai 공식 문서 — Text Generation Inference (TGI)

환경 변수 (Environment Variables)

  • HF_TOKEN(string): gated 모델을 내려받기 위한 읽기 권한 HuggingFace API 토큰이에요. 자세한 내용은 HuggingFace 토큰 문서를 참고하세요.
  • MODEL_ID(string): 추론에 사용할 모델의 ID예요. 지원되는 HuggingFace 모델은 여기에서 확인할 수 있어요.

Vast.ai SDK 설치

vastai pip 패키지가 설치되어 있는지 확인하세요.

pip install vastai

API 키 설정

Vast.ai Serverless API 키를 VAST_API_KEY 환경 변수로 설정해요.

export VAST_API_KEY=<your-api-key>

/generate/ 사용하기

/generate/는 단일 프롬프트에 대한 텍스트 생성을 요청하는 엔드포인트예요. 페이로드의 inputs에 프롬프트를, parameters에 생성 옵션을 담아 보내요.

import asyncio
from vastai import Serverless

MAX_TOKENS = 128

async def main():
    async with Serverless() as client:
        endpoint = await client.get_endpoint(name="my-tgi-endpoint")

        prompt = "Who are you?"

        payload = {
            "inputs": prompt,
            "parameters": {
                "max_new_tokens": MAX_TOKENS,
                "temperature": 0.7,
                "return_full_text": False
            }
        }

        resp = await endpoint.request("/generate", payload, cost=MAX_TOKENS)

if __name__ == "__main__":
    asyncio.run(main())

/generate_stream/ 사용하기

/generate_stream/는 생성되는 대로 토큰을 스트리밍으로 받는 엔드포인트예요. 요청 시 stream=True를 넘기면 응답의 response에서 스트림을 직접 순회하며 토큰을 꺼내 쓸 수 있어요.

import asyncio
from vastai import Serverless

MAX_TOKENS = 1024

def build_prompt(system_prompt: str, user_prompt: str) -> str:
    return (
        f"<<SYS>>\n{system_prompt.strip()}\n<</SYS>>\n\n"
        f"User: {user_prompt.strip()}\n"
        f"Assistant:"
    )

async def main():
    async with Serverless() as client:
        endpoint = await client.get_endpoint(name="my-tgi-endpoint")

        system_prompt = (
            "You are Qwen.\n"
            "You are to only speak in English.\n"
        )
        user_prompt = """
        Critically analyze the extent to which hotdogs are sandwiches.
        """

        prompt = build_prompt(system_prompt, user_prompt)

        payload = {
            "inputs": prompt,
            "parameters": {
                "max_new_tokens": MAX_TOKENS,
                "temperature": 0.7,
                "do_sample": True,
                "return_full_text": False,
            }
        }

        resp = await endpoint.request(
            "/generate_stream",
            payload,
            cost=MAX_TOKENS,
            stream=True,
        )
        stream = resp["response"]

        printed_answer = False
        async for event in stream:
            tok = (event.get("token") or {}).get("text")
            if tok:
                if not printed_answer:
                    printed_answer = True
                    print("Answer:\n")
                print(tok, end="", flush=True)

if __name__ == "__main__":
    asyncio.run(main())

TGI에서 지원하는 추론 파라미터 전체 목록은 TGI의 추론 API 문서를 참고하세요.

더 알아보기 (Learn more)