diff --git a/src/engine.py b/src/engine.py index 3293f59..f80f3c9 100644 --- a/src/engine.py +++ b/src/engine.py @@ -138,7 +138,8 @@ class vLLMEngine: yield batch batch = "" batch_token_counter = 0 - + if batch: + yield batch def _initialize_config(self): quantization = self._get_quantization() diff --git a/src/handler.py b/src/handler.py index 7f82c47..f564e8c 100644 --- a/src/handler.py +++ b/src/handler.py @@ -1,7 +1,12 @@ +import os import runpod from utils import JobInput from engine import vLLMEngine +if not os.getenv("MODEL_NAME"): + os.environ["MODEL_NAME"] = "facebook/opt-125m" + os.environ["CUSTOM_CHAT_TEMPLATE"] = "{{ bos_token }}{% for message in messages %}{% if (message['role'] == 'user') != (loop.index0 % 2 == 0) %}{{ raise_exception('Conversation roles must alternate user/assistant/user/assistant/...') }}{% endif %}{% if message['role'] == 'user' %}{{ '[INST] ' + message['content'] + ' [/INST]' }}{% elif message['role'] == 'assistant' %}{{ message['content'] + eos_token}}{% else %}{{ raise_exception('Only user and assistant roles are supported!') }}{% endif %}{% endfor %}" + vllm_engine = vLLMEngine() async def handler(job):