[object Object]

← back to Exo

Abstract load checkpoint method

6aaea8c74ce61766f8f442a29d48f945588ee727 · 2024-12-08 03:10:31 -0800 · Nel Nibcord

Files touched

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 →