From 3ee2baf71dea039c073a038e645f2e2919852ccc Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Fri, 17 Jul 2026 10:20:01 -0700 Subject: [PATCH] llm: add tool calling support (kimi) (#17061) * llm: add tool calling support * simpler * cls * gpt cleanup * more gpt cleanups * tests for tools calling --- test/null/test_llm_server.py | 92 +++++++++++++++++++++++++++++++++++- tinygrad/llm/cli.py | 89 +++++++++++++++++++++++++++++++--- 2 files changed, 173 insertions(+), 8 deletions(-) diff --git a/test/null/test_llm_server.py b/test/null/test_llm_server.py index 410b9daa84..87a7e28d38 100644 --- a/test/null/test_llm_server.py +++ b/test/null/test_llm_server.py @@ -1,4 +1,4 @@ -import unittest, threading, time +import unittest, threading, time, json from unittest.mock import Mock class TestLLMServer(unittest.TestCase): @@ -156,5 +156,95 @@ class TestLLMServer(unittest.TestCase): self.assertEqual(data["data"][0]["id"], "test-model") self.assertEqual(data["data"][0]["object"], "model") +class TestLLMToolCalls(unittest.TestCase): + """Tool calling through the OpenAI-compatible HTTP API.""" + + @classmethod + def setUpClass(cls): + cls.mock_tok = Mock() + cls.mock_tok.encode = Mock(return_value=[200, 201, 202]) + cls.mock_tok.decode = Mock(return_value="") + cls.mock_tok.preset = "qwen2" + cls.mock_tok.bos_id, cls.mock_tok.eos_id, cls.mock_tok.eot_id = None, 999, None + cls.mock_tok.is_end = Mock(return_value=False) + + cls.mock_model = Mock() + cls.mock_model.get_start_pos = Mock(return_value=0) + + from tinygrad.llm.cli import LLMServer + import jinja2 + # .items() matches tool-aware templates and ensures OpenAI JSON argument strings are normalized before rendering the next turn. + template = jinja2.Template("""{% for m in messages %}{{ m.content or '' }}{% for tc in m.tool_calls or [] %} + {% for key, value in tc.function.arguments.items() %}{{ key }}={{ value }}{% endfor %}{% endfor %}{% endfor %}""") + cls.server = LLMServer(('127.0.0.1', 0), cls.mock_model, "tool-model", cls.mock_tok, template) + cls.port = cls.server.server_address[1] + cls.server_thread = threading.Thread(target=cls.server.serve_forever, daemon=True) + cls.server_thread.start() + time.sleep(0.1) + + from openai import OpenAI + cls.client = OpenAI(base_url=f"http://127.0.0.1:{cls.port}/v1", api_key="test") + + @classmethod + def tearDownClass(cls): + cls.server.shutdown() + cls.server.server_close() + + def set_output(self, text:str): + pieces = dict(enumerate(text, 1)) + self.mock_tok.stream_decoder = Mock(return_value=lambda tid=None: pieces[tid] if tid is not None else "") + self.mock_model.generate = Mock(side_effect=lambda ids, **kwargs: iter(pieces)) + + @staticmethod + def tools(): + return [{"type":"function", "function":{"name":"read", "description":"Read a file", + "parameters":{"type":"object", "properties":{"path":{"type":"string"}}, "required":["path"]}}}] + + def test_streaming_tool_call(self): + self.set_output('before{"name":"read","arguments":{"path":"README.md"}}') + chunks = list(self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Read README.md"}], + tools=self.tools(), stream=True)) + self.assertEqual("".join(c.choices[0].delta.content or "" for c in chunks if c.choices), "before") + calls = [tc for c in chunks if c.choices for tc in c.choices[0].delta.tool_calls or []] + self.assertEqual(len(calls), 1) + self.assertEqual(calls[0].function.name, "read") + self.assertEqual(json.loads(calls[0].function.arguments), {"path":"README.md"}) + self.assertEqual(chunks[-1].choices[0].finish_reason, "tool_calls") + + def test_multiple_xml_tool_calls(self): + self.set_output("\"a\"" + "\"b\"") + response = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Read a and b"}], + tools=self.tools()) + self.assertEqual([json.loads(tc.function.arguments)["path"] for tc in response.choices[0].message.tool_calls], ["a", "b"]) + self.assertEqual(response.choices[0].finish_reason, "tool_calls") + + def test_invalid_tool_call_becomes_content(self): + self.set_output("not a call") + response = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Hello"}], tools=self.tools()) + self.assertEqual(response.choices[0].message.content, "not a call") + self.assertIsNone(response.choices[0].message.tool_calls) + self.assertEqual(response.choices[0].finish_reason, "stop") + + def test_tool_call_in_reasoning_is_not_executed(self): + self.set_output('draft {"name":"wrong","arguments":{}}answer') + response = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Hello"}], tools=self.tools()) + self.assertEqual(response.choices[0].message.content, "answer") + self.assertIsNone(response.choices[0].message.tool_calls) + self.assertEqual(response.choices[0].finish_reason, "stop") + + def test_tool_result_round_trip(self): + self.set_output('{"name":"read","arguments":{"path":"README.md"}}') + first = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Read README.md"}], tools=self.tools()) + call = first.choices[0].message.tool_calls[0] + self.set_output("done") + second = self.client.chat.completions.create(model="tool-model", messages=[ + {"role":"user", "content":"Read README.md"}, + {"role":"assistant", "content":None, "tool_calls":[call.model_dump()]}, + {"role":"tool", "tool_call_id":call.id, "content":"file contents"}, + ], tools=self.tools()) + self.assertEqual(second.choices[0].message.content, "done") + self.assertEqual(second.choices[0].finish_reason, "stop") + if __name__ == '__main__': unittest.main() diff --git a/tinygrad/llm/cli.py b/tinygrad/llm/cli.py index c8d597d583..df67c7ecb0 100644 --- a/tinygrad/llm/cli.py +++ b/tinygrad/llm/cli.py @@ -9,6 +9,58 @@ from tinygrad.llm.model import Transformer if TYPE_CHECKING: import jinja2 +def parse_tool_call(s:str) -> tuple[str, typing.Any]|None: + s = s.strip() + if s.startswith("{"): # hermes JSON format: {"name": ..., "arguments": {...}} + try: + call = json.loads(s) + return call["name"], call.get("arguments", call.get("parameters", {})) + except (json.JSONDecodeError, KeyError): return None + # XML format: \n\nvalue\n... + if (fm := re.match(r"]+)>\s*(.*?)\s*(?:)?$", s, re.DOTALL)): + args = {} + for pm in re.finditer(r"]+)>\s*(.*?)\s*", fm.group(2), re.DOTALL): + try: args[pm.group(1)] = json.loads(pm.group(2)) + except json.JSONDecodeError: args[pm.group(1)] = pm.group(2) + return fm.group(1), args + return None + +def normalize_messages(messages:list[dict]) -> None: + # chat templates expect tool_call arguments as dicts (OpenAI clients send JSON strings) + for m in messages: + for tc in m.get("tool_calls") or []: + if "function" in tc and isinstance(args := tc["function"].get("arguments"), str): + try: tc["function"]["arguments"] = json.loads(args) + except json.JSONDecodeError: pass + +class StreamRouter: + # routes streamed output text to (field, text) deltas, keeping tool_call regions in .buf for the final parse + def __init__(self): + self.buf = "" + self.mode = "undecided" # output inside a think block is sent as reasoning_content + def split(self, tag:str, final:bool) -> tuple[str, bool]: + # split buf on the first full tag, holding back a partial tag at the end unless final + if tag in self.buf: + before, self.buf = self.buf.split(tag, 1) + return before, True + hold = max((i for i in range(1, min(len(self.buf), len(tag))+1) if tag.startswith(self.buf[-i:])), default=0) if not final else 0 + emit, self.buf = self.buf[:len(self.buf)-hold], self.buf[len(self.buf)-hold:] + return emit, False + def route(self, piece:str, final:bool=False) -> typing.Iterator[tuple[str, str]]: + self.buf += piece + if self.mode == "undecided": # decide whether the output starts with a think block + if not final and len(self.buf) < len("") and "".startswith(self.buf): return + self.mode, self.buf = ("reasoning", self.buf[len(""):]) if self.buf.startswith("") else ("content", self.buf) + if self.mode == "reasoning": + emit, done = self.split("", final) + if emit: yield "reasoning_content", emit + if not done: return + self.mode = "content" + if self.mode == "tool": return + emit, found = self.split("", final) + if emit: yield "content", emit + if found: self.mode, self.buf = "tool", "" + self.buf + class SimpleTokenizer: def __init__(self, normal_tokens:dict[str, int], special_tokens:dict[str, int], preset:str="llama3", bos_id:int|None=None, eos_id:int=0, eot_id:int|None=None): @@ -137,20 +189,34 @@ class Handler(HTTPRequestHandler): stderr_log(f"{self.path} {colored('--', 'BLACK')} " f"in:{colored(f'{cache_start_pos:5d}', 'green')} +{len(ids)-cache_start_pos:5d} {colored('--', 'BLACK')} ") tmpl = {"id":f"chatcmpl-{uuid.uuid4().hex[:24]}", "object":"chat.completion.chunk", "created":int(time.time()), "model":model_name} - yield {"choices": [{"index":0, "delta":{"role":"assistant","content":""}, "finish_reason":None}], **tmpl} + def chunk(d:dict): return {"choices": [{"index":0, "delta":d, "finish_reason":None}], **tmpl} + yield chunk({"role":"assistant", "content":""}) out: list[int] = [] finish_reason = "stop" st = time.perf_counter() dec = tok.stream_decoder() + router = StreamRouter() for next_id in model.generate(ids, temperature=temperature): if len(out) == 0: stderr_log(f"prefill:{(len(ids)-cache_start_pos)/((pt:=time.perf_counter())-st):4.0f} tok/s {colored('--', 'BLACK')} ") if tok.is_end(next_id): break out.append(next_id) - yield {"choices": [{"index":0, "delta":{"content":dec(next_id)}, "finish_reason":None}], **tmpl} + for field, delta in router.route(dec(next_id)): yield chunk({field:delta}) if max_tokens is not None and len(out) >= max_tokens: finish_reason = "length" break - if (tail := dec()): yield {"choices": [{"index":0, "delta":{"content":tail}, "finish_reason":None}], **tmpl} + for field, delta in router.route(dec(), final=True): yield chunk({field:delta}) + tool_calls: list[dict] = [] + for m in re.finditer(r"\s*(.*?)\s*(?:|$)", router.buf, re.DOTALL): + if (parsed := parse_tool_call(m.group(1))) is None: + stderr_log(f"failed to parse tool call: {m.group(1)[:200]}") + yield chunk({"content":m.group(0)}) # don't silently drop output the client can't use + else: + name, args = parsed + tool_calls.append({"index":len(tool_calls), "id":f"call_{uuid.uuid4().hex[:24]}", "type":"function", + "function":{"name":name, "arguments":args if isinstance(args, str) else json.dumps(args)}}) + if tool_calls: + yield chunk({"tool_calls":tool_calls}) + if finish_reason == "stop": finish_reason = "tool_calls" yield {"choices": [{"index":0, "delta":{},"finish_reason":finish_reason}], **tmpl} if include_usage: yield {"choices": [], "usage": {"prompt_tokens": len(ids), "completion_tokens": len(out), "total_tokens": len(ids) + len(out)}, **tmpl} @@ -164,6 +230,7 @@ class Handler(HTTPRequestHandler): if DEBUG >= 1: print(json.dumps(body, indent=2)) if self.path == "/v1/chat/completions": # render and tokenize + if not isinstance(self.server.template, FallbackTemplate): normalize_messages(body["messages"]) rendered = self.server.template.render(messages=body["messages"], tools=body.get("tools"), add_generation_prompt=True) ids: list[int] = self.server.tok.encode(rendered) @@ -173,12 +240,20 @@ class Handler(HTTPRequestHandler): max_tokens=max_tokens, temperature=float(body.get("temperature", 0.0))) if body.get("stream"): self.stream_json(chunks) else: - out, finish_reason = [], "stop" + out, reasoning, tool_calls, finish_reason = [], [], [], "stop" for c in chunks: - if c["choices"] and c["choices"][0].get("delta", {}).get("content"): out.append(c["choices"][0]["delta"]["content"]) - if c["choices"] and c["choices"][0].get("finish_reason"): finish_reason = c["choices"][0]["finish_reason"] + if not c["choices"]: continue + choice = c["choices"][0] + if (delta := choice.get("delta", {})): + if delta.get("content"): out.append(delta["content"]) + if delta.get("reasoning_content"): reasoning.append(delta["reasoning_content"]) + tool_calls += [{k:v for k, v in tc.items() if k != "index"} for tc in delta.get("tool_calls", [])] + if choice.get("finish_reason"): finish_reason = choice["finish_reason"] + message: dict[str, typing.Any] = {"role":"assistant", "content":"".join(out) or None} + if reasoning: message["reasoning_content"] = "".join(reasoning) + if tool_calls: message["tool_calls"] = tool_calls self.send_data(json.dumps({**c, "object":"chat.completion", - "choices":[{"index":0, "message":{"role":"assistant","content":"".join(out)}, "finish_reason":finish_reason}]}).encode()) + "choices":[{"index":0, "message":message, "finish_reason":finish_reason}]}).encode()) else: raise RuntimeError(f"unhandled path {self.path}")