← 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
M exo/inference/inference_engine.pyM exo/inference/tinygrad/inference.py
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 →