From 756a936be233c9e1e9c41bddaf5b0329015435bf Mon Sep 17 00:00:00 2001 From: George Hotz Date: Sun, 2 Aug 2026 15:12:29 +0000 Subject: [PATCH] llm: report interrupted stream statistics --- test/null/test_llm_server.py | 22 ++++++++++++ tinygrad/llm/serve.py | 70 ++++++++++++++++++++---------------- tinygrad/viz/serve.py | 2 +- 3 files changed, 63 insertions(+), 31 deletions(-) diff --git a/test/null/test_llm_server.py b/test/null/test_llm_server.py index 2493258bd5..9af50b0296 100644 --- a/test/null/test_llm_server.py +++ b/test/null/test_llm_server.py @@ -109,6 +109,28 @@ class TestLLMServer(unittest.TestCase): self.assertGreater(len(contents), 0) + def test_interrupted_stream_logs_tokens(self): + with patch.object(self.mock_model, "generate", side_effect=lambda ids, **kwargs: iter([300, 301, 999])), \ + patch("tinygrad.llm.serve.stderr_log") as log, patch("tinygrad.llm.serve.colored", side_effect=lambda text, color: text) as color: + stream = self.server.RequestHandlerClass.run_model(Mock(server=self.server), [200, 201, 202], "test") + next(stream) + next(stream) + stream.close() + interrupt = log.call_args.args[0] + self.assertFalse(interrupt.startswith("\n")) + self.assertTrue(interrupt.endswith("\n")) + self.assertIn("gen:", interrupt) + self.assertIn("out: 1", interrupt) + self.assertTrue(any(args[0].startswith("total:") and args[1] == "red" for args, _ in color.call_args_list)) + + def test_stream_disconnect_closes_source(self): + from tinygrad.viz.serve import HTTPRequestHandler + source, handler = Mock(), Mock() + source.__iter__ = Mock(return_value=iter([{}])) + handler.wfile.write.side_effect = BrokenPipeError + HTTPRequestHandler.stream_json(handler, source) + source.close.assert_called_once() + def test_non_streaming(self): resp = self.client.chat.completions.create( model="test-model", diff --git a/tinygrad/llm/serve.py b/tinygrad/llm/serve.py index 768e6da148..77b01f2735 100644 --- a/tinygrad/llm/serve.py +++ b/tinygrad/llm/serve.py @@ -74,40 +74,50 @@ class Handler(HTTPRequestHandler): stderr_log(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} 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() + st = pt = time.perf_counter() dec = tok.stream_decoder() router = StreamRouter(reasoning) - for next_id in model.generate(ids, temperature=temperature): - if len(out) == 0: stderr_log(f"prefill:{(prompt_tokens-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) - 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 - 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": prompt_tokens, "completion_tokens": len(out), - "total_tokens": prompt_tokens + len(out)}, **tmpl} - et = time.perf_counter() - stderr_log(f"gen:{len(out)/(et-pt) if len(out) > 1 else 0:4.0f} tok/s {colored('--', 'BLACK')} " - f"out:{len(out):5d} {colored('--', 'BLACK')} total:{et-st:6.2f}s\n") + def log_stats(interrupted:bool=False): + et = time.perf_counter() + total = f"total:{et-st:6.2f}s" + stderr_log(f"gen:{len(out)/(et-pt) if len(out) > 1 else 0:4.0f} tok/s {colored('--', 'BLACK')} " + f"out:{len(out):5d} {colored('--', 'BLACK')} {colored(total, 'red') if interrupted else total}\n") + completed = False + try: + yield chunk({"role":"assistant", "content":""}) + for next_id in model.generate(ids, temperature=temperature): + if len(out) == 0: + stderr_log(f"prefill:{(prompt_tokens-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) + 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 + 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" + completed = True + yield {"choices": [{"index":0, "delta":{},"finish_reason":finish_reason}], **tmpl} + if include_usage: + yield {"choices": [], "usage": {"prompt_tokens": prompt_tokens, "completion_tokens": len(out), + "total_tokens": prompt_tokens + len(out)}, **tmpl} + log_stats() + except GeneratorExit: + if not completed: log_stats(interrupted=True) + raise def do_POST(self): request_st = time.perf_counter() diff --git a/tinygrad/viz/serve.py b/tinygrad/viz/serve.py index 8574bb1a71..0548f8a3c8 100755 --- a/tinygrad/viz/serve.py +++ b/tinygrad/viz/serve.py @@ -37,7 +37,7 @@ class HTTPRequestHandler(BaseHTTPRequestHandler): self.wfile.flush() self.wfile.write("data: [DONE]\n\n".encode("utf-8")) # pass if client closed connection - except (BrokenPipeError, ConnectionResetError): return + except (BrokenPipeError, ConnectionResetError): source.close() from tinygrad.uop.ops import TrackedGraphRewrite, RewriteTrace, UOp, Ops, GroupOp, srender, sint, sym_infer, range_str, range_start, multirange_str from tinygrad.uop.ops import KernelInfo