[object Object]

← back to Exo

pr suggestion fixes:

8ad70b20b0d761914c03d8b8ae9f7e34a0007d5a · 2024-11-19 01:36:18 -0800 · josh

Files touched

Diff

commit 8ad70b20b0d761914c03d8b8ae9f7e34a0007d5a
Author: josh <eyasufikru567@gmail.com>
Date:   Tue Nov 19 01:36:18 2024 -0800

    pr suggestion fixes:
---
 exo/api/chatgpt_api.py        |  2 +-
 exo/download/hf/hf_helpers.py |  4 ++--
 exo/main.py                   | 12 ++++++++----
 3 files changed, 11 insertions(+), 7 deletions(-)

diff --git a/exo/api/chatgpt_api.py b/exo/api/chatgpt_api.py
index 47f280c4..b3fb71c4 100644
--- a/exo/api/chatgpt_api.py
+++ b/exo/api/chatgpt_api.py
@@ -187,7 +187,7 @@ class ChatGPTAPI:
     self.app.middlewares.append(self.log_request)
   
   async def handle_quit(self, request):
-    print("Received quit signal")
+    if DEBUG>=1: print("Received quit signal")
     response = web.json_response({"detail": "Quit signal received"}, status=200)
     await response.prepare(request)
     await response.write_eof()
diff --git a/exo/download/hf/hf_helpers.py b/exo/download/hf/hf_helpers.py
index a07a060e..acc2f67d 100644
--- a/exo/download/hf/hf_helpers.py
+++ b/exo/download/hf/hf_helpers.py
@@ -103,9 +103,9 @@ def get_repo_root(repo_id: str) -> Path:
   sanitized_repo_id = str(repo_id).replace("/", "--")
   return get_hf_home()/"hub"/f"models--{sanitized_repo_id}"
 
-async def move_models_to_hf():
+async def move_models_to_hf(seed_dir: Union[str, Path]):
   """Move model in resources folder of app to .cache/huggingface/hub"""
-  source_dir = Path(sys.argv[0]).parent
+  source_dir = Path(seed_dir)
   dest_dir = get_hf_home()/"hub"
   await aios.makedirs(dest_dir, exist_ok=True)
   for path in source_dir.iterdir():
diff --git a/exo/main.py b/exo/main.py
index 9ac5bd1e..ccb016e0 100644
--- a/exo/main.py
+++ b/exo/main.py
@@ -36,7 +36,7 @@ parser.add_argument("model_name", nargs="?", help="Model name to run")
 parser.add_argument("--node-id", type=str, default=None, help="Node ID")
 parser.add_argument("--node-host", type=str, default="0.0.0.0", help="Node host")
 parser.add_argument("--node-port", type=int, default=None, help="Node port")
-parser.add_argument("--model-seed-dir", type=str, default=None, help="Model seed directory")
+parser.add_argument("--models-seed-dir", type=str, default=None, help="Model seed directory")
 parser.add_argument("--listen-port", type=int, default=5678, help="Listening port for discovery")
 parser.add_argument("--download-quick-check", action="store_true", help="Quick check local path for model shards download")
 parser.add_argument("--max-parallel-downloads", type=int, default=4, help="Max parallel downloads for model shards download")
@@ -131,10 +131,14 @@ node.on_token.register("update_topology_viz").on_next(
   lambda req_id, tokens, __: topology_viz.update_prompt_output(req_id, inference_engine.tokenizer.decode(tokens)) if topology_viz and hasattr(inference_engine, "tokenizer") else None
 )
 
-if not args.model_seed_dir is None:
+if not args.models_seed_dir is None:
   try:
-    await move_models_to_hf()
-  except:
+    if is_frozen():
+      seed_dir = Path(sys.argv[0]).parent
+      await move_models_to_hf(seed_dir)
+    else:
+      await move_models_to_hf(args.model_seed_dir)
+  except Exception as e:
     print(f"Error moving models to .cache/huggingface: {e}")
 
 def preemptively_start_download(request_id: str, opaque_status: str):

← 06c3f524 removed response return  ·  back to Exo  ·  changes to args 65817ab7 →