[object Object]

← back to Exo

changes to inference engine

6dd2f7ab899168a04f4fdf5a4890186f2d34af5c · 2024-10-22 10:13:07 -0700 · josh

Files touched

Diff

commit 6dd2f7ab899168a04f4fdf5a4890186f2d34af5c
Author: josh <eyasunigussie@Eyasus-MacBook-Air.local>
Date:   Tue Oct 22 10:13:07 2024 -0700

    changes to inference engine
---
 exo/main.py                        | 13 ++++++-------
 exo/orchestration/standard_node.py |  3 +--
 2 files changed, 7 insertions(+), 9 deletions(-)

diff --git a/exo/main.py b/exo/main.py
index 92488912..7faf63c1 100644
--- a/exo/main.py
+++ b/exo/main.py
@@ -168,15 +168,11 @@ async def select_best_inference_engine(node: StandardNode):
           continue
   if any("tinygrad" in engines and len(engines) == 1 for engines in all_supported_engines):
       return "tinygrad"
-  common_engines_across_peers = set.intersection(*all_supported_engines)
-  with open('check_engines.txt', 'w') as f:
-    f.write(common_engines_across_peers)
-    f.close()
-  print(f'common_engines_across_peers:{common_engines_across_peers}')
-  if "mlx" in common_engines_across_peers:
+  common_engine_across_peers = set.intersection(*all_supported_engines)
+  if "mlx" in common_engine_across_peers:
       print('mlx')
       return "mlx"
-  elif "tinygrad" in common_engines_across_peers:
+  elif "tinygrad" in common_engine_across_peers:
       return "tinygrad"
   else:
       raise ValueError("No compatible inference engine found across all nodes")
@@ -221,6 +217,9 @@ async def main():
     loop.add_signal_handler(s, handle_exit)
 
   await node.start(wait_for_peers=args.wait_for_peers)
+  if len(node.peers) > 1:
+    compatible_engine = await select_best_inference_engine(node)
+    node.inference_engine = get_inference_engine(compatible_engine, shard_downloader)
 
   if args.command == "run" or args.run_model:
     model_name = args.model_name or args.run_model
diff --git a/exo/orchestration/standard_node.py b/exo/orchestration/standard_node.py
index 962432fc..54fba341 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -84,8 +84,7 @@ class StandardNode(Node):
       supported_engines.append('tinygrad')
     return supported_engines
 
-  async def broadcast_supported_engines(self):
-    supported_engines = self.get_supported_inference_engines()
+  async def broadcast_supported_engines(self, supported_engines: List):
     await self.broadcast_opaque_status("", json.dumps({
       "type": "supported_inference_engines",
       "node_id": self.id, 

← 3908b97a changes to broadcast func  ·  back to Exo  ·  Moving conditonal apple silicon logic to setup.py ae47fe11 →