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