[object Object]

← back to Exo

Removed ensure_session to clean stuff up. May revisit later

0673d6452c880f9133aeecca123ba919879a0af7 · 2024-12-08 03:12:18 -0800 · Nel Nibcord

Files touched

Diff

commit 0673d6452c880f9133aeecca123ba919879a0af7
Author: Nel Nibcord <blindcrone@tuta.io>
Date:   Sun Dec 8 03:12:18 2024 -0800

    Removed ensure_session to clean stuff up. May revisit later
---
 exo/inference/inference_engine.py   | 6 ------
 exo/inference/tinygrad/inference.py | 5 -----
 2 files changed, 11 deletions(-)

diff --git a/exo/inference/inference_engine.py b/exo/inference/inference_engine.py
index 53dbd923..3a867466 100644
--- a/exo/inference/inference_engine.py
+++ b/exo/inference/inference_engine.py
@@ -33,12 +33,6 @@ class InferenceEngine(ABC):
   async def save_session(self, key, value):
     self.session[key] = value
   
-  async def ensure_session(self, key, check, value_gen, hook=None):
-    if key not in self.session or not check(self.session[key]):
-      await self.save_session(key, value_gen())
-      if hook is not None:
-        hook()
-  
   async def clear_session(self):
     self.session.empty()
   
diff --git a/exo/inference/tinygrad/inference.py b/exo/inference/tinygrad/inference.py
index 95bee691..edf08cea 100644
--- a/exo/inference/tinygrad/inference.py
+++ b/exo/inference/tinygrad/inference.py
@@ -111,8 +111,6 @@ class TinygradDynamicShardInferenceEngine(InferenceEngine):
       Tensor.training = False
       return self.session['loss'](self.model, x, y, l)
     await self.ensure_shard(shard)
-    await self.ensure_session('loss', lambda: loss)
-    await self.ensure_session('jit', lambda: TinyJit(step)) 
     score = await asyncio.get_running_loop().run_in_executor(self.executor, lambda: self.session['jit'](Tensor(inputs), targets, lengths))
     out = score.numpy()
     return out
@@ -126,9 +124,6 @@ class TinygradDynamicShardInferenceEngine(InferenceEngine):
       self.session['opt'].step()
       return score
     await self.ensure_shard(shard)
-    await self.ensure_session('loss', lambda: loss)
-    await self.ensure_session('opt', lambda: opt(nn.state.get_parameters(self.model.model), lr=lr))
-    await self.ensure_session('jit', lambda: TinyJit(step)) 
       
     score = await asyncio.get_running_loop().run_in_executor(self.executor, lambda: self.session['jit'](Tensor(inputs), targets, lengths).realize())
     

← 6aaea8c7 Abstract load checkpoint method  ·  back to Exo  ·  Removed statefulModel stuff from mlx impl too a4313da8 →