Fixed CUDA 11.8 workers
This commit is contained in:
+1
-2
@@ -21,14 +21,13 @@ RUN --mount=type=cache,target=/root/.cache/pip \
|
||||
|
||||
# Install torch and vllm based on CUDA version
|
||||
RUN if [[ "${WORKER_CUDA_VERSION}" == 11.8* ]]; then \
|
||||
python3.11 -m pip install -e git+https://github.com/runpod/vllm-fork-for-sls-worker.git@cuda-11.8#egg=vllm; \
|
||||
python3.11 -m pip install -U --force-reinstall torch==2.1.2 xformers==0.0.23.post1 --index-url https://download.pytorch.org/whl/cu118; \
|
||||
python3.11 -m pip install -e git+https://github.com/runpod/vllm-fork-for-sls-worker.git@cuda-11.8#egg=vllm; \
|
||||
else \
|
||||
python3.11 -m pip install -e git+https://github.com/runpod/vllm-fork-for-sls-worker.git#egg=vllm; \
|
||||
fi && \
|
||||
rm -rf /root/.cache/pip
|
||||
|
||||
|
||||
# Add source files
|
||||
COPY src .
|
||||
|
||||
|
||||
@@ -2,3 +2,4 @@ hf_transfer
|
||||
runpod==1.4.2
|
||||
huggingface-hub
|
||||
packaging
|
||||
pydantic
|
||||
@@ -5,7 +5,7 @@ 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("--download_dir", type=str, default=os.environ.get("MODEL_BASE_PATH"))
|
||||
|
||||
args = parser.parse_args()
|
||||
if not args.model or not args.download_dir:
|
||||
|
||||
Reference in New Issue
Block a user