[object Object]

← back to Exo

fix headless mode with --disable-tui

ae35ada19bb8fd4050f5ddf936923ba9435667a4 · 2024-08-21 19:29:10 +0100 · Alex Cheema

Files touched

Diff

commit ae35ada19bb8fd4050f5ddf936923ba9435667a4
Author: Alex Cheema <alexcheema123@gmail.com>
Date:   Wed Aug 21 19:29:10 2024 +0100

    fix headless mode with --disable-tui
---
 main.py | 10 ++++++----
 1 file changed, 6 insertions(+), 4 deletions(-)

diff --git a/main.py b/main.py
index f8320135..ad969a9a 100644
--- a/main.py
+++ b/main.py
@@ -4,6 +4,8 @@ import signal
 import json
 import time
 import traceback
+import uuid
+from asyncio import CancelledError
 from exo.orchestration.standard_node import StandardNode
 from exo.networking.grpc.grpc_server import GRPCServer
 from exo.networking.grpc.grpc_discovery import GRPCDiscovery
@@ -18,7 +20,6 @@ from exo.inference.tokenizers import resolve_tokenizer
 from exo.orchestration.node import Node
 from exo.models import model_base_shards
 from exo.viz.topology_viz import TopologyViz
-import uuid
 
 # parse args
 parser = argparse.ArgumentParser(description="Initialize GRPC Discovery")
@@ -81,9 +82,9 @@ node = StandardNode(
 )
 server = GRPCServer(node, args.node_host, args.node_port)
 node.server = server
-api = ChatGPTAPI(node, inference_engine.__class__.__name__, response_timeout_secs=args.chatgpt_api_response_timeout_secs, on_chat_completion_request=lambda req_id, __, prompt: topology_viz.update_prompt(req_id, prompt))
+api = ChatGPTAPI(node, inference_engine.__class__.__name__, response_timeout_secs=args.chatgpt_api_response_timeout_secs, on_chat_completion_request=lambda req_id, __, prompt: topology_viz.update_prompt(req_id, prompt) if topology_viz else None)
 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 hasattr(inference_engine, "tokenizer") else tokens)
+    lambda req_id, tokens, __: topology_viz.update_prompt_output(req_id, inference_engine.tokenizer.decode(tokens) if hasattr(inference_engine, "tokenizer") else tokens) if topology_viz else None
 )
 def preemptively_start_download(request_id: str, opaque_status: str):
     try:
@@ -131,7 +132,8 @@ async def run_model_cli(node: Node, inference_engine: InferenceEngine, model_nam
     request_id = str(uuid.uuid4())
     callback_id = f"cli-wait-response-{request_id}"
     callback = node.on_token.register(callback_id)
-    topology_viz.update_prompt(request_id, prompt)
+    if topology_viz:
+        topology_viz.update_prompt(request_id, prompt)
     prompt = tokenizer.apply_chat_template([{"role": "user", "content": prompt}], tokenize=False, add_generation_prompt=True)
 
     try:

← b95916e0 show prompts and outputs in tui  ·  back to Exo  ·  ignore CancelledError when stopping the server 2e123335 →