← back to Exo
fix tokenizer inconsistencies
71e00745cc2e03ea645585072552bb4ea5a7fcf6 · 2024-07-16 22:44:50 -0700 · Alex Cheema
Files touched
M exo/api/chatgpt_api.pyM exo/orchestration/standard_node.pyM exo/topology/device_capabilities.pyM main.py
Diff
commit 71e00745cc2e03ea645585072552bb4ea5a7fcf6
Author: Alex Cheema <alexcheema123@gmail.com>
Date: Tue Jul 16 22:44:50 2024 -0700
fix tokenizer inconsistencies
---
exo/api/chatgpt_api.py | 54 +++++++++++++++++++++++++++++++------
exo/orchestration/standard_node.py | 2 +-
exo/topology/device_capabilities.py | 5 ++--
main.py | 2 +-
4 files changed, 51 insertions(+), 12 deletions(-)
diff --git a/exo/api/chatgpt_api.py b/exo/api/chatgpt_api.py
index fbfe4243..61c75cef 100644
--- a/exo/api/chatgpt_api.py
+++ b/exo/api/chatgpt_api.py
@@ -1,6 +1,7 @@
import uuid
import time
import asyncio
+from transformers import AutoTokenizer
from typing import List
from aiohttp import web
from exo import DEBUG
@@ -8,8 +9,14 @@ from exo.inference.shard import Shard
from exo.orchestration import Node
shard_mappings = {
- "llama-3-8b": Shard(model_id="mlx-community/Meta-Llama-3-8B-Instruct-4bit", start_layer=0, end_layer=0, n_layers=32),
- "llama-3-70b": Shard(model_id="mlx-community/Meta-Llama-3-70B-Instruct-4bit", start_layer=0, end_layer=0, n_layers=80),
+ "llama-3-8b": {
+ "MLXDynamicShardInferenceEngine": Shard(model_id="mlx-community/Meta-Llama-3-8B-Instruct-4bit", start_layer=0, end_layer=0, n_layers=32),
+ "TinygradDynamicShardInferenceEngine": Shard(model_id="llama3-8b-sfr", start_layer=0, end_layer=0, n_layers=32),
+ },
+ "llama-3-70b": {
+ "MLXDynamicShardInferenceEngine": Shard(model_id="mlx-community/Meta-Llama-3-70B-Instruct-4bit", start_layer=0, end_layer=0, n_layers=80),
+ "TinygradDynamicShardInferenceEngine": Shard(model_id="llama3-70b-sfr", start_layer=0, end_layer=0, n_layers=80),
+ },
}
class Message:
@@ -23,25 +30,54 @@ class ChatCompletionRequest:
self.messages = messages
self.temperature = temperature
+def resolve_tinygrad_tokenizer(model_id: str):
+ if model_id == "llama3-8b-sfr":
+ return AutoTokenizer.from_pretrained("TriAiExperiments/SFR-Iterative-DPO-LLaMA-3-8B-R")
+ elif model_id == "llama3-70b-sfr":
+ return AutoTokenizer.from_pretrained("TriAiExperiments/SFR-Iterative-DPO-LLaMA-3-8B-R")
+ else:
+ raise ValueError(f"tinygrad doesnt currently support arbitrary model downloading. unsupported model: {model_id}")
+
+def resolve_tokenizer(model_id: str):
+ try:
+ if DEBUG >= 2: print(f"Trying AutoTokenizer for {model_id}")
+ return AutoTokenizer.from_pretrained(model_id)
+ except:
+ import traceback
+ if DEBUG >= 2: print(traceback.format_exc())
+ if DEBUG >= 2: print(f"Failed to load tokenizer for {model_id}. Falling back to tinygrad tokenizer")
+
+ try:
+ if DEBUG >= 2: print(f"Trying tinygrad tokenizer for {model_id}")
+ return resolve_tinygrad_tokenizer(model_id)
+ except:
+ import traceback
+ if DEBUG >= 2: print(traceback.format_exc())
+ if DEBUG >= 2: print(f"Failed again to load tokenizer for {model_id}. Falling back to mlx tokenizer")
+
+ if DEBUG >= 2: print(f"Trying mlx tokenizer for {model_id}")
+ from exo.inference.mlx.sharded_utils import get_model_path, load_tokenizer
+ return load_tokenizer(get_model_path(model_id))
+
class ChatGPTAPI:
- def __init__(self, node: Node):
+ def __init__(self, node: Node, inference_engine_classname: str):
self.node = node
self.app = web.Application()
self.app.router.add_post('/v1/chat/completions', self.handle_post)
+ self.inference_engine_classname = inference_engine_classname
async def handle_post(self, request):
data = await request.json()
messages = [Message(**msg) for msg in data['messages']]
chat_request = ChatCompletionRequest(data['model'], messages, data['temperature'])
prompt = " ".join([msg.content for msg in chat_request.messages if msg.role == "user"])
- shard = shard_mappings.get(chat_request.model)
+ shard = shard_mappings.get(chat_request.model, {}).get(self.inference_engine_classname)
if not shard:
return web.json_response({'detail': f"Invalid model: {chat_request.model}. Supported: {list(shard_mappings.keys())}"}, status=400)
request_id = str(uuid.uuid4())
- # TODO equivalent for non-mlx since the user can't install this on non-macs even though the tokenizer itself
- from exo.inference.mlx.sharded_utils import get_model_path, load_tokenizer
- tokenizer = load_tokenizer(get_model_path(shard.model_id))
+ tokenizer = resolve_tokenizer(shard.model_id)
+ if DEBUG >= 4: print(f"Resolved tokenizer: {tokenizer}")
prompt = tokenizer.apply_chat_template(
chat_request.messages, tokenize=False, add_generation_prompt=True
)
@@ -63,7 +99,9 @@ class ChatGPTAPI:
continue
await asyncio.sleep(0.1)
if is_finished:
- if result[-1] == tokenizer._tokenizer.eos_token_id:
+ eos_token_id = tokenizer.special_tokens_map.get("eos_token_id") if isinstance(tokenizer._tokenizer, AutoTokenizer) else tokenizer.eos_token_id
+ if DEBUG >= 2: print(f"Checking if end of result {result[-1]=} is {eos_token_id=}")
+ if result[-1] == eos_token_id:
result = result[:-1]
return web.json_response({
"id": f"chatcmpl-{request_id}",
diff --git a/exo/orchestration/standard_node.py b/exo/orchestration/standard_node.py
index c6ede1e4..c554115a 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -178,7 +178,7 @@ class StandardNode(Node):
async def collect_topology(self, visited: set[str] = set(), max_depth: int = 4) -> Topology:
self.topology.update_node(self.id, self.device_capabilities)
- if DEBUG >= 2: print(f"Collecting topoloy {max_depth=} {visited=}")
+ if DEBUG >= 2: print(f"Collecting topology {max_depth=} {visited=}")
prev_visited = visited.copy()
visited.update(p.id() for p in self.peers)
diff --git a/exo/topology/device_capabilities.py b/exo/topology/device_capabilities.py
index 57f97503..8c9787a1 100644
--- a/exo/topology/device_capabilities.py
+++ b/exo/topology/device_capabilities.py
@@ -1,3 +1,4 @@
+from exo import DEBUG
from dataclasses import dataclass
import subprocess
import platform
@@ -41,8 +42,8 @@ def mac_device_capabilities() -> DeviceCapabilities:
def linux_device_capabilities() -> DeviceCapabilities:
import psutil
from tinygrad import Device
-
- print(f"tinygrad {Device.DEFAULT=}")
+
+ if DEBUG >= 2: print(f"tinygrad {Device.DEFAULT=}")
if Device.DEFAULT == "CUDA" or Device.DEFAULT == "NV" or Device.DEFAULT=="GPU":
import pynvml, pynvml_utils
pynvml.nvmlInit()
diff --git a/main.py b/main.py
index 541e99a1..99b1d0c6 100644
--- a/main.py
+++ b/main.py
@@ -38,7 +38,7 @@ node = StandardNode(args.node_id, None, inference_engine, discovery, partitionin
server = GRPCServer(node, args.node_host, args.node_port)
node.server = server
-api = ChatGPTAPI(node)
+api = ChatGPTAPI(node, inference_engine.__class__.__name__)
async def shutdown(signal, loop):
"""Gracefully shutdown the server and close the asyncio loop."""
← c673b4c3 clarify readme
·
back to Exo
·
add stars to readme 365114ec →