[object Object]

← back to Exo

simplify mlx non blocking

874886abc48ecd7eb039c8a650668a6f45137ad2 · 2024-09-05 17:50:54 +0100 · Alex Cheema

Files touched

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 →