← back to Exo
simplify mlx non blocking
874886abc48ecd7eb039c8a650668a6f45137ad2 · 2024-09-05 17:50:54 +0100 · Alex Cheema
Files touched
M exo/inference/mlx/sharded_inference_engine.py
Diff
commit 874886abc48ecd7eb039c8a650668a6f45137ad2
Author: Alex Cheema <alexcheema123@gmail.com>
Date: Thu Sep 5 17:50:54 2024 +0100
simplify mlx non blocking
---
exo/inference/mlx/sharded_inference_engine.py | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/exo/inference/mlx/sharded_inference_engine.py b/exo/inference/mlx/sharded_inference_engine.py
index 5d6378b7..7b920ccc 100644
--- a/exo/inference/mlx/sharded_inference_engine.py
+++ b/exo/inference/mlx/sharded_inference_engine.py
@@ -25,7 +25,7 @@ class MLXDynamicShardInferenceEngine(InferenceEngine):
input_ids = mx.array(inputs["input_ids"])
output_data: np.ndarray = np.array(await loop.run_in_executor(self.executor, self.stateful_sharded_model.step, request_id, input_ids, pixel_values))
else:
- input_ids = await loop.run_in_executor(self.executor, lambda: mx.array(self.tokenizer.encode(prompt)))
+ input_ids = mx.array(await loop.run_in_executor(self.executor, self.tokenizer.encode, prompt))
output_data: np.ndarray = np.array(await loop.run_in_executor(self.executor, self.stateful_sharded_model.step, request_id, input_ids))
return output_data, "", output_data.size == 1 and output_data.item() == self.tokenizer.eos_token_id
← e616d4e8 run realize on the result in tinygrad
·
back to Exo
·
fix broken links in README ca644562 →