← back to Exo
changes to inference engine
6dd2f7ab899168a04f4fdf5a4890186f2d34af5c · 2024-10-22 10:13:07 -0700 · josh
Files touched
M exo/main.pyM exo/orchestration/standard_node.py
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 →