← back to Exo
fix dummy tokenizer
1331ed767922365330b969455a29d652d7ae92ab · 2024-11-25 14:38:03 +0400 · Alex Cheema
Files touched
M exo/inference/dummy_inference_engine.pyM exo/inference/tokenizers.py
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 →