[object Object]

← back to Exo

allow overriding inference_engine and separate flag for TINYGRAD_DEBUG

945f90f676182a751d2ad7bcf20987ab7fe0181e · 2024-07-18 16:09:33 -0700 · Alex Cheema

Files touched

Diff

commit 945f90f676182a751d2ad7bcf20987ab7fe0181e
Author: Alex Cheema <alexcheema123@gmail.com>
Date:   Thu Jul 18 16:09:33 2024 -0700

    allow overriding inference_engine and separate flag for TINYGRAD_DEBUG
---
 main.py | 27 ++++++++++++++++++++++-----
 1 file changed, 22 insertions(+), 5 deletions(-)

diff --git a/main.py b/main.py
index b551bac4..52459f31 100644
--- a/main.py
+++ b/main.py
@@ -4,6 +4,7 @@ import signal
 import uuid
 import platform
 import psutil
+import os
 from typing import List
 from exo.orchestration.standard_node import StandardNode
 from exo.networking.grpc.grpc_server import GRPCServer
@@ -21,16 +22,32 @@ parser.add_argument("--listen-port", type=int, default=5678, help="Listening por
 parser.add_argument("--broadcast-port", type=int, default=5678, help="Broadcast port for discovery")
 parser.add_argument("--wait-for-peers", type=int, default=0, help="Number of peers to wait to connect to before starting")
 parser.add_argument("--chatgpt-api-port", type=int, default=8000, help="ChatGPT API port")
+parser.add_argument("--inference-engine", type=str, default=None, help="Inference engine to use")
 args = parser.parse_args()
 
 print_yellow_exo()
 print(f"Starting exo {platform.system()=} {psutil.virtual_memory()=}")
-if psutil.MACOS:
-    from exo.inference.mlx.sharded_inference_engine import MLXDynamicShardInferenceEngine
-    inference_engine = MLXDynamicShardInferenceEngine()
+if args.inference_engine is None:
+    if psutil.MACOS:
+        from exo.inference.mlx.sharded_inference_engine import MLXDynamicShardInferenceEngine
+        inference_engine = MLXDynamicShardInferenceEngine()
+    else:
+        from exo.inference.tinygrad.inference import TinygradDynamicShardInferenceEngine
+        import tinygrad.helpers
+        tinygrad.helpers.DEBUG.value = int(os.getenv("TINYGRAD_DEBUG", default="0"))
+        inference_engine = TinygradDynamicShardInferenceEngine()
 else:
-    from exo.inference.tinygrad.inference import TinygradDynamicShardInferenceEngine
-    inference_engine = TinygradDynamicShardInferenceEngine()
+    if args.inference_engine == "mlx":
+        from exo.inference.mlx.sharded_inference_engine import MLXDynamicShardInferenceEngine
+        inference_engine = MLXDynamicShardInferenceEngine()
+    elif args.inference_engine == "tinygrad":
+        from exo.inference.tinygrad.inference import TinygradDynamicShardInferenceEngine
+        import tinygrad.helpers
+        tinygrad.helpers.DEBUG.value = int(os.getenv("TINYGRAD_DEBUG", default="0"))
+        inference_engine = TinygradDynamicShardInferenceEngine()
+    else:
+        raise ValueError(f"Inference engine {args.inference_engine} not supported")
+print(f"Using inference engine {inference_engine.__class__.__name__}")
 
 discovery = GRPCDiscovery(args.node_id, args.node_port, args.listen_port, args.broadcast_port)
 node = StandardNode(args.node_id, None, inference_engine, discovery, partitioning_strategy=RingMemoryWeightedPartitioningStrategy())

← 47163d22 broadcast results concurrently fixes #31  ·  back to Exo  ·  reference the code for each feature listed in README 1b194b43 →