[object Object]

← back to Exo

Some stability improvements for tinygrad inference

13572e6a405bdcab2fc8018ef25704bd08e9399f · 2024-11-11 03:59:13 -0800 · Nel Nibcord

Files touched

Diff

commit 13572e6a405bdcab2fc8018ef25704bd08e9399f
Author: Nel Nibcord <blindcrone@tuta.io>
Date:   Mon Nov 11 03:59:13 2024 -0800

    Some stability improvements for tinygrad inference
---
 exo/inference/tinygrad/inference.py    | 6 +++---
 exo/inference/tinygrad/models/llama.py | 5 ++++-
 exo/orchestration/standard_node.py     | 2 +-
 3 files changed, 8 insertions(+), 5 deletions(-)

diff --git a/exo/inference/tinygrad/inference.py b/exo/inference/tinygrad/inference.py
index 3725a82b..9274674c 100644
--- a/exo/inference/tinygrad/inference.py
+++ b/exo/inference/tinygrad/inference.py
@@ -68,7 +68,7 @@ class TinygradDynamicShardInferenceEngine(InferenceEngine):
   async def sample(self, x: np.ndarray) -> np.ndarray:
     logits = x[:, -1, :]
     def sample_wrapper():
-      return sample_logits(Tensor(x).flatten(), TEMPERATURE, 0, 0.8, 0.0, 0.0).realize()
+      return sample_logits(Tensor(logits).flatten(), TEMPERATURE, 0, 0.8, 0.0, 0.0).realize()
     out = await asyncio.get_running_loop().run_in_executor(self.executor, sample_wrapper)
     return out.numpy()
 
@@ -85,7 +85,7 @@ class TinygradDynamicShardInferenceEngine(InferenceEngine):
   async def infer_tensor(self, request_id: str, shard: Shard, input_data: np.ndarray, inference_state: Optional[str] = None) -> np.ndarray:
     await self.ensure_shard(shard)
     start_pos = json.loads(inference_state or "{}").get("start_pos", 0)
-    output_data = await asyncio.get_running_loop().run_in_executor(self.executor, self.model, Tensor(input_data), start_pos)
+    output_data = await asyncio.get_running_loop().run_in_executor(self.executor, lambda: self.model(Tensor(input_data), start_pos).realize())
     return output_data.numpy()
 
   async def ensure_shard(self, shard: Shard):
@@ -96,7 +96,7 @@ class TinygradDynamicShardInferenceEngine(InferenceEngine):
 
     if self.shard != shard:
       parameters = "1B" if "1b" in shard.model_id.lower() else "3B" if "3b" in shard.model_id.lower() else "8B" if "8b" in shard.model_id.lower() else "70B"
-      self.model = await asyncio.get_event_loop().run_in_executor(self.executor, build_transformer, model_path, shard, parameters)
+      self.model = await asyncio.get_running_loop().run_in_executor(self.executor, build_transformer, model_path, shard, parameters)
 
       tokenizer_path = str((model_path if model_path.is_dir() else model_path.parent))
       self.tokenizer = await resolve_tokenizer(tokenizer_path)
diff --git a/exo/inference/tinygrad/models/llama.py b/exo/inference/tinygrad/models/llama.py
index 0d7ff080..1ad6e831 100644
--- a/exo/inference/tinygrad/models/llama.py
+++ b/exo/inference/tinygrad/models/llama.py
@@ -259,7 +259,10 @@ def convert_from_huggingface(weights: Dict[str, Tensor], model: Transformer, n_h
         v = permute(v, n_heads)
       elif "k_proj" in k:
         v = permute(v, n_kv_heads)
-    sd[keymap[k]] = v
+    if k in keymap:
+      sd[keymap[k]] = v
+    else:
+      sd[k] = v
   return sd
 
 
diff --git a/exo/orchestration/standard_node.py b/exo/orchestration/standard_node.py
index f9d221b6..2629268f 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -109,7 +109,7 @@ class StandardNode(Node):
   async def process_result(
     self,
     shard,
-    result,
+    result: np.ndarray,
     request_id: Optional[str] = None,
     inference_state: Optional[str] = None,
   ):

← aefc0d7c I think this is more faithful to how it was originally done  ·  back to Exo  ·  Implemented per-request caching in tinygrad 8205a5ae →