[object Object]

← back to Exo

fix dummy inference

f601a8307053eace28db1030ae2ad0e07c55e96b · 2024-11-25 14:43:21 +0400 · Alex Cheema

Files touched

Diff

commit f601a8307053eace28db1030ae2ad0e07c55e96b
Author: Alex Cheema <alexcheema123@gmail.com>
Date:   Mon Nov 25 14:43:21 2024 +0400

    fix dummy inference
---
 exo/inference/dummy_inference_engine.py | 4 ++--
 exo/inference/tokenizers.py             | 5 +++++
 2 files changed, 7 insertions(+), 2 deletions(-)

diff --git a/exo/inference/dummy_inference_engine.py b/exo/inference/dummy_inference_engine.py
index bdad2b30..34a44d49 100644
--- a/exo/inference/dummy_inference_engine.py
+++ b/exo/inference/dummy_inference_engine.py
@@ -22,7 +22,7 @@ class DummyInferenceEngine(InferenceEngine):
     self.tokenizer = DummyTokenizer()
 
   async def encode(self, shard: Shard, prompt: str) -> np.ndarray:
-    return np.random.randint(1, self.vocab_size, size=(1, len(prompt.split())))
+    return np.array(self.tokenizer.encode(prompt))
   
   async def sample(self, x: np.ndarray) -> np.ndarray:
     if random.random() < 0.1:
@@ -30,7 +30,7 @@ class DummyInferenceEngine(InferenceEngine):
     return np.array([np.random.randint(1, self.vocab_size)])
 
   async def decode(self, shard: Shard, tokens: np.ndarray) -> str:
-    return ' '.join([random_string(np.random.randint(1, 34)) for token in tokens])
+    return self.tokenizer.decode(tokens)
 
   async def infer_tensor(self, request_id: str, shard: Shard, input_data: np.ndarray) -> np.ndarray:
     await self.ensure_shard(shard)
diff --git a/exo/inference/tokenizers.py b/exo/inference/tokenizers.py
index 6248a844..1fbcb839 100644
--- a/exo/inference/tokenizers.py
+++ b/exo/inference/tokenizers.py
@@ -4,6 +4,7 @@ from os import PathLike
 from pathlib import Path
 from typing import Union
 from transformers import AutoTokenizer, AutoProcessor
+import numpy as np
 from exo.download.hf.hf_helpers import get_local_snapshot_dir
 from exo.helpers import DEBUG
 
@@ -11,10 +12,14 @@ from exo.helpers import DEBUG
 class DummyTokenizer:
   def __init__(self):
     self.eos_token_id = 69
+    self.vocab_size = 1000
 
   def apply_chat_template(self, messages, tokenize=True, add_generation_prompt=True):
     return "dummy_tokenized_prompt"
 
+  def encode(self, text):
+    return np.random.randint(1, self.vocab_size, size=(1, len(text.split())))
+
   def decode(self, tokens):
     return "dummy" * len(tokens)
 

← 216e7bff remove redundant jq installation  ·  back to Exo  ·  less strict match on response content 6b28b341 →