Download logic + other changes
This commit is contained in:
@@ -45,8 +45,5 @@ ENV MODEL_NAME=$MODEL_NAME \
|
||||
# Set the entrypoint
|
||||
ENTRYPOINT ["/entrypoint.sh"]
|
||||
|
||||
# Run the Python script to download the model
|
||||
RUN python3.11 -u /download_model.py --model_name $MODEL_NAME --model_revision $MODEL_REVISION --model_base_path $MODEL_BASE_PATH
|
||||
|
||||
# Start the handler
|
||||
CMD STREAMING=$STREAMING MODEL_NAME=$MODEL_NAME MODEL_BASE_PATH=$MODEL_BASE_PATH TOKENIZER=$TOKENIZER QUANTIZATION=$QUANTIZATION python3.11 /handler.py
|
||||
|
||||
@@ -16,10 +16,11 @@
|
||||
- `STREAMING`: Whether to use HTTP Streaming or not.
|
||||
More information on receiving streaming responses from Serverless Endpoints can be found at [Endpoint URLs](https://docs.runpod.io/docs/serverless-endpoint-urls#streamjob_id), and a detailed example at [Llama2 7B Chat | Streaming Token Outputs](https://docs.runpod.io/reference/llama2-7b-chat#streaming-token-outputs).
|
||||
#### Optional:
|
||||
- `HUGGING_FACE_HUB_TOKEN`: Your Hugging Face token to access private or gated models. You can get your token [here](https://huggingface.co/settings/token).
|
||||
- `TOKENIZER`: The specified tokenizer to use. If you want to use the default tokenizer for the model, do not provide this docker argument at all.
|
||||
- `QUANTIZATION`: `awq` to use AWQ Quantization (Base model must be in AWQ format). `squeezellm` for SqueezeLLM quantization - preliminary support.
|
||||
|
||||
### Environment Variables:
|
||||
#### Optional:
|
||||
- `HUGGING_FACE_HUB_TOKEN`: Your Hugging Face token to access private or gated models. You can get your token [here](https://huggingface.co/settings/token).
|
||||
### Compatible Models
|
||||
- LLaMA & LLaMA-2
|
||||
- Mistral
|
||||
|
||||
@@ -1,24 +0,0 @@
|
||||
import os
|
||||
import argparse
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
|
||||
MODEL_NAME = os.environ.get('MODEL_NAME')
|
||||
MODEL_REVISION = os.environ.get('MODEL_REVISION', "main")
|
||||
MODEL_BASE_PATH = os.environ.get('MODEL_BASE_PATH', '/runpod-volume/')
|
||||
HUGGING_FACE_HUB_TOKEN = os.environ.get('HUGGING_FACE_HUB_TOKEN')
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--model_name', type=str, default=MODEL_NAME)
|
||||
parser.add_argument('--model_revision', type=str, default=MODEL_REVISION)
|
||||
parser.add_argument('--model_base_path', type=str, default=MODEL_BASE_PATH)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
snapshot_download(
|
||||
args.model_name,
|
||||
revision=args.model_revision,
|
||||
local_dir=f"{args.model_base_path}{args.model_name.split('/')[1]}",
|
||||
)
|
||||
+8
-4
@@ -12,12 +12,15 @@ import os
|
||||
|
||||
# Prepare the model and tokenizer
|
||||
MODEL_NAME = os.environ.get('MODEL_NAME')
|
||||
MODEL_NAME = MODEL_NAME.replace(".", "_")
|
||||
MODEL_BASE_PATH = os.environ.get('MODEfL_BASE_PATH', '/runpod-volume/')
|
||||
STREAMING = os.environ.get('STREAMING', False) == 'True'
|
||||
TOKENIZER = os.environ.get('TOKENIZER', None)
|
||||
USE_FULL_METRICS = os.environ.get('USE_FULL_METRICS', True)
|
||||
DTYPE = "auto"
|
||||
|
||||
MODEL_BASE_PATH = os.environ.get('MODEL_BASE_PATH')
|
||||
if not os.path.exists(BASE_VOLUME):
|
||||
os.makedirs(BASE_VOLUME)
|
||||
|
||||
USE_FULL_METRICS = os.environ.get('USE_FULL_METRICS', True)
|
||||
MAX_CONCURRENCY = os.environ.get('MAX_CONCURRENCY', 200)
|
||||
TOTAL_RUNNING_JOBS = 0
|
||||
|
||||
@@ -46,7 +49,8 @@ except ValueError:
|
||||
|
||||
# Prepare the engine's arguments
|
||||
engine_args = AsyncEngineArgs(
|
||||
model=f"{MODEL_BASE_PATH}{MODEL_NAME.split('/')[1]}",
|
||||
model=MODEL_NAME,
|
||||
download_dir=MODEL_BASE_PATH,
|
||||
tokenizer=TOKENIZER,
|
||||
tokenizer_mode="auto",
|
||||
tensor_parallel_size=NUM_GPU_SHARD,
|
||||
|
||||
Reference in New Issue
Block a user