Last active
March 17, 2026 17:02
-
-
Save david-andrew/cd8bc7e944b38504f9c7a34aa3ab8f6a to your computer and use it in GitHub Desktop.
toki with local models
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| """ | |
| python dependencies: | |
| - transformers | |
| - torch | |
| - easyrepl | |
| - toki | |
| """ | |
| from threading import Thread | |
| from typing import Generator, Literal, overload | |
| import torch | |
| from easyrepl import REPL | |
| from toki import Agent | |
| from toki.openrouter import OpenRouterMessage, OpenRouterToolResponse, OpenRouterUsageMetadata | |
| from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer | |
| class Model: | |
| """Local `transformers`-backed model with the same interface as `toki.openrouter.Model`.""" | |
| def __init__(self, model: str, allow_parallel_tool_calls: bool = False): | |
| self.model = model | |
| self.allow_parallel_tool_calls = allow_parallel_tool_calls | |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| self.tokenizer = AutoTokenizer.from_pretrained(model) | |
| if self.tokenizer.pad_token_id is None: | |
| self.tokenizer.pad_token = self.tokenizer.eos_token | |
| self._model = AutoModelForCausalLM.from_pretrained(model, torch_dtype="auto") | |
| self._model.to(self.device) | |
| self._model.eval() | |
| self._usage_metadata: OpenRouterUsageMetadata | None = None | |
| @overload | |
| def complete(self, messages: list[OpenRouterMessage], *, stream: Literal[False] = False, tools: None = None, **kwargs) -> str: ... | |
| @overload | |
| def complete(self, messages: list[OpenRouterMessage], *, stream: Literal[False] = False, tools: list, **kwargs) -> str | OpenRouterToolResponse: ... | |
| @overload | |
| def complete(self, messages: list[OpenRouterMessage], *, stream: Literal[True], tools: None = None, **kwargs) -> Generator[str, None, None]: ... | |
| @overload | |
| def complete(self, messages: list[OpenRouterMessage], *, stream: Literal[True], tools: list, **kwargs) -> Generator[str | OpenRouterToolResponse, None, None]: ... | |
| def complete( | |
| self, | |
| messages: list[OpenRouterMessage], | |
| *, | |
| stream: bool = False, | |
| tools: list | None = None, | |
| **kwargs, | |
| ) -> str | OpenRouterToolResponse | Generator[str | OpenRouterToolResponse, None, None]: | |
| if stream: | |
| return self._streaming_complete(messages, tools, **kwargs) | |
| return self._blocking_complete(messages, tools, **kwargs) | |
| @overload | |
| def _blocking_complete(self, messages: list[OpenRouterMessage], tools: None = None, **kwargs) -> str: ... | |
| @overload | |
| def _blocking_complete(self, messages: list[OpenRouterMessage], tools: list, **kwargs) -> str | OpenRouterToolResponse: ... | |
| def _blocking_complete( | |
| self, | |
| messages: list[OpenRouterMessage], | |
| tools: list | None = None, | |
| **kwargs, | |
| ) -> str | OpenRouterToolResponse: | |
| self._reject_tools(tools) | |
| prompt = self._build_prompt(messages) | |
| inputs = self.tokenizer(prompt, return_tensors="pt").to(self.device) | |
| generation_kwargs = self._build_generate_kwargs(**kwargs) | |
| with torch.no_grad(): | |
| outputs = self._model.generate( | |
| **inputs, | |
| **generation_kwargs, | |
| pad_token_id=self.tokenizer.pad_token_id, | |
| ) | |
| generated_tokens = outputs[0, inputs["input_ids"].shape[-1] :] | |
| response = self.tokenizer.decode(generated_tokens, skip_special_tokens=True) | |
| self._set_usage_metadata(prompt_tokens=inputs["input_ids"].shape[-1], completion_text=response) | |
| return response | |
| @overload | |
| def _streaming_complete(self, messages: list[OpenRouterMessage], tools: None = None, **kwargs) -> Generator[str, None, None]: ... | |
| @overload | |
| def _streaming_complete(self, messages: list[OpenRouterMessage], tools: list, **kwargs) -> Generator[str | OpenRouterToolResponse, None, None]: ... | |
| def _streaming_complete( | |
| self, | |
| messages: list[OpenRouterMessage], | |
| tools: list | None = None, | |
| **kwargs, | |
| ) -> Generator[str | OpenRouterToolResponse, None, None]: | |
| self._reject_tools(tools) | |
| prompt = self._build_prompt(messages) | |
| inputs = self.tokenizer(prompt, return_tensors="pt").to(self.device) | |
| streamer = TextIteratorStreamer(self.tokenizer, skip_prompt=True, skip_special_tokens=True) | |
| generation_kwargs = self._build_generate_kwargs(**kwargs) | |
| errors: list[BaseException] = [] | |
| def run_generation() -> None: | |
| try: | |
| with torch.no_grad(): | |
| self._model.generate( | |
| **inputs, | |
| **generation_kwargs, | |
| pad_token_id=self.tokenizer.pad_token_id, | |
| streamer=streamer, | |
| ) | |
| except BaseException as exc: | |
| errors.append(exc) | |
| worker = Thread(target=run_generation, daemon=True) | |
| worker.start() | |
| chunks: list[str] = [] | |
| for chunk in streamer: | |
| chunks.append(chunk) | |
| yield chunk | |
| worker.join() | |
| if errors: | |
| raise errors[0] | |
| response = "".join(chunks) | |
| self._set_usage_metadata(prompt_tokens=inputs["input_ids"].shape[-1], completion_text=response) | |
| def _build_prompt(self, messages: list[OpenRouterMessage]) -> str: | |
| return self.tokenizer.apply_chat_template( | |
| messages, | |
| tokenize=False, | |
| add_generation_prompt=True, | |
| ) | |
| def _build_generate_kwargs(self, **kwargs) -> dict: | |
| max_new_tokens = kwargs.pop("max_new_tokens", 512) | |
| temperature = kwargs.pop("temperature", 0.7) | |
| top_p = kwargs.pop("top_p", 0.95) | |
| generate_kwargs = { | |
| "max_new_tokens": max_new_tokens, | |
| "do_sample": temperature > 0, | |
| **kwargs, | |
| } | |
| if temperature > 0: | |
| generate_kwargs["temperature"] = temperature | |
| generate_kwargs["top_p"] = top_p | |
| return generate_kwargs | |
| def _set_usage_metadata(self, *, prompt_tokens: int, completion_text: str) -> None: | |
| completion_tokens = len(self.tokenizer.encode(completion_text, add_special_tokens=False)) | |
| self._usage_metadata = OpenRouterUsageMetadata( | |
| prompt_tokens=prompt_tokens, | |
| completion_tokens=completion_tokens, | |
| total_tokens=prompt_tokens + completion_tokens, | |
| ) | |
| def _reject_tools(self, tools: list | None) -> None: | |
| if tools: | |
| raise NotImplementedError("This local transformers demo does not implement tool calling.") | |
| def main() -> None: | |
| model = Model("Qwen/Qwen3-0.6B") | |
| agent = Agent(model) | |
| agent.add_system_message("You are a concise, helpful assistant.") | |
| for query in REPL(history=".chat"): | |
| agent.add_user_message(query) | |
| for chunk in agent.execute(stream=True): | |
| print(chunk, end="", flush=True) | |
| print() | |
| if __name__ == "__main__": | |
| main() |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| """ | |
| python dependencies: | |
| - transformers | |
| - torch | |
| - easyrepl | |
| - toki | |
| """ | |
| import json | |
| import logging | |
| import warnings | |
| from collections.abc import Callable | |
| from datetime import UTC, datetime | |
| from threading import Thread | |
| from typing import Generator, Literal, overload | |
| from uuid import uuid4 | |
| import torch | |
| from easyrepl import REPL | |
| from toki import Agent | |
| from toki.openrouter import ( | |
| OpenRouterMessage, | |
| OpenRouterToolCall, | |
| OpenRouterToolResponse, | |
| OpenRouterUsageMetadata, | |
| pretty_tool_call, | |
| ) | |
| from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer | |
| LOG_LEVEL = logging.INFO | |
| logger = logging.getLogger(__name__) | |
| logger.propagate = False | |
| class Model: | |
| """Local `transformers`-backed model with the same interface as `toki.openrouter.Model`.""" | |
| def __init__(self, model: str, allow_parallel_tool_calls: bool = False): | |
| self.model = model | |
| self.allow_parallel_tool_calls = allow_parallel_tool_calls | |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| self.tokenizer = AutoTokenizer.from_pretrained(model) | |
| if self.tokenizer.pad_token_id is None: | |
| self.tokenizer.pad_token = self.tokenizer.eos_token | |
| self._model = AutoModelForCausalLM.from_pretrained(model, torch_dtype="auto") | |
| self._model.to(self.device) | |
| self._model.eval() | |
| self._usage_metadata: OpenRouterUsageMetadata | None = None | |
| @overload | |
| def complete(self, messages: list[OpenRouterMessage], *, stream: Literal[False] = False, tools: None = None, **kwargs) -> str: ... | |
| @overload | |
| def complete(self, messages: list[OpenRouterMessage], *, stream: Literal[False] = False, tools: list, **kwargs) -> str | OpenRouterToolResponse: ... | |
| @overload | |
| def complete(self, messages: list[OpenRouterMessage], *, stream: Literal[True], tools: None = None, **kwargs) -> Generator[str, None, None]: ... | |
| @overload | |
| def complete(self, messages: list[OpenRouterMessage], *, stream: Literal[True], tools: list, **kwargs) -> Generator[str | OpenRouterToolResponse, None, None]: ... | |
| def complete( | |
| self, | |
| messages: list[OpenRouterMessage], | |
| *, | |
| stream: bool = False, | |
| tools: list | None = None, | |
| **kwargs, | |
| ) -> str | OpenRouterToolResponse | Generator[str | OpenRouterToolResponse, None, None]: | |
| if stream: | |
| return self._streaming_complete(messages, tools, **kwargs) | |
| return self._blocking_complete(messages, tools, **kwargs) | |
| @overload | |
| def _blocking_complete(self, messages: list[OpenRouterMessage], tools: None = None, **kwargs) -> str: ... | |
| @overload | |
| def _blocking_complete(self, messages: list[OpenRouterMessage], tools: list, **kwargs) -> str | OpenRouterToolResponse: ... | |
| def _blocking_complete( | |
| self, | |
| messages: list[OpenRouterMessage], | |
| tools: list | None = None, | |
| **kwargs, | |
| ) -> str | OpenRouterToolResponse: | |
| prompt = self._build_prompt(messages, tools=tools) | |
| inputs = self.tokenizer(prompt, return_tensors="pt").to(self.device) | |
| generation_kwargs = self._build_generate_kwargs( | |
| input_token_count=inputs["input_ids"].shape[-1], | |
| **kwargs, | |
| ) | |
| with torch.no_grad(): | |
| outputs = self._model.generate( | |
| **inputs, | |
| **generation_kwargs, | |
| pad_token_id=self.tokenizer.pad_token_id, | |
| ) | |
| generated_tokens = outputs[0, inputs["input_ids"].shape[-1] :] | |
| response = self.tokenizer.decode(generated_tokens, skip_special_tokens=True) | |
| self._set_usage_metadata(prompt_tokens=inputs["input_ids"].shape[-1], completion_text=response) | |
| return self._parse_response(response, tools) | |
| @overload | |
| def _streaming_complete(self, messages: list[OpenRouterMessage], tools: None = None, **kwargs) -> Generator[str, None, None]: ... | |
| @overload | |
| def _streaming_complete(self, messages: list[OpenRouterMessage], tools: list, **kwargs) -> Generator[str | OpenRouterToolResponse, None, None]: ... | |
| def _streaming_complete( | |
| self, | |
| messages: list[OpenRouterMessage], | |
| tools: list | None = None, | |
| **kwargs, | |
| ) -> Generator[str | OpenRouterToolResponse, None, None]: | |
| prompt = self._build_prompt(messages, tools=tools) | |
| inputs = self.tokenizer(prompt, return_tensors="pt").to(self.device) | |
| streamer = TextIteratorStreamer(self.tokenizer, skip_prompt=True, skip_special_tokens=True) | |
| generation_kwargs = self._build_generate_kwargs( | |
| input_token_count=inputs["input_ids"].shape[-1], | |
| **kwargs, | |
| ) | |
| errors: list[BaseException] = [] | |
| def run_generation() -> None: | |
| try: | |
| with torch.no_grad(): | |
| self._model.generate( | |
| **inputs, | |
| **generation_kwargs, | |
| pad_token_id=self.tokenizer.pad_token_id, | |
| streamer=streamer, | |
| ) | |
| except BaseException as exc: | |
| errors.append(exc) | |
| worker = Thread(target=run_generation, daemon=True) | |
| worker.start() | |
| chunks: list[str] = [] | |
| for chunk in streamer: | |
| chunks.append(chunk) | |
| yield chunk | |
| worker.join() | |
| if errors: | |
| raise errors[0] | |
| response = "".join(chunks) | |
| self._set_usage_metadata(prompt_tokens=inputs["input_ids"].shape[-1], completion_text=response) | |
| if tools: | |
| parsed_response = self._parse_response(response, tools) | |
| if isinstance(parsed_response, dict): | |
| yield parsed_response | |
| def _build_prompt(self, messages: list[OpenRouterMessage], *, tools: list | None = None) -> str: | |
| template_kwargs = {} | |
| if tools: | |
| template_kwargs["tools"] = tools | |
| return self.tokenizer.apply_chat_template( | |
| messages, | |
| tokenize=False, | |
| add_generation_prompt=True, | |
| **template_kwargs, | |
| ) | |
| def _build_generate_kwargs(self, *, input_token_count: int, **kwargs) -> dict: | |
| temperature = kwargs.pop("temperature", 0.7) | |
| top_p = kwargs.pop("top_p", 0.95) | |
| max_new_tokens = kwargs.pop("max_new_tokens", None) | |
| generate_kwargs = { | |
| "do_sample": temperature > 0, | |
| **kwargs, | |
| } | |
| if max_new_tokens is not None: | |
| generate_kwargs["max_new_tokens"] = max_new_tokens | |
| else: | |
| generate_kwargs["max_new_tokens"] = max(1, self.tokenizer.model_max_length - input_token_count) | |
| if temperature > 0: | |
| generate_kwargs["temperature"] = temperature | |
| generate_kwargs["top_p"] = top_p | |
| return generate_kwargs | |
| def _set_usage_metadata(self, *, prompt_tokens: int, completion_text: str) -> None: | |
| completion_tokens = len(self.tokenizer.encode(completion_text, add_special_tokens=False)) | |
| self._usage_metadata = OpenRouterUsageMetadata( | |
| prompt_tokens=prompt_tokens, | |
| completion_tokens=completion_tokens, | |
| total_tokens=prompt_tokens + completion_tokens, | |
| ) | |
| def _parse_response(self, response: str, tools: list | None) -> str | OpenRouterToolResponse: | |
| if not tools: | |
| return response | |
| tool_calls = self._extract_tool_calls(response, tools) | |
| if not tool_calls: | |
| return response | |
| thought = response.split("{", 1)[0].strip() | |
| return OpenRouterToolResponse(thought=thought, tool_calls=tool_calls) | |
| def _extract_tool_calls(self, response: str, tools: list) -> list[OpenRouterToolCall]: | |
| allowed_names = {tool["function"]["name"] for tool in tools} | |
| decoder = json.JSONDecoder() | |
| tool_calls: list[OpenRouterToolCall] = [] | |
| idx = 0 | |
| while idx < len(response): | |
| next_open = response.find("{", idx) | |
| if next_open == -1: | |
| break | |
| try: | |
| payload, end = decoder.raw_decode(response[next_open:]) | |
| except json.JSONDecodeError: | |
| idx = next_open + 1 | |
| continue | |
| tool_calls.extend(self._coerce_tool_calls(payload, allowed_names)) | |
| idx = next_open + end | |
| return tool_calls | |
| def _coerce_tool_calls(self, payload: object, allowed_names: set[str]) -> list[OpenRouterToolCall]: | |
| if isinstance(payload, list): | |
| tool_calls: list[OpenRouterToolCall] = [] | |
| for item in payload: | |
| tool_calls.extend(self._coerce_tool_calls(item, allowed_names)) | |
| return tool_calls | |
| if not isinstance(payload, dict): | |
| return [] | |
| if "function_call" in payload: | |
| return self._coerce_tool_calls(payload["function_call"], allowed_names) | |
| if "tool_calls" in payload: | |
| return self._coerce_tool_calls(payload["tool_calls"], allowed_names) | |
| name = payload.get("name") | |
| arguments = payload.get("arguments") | |
| if not isinstance(name, str) or name not in allowed_names or arguments is None: | |
| return [] | |
| if isinstance(arguments, str): | |
| arguments_json = arguments | |
| else: | |
| arguments_json = json.dumps(arguments) | |
| return [ | |
| OpenRouterToolCall( | |
| id=f"local-tool-{uuid4().hex}", | |
| type="function", | |
| function={"name": name, "arguments": arguments_json}, | |
| ) | |
| ] | |
| def get_current_time() -> dict[str, str]: | |
| return {"utc_time": datetime.now(UTC).isoformat()} | |
| def add_numbers(a: float, b: float) -> dict[str, float]: | |
| return {"sum": a + b} | |
| TOOL_SCHEMAS = [ | |
| { | |
| "type": "function", | |
| "function": { | |
| "name": "get_current_time", | |
| "description": "Get the current UTC time.", | |
| "parameters": { | |
| "type": "object", | |
| "properties": {}, | |
| "required": [], | |
| }, | |
| }, | |
| }, | |
| { | |
| "type": "function", | |
| "function": { | |
| "name": "add_numbers", | |
| "description": "Add two numbers together.", | |
| "parameters": { | |
| "type": "object", | |
| "properties": { | |
| "a": {"type": "number", "description": "The first number."}, | |
| "b": {"type": "number", "description": "The second number."}, | |
| }, | |
| "required": ["a", "b"], | |
| }, | |
| }, | |
| }, | |
| ] | |
| TOOLS_BY_NAME: dict[str, Callable[..., object]] = { | |
| "get_current_time": get_current_time, | |
| "add_numbers": add_numbers, | |
| } | |
| def execute_tool_call(tool_call: OpenRouterToolCall) -> str: | |
| function_name = tool_call["function"]["name"] | |
| function = TOOLS_BY_NAME[function_name] | |
| arguments = json.loads(tool_call["function"]["arguments"]) | |
| result = function(**arguments) | |
| result_json = json.dumps(result) | |
| logger.info("Tool result for %s: %s", pretty_tool_call(tool_call), result_json) | |
| return result_json | |
| def run_agent_turn(agent: Agent) -> str: | |
| while True: | |
| result_chunks: list[str] = [] | |
| tool_response: OpenRouterToolResponse | None = None | |
| if agent.tools is None: | |
| for chunk in agent.execute(stream=True): | |
| print(chunk, end="", flush=True) | |
| result_chunks.append(chunk) | |
| print() | |
| return "".join(result_chunks) | |
| for chunk in agent.model.complete(agent.messages, stream=True, tools=agent.tools): | |
| if isinstance(chunk, str): | |
| print(chunk, end="", flush=True) | |
| result_chunks.append(chunk) | |
| else: | |
| tool_response = chunk | |
| if result_chunks: | |
| print() | |
| assistant_text = "".join(result_chunks) | |
| if tool_response is None: | |
| agent.add_assistant_message(assistant_text) | |
| return assistant_text | |
| agent.add_assistant_tool_calls(tool_response["thought"], tool_response["tool_calls"]) | |
| for tool_call in tool_response["tool_calls"]: | |
| logger.info("Tool call: %s", pretty_tool_call(tool_call)) | |
| tool_result = execute_tool_call(tool_call) | |
| agent.add_tool_message(tool_call["id"], tool_result) | |
| def main() -> None: | |
| if not logger.handlers: | |
| handler = logging.StreamHandler() | |
| handler.setFormatter(logging.Formatter("%(levelname)s: %(message)s")) | |
| logger.addHandler(handler) | |
| logger.setLevel(LOG_LEVEL) | |
| for noisy_logger in [ | |
| "httpx", | |
| "httpcore", | |
| "huggingface_hub", | |
| "transformers", | |
| ]: | |
| logging.getLogger(noisy_logger).setLevel(logging.WARNING) | |
| warnings.filterwarnings( | |
| "ignore", | |
| message="Warning: You are sending unauthenticated requests to the HF Hub.*", | |
| ) | |
| model = Model("Qwen/Qwen3-0.6B") | |
| agent = Agent(model, tools=TOOL_SCHEMAS) | |
| agent.add_system_message("You are a concise, helpful assistant. Use tools when they are useful.") | |
| for query in REPL(history=".chat"): | |
| agent.add_user_message(query) | |
| run_agent_turn(agent) | |
| if __name__ == "__main__": | |
| main() |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment