-
Notifications
You must be signed in to change notification settings - Fork 329
Forward declared tools to the DiffusionGemma visual server #864
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
6ed24b0
b8f3c31
12d3153
f6334ad
741b83f
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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")) | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
When an OpenAI-compatible client sends a valid function tool that only includes the required Useful? React with 👍 / 👎.
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. common_chat_tools_parse_oaicompat uses function.value("description", "") and function.value("parameters", {}) at chat.cpp:551-555; only name uses .at(). A name-only tool parses without error, so there is no ERR parse to guard against. There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
When an OpenAI-compatible request includes Useful? React with 👍 / 👎.
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Same as the earlier items on this file: this shim is an internal loopback bridge that Studio spawns, not a public entry point. The Studio backend parses the native tool-call text into structured tool_calls (unslothai/unsloth#6851), so emitting them here as well would double-process the same generation. Forwarding the schemas is what makes the visual server render them into the chat template; the conversion to Chat Completions tool-call fields happens one layer up, not in the shim. There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
When a client sends a forced tool choice (for example Useful? React with 👍 / 👎.
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Same as the earlier forced-tool_choice item: the diffusion visual server samples with no grammar and never reads tool_choice, so required/forced enforcement is not achievable on this path even if the field were forwarded. _tools_for_choice already narrows to the forced function so only that one schema is advertised, which is the most the server can act on here. |
||
| 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) | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
For direct Useful? React with 👍 / 👎.
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. By design: this shim is an internal loopback bridge that Studio spawns, and the Studio backend heals the native tool-call text into structured tool_calls (unslothai/unsloth#6851). On the supported Studio path the call is surfaced as tool_calls; a direct client on the loopback port is not a supported entry point, and emitting tool_calls in the shim would double-process against that healing. There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
When tools are supplied, this new argument lets the visual server render schemas and the model can respond with a tool-call marker, but Useful? React with 👍 / 👎.
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Same as the earlier item: this shim is an internal loopback bridge Studio spawns, and the Studio backend heals the native tool-call text into structured tool_calls (unslothai/unsloth#6851). A direct client on the loopback port is not a supported entry point, and emitting tool_calls here would double-process against that healing. |
||
| 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)) | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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 | ||
|
oobabooga marked this conversation as resolved.
Comment on lines
+227
to
+228
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
When a request uses Useful? React with 👍 / 👎.
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The diffusion visual server renders only the prompt and samples with no grammar (data.grammar is never applied) and never reads tool_choice, so required/forced enforcement is not achievable on this path regardless of forwarding. _tools_for_choice still narrows to the forced function for advertising, so forwarding the choice would be a no-op here. |
||
| with open(self.req, "w", encoding="utf-8") as f: | ||
| json.dump(req, f, ensure_ascii=False) | ||
|
Comment on lines
+228
to
+230
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
On Windows or any non-UTF-8 locale, this new forwarding path serializes tool schemas through Useful? React with 👍 / 👎.
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Fixed in 741b83f: the request file is now opened with encoding="utf-8", so non-ASCII tool descriptions and message content no longer depend on the platform default codec (this also covered the pre-existing message path). Confirmed the same content raises UnicodeEncodeError under cp1252. |
||
| 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,16 +286,16 @@ 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"]: | ||
| raise | ||
| # 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: | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
When a client passes a forced
tool_choicefor one function, this only narrows the advertised schemas and then discards the actual choice before calling the visual server. The downstream request is therefore indistinguishable fromautowith one tool, so the model can still return plain text instead of the required function call; OpenAI-compatible clients that force a tool rely on that call being enforced or rejected, not treated as optional.Useful? React with 👍 / 👎.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Same as the earlier forced-tool_choice item: the diffusion visual server samples with no grammar (data.grammar is never applied) and never reads tool_choice, so required/forced enforcement is not achievable on this path regardless of forwarding. _tools_for_choice still narrows to the forced function for advertising.