diff --git a/unsloth_zoo/diffusion_studio/shim.py b/unsloth_zoo/diffusion_studio/shim.py index 6055b52bf7..2687bfd901 100644 --- a/unsloth_zoo/diffusion_studio/shim.py +++ b/unsloth_zoo/diffusion_studio/shim.py @@ -217,10 +217,25 @@ def health(): return {"status": "ok", "model": MODEL_ID} +def _tools_for_choice(tools, tool_choice): + # Honor tool_choice before advertising: "none" hides tools, a forced function + # narrows to just that one, anything else forwards all. + if isinstance(tool_choice, str) and tool_choice.lower() == "none": + return None + if isinstance(tool_choice, dict): + name = (tool_choice.get("function") or {}).get("name") + if name: + return [t for t in tools or [] + if isinstance(t, dict) and (t.get("function") or {}).get("name") == name] or None + return tools + + @app.post("/v1/chat/completions") async def chat(req: Request): body = await req.json() messages = body.get("messages", []) + # forwarded to the visual server, honoring tool_choice + tools = _tools_for_choice(body.get("tools"), body.get("tool_choice")) stream = bool(body.get("stream", False)) max_blocks = _max_blocks(body) seed = int(body.get("seed", 3407)) @@ -234,7 +249,7 @@ async def chat(req: Request): def work(): with _LOCK: return V.generate_visual(srv, messages, seed=seed, max_blocks=max_blocks, - on_stats=stats_box.update) + on_stats=stats_box.update, tools=tools) try: text = await loop.run_in_executor(None, work) except V.ContextOverflow as exc: @@ -276,7 +291,7 @@ def work(): with _LOCK: full = V.generate_visual(srv, messages, seed=seed, max_blocks=max_blocks, on_frame=on_frame, on_commit=on_commit, - on_stats=stats_box.update) + on_stats=stats_box.update, tools=tools) loop.call_soon_threadsafe(q.put_nowait, ("done", full)) except V.ContextOverflow as exc: # context budget exceeded -> clean user-facing message loop.call_soon_threadsafe(q.put_nowait, ("overflow", exc)) diff --git a/unsloth_zoo/diffusion_studio/visual_engine.py b/unsloth_zoo/diffusion_studio/visual_engine.py index b7e7e13d61..5fbeabdbb9 100644 --- a/unsloth_zoo/diffusion_studio/visual_engine.py +++ b/unsloth_zoo/diffusion_studio/visual_engine.py @@ -216,14 +216,18 @@ def restart(self): pass self._spawn() - def _send(self, messages, n_blocks, seed): + def _send(self, messages, n_blocks, seed, tools=None): # A previous turn may have crashed the decoder; respawn before writing so a dead child does not # strand this and every later turn with a broken pipe. if self.p is None or self.p.poll() is not None: self.restart() - with open(self.req, "w") as f: - json.dump({"seed": int(seed), "n_blocks": int(n_blocks), "messages": messages}, f, - ensure_ascii=False) + req = {"seed": int(seed), "n_blocks": int(n_blocks), "messages": messages} + # Forward tools so the server renders them into the chat template (inputs.tools) + # for schema-correct <|tool_call> args. + if tools: + req["tools"] = tools + with open(self.req, "w", encoding="utf-8") as f: + json.dump(req, f, ensure_ascii=False) try: self.p.stdin.write(self.req + "\n") self.p.stdin.flush() @@ -256,7 +260,7 @@ def _parse_stats(line): return stats -def generate_visual(server, messages, seed=3407, max_blocks=8, on_frame=None, on_commit=None, on_stats=None): +def generate_visual(server, messages, seed=3407, max_blocks=8, on_frame=None, on_commit=None, on_stats=None, tools=None): """Stream one turn through the optimized visual server. on_frame(block, step, total, text): a denoising frame (the current argmax canvas, already decoded). @@ -282,7 +286,7 @@ def _commit(txt): for attempt in (0, 1): try: - return _generate_visual_once(server, messages, seed, max_blocks, _frame, _commit, on_stats) + return _generate_visual_once(server, messages, seed, max_blocks, _frame, _commit, on_stats, tools) except VisualServerCrashed: server.restart() # always bring a fresh server up so the next turn works regardless if attempt == 1 or progressed["emitted"]: @@ -290,8 +294,8 @@ def _commit(txt): # nothing emitted yet -> safe to transparently resend on the fresh server -def _generate_visual_once(server, messages, seed, max_blocks, on_frame, on_commit, on_stats): - server._send(messages, max_blocks, seed) +def _generate_visual_once(server, messages, seed, max_blocks, on_frame, on_commit, on_stats, tools=None): + server._send(messages, max_blocks, seed, tools) full_text = "" while True: