Bug fixes, refactors
This commit is contained in:
+8
-15
@@ -23,7 +23,7 @@ class EngineConfig:
|
|||||||
return quantization if quantization in ["awq", "squeezellm", "gptq"] else None
|
return quantization if quantization in ["awq", "squeezellm", "gptq"] else None
|
||||||
|
|
||||||
def _initialize_config(self):
|
def _initialize_config(self):
|
||||||
return {
|
args = {
|
||||||
"model": self.model_name_or_path,
|
"model": self.model_name_or_path,
|
||||||
"revision": self.model_revision,
|
"revision": self.model_revision,
|
||||||
"download_dir": self.hf_home,
|
"download_dir": self.hf_home,
|
||||||
@@ -36,23 +36,16 @@ class EngineConfig:
|
|||||||
"disable_log_requests": bool(int(os.getenv("DISABLE_LOG_REQUESTS", 1))),
|
"disable_log_requests": bool(int(os.getenv("DISABLE_LOG_REQUESTS", 1))),
|
||||||
"trust_remote_code": bool(int(os.getenv("TRUST_REMOTE_CODE", 0))),
|
"trust_remote_code": bool(int(os.getenv("TRUST_REMOTE_CODE", 0))),
|
||||||
"gpu_memory_utilization": float(os.getenv("GPU_MEMORY_UTILIZATION", 0.95)),
|
"gpu_memory_utilization": float(os.getenv("GPU_MEMORY_UTILIZATION", 0.95)),
|
||||||
"max_parallel_loading_workers": self._get_max_parallel_loading_workers(),
|
"max_parallel_loading_workers": None if device_count() > 1 or not os.getenv("MAX_PARALLEL_LOADING_WORKERS") else int(os.getenv("MAX_PARALLEL_LOADING_WORKERS")),
|
||||||
"max_model_len": self._get_max_model_len(),
|
"max_model_len": int(os.getenv("MAX_MODEL_LENGTH")) if os.getenv("MAX_MODEL_LENGTH") else None,
|
||||||
"tensor_parallel_size": device_count(),
|
"tensor_parallel_size": device_count(),
|
||||||
"seed": int(os.getenv("SEED")),
|
"seed": int(os.getenv("SEED")) if os.getenv("SEED") else None,
|
||||||
"kv_cache_dtype": os.getenv("KV_CACHE_DTYPE"),
|
"kv_cache_dtype": os.getenv("KV_CACHE_DTYPE"),
|
||||||
"block_size": int(os.getenv("BLOCK_SIZE")),
|
"block_size": int(os.getenv("BLOCK_SIZE")) if os.getenv("BLOCK_SIZE") else None,
|
||||||
"swap_space": int(os.getenv("SWAP_SPACE")),
|
"swap_space": int(os.getenv("SWAP_SPACE")) if os.getenv("SWAP_SPACE") else None,
|
||||||
"max_context_len_to_capture": int(os.getenv("MAX_CONTEXT_LEN_TO_CAPTURE")),
|
"max_context_len_to_capture": int(os.getenv("MAX_CONTEXT_LEN_TO_CAPTURE")) if os.getenv("MAX_CONTEXT_LEN_TO_CAPTURE") else None,
|
||||||
"disable_custom_all_reduce": bool(int(os.getenv("DISABLE_CUSTOM_ALL_REDUCE", 0))),
|
"disable_custom_all_reduce": bool(int(os.getenv("DISABLE_CUSTOM_ALL_REDUCE", 0))),
|
||||||
"enforce_eager": bool(int(os.getenv("ENFORCE_EAGER", 0)))
|
"enforce_eager": bool(int(os.getenv("ENFORCE_EAGER", 0)))
|
||||||
}
|
}
|
||||||
|
|
||||||
def _get_max_parallel_loading_workers(self):
|
return {k: v for k, v in args.items() if v is not None}
|
||||||
if device_count() > 1:
|
|
||||||
return None
|
|
||||||
return int(os.getenv("MAX_PARALLEL_LOADING_WORKERS"))
|
|
||||||
|
|
||||||
def _get_max_model_len(self):
|
|
||||||
max_model_len = os.getenv("MAX_MODEL_LENGTH")
|
|
||||||
return int(max_model_len) if max_model_len else None
|
|
||||||
+1
-1
@@ -21,7 +21,7 @@ class vLLMEngine:
|
|||||||
def __init__(self, engine = None):
|
def __init__(self, engine = None):
|
||||||
load_dotenv() # For local development
|
load_dotenv() # For local development
|
||||||
self.config = EngineConfig().config
|
self.config = EngineConfig().config
|
||||||
self.tokenizer = TokenizerWrapper(self.config["tokenizer"], self.config["tokenizer_revision"])
|
self.tokenizer = TokenizerWrapper(self.config.get("tokenizer"), self.config.get("tokenizer_revision"))
|
||||||
self.llm = self._initialize_llm() if engine is None else engine
|
self.llm = self._initialize_llm() if engine is None else engine
|
||||||
self.max_concurrency = int(os.getenv("MAX_CONCURRENCY", DEFAULT_MAX_CONCURRENCY))
|
self.max_concurrency = int(os.getenv("MAX_CONCURRENCY", DEFAULT_MAX_CONCURRENCY))
|
||||||
self.default_batch_size = int(os.getenv("DEFAULT_BATCH_SIZE", DEFAULT_BATCH_SIZE))
|
self.default_batch_size = int(os.getenv("DEFAULT_BATCH_SIZE", DEFAULT_BATCH_SIZE))
|
||||||
|
|||||||
Reference in New Issue
Block a user