Skip to content

Instantly share code, notes, and snippets.

@david-andrew
Last active March 17, 2026 17:02
Show Gist options
  • Select an option

  • Save david-andrew/cd8bc7e944b38504f9c7a34aa3ab8f6a to your computer and use it in GitHub Desktop.

Select an option

Save david-andrew/cd8bc7e944b38504f9c7a34aa3ab8f6a to your computer and use it in GitHub Desktop.
toki with local models
"""
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()
"""
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