[object Object]

← back to Exo

fix dummy tokenizer

1331ed767922365330b969455a29d652d7ae92ab · 2024-11-25 14:38:03 +0400 · Alex Cheema

Files touched

Diff

commit 1331ed767922365330b969455a29d652d7ae92ab
Author: Alex Cheema <alexcheema123@gmail.com>
Date:   Mon Nov 25 14:38:03 2024 +0400

    fix dummy tokenizer
---
 exo/inference/dummy_inference_engine.py | 8 ++++++--
 exo/inference/tokenizers.py             | 6 +++---
 2 files changed, 9 insertions(+), 5 deletions(-)

diff --git a/exo/inference/dummy_inference_engine.py b/exo/inference/dummy_inference_engine.py
index 2712b7d2..bdad2b30 100644
--- a/exo/inference/dummy_inference_engine.py
+++ b/exo/inference/dummy_inference_engine.py
@@ -3,9 +3,10 @@ import numpy as np
 import random
 import string
 import asyncio
-import json
 from exo.inference.inference_engine import InferenceEngine
 from exo.inference.shard import Shard
+from exo.inference.tokenizers import DummyTokenizer
+
 def random_string(length: int):
   return ''.join([random.choice(string.ascii_lowercase) for i in range(length)])
   
@@ -18,12 +19,15 @@ class DummyInferenceEngine(InferenceEngine):
     self.eos_token_id = 0
     self.latency_mean = 0.1
     self.latency_stddev = 0.02
+    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())))
   
   async def sample(self, x: np.ndarray) -> np.ndarray:
-    return np.random.randint(1, self.vocab_size)
+    if random.random() < 0.1:
+      return np.array([self.tokenizer.eos_token_id])
+    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])
diff --git a/exo/inference/tokenizers.py b/exo/inference/tokenizers.py
index 0b8ffa7e..6248a844 100644
--- a/exo/inference/tokenizers.py
+++ b/exo/inference/tokenizers.py
@@ -10,13 +10,13 @@ from exo.helpers import DEBUG
 
 class DummyTokenizer:
   def __init__(self):
-    self.eos_token_id = 0
+    self.eos_token_id = 69
 
   def apply_chat_template(self, messages, tokenize=True, add_generation_prompt=True):
-    return [1, 2, 3]
+    return "dummy_tokenized_prompt"
 
   def decode(self, tokens):
-    return "dummy"
+    return "dummy" * len(tokens)
 
 
 async def resolve_tokenizer(model_id: str):

← 311b4c21 skip inference engine selection if running dummy  ·  back to Exo  ·  use jq to check response content circleci 5d3ac40f →