← 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 →