← back to Exo
Updated unit tests
b787c676de0e525c35f4b223672ee6d501f992a8 · 2024-11-12 18:36:57 -0800 · Nel Nibcord
Files touched
M exo/inference/test_inference_engine.py
Diff
commit b787c676de0e525c35f4b223672ee6d501f992a8
Author: Nel Nibcord <blindcrone@tuta.io>
Date: Tue Nov 12 18:36:57 2024 -0800
Updated unit tests
---
exo/inference/test_inference_engine.py | 4 ++--
1 file changed, 2 insertions(+), 2 deletions(-)
diff --git a/exo/inference/test_inference_engine.py b/exo/inference/test_inference_engine.py
index f86a9085..690ed610 100644
--- a/exo/inference/test_inference_engine.py
+++ b/exo/inference/test_inference_engine.py
@@ -42,7 +42,7 @@ async def test_inference_engine(inference_engine_1: InferenceEngine, inference_e
assert np.array_equal(next_resp_full, resp4)
-asyncio.run(test_inference_engine(MLXDynamicShardInferenceEngine(HFShardDownloader()), MLXDynamicShardInferenceEngine(HFShardDownloader()), "mlx-community/Llama-3.2-1B-Instruct-4bit", 16))
+asyncio.run(test_inference_engine(MLXDynamicShardInferenceEngine(HFShardDownloader()), MLXDynamicShardInferenceEngine(HFShardDownloader()), "llama-3.2-1b", 16))
if os.getenv("RUN_TINYGRAD", default="0") == "1":
import tinygrad
@@ -50,5 +50,5 @@ if os.getenv("RUN_TINYGRAD", default="0") == "1":
from exo.inference.tinygrad.inference import TinygradDynamicShardInferenceEngine
tinygrad.helpers.DEBUG.value = int(os.getenv("TINYGRAD_DEBUG", default="0"))
asyncio.run(
- test_inference_engine(TinygradDynamicShardInferenceEngine(HFShardDownloader()), TinygradDynamicShardInferenceEngine(HFShardDownloader()), "TriAiExperiments/SFR-Iterative-DPO-LLaMA-3-8B-R", 32)
+ test_inference_engine(TinygradDynamicShardInferenceEngine(HFShardDownloader()), TinygradDynamicShardInferenceEngine(HFShardDownloader()), "llama-3-8b", 32)
)
← 6d12deab add better error handling:
·
back to Exo
·
Added a small script to compile grpc 9712d696 →