Compare commits

...
Author SHA1 Message Date
geohot 7b19732a7a gpt fixes 2026-07-17 00:55:50 +00:00
geohot ee19fd0b6a fixes 2026-07-16 17:45:26 -07:00
geohot ce5ae31f5d work 2026-07-16 17:30:14 -07:00
geohot 22bdc7b6c2 rm that 2026-07-16 17:18:34 -07:00
geohot e13e6b3752 add tool calling support to llm 2026-07-16 17:09:35 -07:00
+136 -26
View File
@@ -6,6 +6,26 @@ from tinygrad.helpers import partition, DEBUG, Timing, GlobalCounters, stderr_lo
from tinygrad.viz.serve import TCPServerWithReuse, HTTPRequestHandler
from tinygrad.llm.model import Transformer
def holdback(s:str, tag:str) -> int:
# length of the suffix of s that is a prefix of tag (the tag may be split across streamed pieces)
return max((i for i in range(1, min(len(s), len(tag))+1) if tag.startswith(s[-i:])), default=0)
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: <function=name>\n<parameter=key>\nvalue\n</parameter>...</function>
if (fm := re.match(r"<function=([^>]+)>\s*(.*?)\s*(?:</function>)?$", s, re.DOTALL)):
args = {}
for pm in re.finditer(r"<parameter=([^>]+)>\s*(.*?)\s*</parameter>", 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
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):
@@ -113,26 +133,75 @@ class Handler(HTTPRequestHandler):
def do_GET(self):
if self.path == "/v1/models": self.send_data(json.dumps({"object":"list","data":[{"id":self.server.model_name,"object":"model"}]}).encode())
else: self.send_data((pathlib.Path(__file__).parent / "chat.html").read_bytes(), content_type="text/html")
def run_model(self, ids:list[int], model_name:str, include_usage=False, max_tokens:int|None=None, temperature:float=0.0):
def run_model(self, ids:list[int], model_name:str, include_usage=False, max_tokens:int|None=None, temperature:float=0.0,
parse_tool_calls=False, prefill_think=False):
model, tok = self.server.model, self.server.tok
cache_start_pos = model.get_start_pos(ids)
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()
mode, buf, tool_text = ("reasoning" if prefill_think else "undecided"), "", ""
def route(piece:str, final:bool=False):
nonlocal mode, buf, tool_text
if mode == "undecided": # decide whether the output starts with a think block
buf += piece
if not final and len(buf) < len("<think>") and "<think>".startswith(buf): return
mode, piece, buf = ("reasoning", buf[len("<think>"):], "") if buf.startswith("<think>") else ("content", buf, "")
if mode == "reasoning":
buf += piece
if "</think>" in buf:
before, piece = buf.split("</think>", 1)
if before: yield chunk({"reasoning_content":before})
buf, mode, piece = "", "content", piece.lstrip("\n")
else:
hold = 0 if final else holdback(buf, "</think>")
if (flush := buf[:len(buf)-hold]): yield chunk({"reasoning_content":flush})
buf = buf[len(buf)-hold:]
return
if not parse_tool_calls:
if piece: yield chunk({"content":piece})
else:
tool_text += piece
if tool_text.startswith("<tool_call>"): return
if "<tool_call>" in tool_text:
before, tool_text = tool_text.split("<tool_call>", 1)
if before: yield chunk({"content":before})
tool_text = "<tool_call>" + tool_text
else:
# hold back any suffix that could be the start of a "<tool_call>" tag split across tokens
hold = 0 if final else holdback(tool_text, "<tool_call>")
if (flush := tool_text[:len(tool_text)-hold]): yield chunk({"content":flush})
tool_text = tool_text[len(tool_text)-hold:]
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}
yield from route(dec(next_id))
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}
yield from route(dec(), final=True)
if parse_tool_calls:
tool_calls = []
calls = [(m.group(1), m.group(0)) for m in re.finditer(r"<tool_call>\s*(.*?)\s*</tool_call>", tool_text, re.DOTALL)]
if not calls and tool_text.startswith("<tool_call>"): calls = [(tool_text[len("<tool_call>"):], tool_text)] # unclosed tag
for i, (inner, raw) in enumerate(calls):
if (parsed := parse_tool_call(inner)) is None:
stderr_log(f"failed to parse tool call: {inner[:200]}")
yield chunk({"content":raw}) # don't silently drop output the client can't use
else:
name, args = parsed
tool_calls.append({"index":i, "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}
@@ -146,39 +215,65 @@ class Handler(HTTPRequestHandler):
body: dict[str, typing.Any] = json.loads(raw_body.decode("utf-8"))
if DEBUG >= 1: print(json.dumps(body, indent=2))
if self.path == "/v1/chat/completions":
# extract tokens, last assistant message is treated as prefill
ids: list[int] = tok.prefix()
for i, msg in enumerate(body["messages"]):
ids += tok.role(msg["role"])
content = msg["content"]
if isinstance(content, str): ids += tok.encode(content)
elif isinstance(content, list):
for c in content:
if c["type"] == "text": ids += tok.encode(c["text"])
else: raise RuntimeError(f"unhandled type: {c['type']}")
else: raise RuntimeError(f"unknown content type: {type(content)}")
if msg["role"] == "assistant" and i == len(body["messages"]) - 1: break
ids += tok.end_turn()
else: ids += tok.role("assistant")
messages, tools = body["messages"], body.get("tools")
prefill_think = False
if self.server.template is not None:
# the chat template expects tool_call arguments as dicts, OpenAI clients send them as JSON strings
norm = []
for m in messages:
if m.get("tool_calls"):
m = dict(m)
m["tool_calls"] = [{**tc, "function":{**tc["function"], "arguments":json.loads(a) if isinstance((a:=tc["function"]["arguments"]), str)
else a}} if "function" in tc else tc for tc in m["tool_calls"]]
norm.append(m)
rendered = self.server.template.render(messages=norm, tools=tools, add_generation_prompt=norm[-1]["role"] != "assistant")
prefill_think = rendered.rstrip().endswith("<think>")
ids: list[int] = tok.encode(rendered)
if (prefix := tok.prefix()) and ids[:len(prefix)] != prefix: ids = prefix + ids
if norm[-1]["role"] == "assistant": # last assistant message is treated as prefill, drop its end-of-turn tokens
end = tok.end_turn()
if len(ids) >= len(end) and ids[-len(end):] == end: ids = ids[:-len(end)]
else:
if tools: stderr_log("warning: ignoring tools, install jinja2 to enable tool calling via the model's chat template")
ids = tok.prefix()
for i, msg in enumerate(messages):
ids += tok.role(msg["role"])
content = msg["content"]
if isinstance(content, str): ids += tok.encode(content)
elif isinstance(content, list):
for c in content:
if c["type"] == "text": ids += tok.encode(c["text"])
else: raise RuntimeError(f"unhandled type: {c['type']}")
else: raise RuntimeError(f"unknown content type: {type(content)}")
if msg["role"] == "assistant" and i == len(messages) - 1: break
ids += tok.end_turn()
else: ids += tok.role("assistant")
# reply
max_tokens = body.get("max_completion_tokens") or body.get("max_tokens")
chunks = self.run_model(ids, body["model"], not body.get("stream") or body.get("stream_options",{}).get("include_usage", False),
max_tokens=max_tokens, temperature=float(body.get("temperature", 0.0)))
max_tokens=max_tokens, temperature=float(body.get("temperature", 0.0)), parse_tool_calls=bool(tools),
prefill_think=prefill_think)
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 (delta := c["choices"][0].get("delta", {})):
if delta.get("content"): out.append(delta["content"])
if delta.get("reasoning_content"): reasoning.append(delta["reasoning_content"])
if delta.get("tool_calls"): tool_calls.extend(delta["tool_calls"])
if c["choices"] and c["choices"][0].get("finish_reason"): finish_reason = c["choices"][0]["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"] = [{k:v for k, v in tc.items() if k != "index"} for tc in 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}")
class LLMServer(TCPServerWithReuse):
def __init__(self, server_address:tuple, model:Transformer, model_name:str, tok:SimpleTokenizer):
self.model, self.model_name, self.tok = model, model_name, tok
def __init__(self, server_address:tuple, model:Transformer, model_name:str, tok:SimpleTokenizer, template:typing.Any=None):
self.model, self.model_name, self.tok, self.template = model, model_name, tok, template
super().__init__(server_address, Handler)
def main():
@@ -194,11 +289,26 @@ def main():
model, kv = Transformer.from_gguf(fetch(models.get(args.model, args.model)), args.max_context)
model_name = kv.get('general.name') or kv.get('general.basename') or args.model
file_sizes = [y.nbytes() for y in UOp.sink(*[x.uop for x in nn.state.get_parameters(model)]).toposort() if y.op is Ops.BUFFER]
print(f"using model \"{model_name}\" with {sum(file_sizes):,} bytes and {sum(x.numel() for x in nn.state.get_parameters(model)):,} params")
print(f"using model \"{model_name}\" with {sum(file_sizes):,} bytes and {sum(x.numel() for x in nn.state.get_parameters(model)):,} params, "
f"max context {args.max_context} on {nn.state.get_parameters(model)[0].device}")
# get tokenizer
tok = SimpleTokenizer.from_gguf_kv(kv)
# compile the chat template if jinja2 is available (enables tool calling and model-specific formatting)
template = None
if (ct := kv.get('tokenizer.chat_template')) is not None:
try:
import jinja2
env = jinja2.Environment()
env.filters['tojson'] = lambda obj, **kwargs: json.dumps(obj) # jinja2's tojson escapes <>& for HTML safety
env.globals['raise_exception'] = lambda msg: (_ for _ in ()).throw(RuntimeError(msg))
env.globals['strftime_now'] = lambda fmt: time.strftime(fmt)
for name, key in {'bos_token':'bos', 'eos_token':'eos', 'unk_token':'unknown', 'pad_token':'padding', 'sep_token':'separator'}.items():
if (tid := kv.get(f'tokenizer.ggml.{key}_token_id')) is not None: env.globals[name] = tok.decode([tid])
template = env.from_string(ct)
except ImportError: stderr_log("warning: jinja2 is not installed, the model's chat template is disabled")
# warmup the JIT
if args.warmup or args.serve:
# run 2 tokens through the model twice to capture the JIT before serving
@@ -206,7 +316,7 @@ def main():
for _ in range(2): list(zip(range(2), model.generate([0])))
# start server
if args.serve: LLMServer(('', args.serve), model, model_name, tok).serve_forever()
if args.serve: LLMServer(('', args.serve), model, model_name, tok, template).serve_forever()
# do benchmark
if args.benchmark is not None: