diff --git a/Dockerfile b/Dockerfile index 039e2f1..0a578fe 100644 --- a/Dockerfile +++ b/Dockerfile @@ -16,12 +16,13 @@ RUN --mount=type=cache,target=/root/.cache/pip \ ADD src . ARG MODEL_NAME="" +ENV MODEL_NAME=$MODEL_NAME ARG MODEL_BASE_PATH="" -ARG TOKENIZER="" +ENV MODEL_BASE_PATH=$MODEL_BASE_PATH # Conditionally run download_model.py -RUN if [ -n "$MODEL_NAME" ] && [ -n "$MODEL_BASE_PATH"]; then \ - python3.11 /download_model.py --model $MODEL_NAME --download_dir $MODEL_BASE_PATH --tokenizer $TOKENIZER; \ +RUN if [ -n "$MODEL_NAME" ] && [ -n "$MODEL_BASE_PATH" ]; then \ + python3.11 /download_model.py --model $MODEL_NAME --download_dir $MODEL_BASE_PATH; \ fi # Start the handler diff --git a/src/download_model.py b/src/download_model.py index a9e1e0a..d0a0001 100644 --- a/src/download_model.py +++ b/src/download_model.py @@ -1,22 +1,16 @@ import argparse -from vllm import LLMEngine, SamplingParams, AsyncEngineArgs, utils +from vllm.model_executor.weight_utils import prepare_hf_model_weights if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--model", type=str) parser.add_argument("--download_dir", type=str) - parser.add_argument("--tokenizer", type=str, default=None) args = parser.parse_args() if not args.model or not args.download_dir: raise ValueError("Must specify model and download_dir") - - engine_args = AsyncEngineArgs( - model=args.model, - download_dir=args.download_dir, - tokenizer=args.tokenizer, - dtype="auto" - ) - - llm = LLMEngine.from_engine_args(engine_args) + prepare_hf_model_weights( + model_name_or_path = args.model, + cache_dir=args.download_dir, + ) \ No newline at end of file