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