← back to Exo
Abstract load checkpoint method
6aaea8c74ce61766f8f442a29d48f945588ee727 · 2024-12-08 03:10:31 -0800 · Nel Nibcord
Files touched
M exo/inference/dummy_inference_engine.pyM exo/inference/inference_engine.pyM exo/inference/tinygrad/inference.py
Diff
commit 6aaea8c74ce61766f8f442a29d48f945588ee727
Author: Nel Nibcord <blindcrone@tuta.io>
Date: Sun Dec 8 03:10:31 2024 -0800
Abstract load checkpoint method
---
exo/inference/dummy_inference_engine.py | 3 +++
exo/inference/inference_engine.py | 4 ++++
exo/inference/tinygrad/inference.py | 3 +++
3 files changed, 10 insertions(+)
diff --git a/exo/inference/dummy_inference_engine.py b/exo/inference/dummy_inference_engine.py
index 09026109..daf4b677 100644
--- a/exo/inference/dummy_inference_engine.py
+++ b/exo/inference/dummy_inference_engine.py
@@ -32,3 +32,6 @@ class DummyInferenceEngine(InferenceEngine):
async def ensure_shard(self, shard: Shard):
if self.shard == shard: return
self.shard = shard
+
+ async def load_checkpoint(self, shard: Shard, path: str):
+ await self.ensure_shard(shard)
diff --git a/exo/inference/inference_engine.py b/exo/inference/inference_engine.py
index 7c71e6ef..53dbd923 100644
--- a/exo/inference/inference_engine.py
+++ b/exo/inference/inference_engine.py
@@ -25,6 +25,10 @@ class InferenceEngine(ABC):
@abstractmethod
async def infer_tensor(self, request_id: str, shard: Shard, input_data: np.ndarray) -> np.ndarray:
pass
+
+ @abstractmethod
+ async def load_checkpoint(self, shard: Shard, path: str):
+ pass
async def save_session(self, key, value):
self.session[key] = value
diff --git a/exo/inference/tinygrad/inference.py b/exo/inference/tinygrad/inference.py
index 343deec3..95bee691 100644
--- a/exo/inference/tinygrad/inference.py
+++ b/exo/inference/tinygrad/inference.py
@@ -92,6 +92,9 @@ class TinygradDynamicShardInferenceEngine(InferenceEngine):
tokens = await asyncio.get_running_loop().run_in_executor(self.executor, self.tokenizer.decode, tokens)
return tokens
+ async def load_checkpoint(self, shard: Shard, path: str):
+ await self.ensure_shard(shard)
+
async def infer_tensor(self, request_id: str, shard: Shard, input_data: np.ndarray) -> np.ndarray:
await self.ensure_shard(shard)
def wrap_infer():
← 2a3a2e5e circular include lol
·
back to Exo
·
Removed ensure_session to clean stuff up. May revisit later 0673d645 →