diff --git a/test/unit/test_llm_server.py b/test/unit/test_llm_server.py index b486034643..1afa9ef952 100644 --- a/test/unit/test_llm_server.py +++ b/test/unit/test_llm_server.py @@ -3,10 +3,26 @@ from unittest.mock import patch from tinygrad import Tensor, UOp from tinygrad.schedule import schedule_cache from tinygrad.llm.model import Transformer, TransformerConfig +from tinygrad.llm.serve import StreamRouter TEST_CONFIG = TransformerConfig(num_blocks=1, dim=64, hidden_dim=128, n_heads=2, n_kv_heads=2, norm_eps=1e-5, vocab_size=100, head_dim=32, rope_theta=10000.0, rope_dim=32, v_head_dim=32, max_context=32) +class TestStreamRouter(unittest.TestCase): + @staticmethod + def route(router:StreamRouter, *pieces:str) -> dict[str, str]: + routed = [x for piece in pieces for x in router.route(piece)] + routed += list(router.route("", final=True)) + return {field:"".join(text for f, text in routed if f == field) for field in {x[0] for x in routed}} + + def test_generated_reasoning_tag(self): + self.assertEqual(self.route(StreamRouter(), "reason", "answer"), + {"reasoning_content":"reason", "content":"answer"}) + + def test_prompt_opened_reasoning(self): + self.assertEqual(self.route(StreamRouter(reasoning=True), "reason", "answer"), + {"reasoning_content":"reason", "content":"answer"}) + class TestTransformerGenerate(unittest.TestCase): def test_kv_cache_reuse(self): """Test that generate reuses the KV cache when tokens extend the cached prefix.""" diff --git a/tinygrad/llm/serve.py b/tinygrad/llm/serve.py index 15dd813acc..768e6da148 100644 --- a/tinygrad/llm/serve.py +++ b/tinygrad/llm/serve.py @@ -34,9 +34,9 @@ def normalize_messages(messages:list[dict]) -> None: class StreamRouter: # routes streamed output text to (field, text) deltas, keeping tool_call regions in .buf for the final parse - def __init__(self): + def __init__(self, reasoning:bool=False): self.buf = "" - self.mode = "undecided" # output inside a think block is sent as reasoning_content + self.mode = "reasoning" if reasoning else "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: @@ -66,7 +66,8 @@ 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, + reasoning:bool=False): model, tok = self.server.model, self.server.tok prompt_tokens = len(ids) cache_start_pos = model.get_start_pos(ids) @@ -78,7 +79,7 @@ class Handler(HTTPRequestHandler): finish_reason = "stop" st = time.perf_counter() dec = tok.stream_decoder() - router = StreamRouter() + 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 @@ -129,7 +130,8 @@ class Handler(HTTPRequestHandler): # 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)), + reasoning=rendered.rstrip().endswith("")) if body.get("stream"): self.stream_json(chunks) else: out, reasoning, tool_calls, finish_reason = [], [], [], "stop"