From b42d45ce0f4c4eff71a9f1c665f33aa4e8d727b7 Mon Sep 17 00:00:00 2001 From: alpayariyak Date: Fri, 23 Feb 2024 03:46:47 +0000 Subject: [PATCH] Bug fixes, refactors --- src/config.py | 25 +++++++++---------------- src/engine.py | 2 +- 2 files changed, 10 insertions(+), 17 deletions(-) diff --git a/src/config.py b/src/config.py index c334887..ce70999 100644 --- a/src/config.py +++ b/src/config.py @@ -23,7 +23,7 @@ class EngineConfig: return quantization if quantization in ["awq", "squeezellm", "gptq"] else None def _initialize_config(self): - return { + args = { "model": self.model_name_or_path, "revision": self.model_revision, "download_dir": self.hf_home, @@ -36,23 +36,16 @@ class EngineConfig: "disable_log_requests": bool(int(os.getenv("DISABLE_LOG_REQUESTS", 1))), "trust_remote_code": bool(int(os.getenv("TRUST_REMOTE_CODE", 0))), "gpu_memory_utilization": float(os.getenv("GPU_MEMORY_UTILIZATION", 0.95)), - "max_parallel_loading_workers": self._get_max_parallel_loading_workers(), - "max_model_len": self._get_max_model_len(), + "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": int(os.getenv("MAX_MODEL_LENGTH")) if os.getenv("MAX_MODEL_LENGTH") else None, "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"), - "block_size": int(os.getenv("BLOCK_SIZE")), - "swap_space": int(os.getenv("SWAP_SPACE")), - "max_context_len_to_capture": int(os.getenv("MAX_CONTEXT_LEN_TO_CAPTURE")), + "block_size": int(os.getenv("BLOCK_SIZE")) if os.getenv("BLOCK_SIZE") else None, + "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")) if os.getenv("MAX_CONTEXT_LEN_TO_CAPTURE") else None, "disable_custom_all_reduce": bool(int(os.getenv("DISABLE_CUSTOM_ALL_REDUCE", 0))), "enforce_eager": bool(int(os.getenv("ENFORCE_EAGER", 0))) } - - def _get_max_parallel_loading_workers(self): - 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 \ No newline at end of file + + return {k: v for k, v in args.items() if v is not None} \ No newline at end of file diff --git a/src/engine.py b/src/engine.py index bf6b54d..010c562 100644 --- a/src/engine.py +++ b/src/engine.py @@ -21,7 +21,7 @@ class vLLMEngine: def __init__(self, engine = None): load_dotenv() # For local development 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.max_concurrency = int(os.getenv("MAX_CONCURRENCY", DEFAULT_MAX_CONCURRENCY)) self.default_batch_size = int(os.getenv("DEFAULT_BATCH_SIZE", DEFAULT_BATCH_SIZE))