Compare commits

..

137 Commits

Author SHA1 Message Date
Georg Bauer
b99ce2aa10 fix for stats and model selection
Some checks failed
Weekly OSV dependency audit / dependency-audit (push) Failing after 5s
2026-09-05 08:17:07 +02:00
Georg Bauer
bd6353804b Optimize Qwen3.8 inference on Apple silicon 2026-09-04 23:27:09 +02:00
Georg Bauer
bf82df77cb Add native Qwen MTP speculation 2026-09-03 22:38:52 +02:00
Georg Bauer
714b39f7f3 Implement exact Qwen sparse attention 2026-09-03 21:45:39 +02:00
Georg Bauer
87ccf67d0c Stream Qwen PLE embeddings 2026-09-03 21:12:20 +02:00
Georg Bauer
c414640050 Implement Qwen text core 2026-09-03 20:31:53 +02:00
Georg Bauer
3773cfda2e Integrate Qwen3.8 model intake 2026-09-03 19:56:23 +02:00
Georg Bauer
4c83ac0360 Prepare Qwen3.8 artifact set 2026-09-03 19:15:53 +02:00
Georg Bauer
9556b26c55 Bind sessions to their selected model 2026-09-01 22:52:01 +02:00
Georg Bauer
35f306bf2a Lock tool sessions to their model 2026-09-01 22:37:57 +02:00
Georg Bauer
d40e86e5ef correct thinking modes for each model 2026-09-01 21:03:13 +02:00
Georg Bauer
ee515ee824 Load GLM vision weights on demand 2026-09-01 20:08:36 +02:00
Georg Bauer
46d6a976a5 Add GLM 5.3 Flash support 2026-09-01 19:06:20 +02:00
Georg Bauer
9a33c61ea6 Remove the canceled Mojo spike 2026-08-31 17:30:55 +02:00
Georg Bauer
0dcf181dbe Establish the Mojo 1.0 integration gate 2026-08-31 15:25:04 +02:00
Georg Bauer
2c2c3702b0 Add Cargo dependency audit gates 2026-08-31 14:26:51 +02:00
Georg Bauer
115737302b Move Mojo roadmap to Gitea 2026-08-31 14:22:48 +02:00
Georg Bauer
85138c49d4 Add standalone session exports 2026-08-31 14:07:29 +02:00
Georg Bauer
5c74739f65 Add native glob file discovery 2026-08-31 08:44:44 +02:00
Georg Bauer
8adac261ae Add durable inline SVG presentation 2026-08-31 08:29:24 +02:00
Georg Bauer
0fa15bb68b Remove deprecated Flash preview and legacy MTP 2026-08-31 08:04:20 +02:00
Georg Bauer
7f88e23884 Make F1 open the complete User Guide 2026-08-31 07:43:29 +02:00
Georg Bauer
32c3a1a47a Use login shell PATH for extension hooks
Some checks failed
Weekly OSV dependency audit / dependency-audit (push) Failing after 5s
2026-08-30 22:40:57 +02:00
Georg Bauer
6c4e792c8a Accelerate DeepSeek DSpark inference 2026-08-30 21:45:16 +02:00
Georg Bauer
28231e1faf Add end-to-end MXFP4 Metal support 2026-08-30 20:42:23 +02:00
Georg Bauer
3977261bd3 Add installable agent lifecycle extensions 2026-08-30 16:05:09 +02:00
Georg Bauer
1ac559bbd0 Reach DS4 Metal performance parity 2026-08-30 15:26:27 +02:00
Georg Bauer
79468c65b6 Expand detail views with the sidebar hidden 2026-08-30 13:25:45 +02:00
Georg Bauer
f1b62c65f9 Update agent feature documentation 2026-08-30 13:19:51 +02:00
Georg Bauer
1597ddbc11 Add hierarchical workspace instructions 2026-08-29 21:57:11 +02:00
Georg Bauer
e35f57bc83 Build tagged macOS releases 2026-08-29 21:22:58 +02:00
Georg Bauer
c0681dedb9 Integrate standard agent skills 2026-08-29 21:12:40 +02:00
Georg Bauer
f2772f20c5 Add fresh-agent Ralph loops 2026-08-29 21:05:07 +02:00
Georg Bauer
f12aafcb2e Make agent tools schema-driven 2026-08-29 20:43:56 +02:00
Georg Bauer
f1c177b754 Support DeepSeek V4 Flash 0731 2026-08-29 20:28:50 +02:00
Georg Bauer
ad855b321e workover of the mojo plan 2026-08-29 18:53:28 +02:00
Georg Bauer
5e179a0f2c thinking about mojo 2026-08-29 14:33:03 +02:00
Georg Bauer
02022f4fca Fix dependency audit findings
Some checks failed
Weekly OSV dependency audit / dependency-audit (push) Failing after 4s
2026-08-09 16:41:00 +02:00
Georg Bauer
39ee904b3b Add weekly OSV dependency audit 2026-08-09 16:26:26 +02:00
Georg Bauer
f62a706f12 Trust Dev Brain paths during shell approval 2026-08-02 14:46:26 +02:00
Georg Bauer
e52efb80b0 Store preferences by model profile 2026-08-02 12:38:21 +02:00
Georg Bauer
69c83e7059 Improve chat Markdown exports 2026-08-02 10:38:24 +02:00
Georg Bauer
0b4e4fcc4a Load shell environment once at startup 2026-08-02 10:10:17 +02:00
Georg Bauer
118a9d41e0 feat: record agent approvals 2026-08-02 09:47:28 +02:00
Georg Bauer
94ed73e235 feat: keep track of window size and position 2026-08-01 21:01:43 +02:00
Georg Bauer
c0e4ec13cc Select thinking mode from the chat 2026-08-01 09:49:37 +02:00
Georg Bauer
d0431edaff Use the user login shell for agent commands 2026-08-01 09:40:55 +02:00
Georg Bauer
d1c380f79d feat: DeepSeek V4 Flash 0731 added 2026-08-01 08:40:06 +02:00
Georg Bauer
4175706533 feat: a skill creating skill as default 2026-07-30 13:41:04 +02:00
Georg Bauer
f2fb8d79f7 fix: stricter dev brain validation 2026-07-30 10:21:27 +02:00
Georg Bauer
30b305bc22 fix: better validate messages 2026-07-30 10:05:00 +02:00
Georg Bauer
73d7c8af27 feat: AI based permission checks 2026-07-30 09:47:54 +02:00
Georg Bauer
c101ca1380 chore: removed old plan document 2026-07-29 19:29:00 +02:00
Georg Bauer
92f351ae0b chore: refactorings for more code reorganizing 2026-07-29 19:28:28 +02:00
Georg Bauer
95ba6c3c76 chore: updated refactor plan 2026-07-29 19:14:15 +02:00
6387538da9 Merge pull request 'Refactor application and runtime boundaries' (#61) from linux-refactoring-base into main
Reviewed-on: #61
2026-07-29 16:48:38 +00:00
Georg Bauer
83905690e6 Fix refactor recovery and metrics 2026-07-29 18:47:14 +02:00
Hermes Agent
607aaa2e7a Restore DSML tool-call recovery 2026-07-29 11:00:31 +00:00
Hermes Agent
0f0a1a5409 Refactor application and runtime boundaries 2026-07-29 10:40:39 +00:00
Hermes Agent
902037f947 Support Linux refactoring checks 2026-07-29 10:21:56 +00:00
Hermes Agent
36f1056cf7 Document refactoring opportunities 2026-07-29 07:08:03 +00:00
Georg Bauer
d8b29fb1d1 fix: hopefully fixed a2ui confusion by the model 2026-07-28 22:36:16 +02:00
Georg Bauer
39c4181ae2 chore: removed reference to old license for stuff we don't use anymore 2026-07-28 22:19:40 +02:00
Georg Bauer
9d1c403ae7 Replace C browser client with headless Chrome 2026-07-28 22:17:45 +02:00
Georg Bauer
9683d5d10b make a2ui usage more clear to prevent every-session-spam 2026-07-28 21:37:05 +02:00
Georg Bauer
dbccaf7088 fix: fixed git handling 2026-07-28 21:11:46 +02:00
Georg Bauer
9b5b3a2f6a fix: less fragile dev brain sources revisions 2026-07-28 20:47:44 +02:00
Georg Bauer
ed8f7c4a01 fix: fixes to agent behaviour and llm steering 2026-07-28 20:07:41 +02:00
Georg Bauer
1f028950b3 fix: better sidebar layout 2026-07-28 19:40:31 +02:00
Georg Bauer
eaa2fcfe42 Discard archived session checkpoints 2026-07-28 19:20:10 +02:00
Georg Bauer
74c5948955 Align session titles in the sidebar 2026-07-28 19:11:23 +02:00
Georg Bauer
191ff0b2ce Persist the default Git diff layout 2026-07-28 19:07:59 +02:00
Georg Bauer
07c89094b4 Add current session statistics 2026-07-28 19:04:02 +02:00
Georg Bauer
44d614378f Group chat token counts by thousands 2026-07-28 18:59:07 +02:00
Georg Bauer
2391e2a680 Add verified Dev Brain skills 2026-07-28 18:57:08 +02:00
Georg Bauer
5ec873074b Add A2UI heatmap charts 2026-07-28 18:44:45 +02:00
Georg Bauer
abac978f54 Document managed model storage 2026-07-28 18:39:37 +02:00
Georg Bauer
00af0335ff Use slider toggles for boolean controls 2026-07-28 18:38:30 +02:00
Georg Bauer
9cf1c70a67 Keep runtime files inside application support 2026-07-28 18:34:08 +02:00
Georg Bauer
31f6426eef Replace gix with configurable libgit2 diffs 2026-07-28 18:25:24 +02:00
Georg Bauer
ed2226bc60 Replace Git subprocesses with native APIs 2026-07-28 06:55:05 +02:00
Georg Bauer
90a445cafd fix: better validate for dev_brain and 50k reminder not triggering rebuild 2026-07-27 23:08:47 +02:00
Georg Bauer
d371d9d659 Render tool output as terminal text 2026-07-27 22:07:21 +02:00
Georg Bauer
a7f26ccb71 Fix Git diff rendering 2026-07-27 22:01:24 +02:00
Georg Bauer
6001e7312b fix: better handling of dev brain 2026-07-27 21:42:58 +02:00
Georg Bauer
9a55a49b00 chore: removed old plan that was finished 2026-07-27 21:17:05 +02:00
Georg Bauer
cfc4fce8d7 feat: first cut at dev brain 2026-07-27 21:16:10 +02:00
Georg Bauer
56be87d288 Add pane shortcuts to View menu 2026-07-27 20:14:19 +02:00
Georg Bauer
9d58f79295 Fix Git action bar sizing 2026-07-27 20:04:21 +02:00
Georg Bauer
b9933ae076 Add Git worktree pane 2026-07-27 19:44:12 +02:00
Georg Bauer
aebfe7aeb6 Add project and branch chat controls 2026-07-27 19:20:57 +02:00
Georg Bauer
7d865df20f Run chats independently 2026-07-27 19:12:59 +02:00
Georg Bauer
22a939751e Keep chat scrolling under user control 2026-07-27 18:58:45 +02:00
Georg Bauer
7f14b6008a Order sessions by recent use 2026-07-27 18:53:39 +02:00
Georg Bauer
1017087fc9 fix: better A2UI description 2026-07-27 18:49:06 +02:00
Georg Bauer
4adb49a4d1 Reuse canonical chat context 2026-07-27 18:35:52 +02:00
Georg Bauer
ee7e5696a5 chore: updated plan for dev brain 2026-07-27 17:08:13 +02:00
Georg Bauer
50ae1cd5b9 fix: markdown rendering for chat was broken 2026-07-27 16:38:48 +02:00
Georg Bauer
c4e92f6206 Add preference section navigation 2026-07-27 15:26:25 +02:00
Georg Bauer
c96675c462 feat: more native menus with expected reactions 2026-07-27 15:09:12 +02:00
Georg Bauer
8f02f0934b Improve the default system prompt 2026-07-27 14:26:10 +02:00
Georg Bauer
a90a241ae1 Confirm session deletion 2026-07-27 14:18:32 +02:00
Georg Bauer
8a08cd7b52 Improve chat rendering and generation summaries 2026-07-27 14:04:14 +02:00
Georg Bauer
6a541b51f9 fix: rendering in a2ui surface and viewport aligned 2026-07-27 13:53:01 +02:00
Georg Bauer
f69623ad02 feat: moved a2ui surface to a separate panel 2026-07-27 12:58:46 +02:00
Georg Bauer
ef1c574a7f fix: better hardening of tool calls 2026-07-27 12:12:35 +02:00
Georg Bauer
4a309091b6 prev/next A2UI surface and better dismissal 2026-07-27 11:52:15 +02:00
Georg Bauer
c9c2d8efd5 feat: a2ui interface to enablee the LLM to give structured information 2026-07-27 11:19:57 +02:00
Georg Bauer
27b10ee0e1 chore: more on menus 2026-07-27 07:42:59 +02:00
Georg Bauer
61c2028e27 chore: reduce plan to bare minimum - no issues 2026-07-27 07:36:07 +02:00
Georg Bauer
21b9223706 chore: updated PLAN 2026-07-26 22:28:59 +02:00
Georg Bauer
47b405e72b fix: ssd streaming and deepseek fixes 2026-07-26 21:50:59 +02:00
Georg Bauer
0d80c217c4 Finish DS4 execution parity 2026-07-26 20:44:23 +02:00
Georg Bauer
fd3f8e45dc Upgrade Iced and render Markdown tables 2026-07-26 18:45:33 +02:00
Georg Bauer
4420b81117 Integrate DS4 execution parity in Rust 2026-07-26 17:58:05 +02:00
Georg Bauer
c9f0c3661c Harden local tool execution 2026-07-26 14:36:38 +02:00
Georg Bauer
171b041ba6 Finish long-running agent parity 2026-07-26 13:41:53 +02:00
Georg Bauer
65c9cbfc45 Add GLM 5.2 Metal execution 2026-07-26 12:16:30 +02:00
Georg Bauer
1de954b579 Support DeepSeek V4 Pro execution 2026-07-26 12:02:06 +02:00
Georg Bauer
3c75b8f6c1 Support long-running agent turns 2026-07-26 11:07:28 +02:00
Georg Bauer
2a14b93335 Make compaction transitions recoverable 2026-07-26 10:56:19 +02:00
Georg Bauer
6aa45b2cf0 Add agent context compaction 2026-07-26 10:48:29 +02:00
Georg Bauer
1dcadeb882 Scroll to the focused preference from its unscrolled position 2026-07-26 10:27:04 +02:00
Georg Bauer
c5e812f9b9 Scroll the focused preference into view 2026-07-26 10:22:14 +02:00
Georg Bauer
671724949f Keep the composer out of a dialog's field order 2026-07-26 10:11:40 +02:00
Georg Bauer
a0c3a72f4e Move focus between form fields with tab 2026-07-26 10:07:39 +02:00
Georg Bauer
f2133d561b Chart session checkpoints as their own cache bucket 2026-07-26 09:56:52 +02:00
Georg Bauer
5d441038bc Move preferences from the database to a YAML config file
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-07-26 09:50:39 +02:00
Georg Bauer
3f2c42513f Show the context fill grade as a pie chart
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-07-26 09:33:14 +02:00
Georg Bauer
76a5dd5b26 Name session checkpoints by their session title
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-07-26 09:25:37 +02:00
Georg Bauer
1f27270d66 Reclaim KV checkpoints the index cannot reach
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-07-26 09:25:37 +02:00
Georg Bauer
4fc423314b Center the detail tabs over the chat area
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-07-26 09:16:55 +02:00
Georg Bauer
63e2a74ad5 Manage the KV cache disc usage from preferences and stats
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-07-26 09:10:18 +02:00
Georg Bauer
b182e7007c Let the sidebar divider be dragged 2026-07-26 00:12:04 +02:00
Georg Bauer
9bd648d77b Explain each preference in a hover tooltip 2026-07-26 00:04:57 +02:00
Georg Bauer
ce18a2db9b Collapse a project's sessions by clicking its name 2026-07-25 23:58:21 +02:00
Georg Bauer
e44680541d Reopen the last project with a focused draft chat 2026-07-25 23:52:39 +02:00
153 changed files with 108042 additions and 9578 deletions

21
.cargo/audit.toml Normal file
View File

@@ -0,0 +1,21 @@
[advisories]
ignore = [
# Unmaintained transitive dependencies with no patched version.
"RUSTSEC-2020-0095", # difference via the release-only apple-codesign crate
"RUSTSEC-2024-0436", # paste via Iced's wgpu Metal backend
"RUSTSEC-2025-0134", # rustls-pemfile via apple-codesign's unused remote-signing stack
"RUSTSEC-2026-0192", # ttf-parser via Iced's text and SVG renderers
"RUSTSEC-2026-0206", # rustybuzz via Iced's SVG renderer
# No compatible upgrade exists; Cryoglyph does not call the affected LruCache::pop API.
"RUSTSEC-2026-0253",
# No patch exists; the release packager only creates ad-hoc signatures without RSA keys.
"RUSTSEC-2023-0071",
]
[output]
deny = ["warnings"]
format = "terminal"
quiet = false
show_tree = true

2
.cargo/config.toml Normal file
View File

@@ -0,0 +1,2 @@
[env]
MACOSX_DEPLOYMENT_TARGET = "15.0"

154
.gitea/scripts/release.py Normal file
View File

@@ -0,0 +1,154 @@
#!/usr/bin/env python3
"""Create and publish Gitea releases for tagged DS4Server builds."""
from __future__ import annotations
import json
import os
import re
import subprocess
import sys
import unittest
from urllib.error import HTTPError
from urllib.parse import quote
from urllib.request import Request, urlopen
def issue_numbers(message: str) -> list[int]:
return sorted({int(number) for number in re.findall(r"#(\d+)", message)})
def markdown(text: str) -> str:
return text.replace("\\", "\\\\").replace("[", "\\[").replace("]", "\\]")
def release_body(issues: list[dict[str, object]], server: str, repository: str, sha: str) -> str:
groups = {"Fixes": [], "Improvements": [], "Other changes": []}
for issue in issues:
labels = {label["name"] for label in issue.get("labels", [])} # type: ignore[index]
group = "Fixes" if "bug" in labels else "Improvements" if "enhancement" in labels else "Other changes"
title = " ".join(str(issue["title"]).split())
groups[group].append(f"- [#{issue['number']} {markdown(title)}]({issue['html_url']})")
sections = [f"## {name}\n" + "\n".join(items) for name, items in groups.items() if items]
if not sections:
sections.append("No closed issues were linked from commits in this release.")
sections.append(f"Built from [{sha[:12]}]({server}/{repository}/commit/{sha}).")
return "\n\n".join(sections)
class Gitea:
def __init__(self) -> None:
self.api_url = os.environ["GITEA_API_URL"].rstrip("/")
self.repository = os.environ["GITEA_REPOSITORY"]
self.token = os.environ["GITEA_TOKEN"]
def request(self, method: str, path: str, payload: object | None = None) -> object:
data = None if payload is None else json.dumps(payload).encode()
request = Request(
f"{self.api_url}{path}",
data=data,
method=method,
headers={"Authorization": f"token {self.token}", "Content-Type": "application/json"},
)
with urlopen(request) as response:
content = response.read()
return None if not content else json.loads(content)
def issue(self, number: int) -> dict[str, object]:
return self.request("GET", f"/repos/{self.repository}/issues/{number}") # type: ignore[return-value]
def commit_messages(sha: str) -> str:
try:
previous = subprocess.check_output(
["git", "describe", "--tags", "--abbrev=0", f"{sha}^"],
stderr=subprocess.DEVNULL,
text=True,
).strip()
except subprocess.CalledProcessError:
previous = ""
revision = f"{previous}..{sha}" if previous else sha
return subprocess.check_output(["git", "log", "--format=%s%n%b", revision], text=True)
def prepare() -> None:
gitea = Gitea()
tag = os.environ["GITEA_REF_NAME"]
sha = os.environ["GITEA_SHA"]
server = os.environ["GITEA_SERVER_URL"].rstrip("/")
issues = [gitea.issue(number) for number in issue_numbers(commit_messages(sha))]
closed = [issue for issue in issues if issue["state"] == "closed"]
payload = {
"tag_name": tag,
"target_commitish": sha,
"name": tag,
"body": release_body(closed, server, gitea.repository, sha),
"draft": True,
"prerelease": False,
}
try:
release = gitea.request("GET", f"/repos/{gitea.repository}/releases/tags/{quote(tag, safe='')}")
except HTTPError as error:
if error.code != 404:
raise
release = gitea.request("POST", f"/repos/{gitea.repository}/releases", payload)
else:
if not release["draft"]: # type: ignore[index]
raise RuntimeError(f"release {tag} is already published")
release = gitea.request("PATCH", f"/repos/{gitea.repository}/releases/{release['id']}", payload) # type: ignore[index]
for asset in release["assets"] or []: # type: ignore[index]
gitea.request(
"DELETE",
f"/repos/{gitea.repository}/releases/{release['id']}/assets/{asset['id']}", # type: ignore[index]
)
with open(os.environ["GITHUB_OUTPUT"], "a", encoding="utf-8") as output:
output.write(f"release_id={release['id']}\n") # type: ignore[index]
def publish() -> None:
gitea = Gitea()
release_id = int(os.environ["RELEASE_ID"])
gitea.request("PATCH", f"/repos/{gitea.repository}/releases/{release_id}", {"draft": False})
class ReleaseTests(unittest.TestCase):
def test_notes_are_unique_sorted_grouped_and_linked(self) -> None:
self.assertEqual(issue_numbers("Fix #12, refs #3 and #12"), [3, 12])
body = release_body(
[
{
"number": 12,
"title": "Fix [unlock]",
"html_url": "https://example.test/issues/12",
"labels": [{"name": "bug"}],
},
{
"number": 3,
"title": "Add search",
"html_url": "https://example.test/issues/3",
"labels": [{"name": "enhancement"}],
},
],
"https://example.test",
"owner/repo",
"1234567890abcdef",
)
self.assertIn("## Fixes\n- [#12 Fix \\[unlock\\]](https://example.test/issues/12)", body)
self.assertIn("## Improvements\n- [#3 Add search](https://example.test/issues/3)", body)
self.assertIn("[1234567890ab](https://example.test/owner/repo/commit/1234567890abcdef)", body)
if __name__ == "__main__":
command = sys.argv[1] if len(sys.argv) == 2 else ""
if command == "prepare":
prepare()
elif command == "publish":
publish()
elif command == "test":
unittest.main(argv=[sys.argv[0]])
else:
raise SystemExit("usage: release.py prepare|publish|test")

View File

@@ -0,0 +1,31 @@
name: Weekly OSV dependency audit
on:
schedule:
# Runs every Monday. Change the minute to stagger repositories.
- cron: "41 3 * * 1"
workflow_dispatch:
jobs:
dependency-audit:
runs-on: linux-arm64
timeout-minutes: 20
steps:
- name: Check out repository
uses: actions/checkout@v4
- name: Install cargo-audit
uses: taiki-e/install-action@v2
with:
tool: cargo-audit@0.22.2
- name: Audit Rust dependencies
run: |
cargo audit
cargo audit --file tools/macos-packager/Cargo.lock
- name: Scan dependencies for known vulnerabilities
uses: docker://ghcr.io/google/osv-scanner:v2
with:
args: scan source --recursive .

View File

@@ -0,0 +1,108 @@
name: Tagged release
on:
push:
tags:
- "*"
env:
CARGO_TERM_COLOR: always
jobs:
prepare-release:
runs-on: linux-arm64
outputs:
release_id: ${{ steps.release.outputs.release_id }}
steps:
- name: Check out the tag
uses: actions/checkout@v4
with:
fetch-depth: 0
- name: Install cargo-audit
uses: taiki-e/install-action@v2
with:
tool: cargo-audit@0.22.2
- name: Audit Rust dependencies
run: |
cargo audit
cargo audit --file tools/macos-packager/Cargo.lock
- name: Test release-note generation
run: python3 .gitea/scripts/release.py test
- name: Create draft release
id: release
env:
GITEA_API_URL: ${{ gitea.server_url }}/api/v1
GITEA_REPOSITORY: ${{ gitea.repository }}
GITEA_REF_NAME: ${{ gitea.ref_name }}
GITEA_SERVER_URL: ${{ gitea.server_url }}
GITEA_SHA: ${{ gitea.sha }}
GITEA_TOKEN: ${{ secrets.GITEA_TOKEN }}
run: python3 .gitea/scripts/release.py prepare
build-macos-arm64:
needs: prepare-release
runs-on: linux-arm64
container: ghcr.io/rust-cross/cargo-zigbuild@sha256:82af75c41958c2af2787e8bedd912da7678a9438937e223e9d83d006d747b38b
steps:
- name: Check out the tag
env:
GITEA_REF_NAME: ${{ gitea.ref_name }}
GITEA_REPOSITORY: ${{ gitea.repository }}
GITEA_SERVER_URL: ${{ gitea.server_url }}
GITEA_TOKEN: ${{ secrets.GITEA_TOKEN }}
run: |
set -eu
git init
git -c http.extraHeader="Authorization: token $GITEA_TOKEN" fetch --depth=1 \
"$GITEA_SERVER_URL/$GITEA_REPOSITORY.git" "refs/tags/$GITEA_REF_NAME"
git checkout --detach FETCH_HEAD
- name: Build and package the Apple Silicon application
env:
CARGO_TARGET_AARCH64_APPLE_DARWIN_RUSTFLAGS: -C link-arg=-Wl,-headerpad,0x1000
GITEA_REF_NAME: ${{ gitea.ref_name }}
run: |
set -eu
version=$(printf '%s' "$GITEA_REF_NAME" | sed 's/[^A-Za-z0-9._-]/-/g')
asset="dist/DS4Server-${version}-darwin-arm64.dmg"
cargo zigbuild --release --locked --target aarch64-apple-darwin --bin ds4-server
cargo run --release --locked --manifest-path tools/macos-packager/Cargo.toml -- \
"$GITEA_REF_NAME" target/aarch64-apple-darwin/release/ds4-server "$asset"
- name: Upload the DMG
env:
GITEA_API_URL: ${{ gitea.server_url }}/api/v1
GITEA_REPOSITORY: ${{ gitea.repository }}
GITEA_TOKEN: ${{ secrets.GITEA_TOKEN }}
RELEASE_ID: ${{ needs.prepare-release.outputs.release_id }}
run: |
set -eu
asset=$(find dist -maxdepth 1 -type f -name '*.dmg' -print -quit)
test -n "$asset"
name=$(basename "$asset")
curl --fail --silent --show-error --retry 3 \
-H "Authorization: token $GITEA_TOKEN" \
-H "Content-Type: application/octet-stream" \
--data-binary "@$asset" \
"$GITEA_API_URL/repos/$GITEA_REPOSITORY/releases/$RELEASE_ID/assets?name=$name"
publish-release:
needs:
- prepare-release
- build-macos-arm64
runs-on: linux-arm64
steps:
- name: Check out the tag
uses: actions/checkout@v4
- name: Publish complete release
env:
GITEA_API_URL: ${{ gitea.server_url }}/api/v1
GITEA_REPOSITORY: ${{ gitea.repository }}
GITEA_TOKEN: ${{ secrets.GITEA_TOKEN }}
RELEASE_ID: ${{ needs.prepare-release.outputs.release_id }}
run: python3 .gitea/scripts/release.py publish

View File

@@ -3,6 +3,8 @@
- Prefer simple, idiomatic Rust; reuse existing code and dependencies before adding abstractions or crates.
- Keep changes focused, handle errors explicitly, and add the smallest useful test for non-trivial behavior.
- Preserve `rustfmt` output and keep Clippy warning-free.
- Treat DS4 as the behavioral oracle. Model execution, token processing, context accounting, and KV-cache behavior must remain identical to DS4. Differences are bugs unless they are unavoidable Rust/platform requirements and the user explicitly approves them before implementation; when uncertain, preserve DS4 behavior.
- Keep all application, model-loading, inference-graph, session, and orchestration code in Rust. Never copy, vendor, compile, or call DS4's C implementation; use it only as a behavioral and performance reference. Native code is limited to the Objective-C Metal bridge and `.metal` kernels required to execute GPU work.
- This is not a GitHub project. Use direct `git` commands for version control and the `tea` CLI for forge operations; do not use GitHub tools or workflows.
- Issues are managed with the command "tea" run from the project directory.
@@ -11,6 +13,7 @@
- Follow the Codex macOS dark UI: near-black backgrounds, subtly raised surfaces, quiet borders, rounded corners, and restrained contrast.
- Present overviews as one bordered panel with divided rows, a clear primary label, muted supporting text, and compact trailing actions.
- Reuse the shared surface and action-button styles in `src/app/view.rs`; keep button shape, padding, typography, hover, and disabled states consistent across windows. Destructive actions may differ by color only, while navigation controls may remain flat.
- Render boolean controls as macOS-style slider toggles through the shared `toggle` helper, not checkboxes.
- Prefer generous spacing and clear hierarchy over decoration; avoid one-off colors, card stacks, oversized controls, and screen-specific button styling.
## Commit gates

3144
Cargo.lock generated

File diff suppressed because it is too large Load Diff

View File

@@ -3,6 +3,7 @@ name = "ds4-server"
version = "0.1.0"
edition = "2024"
rust-version = "1.97"
default-run = "ds4-server"
description = "A native macOS coding-agent GUI for DwarfStar"
license = "MIT"
publish = false
@@ -12,19 +13,33 @@ build = "build.rs"
cc = "1.3.0"
[dependencies]
base64 = "0.22.1"
diesel = { version = "2.3.11", features = ["sqlite", "returning_clauses_for_sqlite_3_35", "64-column-tables"] }
diesel_migrations = "2.3.2"
iced = { version = "0.13.1", features = ["advanced", "highlighter", "markdown", "svg", "tokio"] }
git2 = { version = "0.21.0", features = ["https", "vendored-libgit2", "vendored-openssl"] }
headless_chrome = "1.0.22"
iced = { version = "0.14.0", default-features = false, features = ["advanced", "image-without-codecs", "markdown", "svg", "tokio", "wgpu"] }
image = { version = "0.25.10", default-features = false, features = ["gif", "jpeg", "png", "webp"] }
libc = "0.2.186"
memmap2 = "0.9.11"
png = "0.17.16"
rfd = "0.15.4"
pulldown-cmark = "0.12.2"
regex = "1.13.1"
rfd = { version = "0.15.4", default-features = false }
rustc-hash = "2.1.3"
serde = { version = "1.0.228", features = ["derive"] }
serde_json = { version = "1.0.149", features = ["preserve_order", "raw_value"] }
serde_norway = "0.9.42"
sha2 = "0.11.0"
shlex = "2.0.1"
time = { version = "0.3.54", features = ["formatting", "parsing"] }
turbovault-parser = "1.6.0"
ureq = { version = "3.3.0", default-features = false, features = ["rustls"] }
url = "2.5.8"
usvg = "0.45.1"
[target.'cfg(target_os = "macos")'.dependencies]
muda = "0.15.3"
muda = { version = "0.19.3", default-features = false }
[package.metadata.packager]
product-name = "DS4Server"
@@ -32,7 +47,11 @@ identifier = "de.rfc1437.ds4server"
description = "A native macOS coding-agent GUI for DwarfStar"
binaries = [{ path = "ds4-server", main = true }]
icons = ["assets/DS4Server.icns", "assets/app-icon.png"]
resources = ["metal"]
resources = ["metal", "assets/dev-brain"]
[package.metadata.packager.macos]
minimum-system-version = "13.0"
minimum-system-version = "15.0"
[[bin]]
name = "qwen38-artifacts"
path = "tools/qwen38-artifacts.rs"

334
PLAN.md
View File

@@ -1,334 +0,0 @@
# DS4Server implementation plan
Bundle/application identifier: `DS4Server.rfc1437.de`
DS4Server is the native macOS Rust/Iced application counterpart to `../ds4`.
The reference implementation remains normative: `ds4.c`/`ds4.h` define model
and session behavior, `ds4_server.c` defines the HTTP API, and `ds4_agent.c`
defines local-agent behavior and tools. Parity means matching observable
behavior and fixtures, not merely exposing similar controls. All planned
execution targets one self-contained Mac.
## Current state
- The macOS application shell, projects, session rename/pin/archive/delete,
SQLite transcripts, model preferences and downloads, Stats dashboard, native
Application/Edit/Window menu integration, development bundle, and shared lazy
model lifecycle are implemented.
- DeepSeek V4 Flash runs through the Rust/Metal executor with full configured
context, ratio-4 sparse attention, cold and resumed batched prefill, prefix
continuation, durable KV checkpoints, cancellation, and idle unload.
- The configurable localhost endpoint implements all `ds4_server.c` routes,
streaming and non-streaming envelopes, model aliases, CORS, tool replay and
continuation, malformed DSML recovery, transient prefix-aware KV reuse,
disconnect/queue cancellation, and shared model scheduling. The full
automated C/Rust differential corpus remains open.
- Local sessions persist user, assistant, reasoning, and tool roles. The agent
executes the complete `ds4_agent.c` starting tool set and reinjects the same
model-specific tool contract when sessions are opened or continued.
- The remaining baseline gaps are long-running agent behavior and tool safety.
SSD streaming, speculative decoding, steering, GLM 5.2 execution, and
DeepSeek V4 Pro execution are not implemented in the Rust executor. Some
related catalog, validation, and preference plumbing already exists but must
not be treated as runtime support.
## Delivery order
1. **Next:** long-running agent stability and remaining `ds4_agent.c` parity.
2. Tool hardening, approvals, and productive tool presentation.
3. Remaining DS4 execution technology, starting with SSD streaming, then
speculative decoding and the other Metal/runtime parity work.
4. Additional model execution: GLM 5.2 and DeepSeek V4 Pro.
5. Product completion, exhaustive parity verification, and distribution.
6. Optional extensions: Dev Brain and A2UI.
## 1. Next — long-running agent stability and `ds4_agent.c` parity
Goal: a local agent session must be able to run long read/edit/test loops,
cross the context limit repeatedly, survive interruption or relaunch, and
continue with the same durable task state as `ds4-agent`.
### Context compaction
- Port the reference soft trigger: compact before a user turn or tool
continuation at 85% context use, or when at most 8192 tokens remain, with the
free-token threshold capped to one eighth of small contexts.
- Port the hard trigger: before appending a tool result that would leave
insufficient answer room, compact once and retry. If it still does not fit,
return a bounded tool error that tells the model to request less output.
- Use the live model to generate an internal durable task-state summary. The
prompt and generated summary are private compaction work, never ordinary
user/assistant messages, and may not execute tools or retain thinking/DSML
control markup.
- Preserve goals, constraints, files touched, commands and important results,
decisions, known failures, and next steps. Prefer reloadable paths, ranges,
and commands over copying bulky data into the summary.
- Rebuild the model context exactly as the reference does: current system/tool
contract, durable summary, then a recent verbatim tail. Keep up to 10% of the
configured context as the tail, capped at 50000 tokens, and align it to a
user-turn boundary when possible.
- Generate at most the reference summary budget, stop at model control or tool
markers, and never let the private compaction exchange become the reusable
session prefix.
### Durable transition and recovery
- Treat transcript replacement and the new KV checkpoint as one logical
transition. Persist the compacted semantic transcript and checkpoint
metadata only after the rebuilt prefix is valid.
- On cancellation, summary failure, prefill failure, or application exit, keep
the previous durable transcript, invalidate any KV state contaminated by the
private compaction prompt, and make the next turn rebuild safely.
- Reopen a compacted session with the same summary and recent verbatim turns.
If its checkpoint is absent or incompatible, rebuild it from persisted
messages without changing the visible conversation.
- Preserve relevant live tool state across compaction. In particular, append a
compact observation for running shell jobs so the model can still inspect or
stop them after the context rebuild.
- Add a user-visible action equivalent to `/compact`, plus the reference
strip/rebuild behavior: discard a session KV payload without discarding its
transcript, then rebuild on the next use.
- Bind checkpoint compatibility to the model identity, quantization, context,
rendered transcript, and payload ABI. A model/configuration change must
rebuild rather than reuse an invalid prefix.
### Long-turn behavior
- Accept user input while an assistant/tool loop is active. Queue it visibly
and inject it after the current tool result, before the next assistant
continuation, matching `ds4_agent.c` instead of starting a competing turn.
- Preserve the reference rule that there is no arbitrary maximum tool-round
count. Completion, Stop, context pressure, or a real error ends the loop.
- Match the reference date/time context injection and periodic system/tool
prompt reminder so long or reopened sessions do not drift away from the tool
contract.
- Keep Stop cooperative across summarization, compacted-prefix prefill,
generation, and active tool work. A stop must always leave a transcript that
can be reopened.
### Presentation and verification
- Show explicit `Compacting`, rebuilding/prefill, queued-input, stopped, and
failed states without blocking the Iced event loop. Keep the compacted
summary inspectable without presenting the private prompt as user history.
- Add focused tests for soft and hard triggers, tail selection, tool-result
retry, cancellation rollback, checkpoint invalidation, queued input between
tool rounds, running-job preservation, and reopen after compaction.
- Add a reference fixture that runs a long tool loop through compaction, saves,
relaunches, and continues without losing the active task.
Exit criterion: repeat the reference `ds4-agent` long-context scenarios,
including compaction forced by a large tool result, stop during compaction, and
restart after a successful compaction. The same task state, recent turns, tool
contract, and running-job awareness must remain available.
## 2. Tool hardening and safety
Goal: make the existing tool set safe and clear enough for productive daily
use without weakening its ability to inspect, edit, build, and test a project.
- Keep file tools strictly inside the selected project after canonicalization,
including symlinks and nonexistent write targets. Boundary escapes are
rejected, not approved as exceptions.
- Centralize tool approval. Reuse the existing visible-browser approval path
for actions that require consent instead of adding tool-specific modal
systems.
- Require confirmation before destructive shell commands or commands that
affect state outside the project, elevate privileges, launch external apps,
or create unrequested network side effects. Show the exact command and
working directory; support Allow once and Deny. Do not persist broad trust
silently.
- Define the shell environment deliberately so application credentials and
unrelated secrets are not exposed accidentally. Preserve the minimum normal
development environment needed for project commands.
- Render streamed tool calls as compact cards with their real lifecycle:
parsing, awaiting approval, queued, running, completed, failed, or stopped.
Show bounded parameters and results, provide copy/open actions where useful,
and never display raw DSML as ordinary assistant prose.
- Keep background shell jobs cancellable and observable through app stop,
session switches, compaction, and application shutdown. Retain bounded output
files only as long as needed and report cleanup failures.
- Continue enforcing context-aware read/search/bash result limits. Oversized
output must remain available through continuation or a file reference rather
than being inserted wholesale into the prompt.
- Test project-boundary and symlink escapes, destructive-command approval,
denial and cancellation, browser approval, output bounding, and session
switching with a live job.
Exit criterion: use the agent for a real inspect/edit/test cycle while every
side effect is visible, risky actions require consent, Stop works at every
stage, and no file tool can escape the selected project.
## 3. DS4 execution technology parity
Goal: finish the model-independent Metal/runtime capabilities in `ds4.c`
before adding larger model families. Every capability must be shared by local
chat and the HTTP endpoint through the single process-wide model owner.
### 3.1 SSD streaming — first runtime priority
SSD streaming is the capacity prerequisite for larger models and therefore
comes before GLM 5.2 and DeepSeek V4 Pro execution.
- Keep non-routed weights resident while loading routed MoE experts from the
GGUF into a bounded in-memory cache on demand.
- Port DS4's automatic cache sizing from the recommended Metal working set,
non-routed weights, KV/context requirements, graph scratch, and routed-prefill
headroom. Preserve explicit expert-count and GiB-budget modes.
- Port cold start, hot-expert preload, explicit preload count, full resident
layer policy, cache eviction, and the overlapped prefill/decode loading path.
- Keep inference correct under cache misses, short reads, cancellation, memory
pressure, unload, and model replacement. A streaming failure must produce a
bounded request error, not partial model state.
- Report resident bytes, cache capacity, hit/miss and load rates, SSD bytes and
latency, preload progress, and memory-pressure decisions in Stats.
- First prove resident-versus-streaming token parity on DeepSeek V4 Flash. Keep
the cache and I/O layer model-aware so later GLM/Pro milestones add policy and
graph support rather than a second streaming subsystem.
### 3.2 Speculative decoding: legacy MTP and DSpark
- Load and validate the optional Flash legacy-MTP or DSpark support GGUF without
treating either as a standalone model. Preserve exact support-kind and target
compatibility checks.
- Port draft generation, confidence/margin pruning, target-model verification,
accepted-prefix commit, rejected-suffix rollback/replay, cache counters, and
cancellation boundaries. The target model always owns the output stream.
- Honor MTP draft depth/margin and DSpark confidence/strict settings. Match the
reference rules for greedy, sampled, and quality modes rather than applying
speculation where DS4 falls back to target-only decode.
- Match current DS4 compatibility between SSD streaming and DSpark; remove
temporary Rust-only exclusions where the reference supports the combination.
- Expose accepted/drafted token counts, acceptance rate, verifier cost, and
effective speedup in Stats. Do not call the feature complete merely because
it produces correct tokens; it must also preserve checkpoints, tools,
streaming responses, usage accounting, and Stop behavior.
- GLM's in-model MTP path belongs to the GLM milestone, but it should reuse the
verifier/session machinery established here.
### 3.3 Remaining Metal execution controls
- Port directional steering files and exact FFN/attention application,
including DS4 defaults, validation, zero-scale behavior, and checkpoint/model
identity effects.
- Audit every accepted execution preference against the Rust executor. Power,
prefill chunking, quality/exact kernels, warm weights, simulated memory, and
expert profiling must either affect execution as DS4 defines or be disabled
with a clear explanation; persisted no-op settings are not parity.
- Add hardware-backed token/activation fixtures for each mode and keep the
ordinary resident Flash path unchanged when optional features are off.
### 3.4 Single-machine server batching
- Port DS4's resident multi-session batching and server scheduling only after
the serialized path remains the correctness oracle. Preserve per-request
cancellation, finish reasons, usage, and KV ownership while batching prefill
or decode work.
- Keep all scheduling, model state, KV state, and request handling within the
local process. Networked execution and non-Metal backends are outside the
product scope.
Exit criterion: DeepSeek V4 Flash matches DS4 token output and lifecycle in
resident, SSD-streamed, MTP, DSpark, steering, and batched-server
configurations, with optional modes off producing the same baseline behavior
as today.
## 4. Additional model execution
Start these only after the shared capacity and execution technology above is
stable. Catalog entries, settings, tokenizer work, or GGUF validation alone do
not make a model selectable for inference.
### GLM 5.2
- Port the GLM DSA/MLA graph, dense-cache behavior, model-specific tensor and
quantization paths, sampling defaults, reasoning controls, prompt rendering,
and stop tokens.
- Use the already defined GLM tool syntax through the same durable local-agent
loop and expose identical behavior through every HTTP route.
- Port GLM SSD streaming policy, resident full-layer selection, and the MTP
block stored in the main GGUF. Respect GLM restrictions on power, prefill
chunking, steering, and external support models.
- Validate resident and streamed token output against DS4 fixtures before the
Model Manager advertises GLM as runnable.
### DeepSeek V4 Pro
- Generalize the Flash graph only where Pro's dimensions, layers, routed
experts, quantization layouts, or output path actually differ.
- Support resident and SSD-streamed single-machine configurations with explicit
memory admission checks. Never begin a load that cannot leave room for the
configured KV/context and graph working set.
- Match Pro prompt, sampling, checkpoint, HTTP, and agent behavior. Preserve
the reference compatibility matrix for MTP/DSpark rather than assuming Flash
support artifacts work with Pro.
- Validate supported single-file Q2/Q4 configurations against DS4 fixtures
before advertising them as runnable.
Exit criterion: each advertised model passes the same local-agent, checkpoint,
HTTP, SSD-capacity, cancellation, and deterministic token-output matrix as
DeepSeek V4 Flash.
## 5. Product completion and verification
### Reference parity and regression coverage
- Complete the automated black-box C/Rust corpus for Models, Messages, Chat
Completions, Responses, Completions, and OPTIONS in streaming and
non-streaming modes. Cover aliases/defaults, malformed input, context limits,
tools, reasoning, cache reuse, disconnects, queueing, CORS, errors, usage, and
finish reasons.
- Add every new agent/runtime/model capability to both local and endpoint
regression coverage. HTTP conversations remain client-owned and may persist
only opaque content-addressed KV entries, never application transcripts.
- Reuse DS4 prompt, sampling, KV, server, agent, and hardware fixtures. Keep
formatting, Clippy, bundle, and unit tests as commit gates; run hardware
matrices only where the required GGUF and machine are available.
### macOS product and distribution
- Finish the remaining native menu and text behavior: functional File/View/Help
actions, dynamic enabled state, undo/redo parity, and selectable/copyable
transcript text.
- Finish session/model binding and recovery surfaces needed to explain why a
checkpoint is reused or rebuilt. Keep archive as metadata-only and never
delete project contents.
- Add production signing, hardened runtime, required entitlements,
notarization, update delivery, and release automation. Preserve all upstream
DwarfStar/llama.cpp notices for adapted code and kernels.
Exit criterion: a notarized build can be installed on a clean supported Mac,
run the full local-agent and endpoint smoke matrix, restart into its previous
sessions, and update without losing projects, transcripts, models, or KV data.
## 6. Optional future extensions
These are not DS4 baseline parity and must not delay the milestones above.
### Dev Brain
- Let the user opt into one existing Obsidian vault through a native folder
picker and persist a security-scoped bookmark for sandboxed restart access.
- Add one project-independent tool with bounded search, read, and create/append
operations for Markdown notes. Return vault-relative citations and exclude
`.obsidian`, hidden files, deletion, and arbitrary overwrite.
- Reuse the existing file/search implementation with a separate explicit vault
boundary. No vault means no tool and no effect on ordinary sessions.
Exit criterion: store a memory in one project, restart, and retrieve it with a
vault-relative citation from another project without granting broader disk
access.
### A2UI local-chat surfaces
- After the ordinary tool UI is stable, use `../bDS2` as the reference for
typed cards, charts, forms, lists, metrics, mind maps, tables, and tabs.
- Validate and persist versioned structured payloads with their messages.
Render with native Iced/SVG primitives; reject model-authored HTML,
JavaScript, styles, file URLs, and arbitrary commands.
- Keep A2UI local-chat-only until a real endpoint client requires a compatible
protocol. All mutating actions continue through the normal approval path.
Exit criterion: supported surfaces render accessibly, malformed payloads fall
back safely, interactions are validated, and the same surfaces rehydrate after
restart.

158
README.md
View File

@@ -1,51 +1,103 @@
# DS4Server
DS4Server is a native macOS coding-agent application that rewrites the
DwarfStar (`ds4`) inference engine in Rust. It uses Rust and Iced and will combine local model loading, an
OpenAI-compatible localhost endpoint, and project-scoped agent chat in one app.
DS4Server is a native macOS application for running DwarfStar models as a local
coding agent and as a localhost inference server. The inference engine, model
lifecycle, agent loop, and application UI are implemented in Rust with Iced;
the Metal kernels and small native integration layers are adapted from
[DwarfStar (`ds4`)](https://github.com/antirez/ds4).
DS4Server vendors and adapts the Metal kernels, Objective-C Metal glue, and
visible-Chrome web tool runtime from
[DwarfStar (`ds4`)](https://github.com/antirez/ds4). Their copyright and license
notices are retained in [`native/metal/LICENSE`](native/metal/LICENSE) and
[`native/web/LICENSE`](native/web/LICENSE).
## Capabilities
The current milestone provides a Codex-inspired project/session layout. A native
macOS folder picker selects each workspace, then the app asks for its display
name. Projects, sessions, and model preferences are persisted through Diesel in
SQLite. Open Preferences with `Command-,` to configure model, generation,
runtime, local endpoint, and idle-unload settings. The separate Model Manager
(`Shift-Command-M`)
lists local main and DSpark artifacts, their on-disk sizes and state, and lets
you download, resume, validate, or delete them. Rust-native background work
shows live byte progress, speed, and ETA in Model Manager and the app status
bar. Stopping or quitting keeps the partial file; the next Download/Resume
action continues from that exact byte after relaunch. Exact size and SHA-256
verification happen before an artifact becomes usable.
- **Project-based coding chat.** Add local project folders, keep multiple
persistent sessions per project, and rename, retitle, pin, archive, compact,
rebuild, or delete sessions. Chats stream Markdown, reasoning, tool activity,
context use, and generation speed.
- **Local coding tools.** The agent can read, search, list, write, and edit
project files; run and monitor asynchronous shell commands; search Google;
and visit rendered web pages. File access is confined to the project, output
is bounded, risky shell or visible-browser actions require approval, and
tool calls are validated against the same typed schemas shown to the model.
- **Fresh-agent Ralph loops.** For a bounded coding objective, the agent can
run foreground rounds in fresh model contexts while preserving workspace
changes between rounds. A validated structured report is the only
conversational handoff to the next worker; the loop stops on reported
completion, a blocker, failure, cancellation, or its configurable round
limit.
- **Workspace instructions and reusable skills.** A global `AGENTS.md` plus
project and nested `AGENTS.md`/`AGENTS.local.md` files provide hierarchical
guidance that is reconciled as the agent enters or changes a workspace area.
Standard `~/.agents/skills/<name>/SKILL.md` packages are advertised by
metadata and loaded in full only when a task matches.
- **Project-backed Dev Brain.** An optional managed Obsidian vault provides
source-verified wiki pages, ranked search, and on-demand skills. Each session
receives the verified skill names, descriptions, and paths, then loads full
instructions only when a task matches.
- **Long-running sessions.** Transcripts and summaries are stored in SQLite.
Automatic and manual context compaction preserve the complete visible chat,
while durable KV checkpoints make follow-up turns and relaunches resumable.
- **Native Metal inference.** DeepSeek V4 Flash, DeepSeek V4 Pro, and GLM 5.2
run through model-specific Rust/Metal executors with long-context prefill,
sampling, cancellation, idle unload, and a bounded resident session pool.
Supported acceleration includes SSD expert streaming, DeepSeek legacy MTP,
DSpark, GLM MTP, directional steering, and model-specific memory admission.
- **Local HTTP APIs.** A configurable listener on `127.0.0.1:4000` exposes
Models, Chat Completions, Completions, Anthropic Messages, and Responses APIs
with streaming, tool calls, usage accounting, cancellation, and prefix-aware
KV reuse. Endpoint conversations remain client-owned and are not added to
project history.
- **Model and runtime management.** Model Manager downloads, resumes, verifies,
and deletes supported GGUF artifacts with byte progress, speed, and ETA.
Preferences cover generation, runtime, SSD, speculative decoding, steering,
KV storage, the local endpoint, and idle unload. The Stats view reports model,
cache, SSD, speculative, session, and HTTP activity.
The selected DeepSeek V4 Flash model can run directly from a project session.
The model, KV/compressor state, 43-layer graph, sampling, and lifecycle are
owned by Rust; a fixed snapshot of the Objective-C Metal boundary and unchanged
Metal kernels is vendored and built inside this repository. Tokens stream into the chat UI, Stop cancels generation,
follow-up turns reuse durable transcript and KV state, and the model unloads
after the idle timeout. The graph uses the full configured context with the
ratio-4 sparse indexer.
Projects, sessions, and transcripts use SQLite. Settings use a compact YAML file
that contains only values changed from their defaults and can also be edited by
hand.
Project chat includes the native `ds4_agent.c` starting tool set: bounded file
read/continuation, write, anchored edit, search, directory listing, asynchronous
shell jobs, Google search, and rendered page visits. Tool calls and results are
persisted as transcript roles and automatically continue the same model turn.
File tools stay inside the selected project. Web tools ask before starting a
visible Chrome profile.
## A2UI widgets
The app also listens on `127.0.0.1:4000` by default for Models, Chat
Completions, Completions, Anthropic Messages, and Responses APIs. The listener,
port, and opt-in CORS are configurable in Preferences. The
endpoint and local chat share the single model owner. External conversations
are client-managed and never enter the project, session, message, or transcript
database; only opaque content-addressed KV cache files are retained. Model
discovery advertises only supported main artifacts that are fully downloaded
and verified on disk.
DS4Server renders A2UI v1.0 surfaces with these widgets:
- **Content:** Markdown text, images (including avatars), icons, video and audio
players, and dividers.
- **Layout:** rows, columns, lists, cards, modals, and tabs.
- **Controls:** buttons, text fields, checkboxes, sliders, date/time inputs, and
choice pickers.
- **Research:** bar, line, area, stacked-bar, pie, donut, and heatmap charts; tables,
metrics, timelines, maps, mind maps, and forms.
Video posters render in the surface; video and audio playback uses native
macOS AVKit player windows with standard transport, volume, and fullscreen
controls.
The newest A2UI surface is pinned above the transcript in a default 50/50
workspace split. Dismissing it requires confirmation and records a durable
boundary in the chat database, so reopening the session restores only the
newest surface created after the last dismissal. Left and right arrow buttons
reconstruct earlier and later read-only surface epochs from the recorded A2UI
protocol calls; model updates remain routed only to the current surface.
### Live A2UI validation
With DS4Server running and its local endpoint enabled, run the opt-in model
validation suite with:
```sh
cargo run --release -- validate-a2ui
```
It sends natural-language chart, form, choice, and full-composition prompts
through the current A2UI tool description, then replays every returned protocol
message through the real catalog and component-graph validators. Use `--case
pie-natural`, `--case composed-basics`, `--attempts 3`, `--model MODEL_ID`, or
`--endpoint URL` to narrow or repeat runs. This slow, model-dependent suite is
intentionally separate from `cargo test`.
## Build and run
DS4Server targets Apple silicon with macOS 15 or newer and requires Rust 1.97 plus
`cargo-packager` 0.11.8.
```sh
cargo install cargo-packager --locked --version 0.11.8
@@ -53,13 +105,25 @@ make bundle
open target/release/DS4Server.app
```
State is stored at:
Tagged versions are also published as Apple Silicon DMG files on the
[repository releases page](https://git.rfc1437.de/hugo/DS4Server/releases).
## Application data
```text
~/Library/Application Support/DS4Server.rfc1437.de/data.sqlite3
~/Library/Application Support/DS4Server.rfc1437.de/kv-cache/
~/Library/Application Support/de.rfc1437.ds4server/data.sqlite3
~/Library/Application Support/de.rfc1437.ds4server/config.yaml
~/Library/Application Support/de.rfc1437.ds4server/AGENTS.md
~/Library/Application Support/de.rfc1437.ds4server/models/
~/Library/Application Support/de.rfc1437.ds4server/kv-cache/
```
Deleting a project or session removes only DS4Server metadata. It never deletes
the referenced project directory. See [PLAN.md](PLAN.md) for the implementation
roadmap.
Deleting a project or session removes DS4Server metadata and its associated
checkpoint; it never deletes the referenced project directory. Standard agent
skills live separately under `~/.agents/skills/`. See the
[issue tracker](https://git.rfc1437.de/hugo/DS4Server/issues) for remaining
implementation work.
The adapted Metal and browser sources retain their upstream notices in
[`native/metal/LICENSE`](native/metal/LICENSE) and
[`native/web/LICENSE`](native/web/LICENSE).

144
assets/a2ui/catalog.json Normal file
View File

@@ -0,0 +1,144 @@
{
"$schema": "https://json-schema.org/draft/2020-12/schema",
"$id": "https://ds4server.local/a2ui/v1_0/catalog.json",
"title": "DS4Server A2UI Catalog",
"catalogId": "https://ds4server.local/a2ui/v1_0/catalog.json",
"instructions": "Use the Basic Catalog for interaction and layout. Use research components for compact analytical results.",
"components": {
"Text": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Text" },
"Image": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Image" },
"Icon": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Icon" },
"Video": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Video" },
"AudioPlayer": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/AudioPlayer" },
"Divider": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Divider" },
"Row": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Row" },
"Column": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Column" },
"List": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/List" },
"Card": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Card" },
"Modal": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Modal" },
"Tabs": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Tabs" },
"Button": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Button" },
"TextField": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/TextField" },
"CheckBox": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/CheckBox" },
"Slider": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/Slider" },
"DateTimeInput": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/DateTimeInput" },
"ChoicePicker": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/components/ChoicePicker" },
"Chart": { "$ref": "#/$defs/chart" },
"Table": { "$ref": "#/$defs/table" },
"Metric": { "$ref": "#/$defs/metric" },
"Timeline": { "$ref": "#/$defs/timeline" },
"Map": { "$ref": "#/$defs/map" },
"MindMap": { "$ref": "#/$defs/mindMap" },
"Form": { "$ref": "#/$defs/form" }
},
"functions": {
"required": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/required" },
"regex": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/regex" },
"length": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/length" },
"numeric": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/numeric" },
"email": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/email" },
"formatString": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/formatString" },
"formatNumber": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/formatNumber" },
"formatCurrency": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/formatCurrency" },
"formatDate": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/formatDate" },
"pluralize": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/pluralize" },
"openUrl": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/openUrl" },
"and": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/and" },
"or": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/or" },
"not": { "$ref": "https://a2ui.org/specification/v1_0/catalogs/basic/catalog.json#/functions/not" }
},
"$defs": {
"surfaceProperties": {
"type": "object",
"properties": {
"iconUrl": { "type": "string", "format": "uri" },
"agentDisplayName": { "type": "string" }
},
"additionalProperties": true
},
"researchBase": {
"type": "object",
"properties": {
"id": { "type": "string" },
"component": { "type": "string" },
"title": { "type": "string" },
"accessibility": { "type": "object" },
"weight": { "type": "number" },
"checks": { "type": "array" }
},
"required": ["id", "component"]
},
"chart": {
"allOf": [
{ "$ref": "#/$defs/researchBase" },
{ "properties": { "component": { "const": "Chart" }, "chartType": { "enum": ["bar", "line", "area", "stackedBar", "pie", "donut", "heatmap"] }, "series": {} }, "required": ["chartType", "series"] }
],
"unevaluatedProperties": false
},
"table": {
"allOf": [
{ "$ref": "#/$defs/researchBase" },
{ "properties": { "component": { "const": "Table" }, "columns": {}, "rows": {} }, "required": ["columns", "rows"] }
],
"unevaluatedProperties": false
},
"metric": {
"allOf": [
{ "$ref": "#/$defs/researchBase" },
{ "properties": { "component": { "const": "Metric" }, "label": {}, "value": {}, "detail": {}, "trend": {} }, "required": ["label", "value"] }
],
"unevaluatedProperties": false
},
"timeline": {
"allOf": [
{ "$ref": "#/$defs/researchBase" },
{ "properties": { "component": { "const": "Timeline" }, "events": {} }, "required": ["events"] }
],
"unevaluatedProperties": false
},
"map": {
"allOf": [
{ "$ref": "#/$defs/researchBase" },
{ "properties": { "component": { "const": "Map" }, "locations": {} }, "required": ["locations"] }
],
"unevaluatedProperties": false
},
"mindMap": {
"allOf": [
{ "$ref": "#/$defs/researchBase" },
{ "properties": { "component": { "const": "MindMap" }, "nodes": {} }, "required": ["nodes"] }
],
"unevaluatedProperties": false
},
"form": {
"allOf": [
{ "$ref": "#/$defs/researchBase" },
{ "properties": { "component": { "const": "Form" }, "children": { "$ref": "https://a2ui.org/specification/v1_0/common_types.json#/$defs/ChildList" }, "submitLabel": { "type": "string" }, "action": { "$ref": "https://a2ui.org/specification/v1_0/common_types.json#/$defs/Action" } }, "required": ["children"] }
],
"unevaluatedProperties": false
},
"anyComponent": {
"oneOf": [
{ "$ref": "#/components/Text" }, { "$ref": "#/components/Image" }, { "$ref": "#/components/Icon" },
{ "$ref": "#/components/Video" }, { "$ref": "#/components/AudioPlayer" }, { "$ref": "#/components/Divider" },
{ "$ref": "#/components/Row" }, { "$ref": "#/components/Column" }, { "$ref": "#/components/List" },
{ "$ref": "#/components/Card" }, { "$ref": "#/components/Modal" }, { "$ref": "#/components/Tabs" },
{ "$ref": "#/components/Button" }, { "$ref": "#/components/TextField" }, { "$ref": "#/components/CheckBox" },
{ "$ref": "#/components/Slider" }, { "$ref": "#/components/DateTimeInput" }, { "$ref": "#/components/ChoicePicker" },
{ "$ref": "#/components/Chart" }, { "$ref": "#/components/Table" }, { "$ref": "#/components/Metric" },
{ "$ref": "#/components/Timeline" }, { "$ref": "#/components/Map" }, { "$ref": "#/components/MindMap" },
{ "$ref": "#/components/Form" }
],
"discriminator": { "propertyName": "component" }
},
"anyFunction": {
"oneOf": [
{ "$ref": "#/functions/required" }, { "$ref": "#/functions/regex" }, { "$ref": "#/functions/length" },
{ "$ref": "#/functions/numeric" }, { "$ref": "#/functions/email" }, { "$ref": "#/functions/formatString" },
{ "$ref": "#/functions/formatNumber" }, { "$ref": "#/functions/formatCurrency" }, { "$ref": "#/functions/formatDate" },
{ "$ref": "#/functions/pluralize" }, { "$ref": "#/functions/openUrl" }, { "$ref": "#/functions/and" },
{ "$ref": "#/functions/or" }, { "$ref": "#/functions/not" }
]
}
}
}

View File

@@ -0,0 +1,15 @@
# Dev Brain purpose
Dev Brain compiles durable, source-backed knowledge from the registered projects.
## Priorities
- Architecture, behavior, decisions, invariants, workflows, and relationships.
- Small topic pages that answer recurring development questions.
- Visible uncertainty and exact project provenance.
## Recurring questions
- Where does a behavior live, and what must remain invariant when it changes?
- Which decisions constrain the current implementation?
- Which project sources must be rechecked before relying on this page?

View File

@@ -0,0 +1,67 @@
# Dev Brain schema
The registered projects are authoritative. This vault is a derived, human-readable wiki.
## Managed paths
DS4Server manages `purpose.md`, `schema.md`, `index.md`, `skills.md`, `log.md`, and topic pages below `projects/`, `subsystems/`, `concepts/`, `decisions/`, `invariants/`, `workflows/`, and `skills/`. Topic pages must opt in with `dev_brain: true`. Other notes, hidden files, Obsidian settings, attachments, and trash are never modified.
## Topic frontmatter
```yaml
---
dev_brain: true
type: subsystem # project, subsystem, concept, decision, invariant, workflow, or skill
project: Registered project name
status: verified # verified, stale, or needs-review
verified_at: 2026-07-27T12:00:00Z
sources:
- project: Registered project name
path: src/example.rs
symbol: optional_symbol
revision: latest-commit-containing-this-file-content
# For a clean tracked file, use `git log -1 --format=%H -- src/example.rs`,
# or a unique lowercase hex prefix of at least 7 characters. Use hash when
# this file differs from HEAD, is untracked, or the registered project is not Git.
---
```
`project` names a registered project, not this Dev Brain vault; `dev_brain_info` lists the exact registered names and folders. Each source `path` is relative to that registered project's folder; never resolve it inside the vault. Cite only files that support the page's claims, not every dirty file in the project. Each source has exactly one evidence version: `revision` or a lowercase SHA-256 `hash`. A revision is the newest commit that changed that source file, returned by `git log -1 --format=%H -- path`; validation requires that per-file revision and compares its file content with the current clean file. Every source therefore has its own revision; never stamp all sources with the repository's overall HEAD. Use a hash when that specific file differs from HEAD, is untracked, or its registered project is not Git.
## Skills
Skills are on-demand instructions stored as one Markdown file per skill below `skills/`. They use the same provenance and status fields as topic pages, set `type: skill`, and additionally require `name` and `description` frontmatter. Names contain 164 lowercase letters, digits, or hyphens, with no leading, trailing, or consecutive hyphens. Descriptions contain 11024 characters and explain both what the skill does and when to use it. Only skills whose status and current evidence are both `verified` appear in generated `skills.md` and the session system prompt. The prompt exposes each verified skill's name, description, and Markdown path; read the complete matching file on demand before following it. `index.md` lists all managed topic pages, while `skills.md` lists verified skills only. Do not hand-edit either generated index.
Files with `built_in: true` are app-managed skills restored from the DS4Server bundle. Do not create or edit them by hand.
```yaml
---
dev_brain: true
type: skill
name: review-release
description: Verify a release candidate against the project checklist. Use before publishing a release.
project: Registered project name
status: verified
verified_at: 2026-07-27T12:00:00Z
sources:
- project: Registered project name
path: docs/releasing.md
revision: full-or-unique-short-clean-git-revision
---
```
## Compilation
Read the purpose first. Examine high-signal manifests, documentation, schemas, entry points, public interfaces, and tests. Create the smallest coherent topic set that answers the purpose; do not mirror every source file. Use ordinary Obsidian wikilinks and embeds to connect topics.
## Querying
Start with `index.md` or ranked search, then follow links and backlinks. A `verified` page is authoritative only while all recorded evidence is current. Treat `stale` pages as navigation leads and `needs-review` pages as explicit uncertainty; check project sources before making claims from either.
## Refresh and semantic validation
When validation reports drifted sources, handle each listed file independently: re-read it, compare its current behavior with the page's documented findings, update, split, merge, or retire affected knowledge, then update only that source's revision or hash. For a large revision-backed file, `git diff <recorded-revision> -- path` shows the change from the recorded evidence to the current working-tree version; use it to focus the review, then read enough current context to validate the affected findings. Inspect the commits affecting that path since the recorded revision as well. When code was removed, inspect the full commits and search the current project, callers, and tests before concluding the behavior disappeared: it may have been renamed, replaced, or moved into another file that should replace or join the page's cited sources. Repeat for every listed file and re-run validation. Keep `index.md` exact, and append a material update entry to `log.md` with source revisions. Only mark a page `verified` after its claims have been checked against its current evidence. Unsupported conclusions stay `needs-review`.
## Publication
Call `dev_brain_info`, then maintain these pages with the ordinary file tools. Append material updates to `log.md`; `index.md` and `skills.md` are generated and should not be hand-maintained. Finish with `dev_brain_validate`. Structural or provenance errors must be repaired, while broken links are warnings so valid work remains editable instead of being discarded.

View File

@@ -0,0 +1,58 @@
---
dev_brain: true
type: skill
built_in: true
name: create-dev-brain-skill
description: Create or update a source-backed skill in Dev Brain. Use when adding reusable instructions that should activate on demand in future development tasks.
project: DS4Server built-in
status: verified
verified_at: 2026-07-30T00:00:00Z
sources: []
---
# Create a Dev Brain skill
Create a skill when a repeatable development task needs instructions beyond the model's general knowledge. Keep it concise and source-backed.
## File structure
A Dev Brain skill is one Markdown file, not a directory:
```text
skills/
└── skill-name.md
```
Use this structure:
```yaml
---
dev_brain: true
type: skill
name: skill-name
description: State what the skill does and when to use it.
project: Exact registered project name
status: needs-review # change to verified only after checking every claim
verified_at: 2026-07-30T12:00:00Z
sources:
- project: Exact registered project name
path: path/to/authoritative-source.md
revision: latest-commit-that-changed-this-file
---
# Human-readable skill title
Write direct instructions for completing the task. Link to related Dev Brain pages when useful.
```
Use lowercase letters, digits, and hyphens for `name`. Put all trigger conditions in `description`. Cite at least one registered project source, using the exact per-file revision for a clean Git file or its lowercase SHA-256 hash when required by `schema.md`.
## Create and activate
1. Call `dev_brain_info` and read `purpose.md` and `schema.md`.
2. Read the authoritative project sources for the instructions.
3. Write `skills/<name>.md` with current evidence and `status: verified`; use `needs-review` until every claim is checked.
4. Append the material change to `log.md`.
5. Run `dev_brain_validate`.
Validation is the activation step. It verifies the file, rebuilds generated `skills.md`, and makes the skill's name, description, and path available to the agent. Never edit `skills.md` directly.

View File

@@ -0,0 +1,3 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 20 20" fill="none" stroke="#fff" stroke-width="1.7" stroke-linecap="round" stroke-linejoin="round">
<path d="M11.75 4.5 6.25 10l5.5 5.5"/>
</svg>

After

Width:  |  Height:  |  Size: 201 B

View File

@@ -0,0 +1,3 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 20 20" fill="none" stroke="#fff" stroke-width="1.7" stroke-linecap="round" stroke-linejoin="round">
<path d="m8.25 4.5 5.5 5.5-5.5 5.5"/>
</svg>

After

Width:  |  Height:  |  Size: 200 B

View File

@@ -0,0 +1,4 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 20 20" fill="none" stroke="#fff" stroke-width="1.6" stroke-linecap="round" stroke-linejoin="round">
<path d="M2.75 6V5.75A1.75 1.75 0 0 1 4.5 4h3l1.7 1.75h6.3A1.75 1.75 0 0 1 17.25 7.5v.5"/>
<path d="M3.75 8h13.5a1 1 0 0 1 .95 1.3l-1.65 5.1A2.25 2.25 0 0 1 14.4 16H5.25a2.25 2.25 0 0 1-2.2-1.8L2 9.75A1.43 1.43 0 0 1 3.4 8z"/>
</svg>

After

Width:  |  Height:  |  Size: 391 B

3
assets/icons/robot.svg Normal file
View File

@@ -0,0 +1,3 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 20 20" fill="none" stroke="#fff" stroke-width="1.5" stroke-linecap="round" stroke-linejoin="round">
<path d="M10 2.5v2M7.5 2.5h5M4 7.5A2.5 2.5 0 0 1 6.5 5h7A2.5 2.5 0 0 1 16 7.5v6a2.5 2.5 0 0 1-2.5 2.5h-7A2.5 2.5 0 0 1 4 13.5zM4 10H2.5M17.5 10H16M7.5 9.5h.01M12.5 9.5h.01M7.5 13h5"/>
</svg>

After

Width:  |  Height:  |  Size: 346 B

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,115 @@
{
"format": 1,
"source": {
"repository": "Youssofal/Qwen3.8-Flash-Next-MTPLX-Bare-Speed",
"url": "https://huggingface.co/Youssofal/Qwen3.8-Flash-Next-MTPLX-Bare-Speed/tree/74559cdf34fbfc0b593de72d17e93f37fd4f9ea7",
"revision": "74559cdf34fbfc0b593de72d17e93f37fd4f9ea7",
"base_repository": "Qwen/Qwen3.8-Flash-Next",
"base_url": "https://huggingface.co/Qwen/Qwen3.8-Flash-Next/tree/de4b8e4d43b917e7706784d8bb445c9af86a3540",
"base_revision": "de4b8e4d43b917e7706784d8bb445c9af86a3540",
"license": "Qwen Community License 1.0",
"converter": "qwen38-artifacts-v1-identity"
},
"tensor_inventory": "qwen38-flash-next-bare-speed-tensors.tsv",
"tensor_inventory_sha256": "b5731e6febcf865d276a0e7b144da02375f2d7e3129594f271ee9c6c351f4c8f",
"config": {
"/architectures/0": "Qwen4ExpForConditionalGeneration",
"/model_type": "qwen4_exp",
"/quantization_config/bits": 4,
"/quantization_config/group_size": 64,
"/text_config/full_attention_interval": 4,
"/text_config/hc_count": 4,
"/text_config/hc_lowrank": 320,
"/text_config/hidden_size": 2560,
"/text_config/indexer_budget": 2048,
"/text_config/indexer_compress_ratio": 4,
"/text_config/indexer_head_dim": 128,
"/text_config/indexer_kv_heads": 1,
"/text_config/indexer_n_heads": 4,
"/text_config/linear_conv_kernel_dim": 4,
"/text_config/linear_key_head_dim": 128,
"/text_config/linear_num_key_heads": 16,
"/text_config/linear_num_value_heads": 48,
"/text_config/linear_value_head_dim": 128,
"/text_config/max_position_embeddings": 262144,
"/text_config/moe_intermediate_size": 640,
"/text_config/mtp_num_hidden_layers": 1,
"/text_config/ngram_size": 3,
"/text_config/ngram_vocab_size_base": 20000000,
"/text_config/num_attention_heads": 24,
"/text_config/num_experts": 512,
"/text_config/num_experts_per_tok": 10,
"/text_config/num_hidden_layers": 48,
"/text_config/num_key_value_heads": 2,
"/text_config/output_gate_type": "sigmoid",
"/text_config/partial_rotary_factor": 0.25,
"/text_config/ple_conv_kernel_size": 4,
"/text_config/ple_embed_dim": 2560,
"/text_config/rms_norm_eps": 0.000001,
"/text_config/rope_parameters/rope_theta": 10000000,
"/text_config/shared_expert_intermediate_size": 640,
"/text_config/split_ngram_parts": 128,
"/text_config/vocab_size": 248320
},
"runtime": {
"/arch_id": "qwen4-next",
"/mtp_contract/base_hidden_variant": "post_norm",
"/mtp_contract/concat_order": "embedding_hidden",
"/mtp_contract/hidden_variant": "post_norm",
"/mtp_contract/mtp_position_mode": "cache",
"/mtp_contract/mtp_quant_group_size": 64,
"/mtp_contract/mtp_quant_mode": "affine",
"/mtp_depth_max": 3,
"/mtp_sidecar": "prequantized-mlx-affine-qwen-moe-switch-mlx",
"/mtplx_version": "2.9.3",
"/public_model_id": "mtplx-flash-next-bare-speed",
"/sampler/temperature": 1.0,
"/sampler/top_k": 20,
"/sampler/top_p": 0.95
},
"representative_rows": [
{"class":"expert","file":"model-00001-of-00017.safetensors","tensor":"language_model.model.layers.0.mlp.switch_mlp.gate_proj.weight","row":0,"values":64,"bits":4,"group_size":64,"sha256":"0fe4ae597854a1483e2d63f07cada72d4656c4178e1d696ce4311b838ba05964"},
{"class":"dense","file":"model-00016-of-00017.safetensors","tensor":"language_model.model.layers.0.linear_attn.in_proj_qkv.weight","row":0,"values":64,"bits":4,"group_size":64,"sha256":"59d7d7bf84e88823704d9f7b2c84e41b46f481c658ed8a485f2fd063ecf9ca34"},
{"class":"qsa","file":"model-00016-of-00017.safetensors","tensor":"language_model.model.layers.11.self_attn.indexer.index_qk_proj.weight","row":0,"values":64,"bits":8,"group_size":64,"sha256":"c803ea77148621a5d6dfa4060a52c75f57185234a0bec4771671f69336bb8346"},
{"class":"gdn","file":"model-00016-of-00017.safetensors","tensor":"language_model.model.layers.0.linear_attn.A_log","row":0,"values":48,"sha256":"88c53a2a04bda1d96ee1ade6fa7dfa9c49d3245b1cc1e83673f02b7a36e07a85"},
{"class":"ple","file":"ngram-table.safetensors","tensor":"ngram.weight","row":0,"values":64,"bits":4,"group_size":32,"sha256":"2243dd9766046bb80d98e3baf5e59ba246958cfe4f340ebce8c1fff57a2810d9"},
{"class":"mtp","file":"mtp.safetensors","tensor":"mtp.layers.0.mlp.switch_mlp.gate_proj.weight","row":0,"values":64,"bits":4,"group_size":32,"sha256":"393d82675ae7e25275243fd4907d3894146ac1c5662b601aae998c489d10af5c"}
],
"files": [
{"path":"model-00001-of-00017.safetensors","role":"core","size":4666167150,"sha256":"a27232c9434b9d8961f198cf36f44e228b3f23a5f8d874ba4ca8af32a1b23ffc"},
{"path":"model-00002-of-00017.safetensors","role":"core","size":4299165712,"sha256":"6268ae80950f6b935fca672fae5aa455eb3ea4f139730e54e6e5f8581faa94dc"},
{"path":"model-00003-of-00017.safetensors","role":"core","size":4666167168,"sha256":"71a9abbd420dce567f2c06593bfee2d6cea681989d51eaef1b1c5dffd65a844b"},
{"path":"model-00004-of-00017.safetensors","role":"core","size":4299165702,"sha256":"c8112da735bb1c8dbc4757926ab8ec1eba7a707a8767c7e43c6150586cf63a4c"},
{"path":"model-00005-of-00017.safetensors","role":"core","size":4666167169,"sha256":"48258b931e7596c1b43c8b38c044060787ef6d85cca2c1935852bc962801f792"},
{"path":"model-00006-of-00017.safetensors","role":"core","size":4299165708,"sha256":"f357d2d66f7b768ebb19ba464556bbe46c572cf3cd1f8226faddf9216b150d0c"},
{"path":"model-00007-of-00017.safetensors","role":"core","size":4666167168,"sha256":"d3152e808d013e3afa4a9c391f6d8088e74a10303c89d45fb61dad7e457160f6"},
{"path":"model-00008-of-00017.safetensors","role":"core","size":4299165697,"sha256":"1d93bd4872f338fa19416bac64ebeb05aedf2721a07ac48e37c53921252c95c1"},
{"path":"model-00009-of-00017.safetensors","role":"core","size":4666167162,"sha256":"d22d4be580296e3d4fa5f4fe5902c8fc448d8204b6dcde4641f87b9c675930ca"},
{"path":"model-00010-of-00017.safetensors","role":"core","size":4299165708,"sha256":"111d0166c888cabb7d6fc19c4bb7eb42e2b51946f71b4983dd750b031755637e"},
{"path":"model-00011-of-00017.safetensors","role":"core","size":4666167159,"sha256":"73dad882cf94a5486cfe0dcda876b0f740556bc469608d2ced7a07540e699803"},
{"path":"model-00012-of-00017.safetensors","role":"core","size":4299165706,"sha256":"a6057d5a06276461760a4f4e88991550a43abdc5f8afacef06e860f5e7843393"},
{"path":"model-00013-of-00017.safetensors","role":"core","size":4666167174,"sha256":"ff25dc6e1c250c9b9ece83d59111a4c4055dad1b1f160546c614dcc17395b152"},
{"path":"model-00014-of-00017.safetensors","role":"core","size":4299165704,"sha256":"d8c9de3ae28abccc8e3da09dc92fbbed4039d6d65dabb2657edebc11b2b09652"},
{"path":"model-00015-of-00017.safetensors","role":"core","size":4666167146,"sha256":"a5438d468074c709febdeadb9f6a2db54fa9fb30a37c33de78584f28d6fa1191"},
{"path":"model-00016-of-00017.safetensors","role":"core","size":4300831084,"sha256":"4a07409bd568cd1435641495d90c805dcb0d0e5d8dc64409273d604a7d06b342"},
{"path":"model-00017-of-00017.safetensors","role":"core","size":18354282,"sha256":"9c064e0b8dc0054e2b453fe2b8450c287e1ede6e74bb74ecf5aaed1ef9271913"},
{"path":"ngram-table.safetensors","role":"ple","size":32000154008,"sha256":"af7f0880c1798e60353b07cc6b0faa0d41f78734818d858ff013ab7e3bce9de5"},
{"path":"mtp.safetensors","role":"mtp","size":1672575532,"sha256":"1bd4b2fc20a93b53f22b4cab0b14d68bc0fe666bb64e4e2f898872c30d7d26ae"},
{"path":"model.safetensors.index.json","role":"core-index","size":284868,"sha256":"81e77f5b25bb29e25e239ae2f9f324c9cada379ee600441d6c229f1d95a13443"},
{"path":"config.json","role":"config","size":227092,"sha256":"826302af9ff3293cfc9975c9a5a96546fa0aa9f0808505adfc30a62fef806965"},
{"path":"generation_config.json","role":"config","size":202,"sha256":"e70c136c1b78ddc1fb0905bac8e733a4dc448d4f852a5dd75143fffc70be550e"},
{"path":"mtplx_runtime.json","role":"config","size":2676,"sha256":"381c7666872f878542494a48a874bb0784c72a955d18cc077324d3c9e74b8cdb"},
{"path":"chat_template.jinja","role":"tokenizer","size":8952,"sha256":"c3cf9e34abf4f9e36c2d72165aa9c132d3e2a725b6c2586aaa3a8af9d7a81041"},
{"path":"tokenizer.json","role":"tokenizer","size":12809320,"sha256":"0997f410c57a1f4e53b09e4be8f4a172d90edd9564368fb0847030937229b9f3"},
{"path":"tokenizer_config.json","role":"tokenizer","size":17928,"sha256":"b11349aafa7cdc6a320767cf7ceb29ed82f7eda5d65e8e0819e76f0ce947bf27"},
{"path":"vocab.json","role":"tokenizer","size":6722759,"sha256":"ce99b4cb2983d118806ce0a8b777a35b093e2000a503ebde25853284c9dfa003"},
{"path":"merges.txt","role":"tokenizer","size":3353259,"sha256":"a9d356d7bdf1ef4949e3e748e95b8e10ad9d4e2e838eddc38a0a7b6b94d1db8d"},
{"path":"LICENSE","role":"license","size":3235,"sha256":"a0dc422560841fd68e06d974907f8b4c709bca44a67daad2b528437bdf676c08"}
],
"excluded": [
{"path":"model-vision.safetensors","reason":"Vision tensors are outside the DS4Server text artifact.","size":897899145,"sha256":"5aff6983f36f52acbecc1f61f703ae7b927e74525bf3b1133ba276bb7a5b52b1"},
{"path":"preprocessor_config.json","reason":"Vision preprocessing is outside the DS4Server text artifact.","size":390,"sha256":"27225450ac9c6529872ee1924fcb0962ff5634834f817040f444118116f4e516"},
{"path":"processor_config.json","reason":"Vision preprocessing is outside the DS4Server text artifact.","size":367,"sha256":"dcc7ce0121d6d1c514804e68e5673e04c239ad2547339ea6330289d272a748f8"},
{"path":"video_preprocessor_config.json","reason":"Video preprocessing is outside the DS4Server text artifact.","size":385,"sha256":"7768af27c1fafa9cc9011c1dc20067e03f8915e03b63504550e11d5066986d13"}
]
}

View File

@@ -11,14 +11,18 @@ fn main() {
.include("native/metal")
.file(metal)
.flag("-fobjc-arc")
.flag("-ffast-math")
.flag("-mcpu=native")
.opt_level(3)
.compile("ds4_metal");
println!("cargo:rerun-if-changed=native/web");
println!("cargo:rerun-if-changed=native/media");
cc::Build::new()
.include("native/web")
.file("native/web/ds4_web.c")
.opt_level(2)
.compile("ds4_web");
.file("native/media/ds4_media.m")
.flag("-fobjc-arc")
.compile("ds4_media");
println!("cargo:rustc-link-lib=framework=Foundation");
println!("cargo:rustc-link-lib=framework=AppKit");
println!("cargo:rustc-link-lib=framework=AVFoundation");
println!("cargo:rustc-link-lib=framework=AVKit");
println!("cargo:rustc-link-lib=framework=Metal");
}

184
docs/METAL_PARITY.md Normal file
View File

@@ -0,0 +1,184 @@
# Apple Metal parity for issue #78
## GLM 5.3 Flash parity target
- Previous DS4 oracle: `8db89fe083ae4d17c9a2428ccd29803d3ae8f577`
- Current DS4 oracle: `ec7642cdd9ec81d01ad4b1fd8f8a3d1511533748`
- Main artifact: `GLM-5.3-Flash-Q2.gguf`, SHA-256
`e81fd6241c6e55a64e1e14e47a3eab61a173fa8d7e4b5c1d1848827119705b32`
- Vision artifact: `GLM-5.3-Flash-Vision-Encoder.gguf`, SHA-256
`ae23e14c6979e889051b2e4a39351abcdafb161e18e606fae4d8c40095a4bf3a`
- Machine: MacBook Pro `Mac17,7`, Apple M5 Max, 128 GiB unified memory,
resident Q2 weights
The applicable commits after the previous oracle are:
| DS4 commit(s) | Scope carried into DS4Server |
| --- | --- |
| `147109a` | GLM 5.3 shape, GGUF binding, KDA/DSA attention, hyperconnections, routed MoE, BF16 tensors, output head, integrated MTP, and Metal kernels |
| `4771329`, `9f95d9f` | Vision encoder, image preprocessing, prompt overlays, and compact-prefill correctness |
| `b0c31af` | Pooled indexer cache, bounded attention memory, and batched prefill |
| `747eafc`, `1df8ff0` | Fixed-image encoder oracle, replay/logit checks, and multimodal quality gates |
| `8adb221` | Directional attention and FFN steering at the DS4 residual points |
| `639d4eb` | Continued-prefill cache and position behavior |
| `6cf658a` | BF16 decode projection schedule and Metal kernel specialization |
| `01f3316`, `81c8c39`, `5da3206` | Current Metal compact-cache ABI, Q2/full-checkpoint binding, and GLM streaming compatibility |
CUDA- and ROCm-only commits in the range are excluded. The Rust release gates
compare three-run medians against this exact DS4 build, require identical
greedy token hashes, and fail only when Rust is more than 2% slower than DS4;
there is no upper bound when Rust is faster. Text measures an
8192-token frontier plus steady decode, both with and without integrated MTP.
Vision compares the fixed encoder embedding (cosine at least 0.995, mean
absolute error at most 0.001, maximum error at most 0.06), encoder latency,
multimodal prefill, steady decode, image-conditioned logits, and replay.
### GLM 5.3 Flash acceptance results
The text fixture is `ds4.c`, rendered as one no-system chat turn (`-sys ''` in
DS4), truncated to exactly 8192 tokens, with a 32768-token allocation and 64
greedy output tokens. Same-process three-run medians are used because both DS4
and DS4Server exhibit a load-order/thermal artifact with this 96 GB model.
| Capability | Metric | DS4 | DS4Server | Result |
| --- | --- | ---: | ---: | --- |
| Text | Prefill tok/s | 288.91 | 295.21 | Rust 2.18% faster |
| Text | Decode tok/s | 22.58 | 22.78 | Rust 0.89% faster |
| Integrated MTP | Decode tok/s | 26.55 | 26.80 | Rust 0.93% faster |
| Vision encoder | Fixed-image seconds | 2.6255 | 2.4354 | Rust 7.24% faster |
| Vision prompt | Prefill tok/s | 185.81 | 201.51 | Rust 8.45% faster |
| Vision prompt | Decode tok/s | 26.40 | 27.21 | Rust 3.09% faster |
Text and MTP produced the same token SHA-256,
`0b61e613708d8367fc87f5a32fbfd4ba098b40ee920645b1e38381e867131108`.
The 1094-token fixed-image prompt produced the same vision token SHA-256,
`424fdadc10c0a2bd95b405a3985151e84aa844a209d2a2ca72cd01c39dfa53db`.
The Rust vision embedding compared with the DS4 oracle at cosine
`0.999958726`, mean absolute error `0.000036702`, and maximum error
`0.006153620`.
The checkpoint advertises a 1048576-token maximum. On this 128 GiB machine,
32768 is the default manageable context and 50000 is the extended option: the
50000-token run used 3.51 GiB of live runtime tensors in addition to 90.61 GiB
of mapped model wrappers and retained the exact output hash. Larger allocations
leave too little operating-system and application headroom for a reliable
default.
## Frozen comparison
- DS4 oracle: `8db89fe083ae4d17c9a2428ccd29803d3ae8f577`
- DS4Server pre-issue baseline: `79468c65b68c28e1539acf0a2a14bd42e000e365`
- DeepSeek V4 Flash 0731 IQ2XXS/w2Q2K model SHA-256:
`ca22ae2f838e14077c22bc1c1417b71b45b5e5a3687bd96c2ac6e17fdb6261c0`
- DSpark support-model SHA-256:
`7e319924541db3f7a163ed7e11d7532a70d48228ab59d36cb81e1d4511885360`
- Machine: MacBook Pro `Mac17,7`, Apple M5 Max, 128 GiB unified memory,
macOS 26.6.2, power 100%
- Context and prefill allocation: 4096 tokens; quality, steering, speculative
decoding, and DSpark disabled for the throughput gate
- SSD mode: 4096 cached experts, 4096 popularity-preloaded experts, ordinary
(not cold) streaming
- Prompt: `Count from one to two hundred, spelling out every number.` rendered
to 16 tokens; 128 greedy non-EOS tokens are measured
The longer prompt is deliberate: all 128 generated tokens are reachable before
EOS and are identical between DS4 and DS4Server in each mode. A short-response
fixture reached EOS early and made a forced post-EOS speed run compare
unreachable, mode-dependent states.
## DS4 performance-change matrix
The audited range is
`efdadd41e20134af4f3381e1ed90e96fe4faef6f..8db89fe083ae4d17c9a2428ccd29803d3ae8f577`.
Documentation-only, parser/server-only, CUDA, ROCm, tensor-parallel, model
download, and merge commits do not affect the Rust-owned single-session Metal
graph. The applicable performance/runtime commits are accounted for below.
| DS4 commit(s) | DS4 call site or behavior | DS4Server result |
| --- | --- | --- |
| `3ce6777`, `8e0f997` | Native streamed-expert cache pinning and safe fallback | Active in the byte-matched native Metal cache; Rust preserves the cache budget and fallback. |
| `df04887` | `ds4_gpu_attention_indexed_mixed_batch_heads_tensor` | Active through `encode_layer_with_cache_rows`; the long-context fixture crosses the indexed boundary. |
| `74e8f11`, `4591cb1`, `a51e6ec`, `f284165`, `af80694`, `7fb2830`, `0e89a0e`, `769a8ba`, `8d683d6` | DSpark block commit, deterministic scheduling, acceptance replay/direct commit, greedy identity, and exact/opportunistic sampling | Active in `Dspark::propose` and the verifier commit paths; greedy and exact stochastic fixtures cover the frozen support model. |
| `96c3ba4`, `d69a017`, `81f8534`, `ac5cb2c`, `48392e3`, `222b2cb`, `532ec8b` | Metal 4 Q4 projection, compact routed-MoE, aligned Q8, indexed-attention, and routed prefill kernels | Active through the existing `encode_batch_layer` Q4/indexed/routed entry points; no second prefill graph was added. |
| `24903f3` | `ds4_gpu_head_rms_norm_rope_tail_tensor` | Existing Rust Q-head norm/RoPE fusion is retained and used. |
| `4893e0c`, `d14ce35` | Mapped-prefill expert seeding and consistent expert/non-expert maps | Rust seeds with `ds4_gpu_stream_expert_cache_seed_experts_gpu_copy`, falls back to CPU/pread, and restores one cached static decode map with the same non-routed spans. |
| `0dd0d36` | Release speed regression gate | Rust has an ignored release gate with retained raw CSV, DS4 reference threshold, baseline threshold, p50/p95, cache metrics, and token hashes. |
| `4e401a1`, `0e125e6` | Automatic pre-M5/M5 admission for exact decode fusions | Rust binds both device predicates and mirrors per-generation rollback variables without a UI preference. |
| `71d8c2a` | Q/KV pair plus ratio-4/128 compressor projection/store | `try_qkv_pair_compressor_fusion` calls the native combined entry point only for DS4's resident full-decode shapes. |
| `b64d3e0` | Emit-path compressor finalizer | `update_compression_pair` defers both updates only when the resident finalizer is eligible, then requires exact success. |
| `3826f52`, `ad4d05c` | Exact compressor softmax/product and ratio-4 pool | Native kernels are active; Rust's compressor ABI now passes the upstream `decode_one_token` and `defer_finalize` flags exactly. |
| `92d83e6` | HC norm/mix producer clustering | `decode_hc_mix` uses split-norm producer fusion, with the established norm/mix/split sequence as fallback. |
| `1be41c8` | Packed exact split-K attention reduction | Active in the vendored native attention kernel reached by Rust's decode attention calls. |
| `0e8fb54`, `e6a6ab5` | Router/select, shared gate/up, QKV norm/KV RoPE/store, parallel full FFN, inverse-attention RoPE, and pre-M5 ports | All applicable entry points are called with DS4 shape, tensor, quality, steering, profile, SSD, and device guards. Parallel FFN aborts on partial failure; inverse RoPE falls back when the attention kernel does not consume it. |
| `0658dac` | Long-context prefill/decode correctness | The 4096-crossing ignored fixture covers compressed/indexed state and the next decode token. |
| `023614e`, `0ad494e`, `42033ee` | Remove stream-512 experiment; make exact paths automatic; verify DFlash pipeline | No stream-512 or opt-in experiment remains. Feature rollbacks retain identical token hashes. |
| `24fa85e` | Keep successfully loaded disk KV checkpoints | Existing Rust checkpoint load/round-trip behavior is verified; only invalid checkpoints follow deletion policy. |
| `7694112` | Reuse a shorter exact GLM live prefix | `live_prefix_rewind_target` rewinds to `prompt_len - 1`, clears pending speculation, and reevaluates the last prompt token. |
| `5186e2c` | Keep Metal decode changes portable | Device checks and unavailable returns preserve non-M5 and unfused execution. |
The Metal MXFP4 commits (`725b084`, `1f862bd`, `7bec128`, `bcd1e57`,
`f75a065`, and their merge/format work) are intentionally excluded: MXFP4 is
issue #79. `ds4_gpu_set_decode_pipeline_fast_lookup` is reset to disabled,
matching DS4 because that cache is enabled only for the pre-M5 MXFP4 path.
## Results
Three implementations were run sequentially, three times per mode, with a
five-second cooldown. Medians are from
[`metal-parity-issue-78.csv`](metal-parity-issue-78.csv).
| Implementation | Mode | Prefill tok/s | TTFT ms | First decode ms | Stable tok/s | p50 ms | p95 ms |
| --- | --- | ---: | ---: | ---: | ---: | ---: | ---: |
| DS4Server baseline | resident | 69.63 | 256.07 | 26.37 | 39.53 | 24.88 | 26.61 |
| DS4 | resident | 56.49 | 305.94 | 22.71 | 45.60 | 21.91 | 22.32 |
| DS4Server issue #78 | resident | 69.76 | 252.01 | 22.66 | 45.60 | 21.94 | 22.33 |
| DS4Server baseline | SSD | 6.81 | 2476.94 | 125.30 | 8.63 | 114.93 | 124.28 |
| DS4 | SSD | 11.88 | 1393.54 | 46.75 | 26.11 | 35.62 | 46.86 |
| DS4Server issue #78 | SSD | 11.49 | 1440.79 | 48.84 | 25.07 | 37.70 | 48.85 |
TTFT is the per-run prompt prefill duration plus first decode latency; its
median is reported separately from the first decode kernel latency.
Resident stable decode improves 15.4% over the baseline and matches DS4.
SSD stable decode improves 190.3% over the baseline and is 4.0% below
DS4, inside the required 5% boundary. SSD cache accounting is identical to
DS4 for the measured token sequence: 4096 entries, 32,958 hits, 4,194 misses,
and 29,684,662,272 pread bytes. The resident and SSD token SHA-256 values are,
respectively, `ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d`
and `2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052`;
DS4 emitted the same 128-token sequences in the corresponding modes.
Representative `/usr/bin/time -l` measurements report:
| Implementation | Mode | Maximum RSS bytes | Peak footprint bytes | Process swaps |
| --- | --- | ---: | ---: | ---: |
| DS4 | resident | 78,790,656 | 4,746,511,440 | 0 |
| DS4Server issue #78 | resident | 129,024,000 | 104,071,816 | 0 |
| DS4 | SSD | 29,066,280,960 | 29,595,620,664 | 0 |
| DS4Server issue #78 | SSD | 29,098,393,600 | 104,350,344 | 0 |
The post-run system memory-pressure reading was 92% free. Every measured
process reported zero swaps; the post-run host counters were 0 swap-ins and 20
historical swap-outs.
## Reproduction
DS4 uses the exact 16 Rust-rendered prompt token IDs with its speed-benchmark
format. DS4Server uses the installed-model fixture:
```sh
DS4_REFERENCE_TPS=45.60 \
DS4SERVER_BASELINE_TPS=39.526849 \
cargo test --release --all-features flash_0731_m5_decode_performance_gate \
-- --ignored --nocapture --test-threads=1
DS4SERVER_BENCH_SSD=1 \
DS4_REFERENCE_TPS=26.11 \
DS4SERVER_BASELINE_TPS=8.634659 \
cargo test --release --all-features flash_0731_m5_decode_performance_gate \
-- --ignored --nocapture --test-threads=1
```
The gate fails unless throughput improves over the frozen DS4Server baseline
and remains at least 95% of DS4. `DS4SERVER_BENCH_TOKENS=1` prints the prompt
and generated sequences for direct DS4 comparison.

104
docs/QWEN38_ARTIFACTS.md Normal file
View File

@@ -0,0 +1,104 @@
# Qwen3.8 Flash Next Bare Speed artifacts
DS4Server uses the text artifacts published at
`Youssofal/Qwen3.8-Flash-Next-MTPLX-Bare-Speed` revision
`74559cdf34fbfc0b593de72d17e93f37fd4f9ea7`. The upstream base is
`Qwen/Qwen3.8-Flash-Next` revision
`de4b8e4d43b917e7706784d8bb445c9af86a3540`, under the Qwen Community
License 1.0 included in the artifact set.
The source already separates the 71,742,682,599-byte resident text core,
32,000,154,008-byte PLE table, and 1,672,575,532-byte MTP head. DS4Server's
artifact conversion is therefore deliberately byte-preserving: the Rust tool
validates and materializes the selected files without changing their tensor
containers or numerical interpretation. The complete selected download is
105,438,842,430 bytes.
The manifest at `assets/models/qwen38-flash-next-bare-speed.json` pins every
selected file by byte size and SHA-256. Its companion TSV records all 2,527
text, PLE, and MTP tensors with their source file, name, dtype, shape,
quantization parameters, and byte range. Validation also compares the 2,466
core tensors with the source index, rejects duplicate names, invalid byte
ranges, and precision-map drift, and checks the fixed `qwen4_exp` architecture
and MTPLX serving contract before any payload is accepted.
The precision contract follows the pinned files and their per-tensor
`quantization_config`: affine 4-bit matrices use 64-weight groups; router
matrices and QSA `index_qk_proj` use the source's declared 8-bit groups; GDN
convolution and recurrent parameters, norms, hyper-connection parameters, and
the MTP mixing/norm parameters remain BF16. The MTP layer's attention and MoE
matrices remain in their published 2-, 4-, or 8-bit affine representation;
they are not up-converted to satisfy a broader prose description. The PLE
sidecar's own header pins its 4-bit rows to 32-weight groups. No Python, MLX,
llama.cpp runtime, or foreign converter is part of the tool or product.
`model-vision.safetensors` and its image/video processor configs are explicitly
excluded. The full config is retained because its `text_config` is the loader
authority. llama.cpp commit `b356fa2624643b6d5753162ae43efff8cdd4d8cb`
is not used as a converter because its reviewed path did not preserve this
separate MTP artifact.
Download and authenticate every selected file without loading tensors:
```sh
cargo run --release --bin qwen38-artifacts -- \
fetch assets/models/qwen38-flash-next-bare-speed.json OUTPUT_DIRECTORY
```
Interrupted downloads remain as resumable `.part` files. The fetcher uses the
manifest's immutable revision URLs, validates resumed ranges, authenticates
each completed file, and then runs the complete inventory and precision
verifier.
Verify a pinned checkout:
```sh
cargo run --release --bin qwen38-artifacts -- \
verify assets/models/qwen38-flash-next-bare-speed.json SOURCE_DIRECTORY
```
Materialize a clean, authenticated artifact directory:
```sh
cargo run --release --bin qwen38-artifacts -- \
materialize assets/models/qwen38-flash-next-bare-speed.json \
SOURCE_DIRECTORY OUTPUT_DIRECTORY
```
The output directory must not already exist. Materialization uses hard links
when possible and otherwise copies bytes, then re-runs the same full verifier
against the output.
## Application intake
The native Model Manager exposes the complete 29-file set as one resumable
artifact. It downloads from the immutable source revision above, verifies each
file before promotion, retains interrupted `.part` files, and removes the core,
PLE, MTP, tokenizer, metadata, checksum markers, and partials together.
Model selection is persisted as `qwen3.8-flash-next`; intake recognizes the
artifact from the pinned `qwen4_exp` metadata rather than its directory name.
Before Metal allocation, the Rust loader checks every configuration and runtime
contract in the manifest and all 2,527 safetensors records against the frozen
inventory. The tokenizer uses the pinned byte-level BPE, ChatML special tokens,
official low/medium/xhigh/direct formatting, grouped tool responses, and the
official temperature 1.0, top-p 0.95, and top-k 20 defaults.
At the native 262,144-token ceiling, admission reports the 71,742,682,599-byte
resident core, the separately mapped 32,000,154,008-byte PLE table, the
1,672,575,532-byte optional MTP artifact, 6,557,908,992 bytes of KV/recurrent
state, and 27,262,976 bytes of prefill transient storage. PLE mapping is not
counted as resident allocation, and MTP is counted only when enabled. Execution
still fails with an explicit backend-unavailable error until the Rust Metal
backend is delivered by issue #95; it never falls through to a GGUF executor.
## Verification receipt
On 2026-09-03 the pinned source passed full authentication and inventory,
precision, and representative-row verification. Two separately empty output
directories were then materialized and independently verified. Each contained
the same 29 selected files and 105,438,842,430 artifact bytes; normalized
SHA-256 listings for all 31 output files (artifacts, manifest, and inventory)
were byte-identical and had list digest
`3e0233c2356b26859de43b25b2a3df08169dd8263071432bdf2fc497da74007c`.
The six dequantized-row receipts are pinned directly in the manifest.

146
docs/QWEN38_PERFORMANCE.md Normal file
View File

@@ -0,0 +1,146 @@
# Qwen3.8 Flash Next performance and parity receipt
Issue #99 optimizes the pinned `Youssofal/Qwen3.8-Flash-Next-MTPLX-Bare-Speed`
text model without changing its artifact precision, sampler, or token behavior.
All application and inference work remains Rust; native additions are confined to
the existing Objective-C Metal bridge and Qwen `.metal` kernels. Integrated MTP
is the Qwen production default; an explicitly stored user opt-out remains
authoritative.
## Result
The acceptance B1 median is **60.825779 tok/s MTP**. A later B2 confirmation,
after adding complete repair telemetry and under the then-current machine load,
measured **58.019436 tok/s**. Both sequences produced the same 64-token SHA-256
`46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72`, with 44
accepted drafts from 60 proposals in 20 verifier passes. The fresh matched MTPLX
median was 59.751170 tok/s, putting B1 1.8% above and B2 2.9% below the oracle;
both are inside the issue's 5% parity band. Raw per-run data is in
`docs/qwen38-issue-99.csv`.
Plain AR reached a 50.671077 tok/s median with stable token hash
`15373689bcb1c667123accfcfbe49736466f28ce7beb89d2cbe5e4120416056d`.
The published 47 tok/s MTPLX figure is contextual rather than a fresh matched
AR baseline: after the valid MTP oracle run, the installed MTPLX package and
checked-out source drifted apart. The current CLI rejects this Flash Next pack
as unsupported and the source path fails at the missing
`DecoderLayer.input_layernorm`, so no invalid number is presented as a fresh AR
comparison.
The published 75.9 tok/s MTP result is from a different MTPLX runtime/workload.
With identical pinned artifacts, 4,096 context, prompt, sampling, seed, and
64-token measurement window, the current usable oracle result is 59.751170
tok/s. The remaining gap is therefore not explained by quantization drift or a
different output distribution.
Maximum fan control could not be established because neither ThermalForge nor
TG Pro is installed; MTPLX reports its fan boost as unavailable. Both compared
engines ran under the same no-boost constraint. No system utility was installed
or substituted for the missing control.
## Fixed workload
- Hardware: MacBook Pro `Mac17,7`, Apple M5 Max, 18 CPU cores, 128 GiB unified
memory.
- OS: macOS 26.6.2, build 25G83, arm64.
- Model revision: `74559cdf34fbfc0b593de72d17e93f37fd4f9ea7`.
- Model manifest contract SHA-256:
`6f1172de47fa30b9602e13fc7ad14e578a813bb5b1fab8f0ef320041aac6c19a`.
- Tensor inventory SHA-256:
`b5731e6febcf865d276a0e7b144da02375f2d7e3129594f271ee9c6c351f4c8f`.
- `config.json` SHA-256:
`826302af9ff3293cfc9975c9a5a96546fa0aa9f0808505adfc30a62fef806965`.
- `mtplx_runtime.json` SHA-256:
`381c7666872f878542494a48a874bb0784c72a955d18cc077324d3c9e74b8cdb`.
- Prompt: `Write a concise Rust implementation of binary search and explain its invariants.`
- Prompt tokens: 26; context: 4,096; generated tokens: 64.
- Sampling: temperature 1.0, top-p 0.95, top-k 20, min-p 0, seed 12,345.
- One process, three runs: first run cold, following runs warm.
The exact DS4Server commands were:
```sh
DS4SERVER_QWEN38_SOURCE=/private/tmp/qwen38-bare-source \
DS4SERVER_QWEN38_MEASURED=64 DS4SERVER_QWEN38_RUNS=3 \
cargo test --release qwen_m5_performance_gate -- --ignored --nocapture
DS4SERVER_QWEN38_SOURCE=/private/tmp/qwen38-bare-source \
DS4SERVER_QWEN38_MTP=1 DS4SERVER_QWEN38_MEASURED=64 \
DS4SERVER_QWEN38_RUNS=3 \
cargo test --release qwen_m5_performance_gate -- --ignored --nocapture
DS4SERVER_QWEN38_SOURCE=/private/tmp/qwen38-bare-source \
cargo test --release qwen_mtp_matches_target_and_restores_its_cache \
-- --ignored --nocapture
```
The MTPLX comparison used its normal `serve` path with the same local Bare Speed
directory, generation mode, prompt, context, sampler, seed, and single request
stream. A later retry is not a valid benchmark because the local MTPLX
source/package compatibility gate fails as described above.
## MTPLX logic review
The production path now follows the MTPLX/MLX graph rather than retaining the
temporary Qwen experiments used while earlier issues established correctness:
- PLE uses the artifact's 4-bit, group-32 rows, streams them with `pread`, keeps
a bounded hot cache, and publishes hit, miss, eviction, byte, and read-time
counters. The warm MTP runs had 2,080 hits, no misses, and no reads.
- GDN uses fused input projection, MTPLX-equivalent conv+SiLU+L2 normalization,
BF16 rounding points, decay/delta recurrence, gated norm, and direct verifier
capture buffers. Rejected rows restore or advance authoritative recurrence
rather than inventing state.
- QSA uses the MTPLX `qsa_flash_skip` row ownership, online-softmax merge, and
native 256-thread/eight-SIMD layout. Long-context state remains linear and
admitted through 262,144 tokens without a gather-copy fallback.
- MoE uses the exact router ownership and merged gate/up projection. The custom
verifier MoE experiment and NAX path were removed because they are not the
oracle's promoted path.
- MTP uses the same checkpoint, history, proposal, verifier, rejection-repair,
full-accept bonus, and NumPy-compatible PCG64 sampling rules. The exact cache
restoration test passes after the final telemetry change.
- Sparse top-k sampling stays on device for the supported official sampler and
uses the existing exact CPU path only when the requested sampler is outside
that kernel's contract. Full production logits are not copied to the CPU.
- Capability fallbacks remain exact implementations, not alternate Qwen
behavior. Experiments that changed tokens (normal Metal math, wider Tensor
QMM, and recurrence-only repair) were removed rather than left dormant.
The 4-bit/group-64 core, declared 8-bit router/QSA matrices, BF16 recurrent and
normalization tensors, 4-bit/group-32 PLE, and published 2/4/8-bit MTP tensors
are consumed as stored. No blanket up-conversion or replacement quantization is
used. Peak RSS in the final MTP sequence was 52.61 GB, down from the original
102.65 GB path; the 105.44 GB artifact size includes a 32.00 GB file-backed PLE
table that is not resident wholesale.
## A/B evidence
The optimization work used exact hashes or tensor comparisons before promotion:
| Area | Before | After | Decision |
| --- | ---: | ---: | --- |
| Resident memory | 102.65 GB RSS | 52.61 GB RSS | Reused staging and streamed PLE promoted |
| Router | 49.7 ms GPU | 18.7 ms GPU, 23.4 ms wall | Row-owned reduced router promoted |
| 16K QSA AR | 2.6 tok/s | 45.7 tok/s | MTPLX flash-skip rewrite promoted |
| Shared verifier gate/up | 54.7 tok/s | 58.3 tok/s | Exact fused projection promoted |
| Direct GDN verifier capture | 58.5/58.7 tok/s | 60.8 tok/s | Direct capture buffers promoted |
The direct-capture B1 acceptance was followed by the B2 confirmation in the CSV;
both retained the same token hash. Normal Metal math, Tensor QMM at verifier
width, and shortened recurrence repair failed exact-oracle checks and were
removed. Custom MoE verify, fused GDN output projection, and async AR were not
promoted because the MTPLX evidence did not show a win and DS4Server profiles did
not justify overriding that result.
## Context and verification
The long-context gate exercised 2,048, 16,384, 65,536, 131,072, and the native
262,144-token maximum. AR remained 43.488 tok/s at the maximum, and live state
grew from 58,195,968 bytes at 2K to 7,449,083,904 bytes at 262K without collapse
or over-admission. The CSV records every context result and state size.
Final issue verification comprises the exact MTP cache/token oracle, affine
QMV/QMM MLX reference comparisons, focused Qwen tests, the long-context state
gate, the release performance gate, formatting, Clippy with warnings denied, the
macOS application bundle, and the full all-features Rust test suite.

295
docs/USER_GUIDE.md Normal file
View File

@@ -0,0 +1,295 @@
# DS4Server User Guide
DS4Server is a native macOS application for running supported DwarfStar models
locally. Model files, chats, project references, settings, checkpoints, and
agent-extension data stay on this Mac.
## Getting started
1. Open **View > Model Manager** (`⇧⌘M`). Download a supported model and any
acceleration artifact you intend to use, or verify an artifact already on
disk. A model cannot run until its complete file passes verification.
2. Choose **Add project** and select the folder the agent should work in. The
project remains in its original location; DS4Server stores only its reference
and chat data.
3. Choose **File > New Chat** (`⌘N`). Select a thinking profile and a shell
permission mode in the composer.
4. Enter a request and send it. The choices match the selected model: DeepSeek
offers **Direct**, **Think Low**, **Think High**, and **Think Max**; GLM 5.2
offers **Direct**, **Think High**, and **Think Max**; GLM 5.3 Flash always
thinks and offers **Think Low**, **Think High**, and **Think Max**.
The selected model is shared by the app and the optional local HTTP endpoint.
Preferences are saved locally and generation settings are scoped by model and
thinking profile where the controls indicate that.
## Projects, chats, and context
The sidebar switches projects and chats. Project controls add, rename, or delete
a project reference. Deleting one removes its DS4Server sessions and
checkpoints, but never deletes the referenced folder. An unsaved draft can move
to another project; a saved chat remains attached to its original project.
Chat controls can rename, pin, archive, restore, compact, rebuild, export, or
delete a conversation. Archived chats remain under the project's archive until
restored or deleted. Deleting a chat permanently removes its transcript and KV
checkpoint. **File > Export Chat as Markdown** (`⇧⌘S`) or **Export Chat as
HTML** saves the visible conversation, including clearly labeled user,
assistant, thinking, embedded content, and tool sections while excluding
internal system messages. Markdown keeps embedded content as source code;
HTML renders validated SVG figures and A2UI surfaces inline in a self-contained
dark document. Tool-call JSON is collapsed by default and can be expanded from
its disclosure row.
Send another prompt while a response is active to queue it for that chat.
**Stop** cancels current model and agent work. You can work in several chats at
once: each keeps its own transcript, prompt queue, tools, approvals, context,
and interactive surfaces. Independent tool work may overlap, but model
inference is serialized through one shared runtime so weights are not duplicated.
Quitting with active chats asks whether to stop them first.
The context indicator shows used and available tokens. Near the limit,
DS4Server compacts older model context while retaining the full visible
transcript. Saved chats also keep local KV checkpoints for fast continuation.
A checkpoint belongs to the exact model, quantization, context size, and runtime
configuration that created it. After one of those changes, DS4Server reports a
context rebuild, prefills the saved transcript, and writes a compatible
replacement without losing messages. **Rebuild context on next use** requests
the same safe rebuild manually.
## Agent capabilities and safety
The coding agent can use native tools to list a directory, read files, search
file contents, write or edit files, inspect Git state, and maintain Dev Brain.
It can also run shell commands and use web search or visit pages when the task
needs them. The native glob tool finds files recursively by basename or relative
path pattern without a shell; content matching remains the search tool's job,
and a directory listing remains a non-recursive view of one folder.
Tool calls are checked against typed schemas before execution. Unknown tools,
missing or extra fields, malformed calls, and invalid values return structured
errors instead of running. Results appear in the transcript. Copy actions expose
their complete text, and large output may also be saved to a local file rather
than filling the transcript. **Stop** cancels active generation and cancellable
tool work.
The shell permission mode is stored with each chat:
- **Heuristic** applies DS4Server's built-in command-risk checks.
- **AI based** asks the local model to classify a shell command in a separate
one-shot request. A risky result shows its reason. If classification fails or
is invalid, approval is required.
An approval dialog shows the command and working directory. **Allow once** runs
that operation; **Deny** returns the refusal to the agent. Destructive, network,
privileged, application-control, credential, and outside-trusted-path commands
require approval as applicable. Dev Brain's configured vault is trusted like
the project for ordinary path checks, but does not bypass those risk classes.
### Ralph loops
Ask the agent to use a Ralph loop when a bounded task benefits from repeated
autonomous passes, for example, “Use a Ralph loop with at most 12 rounds to fix
and verify this parser.” Ralph is an agent tool, not a view or slash command.
Its default is 8 rounds and its allowed range is 164.
Each round receives the same objective in a fresh model context, inspects the
current project, and leaves file changes on disk. The next round receives only
the previous validated report, applicable workspace instructions, and available
skills. The current model settings and permission mode still apply. Ralph stops
on completion, a reported blocker, a failed round, cancellation, or the round
limit; a worker's completion report is evidence to review, not independent
certification. Ralph workers cannot start nested Ralph loops.
### Workspace instructions and standard skills
DS4Server reads global instructions from
`~/Library/Application Support/de.rfc1437.ds4server/AGENTS.md`, then project
instructions from root to the working directory. In each directory,
`AGENTS.local.md` overlays `AGENTS.md`. Deeper guidance wins over broader
guidance, while system, developer, and direct user instructions remain higher
priority. Additions, changes, and removals are tracked during a chat and restored
after compaction; an unreadable changed file leaves its last readable version in
effect with a notice.
Manually installed standard skills live at
`~/.agents/skills/<name>/SKILL.md`. DS4Server initially gives the agent each
valid skill's name, description, and location; the agent reads the full skill
only when the task matches. These roots are readable but not writable through
agent file tools. Start a fresh chat after changing the manually installed skill
set so its prompt and tool roots agree.
### Dev Brain
Dev Brain is a project-backed Obsidian wiki for durable, source-cited engineering
knowledge. In **Preferences > Dev Brain**, select an existing vault containing
`.obsidian`, enable it, and save. Ask the agent to build, refresh, search, or
validate the wiki. DS4Server manages its contract files and explicitly marked
topic pages, not Obsidian settings, hidden files, attachments, or unrelated
notes. Verified pages can support answers; stale or needs-review pages are leads
that the agent must recheck against project source. Disabling Dev Brain removes
its tools and guidance without deleting the vault.
## Agent extensions
**Preferences > Agent Extensions** installs and manages portable Codex plugins:
1. Enter an HTTPS Git repository URL and, optionally, a branch, tag, or commit.
2. Choose **Install**. The repository must contain
`.codex-plugin/plugin.json` and provide at least a skills directory or a
supported lifecycle-hooks manifest.
3. Inspect the displayed description, source URL, requested ref, resolved
commit, author, version, skill count, and hook names.
4. Enable the extension. An extension with command hooks first asks for explicit
trust. Its programs run with your user permissions even though DS4Server
invokes them without a shell and limits their environment, runtime, input,
and output.
5. Use **Update** to fetch the configured source/ref, disable the toggle to stop
future use, or **Uninstall** to remove the installed package and its stored
extension data. Hook failures appear on the extension row.
Enable, disable, and update changes apply to subsequent matching hook events.
A current chat does not retroactively receive a `SessionStart` event. Because a
chat's native tool roots are initialized for that session, start a fresh chat
after adding, updating, enabling, or disabling extension skills. An app restart
is not otherwise required.
Extension skills are versioned inside their plugin and are managed with its
toggle, update, and uninstall actions. They are separate from manually installed
standard skills under `~/.agents/skills`. Installed packages, per-extension
data, and `registry.json` live under
`~/Library/Application Support/de.rfc1437.ds4server/extensions/`.
## A2UI interactive surfaces
A2UI is an interactive or graphical response the agent can generate when it is
more useful than prose. It is not another model and is not a tool you operate
directly. Enable or disable it in Preferences. Ask for an interactive chart,
table, form, map, or other surface, then open the **A2UI** view (`⌘2`) if the app
does not switch there automatically.
Buttons can send events back to the agent. Text fields, checkboxes, sliders,
date/time inputs, and choice pickers update the surface's data. Tabs and modals
organize content. Previous and next actions navigate saved surface history;
**Dismiss** removes the current saved surface after confirmation. Ordinary model
prose remains in the Chat transcript. Video and audio controls open native
playback windows.
The embedded catalog currently supports:
- Content: Markdown/text, images and avatars, icons, video, audio, and dividers.
- Layout: rows, columns, lists, cards, modals, and tabs.
- Controls: buttons, text fields, checkboxes, sliders, date/time inputs, and
single- or multiple-choice pickers.
- Visual and research components: bar, line, area, stacked-bar, pie, donut, and
heatmap charts; tables; metrics; timelines; maps; mind maps; and forms.
## The rest of the app
### Models and preferences
**Model Manager** downloads, resumes, verifies, and removes managed model and
acceleration artifacts. Removing an artifact does not remove projects or chats.
Preferences select the active model and its default thinking profile. Generation
values are stored per model and thinking mode; acceleration and SSD settings are
stored per model; the custom system prompt is shared across profiles. Other
sections control A2UI, permission defaults, endpoint settings, Git diff display,
Dev Brain, extensions, checkpoint storage, and diagnostics.
A new draft uses the currently selected model when its first turn is persisted.
From then on the session keeps that exact model, so opening it immediately selects
the same model and that model's thinking choices. Legacy sessions recover the
exact model from their checkpoint when available; a legacy session whose
exact model can no longer be identified remains viewable but cannot be continued.
### Main views, sidebar, and branches
**Chat** shows the conversation, **A2UI** shows interactive surfaces, **Git**
shows the current worktree, and **Stats** reports model, generation, HTTP, SSD,
and KV-cache activity. **View > Show Sidebar** (`⌘B`) hides or restores the
project sidebar. The composer status shows the active project and its local Git
branch. Branch switching is available while that project's chats are idle.
### Git
The Git view separates staged and worktree changes. Select files to stage,
unstage, or commit; selecting a file opens a colored unified or side-by-side
diff. Committing with no selected or staged files asks before staging and
committing every change. Fetch, pull, and push use the `origin` remote and report
authentication, hook, conflict, and other Git errors.
### Local API
Preferences can enable an OpenAI- and Anthropic-compatible endpoint on
`127.0.0.1`, port `4000` by default. Endpoint conversations belong to their
client and do not appear in the sidebar. CORS is off by default; enable it only
for trusted browser clients on this Mac. Saving changed endpoint settings
restarts the local listener.
`reasoning_effort` accepts only the selected model's advertised values from
`GET /v1/models`: DeepSeek accepts `none`, `low`, `high`, or `max`; GLM 5.2
accepts `none`, `high`, or `max`; GLM 5.3 Flash accepts `low`, `high`, or `max`.
Omitting it uses the upstream model default: `low` for DeepSeek and `max` for
both GLM models. Unsupported values and conflicting thinking controls return a
400 error instead of being converted to another effort. This per-model thinking
selection is intentionally part of the DS4Server agent harness, not a DS4 parity
surface; do not replace it with mapped, hidden, or shared choices.
### Data and recovery
Application data is under
`~/Library/Application Support/de.rfc1437.ds4server/`: `config.yaml` stores
preferences, `data.sqlite3` stores projects and chats, `models/` stores managed
artifacts, `kv-cache/` stores checkpoints and metrics data, `extensions/` stores
plugins, and `browser/` stores the agent browser profile. Project source remains
where you selected it.
Removing a model can be recovered by downloading or restoring and verifying the
same artifact. A context incompatible with changed runtime settings is rebuilt
from the saved transcript. Deleted project references, chats, extension data,
and their checkpoints are not restored by DS4Server; use your own backup if you
need recovery from deletion.
## Troubleshooting
- **No model selected or model unavailable:** select a model in Preferences,
then download and verify all required artifacts in Model Manager.
- **Verification failed:** retry verification or download. DS4Server will not
load incomplete or hash-mismatched weights.
- **Context is rebuilding:** the saved checkpoint no longer matches current
model/runtime settings. Let the transcript prefill finish; messages remain.
- **Extension install failed:** use an HTTPS Git URL and verify the repository
contains a valid portable plugin manifest plus skills or supported hooks.
- **Extension hook failed:** read the error on its Preferences row, repair or
update the extension, then retry the matching event. Trust is required before
command hooks can run.
- **Tool approval denied:** the operation was not run. Send a safer request or
repeat it and approve only after checking the command and working directory.
- **A2UI is unavailable:** enable A2UI in Preferences and ask the agent for an
interactive surface in a new prompt.
## Keyboard shortcuts
These are the shortcuts registered by the native DS4Server menu. On Macs whose
top row controls brightness or media, use `fn-F1` when macOS requires it;
DS4Server does not override the system keyboard mapping.
| Shortcut | Action |
| --- | --- |
| `F1` | DS4Server User Guide |
| `⌘,` | Preferences |
| `⇧⌘M` | Model Manager |
| `⌘N` | New Chat |
| `⇧⌘S` | Export Chat as Markdown |
| `⌘B` | Show or hide the sidebar |
| `⌘1` | Chat view |
| `⌘2` | A2UI view |
| `⌘3` | Git view |
| `⌘4` | Stats view |
| `⌘Z` | Undo |
| `⇧⌘Z` | Redo |
| `⌘X` | Cut |
| `⌘C` | Copy |
| `⌘V` | Paste |
| `⌘A` | Select All |
| `⌘Q` | Quit DS4Server |

View File

@@ -0,0 +1,19 @@
implementation,mode,run,prompt_tokens,decode_tokens,prefill_tps,first_token_ms,steady_tps,p50_ms,p95_ms,cache_entries,cache_hits,cache_misses,pread_bytes,token_sha256
ds4server-baseline-79468c65,resident,1,16,128,69.627264,26.270667,39.534586,24.875959,26.607167,0,0,0,0,
ds4server-baseline-79468c65,resident,2,16,128,69.799133,26.365458,39.505867,24.905750,26.624667,0,0,0,0,
ds4server-baseline-79468c65,resident,3,16,128,69.618099,26.454750,39.526849,24.877584,26.585875,0,0,0,0,
ds4server-baseline-79468c65,ssd,1,16,128,6.824429,122.398208,8.638560,114.928000,124.375958,4096,34040,7208,51017416704,
ds4server-baseline-79468c65,ssd,2,16,128,6.805883,126.033250,8.634659,114.667833,124.220583,4096,34040,7208,51017416704,
ds4server-baseline-79468c65,ssd,3,16,128,6.726701,125.304417,8.532966,116.431375,124.278709,4096,34040,7208,51017416704,
ds4-8db89fe0,resident,1,16,128,56.49,22.706,45.67,21.905,22.297,0,0,0,0,ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
ds4-8db89fe0,resident,2,16,128,56.49,22.719,45.59,21.908,22.320,0,0,0,0,ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
ds4-8db89fe0,resident,3,16,128,56.51,22.638,45.60,21.890,22.342,0,0,0,0,ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
ds4-8db89fe0,ssd,1,16,128,12.07,46.746,25.73,36.461,46.861,4096,32958,4194,29684662272,2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
ds4-8db89fe0,ssd,2,16,128,11.88,46.737,26.11,35.622,46.982,4096,32958,4194,29684662272,2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
ds4-8db89fe0,ssd,3,16,128,11.16,47.092,26.20,35.513,46.554,4096,32958,4194,29684662272,2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
ds4server-issue-78,resident,1,16,128,69.930426,22.547583,45.683038,21.889833,22.275042,0,0,0,0,ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
ds4server-issue-78,resident,2,16,128,69.760230,22.657125,45.602401,21.936042,22.331375,0,0,0,0,ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
ds4server-issue-78,resident,3,16,128,69.727346,22.663208,45.361642,22.052833,22.459084,0,0,0,0,ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
ds4server-issue-78,ssd,1,16,128,11.279689,48.890291,24.849723,37.847292,48.853791,4096,32958,4194,29684662272,2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
ds4server-issue-78,ssd,2,16,128,11.520460,48.567417,25.076916,37.696334,48.779625,4096,32958,4194,29684662272,2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
ds4server-issue-78,ssd,3,16,128,11.494703,48.844875,25.066058,37.616542,48.911333,4096,32958,4194,29684662272,2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
1 implementation mode run prompt_tokens decode_tokens prefill_tps first_token_ms steady_tps p50_ms p95_ms cache_entries cache_hits cache_misses pread_bytes token_sha256
2 ds4server-baseline-79468c65 resident 1 16 128 69.627264 26.270667 39.534586 24.875959 26.607167 0 0 0 0
3 ds4server-baseline-79468c65 resident 2 16 128 69.799133 26.365458 39.505867 24.905750 26.624667 0 0 0 0
4 ds4server-baseline-79468c65 resident 3 16 128 69.618099 26.454750 39.526849 24.877584 26.585875 0 0 0 0
5 ds4server-baseline-79468c65 ssd 1 16 128 6.824429 122.398208 8.638560 114.928000 124.375958 4096 34040 7208 51017416704
6 ds4server-baseline-79468c65 ssd 2 16 128 6.805883 126.033250 8.634659 114.667833 124.220583 4096 34040 7208 51017416704
7 ds4server-baseline-79468c65 ssd 3 16 128 6.726701 125.304417 8.532966 116.431375 124.278709 4096 34040 7208 51017416704
8 ds4-8db89fe0 resident 1 16 128 56.49 22.706 45.67 21.905 22.297 0 0 0 0 ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
9 ds4-8db89fe0 resident 2 16 128 56.49 22.719 45.59 21.908 22.320 0 0 0 0 ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
10 ds4-8db89fe0 resident 3 16 128 56.51 22.638 45.60 21.890 22.342 0 0 0 0 ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
11 ds4-8db89fe0 ssd 1 16 128 12.07 46.746 25.73 36.461 46.861 4096 32958 4194 29684662272 2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
12 ds4-8db89fe0 ssd 2 16 128 11.88 46.737 26.11 35.622 46.982 4096 32958 4194 29684662272 2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
13 ds4-8db89fe0 ssd 3 16 128 11.16 47.092 26.20 35.513 46.554 4096 32958 4194 29684662272 2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
14 ds4server-issue-78 resident 1 16 128 69.930426 22.547583 45.683038 21.889833 22.275042 0 0 0 0 ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
15 ds4server-issue-78 resident 2 16 128 69.760230 22.657125 45.602401 21.936042 22.331375 0 0 0 0 ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
16 ds4server-issue-78 resident 3 16 128 69.727346 22.663208 45.361642 22.052833 22.459084 0 0 0 0 ed5461476a82824c3e0440e943a0b2f6b8a8f19dc26a550a1bc255003851ab8d
17 ds4server-issue-78 ssd 1 16 128 11.279689 48.890291 24.849723 37.847292 48.853791 4096 32958 4194 29684662272 2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
18 ds4server-issue-78 ssd 2 16 128 11.520460 48.567417 25.076916 37.696334 48.779625 4096 32958 4194 29684662272 2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052
19 ds4server-issue-78 ssd 3 16 128 11.494703 48.844875 25.066058 37.616542 48.911333 4096 32958 4194 29684662272 2bc33a04bb88f9f3c56f8bae0124b7bf25fd2f32aee81a1bdaf6e979448da052

18
docs/qwen38-issue-99.csv Normal file
View File

@@ -0,0 +1,18 @@
date,engine,mode,sequence,run,state,context,prompt_tokens,generated_tokens,temperature,top_p,top_k,seed,prefill_tps,first_token_ms,steady_tps,decode_tps,p50_cycle_ms,p95_cycle_ms,peak_rss_bytes,token_sha256,cycles,drafted,accepted,verifier_passes,verifier_ms,repair_ms,ple_entries,ple_hits,ple_misses,ple_evictions,ple_pread_bytes,ple_pread_ms,notes
2026-09-04,DS4Server,MTP,B1,1,cold,4096,26,64,1.0,0.95,20,12345,6.360708,59.484667,57.518099,58.042801,54.731042,57.838709,52720156672,46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72,20,60,44,20,1000,0,1648,432,1648,0,164800,,repair timing counter not yet wired
2026-09-04,DS4Server,MTP,B1,2,warm,4096,26,64,1.0,0.95,20,12345,51.095361,52.229584,60.307983,61.119864,52.427709,54.273459,52720156672,46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72,20,60,44,20,946,0,1648,2080,0,0,0,0,repair timing counter not yet wired
2026-09-04,DS4Server,MTP,B1,3,warm,4096,26,64,1.0,0.95,20,12345,51.152493,52.422375,60.014218,60.825779,52.679334,54.147250,52720156672,46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72,20,60,44,20,951,0,1648,2080,0,0,0,0,acceptance median 60.825779
2026-09-04,DS4Server,MTP,B2,1,cold,4096,26,64,1.0,0.95,20,12345,6.527240,62.605750,54.320764,54.834156,57.909625,60.828916,52605190144,46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72,20,60,44,20,1056,30,1648,432,1648,0,164800,63,confirmation after repair telemetry
2026-09-04,DS4Server,MTP,B2,2,warm,4096,26,64,1.0,0.95,20,12345,50.178762,54.061500,57.196397,58.019436,55.002750,57.242875,52605190144,46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72,20,60,44,20,993,29,1648,2080,0,0,0,0,confirmation median 58.019436
2026-09-04,DS4Server,MTP,B2,3,warm,4096,26,64,1.0,0.95,20,12345,49.976289,55.289000,57.556598,58.301539,54.853000,56.619333,52605190144,46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72,20,60,44,20,988,30,1648,2080,0,0,0,0,confirmation after repair telemetry
2026-09-04,DS4Server,AR,A1,1,cold,4096,26,64,1.0,0.95,20,12345,7.796453,21.087125,49.845055,49.805297,18.894042,19.426917,53962735616,15373689bcb1c667123accfcfbe49736466f28ce7beb89d2cbe5e4120416056d,0,0,0,0,0,0,1648,16,1424,0,,,plain autoregressive
2026-09-04,DS4Server,AR,A1,2,warm,4096,26,64,1.0,0.95,20,12345,51.386057,18.606375,51.034286,51.074537,18.392375,18.953542,53962735616,15373689bcb1c667123accfcfbe49736466f28ce7beb89d2cbe5e4120416056d,0,0,0,0,0,0,1648,,,,,,plain autoregressive
2026-09-04,DS4Server,AR,A1,3,warm,4096,26,64,1.0,0.95,20,12345,51.424682,18.256166,50.610874,50.671077,18.575958,19.129667,53962735616,15373689bcb1c667123accfcfbe49736466f28ce7beb89d2cbe5e4120416056d,0,0,0,0,0,0,1648,,,,,,plain autoregressive median 50.671077
2026-09-04,DS4Server,AR,LONG,1,warm,2048,,1,,,,,0.778000,,,29.259000,,,52158349312,b6d28407,,,,,,,,,,,,,synthetic long-context state gate token=5 live_state_bytes=58195968
2026-09-04,DS4Server,AR,LONG,1,warm,16384,,1,,,,,12.101000,,,46.486000,,,,79d71772,,,,,,,,,,,,,synthetic long-context state gate token=5 live_state_bytes=465567744
2026-09-04,DS4Server,AR,LONG,1,warm,65536,,1,,,,,10.037000,,,45.739000,,,,0a1131fe,,,,,,,,,,,,,synthetic long-context state gate token=5 live_state_bytes=1862270976
2026-09-04,DS4Server,AR,LONG,1,warm,131072,,1,,,,,3.277000,,,45.145000,,,,fb678d10,,,,,,,,,,,,,synthetic long-context state gate token=5 live_state_bytes=3724541952
2026-09-04,DS4Server,AR,LONG,1,warm,262144,,1,,,,,0.777000,,,43.488000,,,,01ae755c,,,,,,,,,,,,,synthetic native-maximum state gate token=5 live_state_bytes=7449083904
2026-09-04,MTPLX,MTP,ORACLE,1,warm,4096,26,64,1.0,0.95,20,12345,,,,59.170000,,,,,20,58,44,20,950,,,,,,,,fresh matched run before local source-package drift
2026-09-04,MTPLX,MTP,ORACLE,2,warm,4096,26,64,1.0,0.95,20,12345,,,,59.820000,,,,,20,58,44,20,950,,,,,,,,fresh matched run before local source-package drift
2026-09-04,MTPLX,MTP,ORACLE,3,warm,4096,26,64,1.0,0.95,20,12345,,,,59.751170,,,,,20,58,44,20,950,,,,,,,,matched median 59.751170
1 date engine mode sequence run state context prompt_tokens generated_tokens temperature top_p top_k seed prefill_tps first_token_ms steady_tps decode_tps p50_cycle_ms p95_cycle_ms peak_rss_bytes token_sha256 cycles drafted accepted verifier_passes verifier_ms repair_ms ple_entries ple_hits ple_misses ple_evictions ple_pread_bytes ple_pread_ms notes
2 2026-09-04 DS4Server MTP B1 1 cold 4096 26 64 1.0 0.95 20 12345 6.360708 59.484667 57.518099 58.042801 54.731042 57.838709 52720156672 46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72 20 60 44 20 1000 0 1648 432 1648 0 164800 repair timing counter not yet wired
3 2026-09-04 DS4Server MTP B1 2 warm 4096 26 64 1.0 0.95 20 12345 51.095361 52.229584 60.307983 61.119864 52.427709 54.273459 52720156672 46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72 20 60 44 20 946 0 1648 2080 0 0 0 0 repair timing counter not yet wired
4 2026-09-04 DS4Server MTP B1 3 warm 4096 26 64 1.0 0.95 20 12345 51.152493 52.422375 60.014218 60.825779 52.679334 54.147250 52720156672 46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72 20 60 44 20 951 0 1648 2080 0 0 0 0 acceptance median 60.825779
5 2026-09-04 DS4Server MTP B2 1 cold 4096 26 64 1.0 0.95 20 12345 6.527240 62.605750 54.320764 54.834156 57.909625 60.828916 52605190144 46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72 20 60 44 20 1056 30 1648 432 1648 0 164800 63 confirmation after repair telemetry
6 2026-09-04 DS4Server MTP B2 2 warm 4096 26 64 1.0 0.95 20 12345 50.178762 54.061500 57.196397 58.019436 55.002750 57.242875 52605190144 46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72 20 60 44 20 993 29 1648 2080 0 0 0 0 confirmation median 58.019436
7 2026-09-04 DS4Server MTP B2 3 warm 4096 26 64 1.0 0.95 20 12345 49.976289 55.289000 57.556598 58.301539 54.853000 56.619333 52605190144 46ce63601c44d80b5ef8297eaa49a636f6680040b916fe30cf7a154c05daef72 20 60 44 20 988 30 1648 2080 0 0 0 0 confirmation after repair telemetry
8 2026-09-04 DS4Server AR A1 1 cold 4096 26 64 1.0 0.95 20 12345 7.796453 21.087125 49.845055 49.805297 18.894042 19.426917 53962735616 15373689bcb1c667123accfcfbe49736466f28ce7beb89d2cbe5e4120416056d 0 0 0 0 0 0 1648 16 1424 0 plain autoregressive
9 2026-09-04 DS4Server AR A1 2 warm 4096 26 64 1.0 0.95 20 12345 51.386057 18.606375 51.034286 51.074537 18.392375 18.953542 53962735616 15373689bcb1c667123accfcfbe49736466f28ce7beb89d2cbe5e4120416056d 0 0 0 0 0 0 1648 plain autoregressive
10 2026-09-04 DS4Server AR A1 3 warm 4096 26 64 1.0 0.95 20 12345 51.424682 18.256166 50.610874 50.671077 18.575958 19.129667 53962735616 15373689bcb1c667123accfcfbe49736466f28ce7beb89d2cbe5e4120416056d 0 0 0 0 0 0 1648 plain autoregressive median 50.671077
11 2026-09-04 DS4Server AR LONG 1 warm 2048 1 0.778000 29.259000 52158349312 b6d28407 synthetic long-context state gate token=5 live_state_bytes=58195968
12 2026-09-04 DS4Server AR LONG 1 warm 16384 1 12.101000 46.486000 79d71772 synthetic long-context state gate token=5 live_state_bytes=465567744
13 2026-09-04 DS4Server AR LONG 1 warm 65536 1 10.037000 45.739000 0a1131fe synthetic long-context state gate token=5 live_state_bytes=1862270976
14 2026-09-04 DS4Server AR LONG 1 warm 131072 1 3.277000 45.145000 fb678d10 synthetic long-context state gate token=5 live_state_bytes=3724541952
15 2026-09-04 DS4Server AR LONG 1 warm 262144 1 0.777000 43.488000 01ae755c synthetic native-maximum state gate token=5 live_state_bytes=7449083904
16 2026-09-04 MTPLX MTP ORACLE 1 warm 4096 26 64 1.0 0.95 20 12345 59.170000 20 58 44 20 950 fresh matched run before local source-package drift
17 2026-09-04 MTPLX MTP ORACLE 2 warm 4096 26 64 1.0 0.95 20 12345 59.820000 20 58 44 20 950 fresh matched run before local source-package drift
18 2026-09-04 MTPLX MTP ORACLE 3 warm 4096 26 64 1.0 0.95 20 12345 59.751170 20 58 44 20 950 matched median 59.751170

View File

@@ -197,33 +197,6 @@ kernel void kernel_mul_mv_q8_0_f32(
kernel_mul_mv_q8_0_f32_impl<N_R0_Q8_0, constant ds4_metal_args_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
[[host_name("kernel_mul_mv_q8_0_f32_r4")]]
kernel void kernel_mul_mv_q8_0_f32_r4(
constant ds4_metal_args_mul_mv & args,
device const char * src0,
device const char * src1,
device char * dst,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_mul_mv_q8_0_f32_impl<4, constant ds4_metal_args_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
// Output projection alias used by the optimized host dispatch.
[[host_name("kernel_mul_mv_q8_0_f32_nr4")]]
kernel void kernel_mul_mv_q8_0_f32_nr4(
constant ds4_metal_args_mul_mv & args,
device const char * src0,
device const char * src1,
device char * dst,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_mul_mv_q8_0_f32_impl<4, constant ds4_metal_args_mul_mv &>(
args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
// Decode Q-A/KV pair. Both projections consume the same activation row but
// have independent weight ranges and output extents. Keep the standalone Q8_0
@@ -497,12 +470,22 @@ kernel void kernel_dsv4_shared_gate_up_swiglu_q8_0(
clamp_value, shmem, tgpig, tiisg, sgitg);
}
[[host_name("kernel_dsv4_shared_gate_up_swiglu_q8_0_r4")]]
kernel void kernel_dsv4_shared_gate_up_swiglu_q8_0_r4(
// Decode-only fusion of the router logits matvec (F16, embd -> n_expert)
// with the shared-expert gate/up SwiGLU (Q8_0, embd -> shared). Both read
// the same normalized FFN input back to back; one dispatch removes one
// launch per decode layer. Router threadgroups replicate
// kernel_mul_mv_f16_f32_4 (nsg=8, nr0=2); shared threadgroups host two
// virtual 4-simdgroup cohorts replicating
// kernel_dsv4_shared_gate_up_swiglu_q8_0 (nsg=4, nr0=2), including its
// per-row simd/shmem reduction trees. Bit-exact by construction.
kernel void kernel_dsv4_router_shared_gate_up_q8_0(
constant ds4_metal_args_mul_mv & args,
constant ds4_metal_args_mul_mv & sargs,
device const char * src0_router,
device const char * src0_gate,
device const char * src0_up,
device const char * src1,
device char * dst_router,
device char * dst_gate,
device char * dst_up,
device char * dst_mid,
@@ -511,11 +494,166 @@ kernel void kernel_dsv4_shared_gate_up_swiglu_q8_0_r4(
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
kernel_dsv4_shared_gate_up_swiglu_q8_0_impl<4, true>(
args, src0_gate, src0_up, src1, dst_gate, dst_up, dst_mid,
clamp_value, shmem, tgpig, tiisg, sgitg);
constexpr short NW = N_SIMDWIDTH;
const uint router_tgs = ((uint)args.ne01 + 1u) / 2u;
if (tgpig.x < router_tgs) {
// Exact replica of kernel_mul_mv_f16_f32_4 with NSG=8, NR0=2.
constexpr short NSG = 8;
constexpr short NR0 = 2;
constexpr short NB = 32;
constexpr short NF = 16;
constexpr short NF4 = NF/4;
const int nb = args.ne00/NB;
const int r0 = tgpig.x*NR0;
device const float4 * y4 = (device const float4 *) src1;
device const half4 * ax4[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
ax4[row] = (device const half4 *)
(src0_router + (uint64_t)(r0 + row)*args.nb01);
}
float sumf[NR0] = { 0.f };
const short ix = tiisg/(NW/NF);
const short il = tiisg%(NW/NF);
const int ib0 = sgitg*NF + ix;
device const float4 * yb4 = y4 + (ib0*NB + il*NF)/4;
for (int ib = ib0; ib < nb; ib += NSG*NF) {
float4 yl4[NF4];
FOR_UNROLL (short i = 0; i < NF4; ++i) {
yl4[i] = yb4[i];
}
FOR_UNROLL (short row = 0; row < NR0; row++) {
device const half4 * xb4 = ax4[row] + (ib*NB + il*NF)/4;
float sumq = 0.f;
FOR_UNROLL (short i = 0; i < NF4; ++i) {
sumq += dot(float4(xb4[i]), yl4[i]);
}
sumf[row] += sumq;
}
yb4 += NSG*NF*NW/4;
}
device float * dst_f32 = (device float *) dst_router;
helper_mv_reduce_and_write<NR0>(dst_f32, sumf, r0, args.ne01,
tiisg, sgitg, shmem);
return;
}
// Shared-expert part: two virtual nsg=4 cohorts per threadgroup, each an
// exact replica of kernel_dsv4_shared_gate_up_swiglu_q8_0 (NR0=2).
constexpr short NSG = 4;
constexpr short NR0 = 2;
constexpr short NQ = 8;
const uint cohort = sgitg >> 2;
const ushort vsg = sgitg & 3u;
const uint vt = (tgpig.x - router_tgs) * 2u + cohort;
const int nb = sargs.ne00 / QK8_0;
const int r0 = vt * NR0;
device const float *y = (device const float *) src1;
device const block_q8_0 *ag[NR0];
device const block_q8_0 *au[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
const uint64_t offset0 = (uint64_t)(r0 + row) * sargs.nb01;
ag[row] = (device const block_q8_0 *)(src0_gate + offset0);
au[row] = (device const block_q8_0 *)(src0_up + offset0);
}
float sumg[NR0] = { 0.f };
float sumu[NR0] = { 0.f };
const short ix = tiisg / (NW / NQ);
const short il = tiisg % (NW / NQ);
const int ib0 = vsg * NQ + ix;
float yl[NQ];
device const float *yb = y + ib0 * QK8_0 + il * NQ;
for (int ib = ib0; ib < nb; ib += NSG * NQ) {
FOR_UNROLL (short i = 0; i < NQ; ++i) {
yl[i] = yb[i];
}
FOR_UNROLL (short row = 0; row < NR0; ++row) {
device const int8_t *qg = ag[row][ib].qs + il * NQ;
device const int8_t *qu = au[row][ib].qs + il * NQ;
float sg = 0.f;
float su = 0.f;
FOR_UNROLL (short i = 0; i < NQ; ++i) {
sg += qg[i] * yl[i];
su += qu[i] * yl[i];
}
sumg[row] += sg * ag[row][ib].d;
sumu[row] += su * au[row][ib].d;
}
yb += NSG * NQ * QK8_0;
}
threadgroup float *shmem_f32 = (threadgroup float *)shmem + cohort * (2*NR0*NW);
threadgroup float *sh_gate[NR0];
threadgroup float *sh_up[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
sh_gate[row] = shmem_f32 + NW * row;
sh_up[row] = shmem_f32 + NW * (NR0 + row);
if (vsg == 0) {
sh_gate[row][tiisg] = 0.0f;
sh_up[row][tiisg] = 0.0f;
}
sumg[row] = simd_sum(sumg[row]);
sumu[row] = simd_sum(sumu[row]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
FOR_UNROLL (short row = 0; row < NR0; ++row) {
if (tiisg == 0) {
sh_gate[row][vsg] = sumg[row];
sh_up[row][vsg] = sumu[row];
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
device float *gate_f32 = (device float *)dst_gate;
device float *up_f32 = (device float *)dst_up;
device float *mid_f32 = (device float *)dst_mid;
FOR_UNROLL (short row = 0; row < NR0 && r0 + row < sargs.ne01; ++row) {
const float gate = simd_sum(sh_gate[row][tiisg]);
const float up = simd_sum(sh_up[row][tiisg]);
if (tiisg == 0 && vsg == 0) {
const uint out_row = r0 + row;
gate_f32[out_row] = gate;
up_f32[out_row] = up;
float g = gate;
float u = up;
if (clamp_value > 1.0e-6f) {
g = min(g, clamp_value);
u = clamp(u, -clamp_value, clamp_value);
}
const float silu = g / (1.0f + exp(-g));
mid_f32[out_row] = silu * u;
}
}
}
[[host_name("kernel_dsv4_shared_mid_swiglu_q8_0")]]
kernel void kernel_dsv4_shared_mid_swiglu_q8_0(
constant ds4_metal_args_mul_mv & args,
@@ -535,24 +673,6 @@ kernel void kernel_dsv4_shared_mid_swiglu_q8_0(
clamp_value, shmem, tgpig, tiisg, sgitg);
}
[[host_name("kernel_dsv4_shared_mid_swiglu_q8_0_r4")]]
kernel void kernel_dsv4_shared_mid_swiglu_q8_0_r4(
constant ds4_metal_args_mul_mv & args,
device const char * src0_gate,
device const char * src0_up,
device const char * src1,
device char * dst_gate,
device char * dst_up,
device char * dst_mid,
constant float &clamp_value,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_dsv4_shared_gate_up_swiglu_q8_0_impl<4, false>(
args, src0_gate, src0_up, src1, dst_gate, dst_up, dst_mid,
clamp_value, shmem, tgpig, tiisg, sgitg);
}
template<typename T0, typename T1, short NR0, typename args_t>
void kernel_mul_mv_t_t_impl(
@@ -974,6 +1094,311 @@ kernel void kernel_mul_mv_f16_f32_pair_compressor_store_4(
state_score[dst] = projected_score[col] + ape_v;
}
// Decode compressor + indexer-compressor projection in one dispatch. Both
// pairs read the same normalized activation with the same F16 matvec shape,
// so one launch covers all four matrices: threadgroups below the first
// range boundary run the exact paired matvec + state store of
// kernel_mul_mv_f16_f32_pair_compressor_store_4 for the attention
// compressor, the rest for the indexer compressor. Per-row reduction trees
// and the per-threadgroup state stores are unchanged, keeping the fused
// result bit-identical to the two separate dispatches while removing one
// dispatch per decode layer.
kernel void kernel_mul_mv_f16_f32_quad_compressor_store_4(
constant ds4_metal_args_mul_mv & args,
constant ds4_metal_args_compressor_pair_store & store0,
constant ds4_metal_args_compressor_pair_store & store1,
device const char * src0_a0,
device const char * src0_b0,
device const char * src0_a1,
device const char * src0_b1,
device const char * src1,
device char * dst_a0,
device char * dst_b0,
device char * dst_a1,
device char * dst_b1,
device const char * ape0,
device const char * ape1,
device float * state0_kv,
device float * state0_score,
device float * state1_kv,
device float * state1_score,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiitg [[thread_index_in_threadgroup]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constexpr short NR0 = 2;
const uint tgs0 = ((uint)store0.width + NR0 - 1u) / NR0;
const bool second = tgpig.x >= tgs0;
uint3 local_tgpig = tgpig;
if (second) local_tgpig.x = tgpig.x - tgs0;
ds4_metal_args_mul_mv largs = args;
largs.nr0 = NR0;
largs.ne01 = second ? (int32_t)store1.width : (int32_t)store0.width;
if (!second) {
kernel_mul_mv_f16_f32_pair_4_impl<NR0>(
largs, src0_a0, src0_b0, src1, dst_a0, dst_b0,
shmem, local_tgpig, tiisg, sgitg);
} else {
kernel_mul_mv_f16_f32_pair_4_impl<NR0>(
largs, src0_a1, src0_b1, src1, dst_a1, dst_b1,
shmem, local_tgpig, tiisg, sgitg);
}
threadgroup_barrier(mem_flags::mem_device);
// State append: identical to the paired store kernel, scoped to the
// range this threadgroup just projected (its own outputs only).
constant ds4_metal_args_compressor_pair_store & store = second ? store1 : store0;
if (tiitg >= NR0 || store.width == 0u || store.ratio == 0u) {
return;
}
const uint col = local_tgpig.x * (uint)NR0 + tiitg;
if (col >= store.width) return;
const uint pos_mod = store.pos % store.ratio;
const uint dst_row = store.ratio == 4u ? store.ratio + pos_mod : pos_mod;
const uint dst = dst_row * store.width + col;
const uint ape_i = pos_mod * store.width + col;
device volatile const float * projected_kv = second
? (device volatile const float *)dst_a1
: (device volatile const float *)dst_a0;
device volatile const float * projected_score = second
? (device volatile const float *)dst_b1
: (device volatile const float *)dst_b0;
device const char * ape = second ? ape1 : ape0;
device float * state_kv = second ? state1_kv : state0_kv;
device float * state_score = second ? state1_score : state0_score;
float ape_v;
if (store.ape_type == 1u) {
ape_v = (float)(((device const half *)ape)[ape_i]);
} else {
ape_v = ((device const float *)ape)[ape_i];
}
state_kv[dst] = projected_kv[col];
state_score[dst] = projected_score[col] + ape_v;
}
/* Decode-only fusion: one dispatch covers the q_a/kv Q8 pair projection and
* the four F16 compressor projections (attention + indexer) with their
* state-store epilogue. Both stages read the same normalized attention
* input and write disjoint outputs. The q_a/kv range hosts two virtual
* NSG=4 cohorts per threadgroup, each an exact replica of
* kernel_mul_mv_q8_0_f32_pair (same per-lane K walk and reduction tree, cf.
* kernel_dsv4_router_shared_gate_up_q8_0); the compressor ranges run
* kernel_mul_mv_f16_f32_pair_4_impl<2> and the paired store epilogue
* verbatim, so every output bit matches the two separate dispatches. */
kernel void kernel_dsv4_qkv_pair_quad_compressor_store_q8_0(
constant ds4_metal_args_mul_mv & args0,
constant ds4_metal_args_mul_mv & args1,
constant ds4_metal_args_mul_mv & cargs,
constant ds4_metal_args_compressor_pair_store & store0,
constant ds4_metal_args_compressor_pair_store & store1,
constant uint & pair_vtgs,
device const char * qw0,
device const char * qw1,
device const char * cw0a,
device const char * cw0b,
device const char * cw1a,
device const char * cw1b,
device const char * src1,
device char * dst0,
device char * dst1,
device char * cdst_a0,
device char * cdst_b0,
device char * cdst_a1,
device char * cdst_b1,
device const char * ape0,
device const char * ape1,
device float * state0_kv,
device float * state0_score,
device float * state1_kv,
device float * state1_score,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiitg [[thread_index_in_threadgroup]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constexpr short NW = N_SIMDWIDTH;
const uint pair_ctgs = (pair_vtgs + 1u) / 2u;
if (tgpig.x < pair_ctgs) {
/* Q8 pair range: cohort c of threadgroup t runs virtual pair
* threadgroup 2t+c with the original NSG=4 mapping. */
constexpr short NSG = 4;
constexpr short NQ = 8;
constexpr short NR0 = 2;
const uint cohort = sgitg >> 2;
const ushort vsg = sgitg & 3u;
const uint vt = tgpig.x * 2u + cohort;
const bool valid = vt < pair_vtgs;
const int r0 = vt * NR0;
const bool active_a = valid && r0 < args0.ne01;
const bool active_b = valid && r0 < args1.ne01;
const int nb = args0.ne00 / QK8_0;
device const float *y = (device const float *)src1;
device const block_q8_0 *ax_a[NR0];
device const block_q8_0 *ax_b[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
const int out_row = r0 + row;
ax_a[row] = active_a && out_row < args0.ne01
? (device const block_q8_0 *)(qw0 + (uint64_t)out_row * args0.nb01)
: (device const block_q8_0 *)qw0;
ax_b[row] = active_b && out_row < args1.ne01
? (device const block_q8_0 *)(qw1 + (uint64_t)out_row * args1.nb01)
: (device const block_q8_0 *)qw1;
}
float suma[NR0] = { 0.f };
float sumb[NR0] = { 0.f };
const short ix = tiisg / (NW / NQ);
const short il = tiisg % (NW / NQ);
const int ib0 = vsg * NQ + ix;
float yl[NQ];
device const float *yb = y + ib0 * QK8_0 + il * NQ;
if (valid) {
for (int ib = ib0; ib < nb; ib += NSG * NQ) {
FOR_UNROLL (short i = 0; i < NQ; ++i) {
yl[i] = yb[i];
}
FOR_UNROLL (short row = 0; row < NR0; ++row) {
const int out_row = r0 + row;
if (active_a && out_row < args0.ne01) {
device const int8_t *qs = ax_a[row][ib].qs + il * NQ;
float sumq = 0.f;
FOR_UNROLL (short i = 0; i < NQ; ++i) {
sumq += qs[i] * yl[i];
}
suma[row] += sumq * ax_a[row][ib].d;
}
if (active_b && out_row < args1.ne01) {
device const int8_t *qs = ax_b[row][ib].qs + il * NQ;
float sumq = 0.f;
FOR_UNROLL (short i = 0; i < NQ; ++i) {
sumq += qs[i] * yl[i];
}
sumb[row] += sumq * ax_b[row][ib].d;
}
}
yb += NSG * NQ * QK8_0;
}
}
threadgroup float *shared =
(threadgroup float *)shmem + cohort * (2 * NR0 * NW);
threadgroup float *sha[NR0];
threadgroup float *shb[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
sha[row] = shared + NW * row;
shb[row] = shared + NW * (NR0 + row);
if (vsg == 0) {
sha[row][tiisg] = 0.0f;
if (active_b) shb[row][tiisg] = 0.0f;
}
suma[row] = simd_sum(suma[row]);
if (active_b) sumb[row] = simd_sum(sumb[row]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
FOR_UNROLL (short row = 0; row < NR0; ++row) {
if (tiisg == 0) {
sha[row][vsg] = suma[row];
if (active_b) shb[row][vsg] = sumb[row];
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
device float *out_a = (device float *)dst0;
device float *out_b = (device float *)dst1;
FOR_UNROLL (short row = 0; row < NR0; ++row) {
const float total_a = simd_sum(sha[row][tiisg]);
if (tiisg == 0 && vsg == 0) {
const int out_row = r0 + row;
if (active_a && out_row < args0.ne01) out_a[out_row] = total_a;
}
if (active_b) {
const float total_b = simd_sum(shb[row][tiisg]);
if (tiisg == 0 && vsg == 0) {
const int out_row = r0 + row;
if (out_row < args1.ne01) out_b[out_row] = total_b;
}
}
}
return;
}
/* Compressor quad range: verbatim body of
* kernel_mul_mv_f16_f32_quad_compressor_store_4 on the shifted grid. */
constexpr short NR0 = 2;
const uint lx = tgpig.x - pair_ctgs;
const uint tgs0 = ((uint)store0.width + NR0 - 1u) / NR0;
const bool second = lx >= tgs0;
uint3 local_tgpig = tgpig;
local_tgpig.x = second ? lx - tgs0 : lx;
ds4_metal_args_mul_mv largs = cargs;
largs.nr0 = NR0;
largs.ne01 = second ? (int32_t)store1.width : (int32_t)store0.width;
if (!second) {
kernel_mul_mv_f16_f32_pair_4_impl<NR0>(
largs, cw0a, cw0b, src1, cdst_a0, cdst_b0,
shmem, local_tgpig, tiisg, sgitg);
} else {
kernel_mul_mv_f16_f32_pair_4_impl<NR0>(
largs, cw1a, cw1b, src1, cdst_a1, cdst_b1,
shmem, local_tgpig, tiisg, sgitg);
}
threadgroup_barrier(mem_flags::mem_device);
// State append: identical to the paired store kernel, scoped to the
// range this threadgroup just projected (its own outputs only).
constant ds4_metal_args_compressor_pair_store & store = second ? store1 : store0;
if (tiitg >= NR0 || store.width == 0u || store.ratio == 0u) {
return;
}
const uint col = local_tgpig.x * (uint)NR0 + tiitg;
if (col >= store.width) return;
const uint pos_mod = store.pos % store.ratio;
const uint dst_row = store.ratio == 4u ? store.ratio + pos_mod : pos_mod;
const uint dst = dst_row * store.width + col;
const uint ape_i = pos_mod * store.width + col;
device volatile const float * projected_kv = second
? (device volatile const float *)cdst_a1
: (device volatile const float *)cdst_a0;
device volatile const float * projected_score = second
? (device volatile const float *)cdst_b1
: (device volatile const float *)cdst_b0;
device const char * ape = second ? ape1 : ape0;
device float * state_kv = second ? state1_kv : state0_kv;
device float * state_score = second ? state1_score : state0_score;
float ape_v;
if (store.ape_type == 1u) {
ape_v = (float)(((device const half *)ape)[ape_i]);
} else {
ape_v = ((device const float *)ape)[ape_i];
}
state_kv[dst] = projected_kv[col];
state_score[dst] = projected_score[col] + ape_v;
}
template<typename T0, typename T1, typename args_t>
void kernel_mul_mv_t_t_short_impl(
args_t args,
@@ -1476,125 +1901,6 @@ constant bool FC_mul_mm_bc_inp [[function_constant(FC_MUL_MM + 0)]];
constant bool FC_mul_mm_bc_out [[function_constant(FC_MUL_MM + 1)]];
#ifdef DS4_METAL_HAS_TENSOR
template<
short NR0, short NR1,
typename SA, typename SA_4x4, typename block_q, short nl,
void (*dequantize_func)(device const block_q *, short, thread SA_4x4 &),
typename T0, typename T0_4x4, typename T1>
kernel void kernel_mul_mm_mpp(
constant ds4_metal_args_mul_mm & args,
device const char * srcA,
device const char * srcB,
device char * dst,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiitg [[thread_index_in_threadgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
(void) sgitg;
constexpr int NK = 32;
constexpr int NL = NK/16;
constexpr int NUM_THREADS = 128;
const int K = args.ne00;
const int M = args.ne0;
const int N = args.ne1;
const int im = tgpig.z;
const int i12 = im%args.ne12;
const int i13 = im/args.ne12;
const int r0 = tgpig.y*NR0;
const int r1 = tgpig.x*NR1;
const uint64_t offset0 = (i12/args.r2)*args.nb02 + (i13/args.r3)*args.nb03;
threadgroup SA *sa = (threadgroup SA *)shmem;
threadgroup SA *sb = sa + NR0*NK;
auto tA = tensor(sa, dextents<int32_t, 2>(NK, NR0));
auto tB = tensor(sb, dextents<int32_t, 2>(NK, NR1));
device const T1 *ptrB = (device const T1 *)(srcB + args.nb12*i12 + args.nb13*i13);
const int strideB = args.nb11/sizeof(T1);
matmul2d<
matmul2d_descriptor(NR1, NR0, NK, false, true, false,
matmul2d_descriptor::mode::multiply_accumulate),
execution_simdgroups<4>> mm;
auto cT = mm.template get_destination_cooperative_tensor<decltype(tB), decltype(tA), float>();
#pragma unroll
for (uint16_t i = 0; i < cT.get_capacity(); ++i) {
if (cT.is_valid_element(i)) {
cT[i] = 0.0f;
}
}
for (int loop_k = 0; loop_k < K; loop_k += NK) {
for (int work = tiitg; work < NR0*NL; work += NUM_THREADS) {
const int row = work/NL;
const int k_chunk = work%NL;
const int k_pos = loop_k + k_chunk*16;
const short k_base = k_chunk*16;
if (!FC_mul_mm_bc_out || r0 + row < M) {
if (is_same<T0_4x4, block_q>::value && FC_mul_mm_bc_inp) {
device const T0 *row_ptr = (device const T0 *)(srcA + args.nb01*(r0 + row) + offset0);
FOR_UNROLL (short i = 0; i < 16; i++) {
sa[row*NK + k_base + i] = (k_pos + i < K) ? (SA)row_ptr[k_pos + i] : (SA)0;
}
} else {
const int block_idx = k_pos/(16*nl);
const short il = (k_pos/16)%nl;
device const block_q *row_ptr = (device const block_q *)(srcA + args.nb01*(r0 + row) + offset0);
SA_4x4 temp_a;
dequantize_func(row_ptr + block_idx, il, temp_a);
FOR_UNROLL (short i = 0; i < 16; i++) {
sa[row*NK + k_base + i] = (k_pos + i < K) ? temp_a[i/4][i%4] : (SA)0;
}
}
} else {
FOR_UNROLL (short i = 0; i < 16; i++) {
sa[row*NK + k_base + i] = (SA)0;
}
}
}
for (int work = tiitg; work < NK*NR1; work += NUM_THREADS) {
const int col = work/NK;
const int k = work%NK;
if ((!FC_mul_mm_bc_out && !FC_mul_mm_bc_inp) ||
(r1 + col < N && loop_k + k < K)) {
sb[col*NK + k] = (SA)ptrB[(uint64_t)(r1 + col)*strideB + loop_k + k];
} else {
sb[col*NK + k] = (SA)0;
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
auto mA = tA.slice(0, 0);
auto mB = tB.slice(0, 0);
mm.run(mB, mA, cT);
threadgroup_barrier(mem_flags::mem_threadgroup);
}
device float *dst_batch = (device float *)dst + im*N*M;
if (!FC_mul_mm_bc_out) {
device float *dst_tile = dst_batch + r0 + (uint64_t)r1*M;
auto tD = tensor(dst_tile, dextents<int32_t, 2>(NR0, NR1), array<int, 2>({1, M}));
cT.store(tD);
} else {
auto tD = tensor(dst_batch, dextents<int32_t, 2>(M, N), array<int, 2>({1, M}));
auto mD = tD.slice(r0, r1);
cT.store(mD);
}
}
typedef decltype(kernel_mul_mm_mpp<64, 32, half, half4x4, float4x4, 1, dequantize_f32, float, float4x4, float>) mul_mm_mpp_t;
template [[host_name("kernel_mul_mm_f16_f32_mpp")]] kernel mul_mm_mpp_t kernel_mul_mm_mpp<64, 32, half, half4x4, half4x4, 1, dequantize_f16, half, half4x4, float>;
// Retained Metal4/TensorOps dense prefill kernel. The legacy MPP prototype
// staged both operands in threadgroup memory; this version stages only the
// model weight tile and lets MPP read the dense RHS activation matrix directly
@@ -2144,242 +2450,6 @@ kernel void kernel_mul_mm_f16_f32_scaled(
}
}
kernel void kernel_mul_mm_f16_f32_pair(
constant ds4_metal_args_mul_mm & args,
device const char * src0_a,
device const char * src0_b,
device const char * src1,
device char * dst_a,
device char * dst_b,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiitg[[thread_index_in_threadgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
threadgroup half * sa_a = (threadgroup half *)(shmem);
threadgroup half * sa_b = (threadgroup half *)(shmem + 4096);
threadgroup half * sb = (threadgroup half *)(shmem + 8192);
constexpr int NR0 = 64;
constexpr int NR1 = 32;
constexpr int NK = 32;
constexpr int NL0 = NK/16;
constexpr int NL1 = NK/8;
const int im = tgpig.z;
const int r0 = tgpig.y*NR0;
const int r1 = tgpig.x*NR1;
const short nr0 = (args.ne0 - r0 < NR0) ? (args.ne0 - r0) : NR0;
const short nr1 = (args.ne1 - r1 < NR1) ? (args.ne1 - r1) : NR1;
const short lr0 = ((short)tiitg/NL0) < nr0 ? ((short)tiitg/NL0) : nr0 - 1;
const short lr1 = ((short)tiitg/NL1) < nr1 ? ((short)tiitg/NL1) : nr1 - 1;
const short il0 = (tiitg % NL0);
short il = il0;
const int i12 = im%args.ne12;
const int i13 = im/args.ne12;
const uint64_t offset0 = (i12/args.r2)*args.nb02 + (i13/args.r3)*args.nb03;
const short offset1 = il0;
device const half4x4 * xa = (device const half4x4 *)(src0_a + args.nb01*(r0 + lr0) + offset0) + offset1;
device const half4x4 * xb = (device const half4x4 *)(src0_b + args.nb01*(r0 + lr0) + offset0) + offset1;
const short iy = 8*(tiitg % NL1);
device const float * y = (device const float *)(src1
+ args.nb13*i13
+ args.nb12*i12
+ args.nb11*(r1 + lr1)
+ args.nb10*iy);
simdgroup_half8x8 ma[4];
simdgroup_half8x8 mb[2];
simdgroup_float8x8 mc_a[8];
simdgroup_float8x8 mc_b[8];
for (short i = 0; i < 8; i++) {
mc_a[i] = make_filled_simdgroup_matrix<float, 8>(0.f);
mc_b[i] = make_filled_simdgroup_matrix<float, 8>(0.f);
}
for (int loop_k = 0; loop_k < args.ne00; loop_k += NK) {
half4x4 temp_a;
half4x4 temp_b;
dequantize_f16(xa, il, temp_a);
dequantize_f16(xb, il, temp_b);
threadgroup_barrier(mem_flags::mem_threadgroup);
FOR_UNROLL (short i = 0; i < 16; i++) {
const short sx = 2*il0 + i/8;
const short sy = (tiitg/NL0)/8;
const short lx = (tiitg/NL0)%8;
const short ly = i%8;
const short ib = 8*sx + sy;
*(sa_a + 64*ib + 8*ly + lx) = temp_a[i/4][i%4];
*(sa_b + 64*ib + 8*ly + lx) = temp_b[i/4][i%4];
}
if (FC_mul_mm_bc_inp) {
for (short i = 0; i < 8; ++i) {
const short sx = (tiitg%NL1);
const short sy = (tiitg/NL1)/8;
const short lx = i;
const short ly = (tiitg/NL1)%8;
const short ib = 4*sx + sy;
*(sb + 64*ib + 8*ly + lx) = loop_k + iy + i < args.ne00 ? (half) *((device float *) y + i) : 0;
}
} else {
const short sx = (tiitg%NL1);
const short sy = (tiitg/NL1)/8;
const short ly = (tiitg/NL1)%8;
const short ib = 4*sx + sy;
*(threadgroup half2x4 *)(sb + 64*ib + 8*ly) = (half2x4)(*((device float2x4 *) y));
}
il = (il + 2 < 1) ? il + 2 : il % 2;
xa = (il < 2) ? xa + 2 : xa;
xb = (il < 2) ? xb + 2 : xb;
y += NK;
threadgroup_barrier(mem_flags::mem_threadgroup);
threadgroup const half * lsma_a = (sa_a + 4*64*(sgitg%2));
threadgroup const half * lsma_b = (sa_b + 4*64*(sgitg%2));
threadgroup const half * lsmb = (sb + 2*64*(sgitg/2));
FOR_UNROLL (short ik = 0; ik < NK/8; ik++) {
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 2; i++) {
simdgroup_load(mb[i], lsmb + 64*i, 8, 0, false);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 4; i++) {
simdgroup_load(ma[i], lsma_a + 64*i, 8, 0, false);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 8; i++) {
simdgroup_multiply_accumulate(mc_a[i], mb[i/4], ma[i%4], mc_a[i]);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 4; i++) {
simdgroup_load(ma[i], lsma_b + 64*i, 8, 0, false);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 8; i++) {
simdgroup_multiply_accumulate(mc_b[i], mb[i/4], ma[i%4], mc_b[i]);
}
lsma_a += 8*64;
lsma_b += 8*64;
lsmb += 4*64;
}
}
if (!FC_mul_mm_bc_out || (r0 + NR0 <= args.ne0 && r1 + NR1 <= args.ne1)) {
device float * C_a = (device float *) dst_a +
(r0 + 32*(sgitg & 1)) +
(r1 + 16*(sgitg >> 1)) * args.ne0 + im*args.ne1*args.ne0;
device float * C_b = (device float *) dst_b +
(r0 + 32*(sgitg & 1)) +
(r1 + 16*(sgitg >> 1)) * args.ne0 + im*args.ne1*args.ne0;
for (short i = 0; i < 8; i++) {
simdgroup_store(mc_a[i], C_a + 8*(i%4) + 8*args.ne0*(i/4), args.ne0, 0, false);
simdgroup_store(mc_b[i], C_b + 8*(i%4) + 8*args.ne0*(i/4), args.ne0, 0, false);
}
} else {
threadgroup_barrier(mem_flags::mem_threadgroup);
threadgroup float * temp_str = (threadgroup float *) shmem;
for (short i = 0; i < 8; i++) {
simdgroup_store(mc_a[i],
temp_str + 32*(sgitg&1) + (16*(sgitg >> 1))*NR0 + 8*(i%4) + 8*NR0*(i/4),
NR0,
0,
false);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
if (sgitg == 0) {
for (int j = tiitg; j < nr1; j += NR1) {
device float * D = (device float *) dst_a + r0 + (r1 + j)*args.ne0 + im*args.ne1*args.ne0;
device float4 * D4 = (device float4 *) D;
threadgroup float * C = temp_str + (j*NR0);
threadgroup float4 * C4 = (threadgroup float4 *) C;
int i = 0;
for (; i < nr0/4; i++) {
*(D4 + i) = *(C4 + i);
}
i *= 4;
for (; i < nr0; i++) {
*(D + i) = *(C + i);
}
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
for (short i = 0; i < 8; i++) {
simdgroup_store(mc_b[i],
temp_str + 32*(sgitg&1) + (16*(sgitg >> 1))*NR0 + 8*(i%4) + 8*NR0*(i/4),
NR0,
0,
false);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
if (sgitg == 0) {
for (int j = tiitg; j < nr1; j += NR1) {
device float * D = (device float *) dst_b + r0 + (r1 + j)*args.ne0 + im*args.ne1*args.ne0;
device float4 * D4 = (device float4 *) D;
threadgroup float * C = temp_str + (j*NR0);
threadgroup float4 * C4 = (threadgroup float4 *) C;
int i = 0;
for (; i < nr0/4; i++) {
*(D4 + i) = *(C4 + i);
}
i *= 4;
for (; i < nr0; i++) {
*(D + i) = *(C + i);
}
}
}
}
}
typedef decltype(kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, float4x4, 1, dequantize_f32, float, float4x4, float, float2x4>) mul_mm_t;
// Host-visible prefill matmul variants for F16 and Q8_0 weights.

View File

@@ -21,20 +21,6 @@ struct ds4_metal_args_dsv4_hc_weighted_sum {
uint64_t nb1;
};
struct ds4_metal_args_dsv4_hc_weighted_sum_norm {
int64_t n_embd;
int64_t n_hc;
int64_t n_tokens;
uint64_t nb_x0;
uint64_t nb_x1;
uint64_t nb_x2;
uint64_t nb_w0;
uint64_t nb_w1;
uint64_t nb0;
uint64_t nb1;
uint64_t nb_norm1;
float norm_eps;
};
struct ds4_metal_args_dsv4_output_hc_weights4 {
float post_scale;
@@ -411,6 +397,68 @@ kernel void kernel_dsv4_hc_split_weighted_sum(
// kernel_dsv4_hc_split_weighted_sum, stores the HC-pre row for diagnostics, and
// reuses the just-collapsed values from threadgroup memory for the RMSNorm
// reduction.
static __attribute__((always_inline)) inline void ds4_hc_comb_weights4_exact(
constant ds4_metal_args_dsv4_hc_split_weighted_sum_norm & args,
device volatile const float *mix,
device const float *scale,
device const float *base,
device float *out) {
const float epsv = args.eps;
const float comb_scale = scale[2];
float4 r0 =
*((device volatile const float4 *)(mix + 8)) * comb_scale +
*((device const float4 *)(base + 8));
float4 r1 =
*((device volatile const float4 *)(mix + 12)) * comb_scale +
*((device const float4 *)(base + 12));
float4 r2 =
*((device volatile const float4 *)(mix + 16)) * comb_scale +
*((device const float4 *)(base + 16));
float4 r3 =
*((device volatile const float4 *)(mix + 20)) * comb_scale +
*((device const float4 *)(base + 20));
const float m0 = max(max(r0.x, r0.y), max(r0.z, r0.w));
const float m1 = max(max(r1.x, r1.y), max(r1.z, r1.w));
const float m2 = max(max(r2.x, r2.y), max(r2.z, r2.w));
const float m3 = max(max(r3.x, r3.y), max(r3.z, r3.w));
r0 = exp(r0 - m0);
r1 = exp(r1 - m1);
r2 = exp(r2 - m2);
r3 = exp(r3 - m3);
r0 = r0 * (1.0f / (r0.x + r0.y + r0.z + r0.w)) + epsv;
r1 = r1 * (1.0f / (r1.x + r1.y + r1.z + r1.w)) + epsv;
r2 = r2 * (1.0f / (r2.x + r2.y + r2.z + r2.w)) + epsv;
r3 = r3 * (1.0f / (r3.x + r3.y + r3.z + r3.w)) + epsv;
float4 col_inv = 1.0f / (r0 + r1 + r2 + r3 + epsv);
r0 *= col_inv;
r1 *= col_inv;
r2 *= col_inv;
r3 *= col_inv;
for (int iter = 1; iter < args.sinkhorn_iters; ++iter) {
r0 *= 1.0f / (r0.x + r0.y + r0.z + r0.w + epsv);
r1 *= 1.0f / (r1.x + r1.y + r1.z + r1.w + epsv);
r2 *= 1.0f / (r2.x + r2.y + r2.z + r2.w + epsv);
r3 *= 1.0f / (r3.x + r3.y + r3.z + r3.w + epsv);
col_inv = 1.0f / (r0 + r1 + r2 + r3 + epsv);
r0 *= col_inv;
r1 *= col_inv;
r2 *= col_inv;
r3 *= col_inv;
}
*((device float4 *)(out + 8)) = r0;
*((device float4 *)(out + 12)) = r1;
*((device float4 *)(out + 16)) = r2;
*((device float4 *)(out + 20)) = r3;
}
kernel void kernel_dsv4_hc_split_weighted_sum_norm4(
constant ds4_metal_args_dsv4_hc_split_weighted_sum_norm & args,
device const char * mixes,
@@ -519,7 +567,6 @@ kernel void kernel_dsv4_hc_split_weighted_sum_norm4(
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float sumf = 0.0f;
for (uint i = tid; i < n4; i += ntg) {
device const float4 *x0 = (device const float4 *)(x + 0 * args.nb_x1 + (uint64_t)row * args.nb_x2);
@@ -884,6 +931,119 @@ kernel void kernel_dsv4_q8_hc_expand4_q8_0(
}
}
kernel void kernel_dsv4_q8_hc_expand4_q8_0_vec_hc(
constant ds4_metal_args_mul_mv & mv,
constant ds4_metal_args_dsv4_hc_expand & hc,
device const char * weight,
device const char * input,
device char * block_out,
device const char * residual,
device const char * post,
device const char * comb,
device char * dst,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
if (hc.n_hc != 4 || hc.n_tokens != 1) {
return;
}
const short NSG = FC_mul_mv_nsg;
constexpr short NW = N_SIMDWIDTH;
constexpr short NQ = 8;
constexpr short NR0 = N_R0_Q8_0;
const int nb = mv.ne00 / QK8_0;
const int row0 = tgpig.x * NR0;
const short ix = tiisg / (NW / NQ);
const short il = tiisg % (NW / NQ);
const int ib0 = sgitg * NQ + ix;
device const float *y = (device const float *)(input);
device const float *yb = y + ib0 * QK8_0 + il * NQ;
device const block_q8_0 *ax[NR0];
FOR_UNROLL(short row = 0; row < NR0; ++row) {
const uint64_t off0 = (uint64_t)(row0 + row) * mv.nb01;
ax[row] = (device const block_q8_0 *)(weight + off0);
}
float sumf[NR0] = { 0.0f };
float yl[NQ];
for (int ib = ib0; ib < nb; ib += NSG * NQ) {
FOR_UNROLL(short i = 0; i < NQ; ++i) {
yl[i] = yb[i];
}
FOR_UNROLL(short row = 0; row < NR0; ++row) {
device const int8_t *qs = ax[row][ib].qs + il * NQ;
float sumq = 0.0f;
FOR_UNROLL(short i = 0; i < NQ; ++i) {
sumq += qs[i] * yl[i];
}
sumf[row] += sumq * ax[row][ib].d;
}
yb += NSG * NQ * QK8_0;
}
threadgroup float *shmem_f32[NR0];
FOR_UNROLL(short row = 0; row < NR0; ++row) {
shmem_f32[row] = (threadgroup float *)shmem + NW * row;
if (sgitg == 0) {
shmem_f32[row][tiisg] = 0.0f;
}
sumf[row] = simd_sum(sumf[row]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
FOR_UNROLL(short row = 0; row < NR0; ++row) {
if (tiisg == 0) {
shmem_f32[row][sgitg] = sumf[row];
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
FOR_UNROLL(short row = 0; row < NR0; ++row) {
const int d = row0 + row;
if (d >= mv.ne01) {
continue;
}
const float block_v = simd_sum(shmem_f32[row][tiisg]);
if (tiisg == 0 && sgitg == 0) {
*((device float *)(block_out + (uint64_t)d * sizeof(float))) = block_v;
const float r0 = *((device const float *)(residual + (uint64_t)d * hc.nb_res0 + 0 * hc.nb_res1));
const float r1 = *((device const float *)(residual + (uint64_t)d * hc.nb_res0 + 1 * hc.nb_res1));
const float r2 = *((device const float *)(residual + (uint64_t)d * hc.nb_res0 + 2 * hc.nb_res1));
const float r3 = *((device const float *)(residual + (uint64_t)d * hc.nb_res0 + 3 * hc.nb_res1));
const float4 post4 = *((device const float4 *)post);
const float4 comb0 = *((device const float4 *)(comb + 0 * hc.nb_comb1));
const float4 comb1 = *((device const float4 *)(comb + 1 * hc.nb_comb1));
const float4 comb2 = *((device const float4 *)(comb + 2 * hc.nb_comb1));
const float4 comb3 = *((device const float4 *)(comb + 3 * hc.nb_comb1));
float4 acc = block_v * post4;
acc += comb0 * r0;
acc += comb1 * r1;
acc += comb2 * r2;
acc += comb3 * r3;
FOR_UNROLL (short dst_hc = 0; dst_hc < 4; ++dst_hc) {
*((device float *)(dst + (uint64_t)d * hc.nb0 +
(uint64_t)dst_hc * hc.nb1)) = acc[dst_hc];
}
}
}
}
// Reduces HC channels to a normal embedding row with the learned pre weights.
// This is the input adapter before the attention block and before the FFN block.
kernel void kernel_dsv4_hc_weighted_sum(
@@ -910,76 +1070,6 @@ kernel void kernel_dsv4_hc_weighted_sum(
*((device float *) (dst + d*args.nb0 + t*args.nb1)) = acc;
}
// The one-row output head immediately applies a learned RMSNorm after reducing
// its four HC streams. Preserve the standalone scalar HC accumulation, write
// the collapsed row for diagnostics, then reload its F32 values from
// threadgroup memory using the standalone RMSNorm's float4 reduction mapping.
kernel void kernel_dsv4_hc_weighted_sum_norm4(
constant ds4_metal_args_dsv4_hc_weighted_sum_norm & args,
device const char * x,
device const char * weights,
device char * dst,
device const char * norm_weight,
device char * norm_dst,
threadgroup float * shared [[threadgroup(0)]],
ushort tid [[thread_position_in_threadgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort ntg [[threads_per_threadgroup]]) {
if (args.n_tokens != 1 || args.n_hc != 4 ||
args.n_embd <= 0 || (args.n_embd & 3) != 0) {
return;
}
const uint n_embd = uint(args.n_embd);
const uint n4 = n_embd >> 2;
threadgroup float *row_shmem = shared;
threadgroup float *sum_shmem = shared + n_embd;
if (sgitg == 0) {
sum_shmem[tiisg] = 0.0f;
}
for (uint d = tid; d < n_embd; d += ntg) {
float acc = 0.0f;
for (int64_t h = 0; h < args.n_hc; ++h) {
const float xv = *((device const float *)(
x + (uint64_t)d*args.nb_x0 + (uint64_t)h*args.nb_x1));
const float wv = *((device const float *)(
weights + (uint64_t)h*args.nb_w0));
acc += xv * wv;
}
row_shmem[d] = acc;
*((device float *)(dst + (uint64_t)d*args.nb0)) = acc;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
threadgroup const float4 *row4 =
(threadgroup const float4 *)row_shmem;
float sumf = 0.0f;
for (uint i = tid; i < n4; i += ntg) {
sumf += dot(row4[i], row4[i]);
}
sumf = simd_sum(sumf);
threadgroup_barrier(mem_flags::mem_threadgroup);
if (tiisg == 0) {
sum_shmem[sgitg] = sumf;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
sumf = sum_shmem[tiisg];
sumf = simd_sum(sumf);
const float mean = sumf/args.n_embd;
const float scale = 1.0f/sqrt(mean + args.norm_eps);
device const float4 *w4 = (device const float4 *)norm_weight;
device float4 *norm4 = (device float4 *)norm_dst;
for (uint i = tid; i < n4; i += ntg) {
norm4[i] = (row4[i]*scale)*w4[i];
}
}
// The one-row HC=4 output head historically materializes four device-F32
// stages across separate launches. Collapse those launches into one tiny
// two-thread group while preserving the scalar/vector lane mapping and every
@@ -1015,3 +1105,440 @@ kernel void kernel_dsv4_output_hc_weights4(
args.post_scale * x + args.eps;
}
}
struct ds4_metal_args_hc_norm_mix {
int32_t n;
int32_t out_dim;
float eps;
};
// Fused unweighted RMSNorm + F16 HC-mix projection for DS4 decode HC-pre.
// The standalone decode path runs kernel_rms_norm_f32_4 over the flattened
// 4*embd HC row (1024 threads, one threadgroup) and then
// kernel_mul_mv_f16_f32_4 (nsg=8, nr0=2) over the normalized row. Both
// stages are reproduced bit-exactly in one dispatch: every threadgroup
// redundantly recomputes the norm partials with the original 1024-thread
// mapping (each real lane covers one virtual thread of each 256-thread
// slice, preserving every simd_sum tree), and the matvec keeps the original
// per-row accumulation order with y = x*scale computed on the fly, which
// rounds identically to the materialized normalized row. The host wrapper
// gates this to n == 16384 && out_dim == 24, where the virtual-thread count
// is exactly 1024 and the mv tail loop is empty.
kernel void kernel_dsv4_hc_rms_norm_mix_f16(
constant ds4_metal_args_hc_norm_mix & args,
device const char * x,
device const char * weight,
device char * dst,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constexpr short NSG = 8; // ds4_gpu_make_plain_mv_dispatch(16384)
constexpr short NW = N_SIMDWIDTH;
constexpr short NR0 = 2; // plain mv nr0
constexpr short NB = 32;
constexpr short NF = 16;
constexpr short NF4 = NF/4;
constexpr uint VTHREADS = 1024u; // rms norm threads at n == 16384
constexpr short VSLICES = VTHREADS/(NSG*NW); // virtual 256-thread slices
const uint n = (uint)args.n;
const uint n4 = n >> 2;
device const float4 *x4 = (device const float4 *)x;
threadgroup float *norm_shmem = (threadgroup float *)shmem; // NW slots
threadgroup float *mv_shmem = (threadgroup float *)shmem + NW; // NW*NR0 slots
// Phase A: exact replica of kernel_rms_norm_f32_4's reduction tree with
// the 1024 virtual threads folded onto this threadgroup's 8 simdgroups.
for (short v = 0; v < VSLICES; ++v) {
const uint vt = (uint)(sgitg + NSG*v)*NW + tiisg;
float sumf = 0.0f;
for (uint i00 = vt; i00 < n4; i00 += VTHREADS) {
sumf += dot(x4[i00], x4[i00]);
}
sumf = simd_sum(sumf);
if (tiisg == 0) {
norm_shmem[sgitg + NSG*v] = sumf;
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float total = norm_shmem[tiisg];
total = simd_sum(total);
const float mean = total/(float)args.n;
const float scale = 1.0f/sqrt(mean + args.eps);
// Phase B: exact replica of kernel_mul_mv_f16_f32_4 (nsg=8, nr0=2) with
// the normalized operand recomputed as x*scale instead of reloaded.
const int nb = args.n/NB;
const int r0 = tgpig.x*NR0;
device const half4 * ax4[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
ax4[row] = (device const half4 *)
(weight + (uint64_t)(r0 + row)*(uint64_t)n*sizeof(half));
}
float sumf_mv[NR0] = { 0.f };
const short ix = tiisg/(NW/NF);
const short il = tiisg%(NW/NF);
const int ib0 = sgitg*NF + ix;
for (int ib = ib0; ib < nb; ib += NSG*NF) {
float4 yl4[NF4];
FOR_UNROLL (short i = 0; i < NF4; ++i) {
yl4[i] = x4[(ib*NB + il*NF)/4 + i]*scale;
}
FOR_UNROLL (short row = 0; row < NR0; row++) {
device const half4 * xb4 = ax4[row] + (ib*NB + il*NF)/4;
float sumq = 0.f;
FOR_UNROLL (short i = 0; i < NF4; ++i) {
sumq += dot(float4(xb4[i]), yl4[i]);
}
sumf_mv[row] += sumq;
}
}
// n == 16384 makes the scalar tail loop of the original empty.
device float * dst_f32 = (device float *) dst;
helper_mv_reduce_and_write<NR0>(dst_f32, sumf_mv, r0, args.out_dim,
tiisg, sgitg, (threadgroup char *)mv_shmem);
}
// M5 specialization: pack two exact NR0=2 HC-mix producer groups into one
// 512-thread group. Two independent eight-simdgroup clusters retain the
// matvec reductions while the exact RMS scale is redundantly formed six,
// rather than twelve, times.
kernel void kernel_dsv4_hc_rms_norm_mix_f16_cluster2(
constant ds4_metal_args_hc_norm_mix & args,
device const char * x,
device const char * weight,
device char * dst,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constexpr short NSG_CLUSTER = 8;
constexpr short NCLUSTER = 2;
constexpr short NSG_TOTAL = NSG_CLUSTER * NCLUSTER;
constexpr short NW = N_SIMDWIDTH;
constexpr short NR0 = 2;
constexpr short NB = 32;
constexpr short NF = 16;
constexpr short NF4 = NF/4;
constexpr uint VTHREADS = 1024u;
constexpr short VSLICES = VTHREADS/(NSG_TOTAL*NW);
const uint n = (uint)args.n;
const uint n4 = n >> 2;
device const float4 *x4 = (device const float4 *)x;
threadgroup float *norm_shmem = (threadgroup float *)shmem;
threadgroup float *mv_shmem = norm_shmem + NW;
// Exact 1024-virtual-thread RMS reduction, now folded two ways over
// the 16 physical simdgroups instead of four ways over eight.
for (short v = 0; v < VSLICES; ++v) {
const uint vt = (uint)(sgitg + NSG_TOTAL*v)*NW + tiisg;
float sumf = 0.0f;
for (uint i00 = vt; i00 < n4; i00 += VTHREADS) {
sumf += dot(x4[i00], x4[i00]);
}
sumf = simd_sum(sumf);
if (tiisg == 0) {
norm_shmem[sgitg + NSG_TOTAL*v] = sumf;
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float total = norm_shmem[tiisg];
total = simd_sum(total);
const float mean = total/(float)args.n;
const float scale = 1.0f/sqrt(mean + args.eps);
// Two independent eight-simdgroup clusters reproduce two original
// NR0=2 matvec threadgroups inside this 512-thread threadgroup.
const short cluster = sgitg / NSG_CLUSTER;
const short local_sg = sgitg - cluster*NSG_CLUSTER;
const int nb = args.n/NB;
const int r0 = (int)tgpig.x*(NCLUSTER*NR0) + cluster*NR0;
device const half4 *ax4[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
ax4[row] = (device const half4 *)
(weight + (uint64_t)(r0 + row)*(uint64_t)n*sizeof(half));
}
float sumf_mv[NR0] = { 0.f };
const short ix = tiisg/(NW/NF);
const short il = tiisg%(NW/NF);
const int ib0 = local_sg*NF + ix;
for (int ib = ib0; ib < nb; ib += NSG_CLUSTER*NF) {
float4 yl4[NF4];
FOR_UNROLL (short i = 0; i < NF4; ++i) {
yl4[i] = x4[(ib*NB + il*NF)/4 + i]*scale;
}
FOR_UNROLL (short row = 0; row < NR0; ++row) {
device const half4 *xb4 = ax4[row] + (ib*NB + il*NF)/4;
float sumq = 0.f;
FOR_UNROLL (short i = 0; i < NF4; ++i) {
sumq += dot(float4(xb4[i]), yl4[i]);
}
sumf_mv[row] += sumq;
}
}
threadgroup float *cluster_shmem[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
cluster_shmem[row] = mv_shmem +
((uint)cluster*NR0 + row)*NW;
if (local_sg == 0) {
cluster_shmem[row][tiisg] = 0.0f;
}
sumf_mv[row] = simd_sum(sumf_mv[row]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
FOR_UNROLL (short row = 0; row < NR0; ++row) {
if (tiisg == 0) {
cluster_shmem[row][local_sg] = sumf_mv[row];
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
device float *mixes_f32 = (device float *)dst;
FOR_UNROLL (short row = 0; row < NR0; ++row) {
const float tot = simd_sum(cluster_shmem[row][tiisg]);
if (tiisg == 0 && local_sg == 0 && r0 + row < args.out_dim) {
mixes_f32[r0 + row] = tot;
}
}
}
kernel void kernel_dsv4_hc_rms_norm_mix_f16_cluster2_pre_norm(
constant ds4_metal_args_hc_norm_mix & args,
constant ds4_metal_args_dsv4_hc_split_weighted_sum_norm & split_args,
device const char * x,
device const char * weight,
device char * dst,
device const float * hc_scale,
device const float * hc_base,
device char * split,
device char * collapse_dst,
device const char * norm_weight,
device char * norm_dst,
device atomic_uint * completion,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constexpr short NSG_CLUSTER = 8;
constexpr short NCLUSTER = 2;
constexpr short NSG_TOTAL = NSG_CLUSTER * NCLUSTER;
constexpr short NW = N_SIMDWIDTH;
constexpr short NR0 = 2;
constexpr short NB = 32;
constexpr short NF = 16;
constexpr short NF4 = NF/4;
constexpr uint VTHREADS = 1024u;
constexpr short VSLICES = VTHREADS/(NSG_TOTAL*NW);
const uint n = (uint)args.n;
const uint n4 = n >> 2;
device const float4 *x4 = (device const float4 *)x;
threadgroup float *norm_shmem = (threadgroup float *)shmem;
threadgroup float *mv_shmem = norm_shmem + NW;
// Exact 1024-virtual-thread RMS reduction, now folded two ways over
// the 16 physical simdgroups instead of four ways over eight.
for (short v = 0; v < VSLICES; ++v) {
const uint vt = (uint)(sgitg + NSG_TOTAL*v)*NW + tiisg;
float sumf = 0.0f;
for (uint i00 = vt; i00 < n4; i00 += VTHREADS) {
sumf += dot(x4[i00], x4[i00]);
}
sumf = simd_sum(sumf);
if (tiisg == 0) {
norm_shmem[sgitg + NSG_TOTAL*v] = sumf;
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float total = norm_shmem[tiisg];
total = simd_sum(total);
const float mean = total/(float)args.n;
const float scale = 1.0f/sqrt(mean + args.eps);
// Two independent eight-simdgroup clusters reproduce two original
// NR0=2 matvec threadgroups inside this 512-thread threadgroup.
const short cluster = sgitg / NSG_CLUSTER;
const short local_sg = sgitg - cluster*NSG_CLUSTER;
const int nb = args.n/NB;
const int r0 = (int)tgpig.x*(NCLUSTER*NR0) + cluster*NR0;
device const half4 *ax4[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
ax4[row] = (device const half4 *)
(weight + (uint64_t)(r0 + row)*(uint64_t)n*sizeof(half));
}
float sumf_mv[NR0] = { 0.f };
const short ix = tiisg/(NW/NF);
const short il = tiisg%(NW/NF);
const int ib0 = local_sg*NF + ix;
for (int ib = ib0; ib < nb; ib += NSG_CLUSTER*NF) {
float4 yl4[NF4];
FOR_UNROLL (short i = 0; i < NF4; ++i) {
yl4[i] = x4[(ib*NB + il*NF)/4 + i]*scale;
}
FOR_UNROLL (short row = 0; row < NR0; ++row) {
device const half4 *xb4 = ax4[row] + (ib*NB + il*NF)/4;
float sumq = 0.f;
FOR_UNROLL (short i = 0; i < NF4; ++i) {
sumq += dot(float4(xb4[i]), yl4[i]);
}
sumf_mv[row] += sumq;
}
}
threadgroup float *cluster_shmem[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
cluster_shmem[row] = mv_shmem +
((uint)cluster*NR0 + row)*NW;
if (local_sg == 0) {
cluster_shmem[row][tiisg] = 0.0f;
}
sumf_mv[row] = simd_sum(sumf_mv[row]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
FOR_UNROLL (short row = 0; row < NR0; ++row) {
if (tiisg == 0) {
cluster_shmem[row][local_sg] = sumf_mv[row];
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
device volatile float *mixes_f32 =
(device volatile float *)dst;
if (local_sg == 0) {
FOR_UNROLL (short row = 0; row < NR0; ++row) {
const float tot = simd_sum(cluster_shmem[row][tiisg]);
if (tiisg == 0 && r0 + row < args.out_dim) {
mixes_f32[r0 + row] = tot;
}
}
}
// The first producer group owns mix[0:4]. After materializing and
// reloading those values, fold the established 1024-thread HC collapse
// and RMS reduction over this group's 512 physical threads as two
// independent virtual slices. This retains the original 32-partial tree.
threadgroup_barrier(mem_flags::mem_device_and_threadgroup);
const uint tid = (uint)sgitg * (uint)NW + (uint)tiisg;
threadgroup float *pre_shmem = norm_shmem + 32u + 4u*NW;
threadgroup float *sum_shmem = pre_shmem + 4;
if (tgpig.x == 0) {
device float *out = (device float *)split;
if (tid == 0) {
const float4 pre_z =
*((device volatile const float4 *)mixes_f32) * hc_scale[0] +
*((device const float4 *)hc_base);
const float4 pre =
1.0f / (1.0f + exp(-pre_z)) + split_args.eps;
*((device float4 *)out) = pre;
pre_shmem[0] = pre.x;
pre_shmem[1] = pre.y;
pre_shmem[2] = pre.z;
pre_shmem[3] = pre.w;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
const uint n4_collapse = uint(split_args.n_embd) >> 2;
const uint i0 = tid;
const uint i1 = tid + 512u;
device const float4 *x0 = (device const float4 *)(
x + 0 * split_args.nb_x1);
device const float4 *x1 = (device const float4 *)(
x + 1 * split_args.nb_x1);
device const float4 *x2 = (device const float4 *)(
x + 2 * split_args.nb_x1);
device const float4 *x3 = (device const float4 *)(
x + 3 * split_args.nb_x1);
float4 v0 = 0.0f;
v0 += x0[i0] * pre_shmem[0];
v0 += x1[i0] * pre_shmem[1];
v0 += x2[i0] * pre_shmem[2];
v0 += x3[i0] * pre_shmem[3];
float sum0 = simd_sum(dot(v0, v0));
float4 v1 = 0.0f;
if (i1 < n4_collapse) {
v1 += x0[i1] * pre_shmem[0];
v1 += x1[i1] * pre_shmem[1];
v1 += x2[i1] * pre_shmem[2];
v1 += x3[i1] * pre_shmem[3];
}
float sum1 = simd_sum(dot(v1, v1));
if (tiisg == 0) {
sum_shmem[sgitg] = sum0;
sum_shmem[sgitg + 16] = sum1;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float sumf = sum_shmem[tiisg];
sumf = simd_sum(sumf);
const float norm_arg =
sumf / float(split_args.n_embd) + split_args.norm_eps;
const float norm_scale = rsqrt(norm_arg);
device float4 *dst4 = (device float4 *)collapse_dst;
device const float4 *w4 = (device const float4 *)norm_weight;
device float4 *norm4 = (device float4 *)norm_dst;
dst4[i0] = v0;
norm4[i0] = (v0 * norm_scale) * w4[i0];
if (i1 < n4_collapse) {
dst4[i1] = v1;
norm4[i1] = (v1 * norm_scale) * w4[i1];
}
} else if (tgpig.x == 1 && tid == 0) {
device float *out = (device float *)split;
const float4 post_z =
*((device volatile const float4 *)(mixes_f32 + 4)) * hc_scale[1] +
*((device const float4 *)(hc_base + 4));
*((device float4 *)(out + 4)) = 2.0f / (1.0f + exp(-post_z));
}
// Groups 2..5 own exactly the comb range consumed by the
// continuation. Their four-way completion overlaps TG0's independent
// pre-collapse/RMS epilogue. Every writer crosses the uniform publish
// fence; only lane zero then participates in the completion protocol.
atomic_thread_fence(mem_flags::mem_device,
memory_order_seq_cst,
thread_scope_device);
if (tgpig.x < 2 || tid != 0) {
return;
}
const uint old = atomic_fetch_add_explicit(
completion, 1u, memory_order_relaxed);
if (old + 1u != 4u) {
return;
}
atomic_thread_fence(mem_flags::mem_device,
memory_order_seq_cst,
thread_scope_device);
ds4_hc_comb_weights4_exact(
split_args, mixes_f32, hc_scale, hc_base,
(device float *)split);
atomic_thread_fence(mem_flags::mem_device,
memory_order_seq_cst,
thread_scope_device);
atomic_store_explicit(completion, 0u, memory_order_relaxed);
}

View File

@@ -319,6 +319,228 @@ kernel void kernel_dsv4_compressor_pack_ratio4(
}
}
// Decode already holds the complete ratio-4 recurrent window in state layout:
// eight rows of two head_dim planes. Pack the previous plane from rows 0..3
// and the current plane from rows 4..7 directly into the transposed [head_dim,
// 8] layout consumed by the exact GGML softmax/multiply/sum sequence. KV and
// score move together; no arithmetic or reduction order changes.
kernel void kernel_dsv4_compressor_pack_ratio4_decode_ggml(
constant ds4_metal_args_dsv4_compressor_pack_ratio4 & args,
device const uint * state_kv,
device const uint * state_score,
device uint * packed_kv,
device uint * packed_score,
uint row [[threadgroup_position_in_grid]],
uint tid [[thread_index_in_threadgroup]]) {
if (row >= 8u || args.head_dim == 0u || args.n_threads == 0u) {
return;
}
const uint64_t state_row_stride = 2ull * args.head_dim;
const uint64_t src_plane = row >= 4u ? args.head_dim : 0u;
for (uint col = tid; col < args.head_dim; col += args.n_threads) {
const uint64_t src = (uint64_t)row * state_row_stride +
src_plane + col;
const uint64_t dst = (uint64_t)col * 8u + row;
packed_kv[dst] = state_kv[src];
packed_score[dst] = state_score[src];
}
}
// Exact decode specialization for the first two operations in GGML's
// softmax -> multiply -> sum_rows compressor reduction. The normalized
// softmax values are deliberately materialized in device memory and reloaded
// after a device barrier before the in-place product, preserving the dispatch
// boundary's float store/load semantics. The final sum remains the standalone
// eight-thread sum_rows kernel: changing a 32-thread group into the original
// eight-thread reduction inside this kernel would make its threadgroup
// barriers non-uniform or alter simd_sum's active-lane topology.
kernel void kernel_dsv4_compressor_exact_softmax_product_ratio4(
constant ds4_metal_args_dsv4_compressor_pack_ratio4 & args,
device const float * packed_kv,
device const float * packed_score,
device float * softmax,
device float * product,
threadgroup float * softmax_scratch [[threadgroup(0)]],
uint row [[threadgroup_position_in_grid]],
uint tid [[thread_position_in_threadgroup]]) {
if (row >= args.head_dim || args.n_comp != 1u ||
args.n_threads != 32u) {
return;
}
device const float4 * score4 =
(device const float4 *)(packed_score + (uint64_t)row * 8u);
device float4 * softmax4 =
(device float4 *)(softmax + (uint64_t)row * 8u);
const float scale = (float)args.replay;
const float zero = (float)(args.n_comp - 1u);
// Match kernel_soft_max_f32_4(width=8, nth=32) literally. Only lanes zero
// and one own float4s, while all 32 lanes participate in both reductions.
float4 lmax4 = -INFINITY;
for (int i00 = (int)tid; i00 < 2; i00 += 32) {
lmax4 = fmax(lmax4, score4[i00] * scale + (float4)zero);
}
const float lmax =
MAX(MAX(lmax4[0], lmax4[1]), MAX(lmax4[2], lmax4[3]));
const float max_val = simd_max(lmax);
float4 lsum4 = 0.0f;
for (int i00 = (int)tid; i00 < 2; i00 += 32) {
const float4 exp_score4 =
exp((score4[i00] * scale + (float4)zero) - max_val);
lsum4 += exp_score4;
softmax4[i00] = exp_score4;
}
const float lsum =
lsum4[0] + lsum4[1] + lsum4[2] + lsum4[3];
threadgroup_barrier(mem_flags::mem_none);
const float sum = simd_sum(lsum);
const float inv_sum = 1.0f / sum;
for (int i00 = (int)tid; i00 < 2; i00 += 32) {
softmax4[i00] *= inv_sum;
}
// Force the same normalized-softmax device store/reload boundary that the
// separate multiply dispatch observes.
threadgroup_barrier(mem_flags::mem_device);
device volatile const float * reloaded_softmax =
(device volatile const float *)(softmax + (uint64_t)row * 8u);
device const float * kv_row = packed_kv + (uint64_t)row * 8u;
device float * product_row = product + (uint64_t)row * 8u;
// Match kernel_bin_fuse_f32_f32_f32(width=8, nth=4): four lanes each
// process their low element followed by the element four positions later.
if (tid < 4u) {
for (uint i0 = tid; i0 < 8u; i0 += 4u) {
float value = kv_row[i0];
value *= reloaded_softmax[i0];
product_row[i0] = value;
}
}
// All 32 lanes reach the final device barrier. The following standalone
// sum_rows dispatch performs the required global reload and exact TG8
// two-stage simd_sum topology.
threadgroup_barrier(mem_flags::mem_device);
(void)softmax_scratch;
}
// Exact one-dispatch ratio-4 decode pool. This specializes the three-dispatch
// pack -> exact softmax/product -> sum_rows chain above without changing any
// floating-point operation or reduction topology. The normalized softmax and
// product are still materialized and volatile-reloaded through device memory.
// The two simd_sum calls in the final reduction execute under an eight-lane
// active mask, exactly matching kernel_sum_rows_f32_f32's original TG8.
kernel void kernel_dsv4_compressor_exact_pool_ratio4_decode_ggml(
constant ds4_metal_args_dsv4_compressor_pack_ratio4 & args,
device const float * state_kv,
device const float * state_score,
device float * softmax,
device float * product,
device float * dst,
threadgroup float * sum_scratch [[threadgroup(0)]],
uint col [[threadgroup_position_in_grid]],
uint tid [[thread_position_in_threadgroup]]) {
if (col >= args.head_dim || args.n_comp != 1u ||
args.n_threads != 32u) {
return;
}
const uint64_t state_row_stride = 2ull * args.head_dim;
const float scale = (float)args.replay;
const float zero = (float)(args.n_comp - 1u);
// Match the packed float4 ownership: lane 0 owns rows 0..3 and lane 1
// rows 4..7. The gather itself is an integer-addressed bit-preserving load.
float4 score_values = -INFINITY;
if (tid < 2u) {
const uint row0 = 4u * tid;
for (uint j = 0u; j < 4u; ++j) {
const uint row = row0 + j;
const uint64_t src = (uint64_t)row * state_row_stride +
(row >= 4u ? args.head_dim : 0u) + col;
score_values[j] = state_score[src];
}
}
const uint64_t scratch_base = (uint64_t)col * 8u;
device float4 * softmax4 =
(device float4 *)(softmax + scratch_base);
// Verbatim kernel_soft_max_f32_4(width=8, nth=32) arithmetic.
float4 lmax4 = -INFINITY;
for (int i00 = (int)tid; i00 < 2; i00 += 32) {
lmax4 = fmax(lmax4, score_values * scale + (float4)zero);
}
const float lmax =
MAX(MAX(lmax4[0], lmax4[1]), MAX(lmax4[2], lmax4[3]));
const float max_val = simd_max(lmax);
float4 lsum4 = 0.0f;
for (int i00 = (int)tid; i00 < 2; i00 += 32) {
const float4 exp_score4 =
exp((score_values * scale + (float4)zero) - max_val);
lsum4 += exp_score4;
softmax4[i00] = exp_score4;
}
const float lsum =
lsum4[0] + lsum4[1] + lsum4[2] + lsum4[3];
threadgroup_barrier(mem_flags::mem_none);
const float sum = simd_sum(lsum);
const float inv_sum = 1.0f / sum;
for (int i00 = (int)tid; i00 < 2; i00 += 32) {
softmax4[i00] *= inv_sum;
}
threadgroup_barrier(mem_flags::mem_device);
device volatile const float * reloaded_softmax =
(device volatile const float *)(softmax + scratch_base);
device float * product_row = product + scratch_base;
// Verbatim width=8, TG4 multiply ownership: low element, then +4.
if (tid < 4u) {
for (uint i0 = tid; i0 < 8u; i0 += 4u) {
const uint64_t src = (uint64_t)i0 * state_row_stride +
(i0 >= 4u ? args.head_dim : 0u) + col;
float value = state_kv[src];
value *= reloaded_softmax[i0];
product_row[i0] = value;
}
}
// Preserve the product dispatch's device store/reload boundary.
threadgroup_barrier(mem_flags::mem_device);
device volatile const float * reloaded_product =
(device volatile const float *)product_row;
// Reproduce kernel_sum_rows_f32_f32(width=8, TG8) literally. MSL defines
// simdgroup collectives over active lanes, so the branch recreates the
// original eight-lane partial SIMD group inside this 32-thread group.
sum_scratch[tid] = 0.0f;
float row_sum = 0.0f;
if (tid < 8u) {
row_sum += reloaded_product[tid];
row_sum = simd_sum(row_sum);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
if (tid == 0u) {
sum_scratch[0] = row_sum;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
if (tid < 8u) {
row_sum = sum_scratch[tid];
row_sum = simd_sum(row_sum);
if (tid == 0u) {
dst[col] = row_sum;
}
}
}
// Ratio-4 compression keeps two 4-row halves of recurrent state. After an
// emitted compressed row, the second half becomes the next window's previous
// half. The old encoder expressed this as four generic copies; this DS4-specific

View File

@@ -70,7 +70,7 @@ struct ds4_metal_args_dsv4_indexed_attention {
uint32_t window;
uint32_t ratio;
uint32_t comp_kv_f16;
uint32_t pad0;
uint32_t n_splits;
uint64_t q_token_stride;
uint64_t q_head_stride;
uint64_t raw_row_stride;
@@ -195,6 +195,17 @@ struct ds4_metal_args_glm_store_indexer_k {
float pad1;
};
struct ds4_metal_args_glm53_indexer_pool_update {
uint32_t pos0;
uint32_t n_tokens;
uint32_t cache_cap;
uint32_t head_dim;
uint32_t pool_size;
uint32_t cache_f16;
float eps;
uint32_t pad0;
};
struct ds4_metal_args_glm_attention_full {
uint32_t pos0;
uint32_t n_tokens;
@@ -221,6 +232,15 @@ struct ds4_metal_args_glm_fill_selected_range_batch {
uint32_t pad_row;
};
struct ds4_metal_args_glm53_expand_pool_selection {
uint32_t n_tokens;
uint32_t pos0;
uint32_t selected_pools;
uint32_t index_topk;
uint32_t pool_size;
uint32_t output_width;
};
struct ds4_metal_args_glm_indexer_rope_tail {
uint32_t n_tokens;
uint32_t n_head;
@@ -252,6 +272,8 @@ struct ds4_metal_args_glm_indexer_scores_batch {
uint32_t head_dim;
uint32_t pos0;
uint32_t cache_f16;
uint32_t row_group_size;
uint32_t pad0;
uint64_t q_token_stride;
uint64_t q_head_stride;
uint64_t weights_token_stride;
@@ -259,6 +281,13 @@ struct ds4_metal_args_glm_indexer_scores_batch {
float scale;
};
static inline uint glm_indexer_batch_visible_rows(
constant ds4_metal_args_glm_indexer_scores_batch &args,
uint token) {
const uint group = max(args.row_group_size, 1u);
return min((args.pos0 + token + 1u) / group, args.n_rows);
}
struct ds4_metal_args_glm_qk_lowrank {
uint32_t n_head;
uint32_t kv_lora_dim;
@@ -933,6 +962,117 @@ kernel void kernel_glm_store_indexer_k(
}
}
static inline float glm53_pool_bf16_to_f32(ushort value) {
return as_type<float>((uint)value << 16);
}
kernel void kernel_glm53_indexer_pool_update(
constant ds4_metal_args_glm53_indexer_pool_update &args,
device const char *raw_k,
device const char *gate,
device const float *norm_weight,
device const float *norm_bias,
device const ushort *ape,
device char *pool_cache,
device float *tail_k,
device float *tail_gate,
threadgroup float *shared [[threadgroup(0)]],
uint tid [[thread_index_in_threadgroup]],
uint3 tgpig [[threadgroup_position_in_grid]]) {
if (args.head_dim == 0u || args.pool_size == 0u ||
tid >= args.head_dim || args.n_tokens == 0u) return;
const uint pool = args.pos0 / args.pool_size + tgpig.x;
const uint pool_start = pool * args.pool_size;
const uint input_end = args.pos0 + args.n_tokens;
if (pool_start >= input_end || pool_start + args.pool_size <= args.pos0) return;
threadgroup float *rows = shared;
threadgroup float *mean = rows + args.pool_size * args.head_dim;
threadgroup float *inv = mean + args.pool_size;
const bool complete = pool_start + args.pool_size <= input_end;
for (uint r = 0; r < args.pool_size; r++) {
const uint pos = pool_start + r;
float k_value = 0.0f;
float gate_value = 0.0f;
if (pos >= args.pos0 && pos < input_end) {
const uint src_row = pos - args.pos0;
k_value = ((device const float *)raw_k)[
(uint64_t)src_row * args.head_dim + tid];
gate_value = ((device const float *)gate)[
(uint64_t)src_row * args.head_dim + tid];
if (!complete) {
tail_k[(uint64_t)r * args.head_dim + tid] = k_value;
tail_gate[(uint64_t)r * args.head_dim + tid] = gate_value;
}
} else {
k_value = tail_k[(uint64_t)r * args.head_dim + tid];
gate_value = tail_gate[(uint64_t)r * args.head_dim + tid];
}
rows[(uint64_t)r * args.head_dim + tid] = k_value;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
if (!complete || pool >= (args.cache_cap + args.pool_size - 1u) / args.pool_size) {
return;
}
if (tid < args.pool_size) {
const uint r = tid;
float sum = 0.0f;
for (uint d = 0; d < args.head_dim; d++) {
sum += rows[(uint64_t)r * args.head_dim + d];
}
const float m = sum / (float)args.head_dim;
float ss = 0.0f;
for (uint d = 0; d < args.head_dim; d++) {
const float delta = rows[(uint64_t)r * args.head_dim + d] - m;
ss += delta * delta;
}
mean[r] = m;
inv[r] = rsqrt(ss / (float)args.head_dim + args.eps);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float max_logit = -INFINITY;
float logits[4];
for (uint r = 0; r < args.pool_size; r++) {
const uint pos = pool_start + r;
float gate_value;
if (pos >= args.pos0) {
const uint src_row = pos - args.pos0;
gate_value = ((device const float *)gate)[
(uint64_t)src_row * args.head_dim + tid];
} else {
gate_value = tail_gate[(uint64_t)r * args.head_dim + tid];
}
logits[r] = gate_value +
glm53_pool_bf16_to_f32(ape[(uint64_t)r * args.head_dim + tid]);
max_logit = max(max_logit, logits[r]);
}
float denom = 0.0f;
for (uint r = 0; r < args.pool_size; r++) {
logits[r] = exp(logits[r] - max_logit);
denom += logits[r];
}
float pooled = 0.0f;
for (uint r = 0; r < args.pool_size; r++) {
const float normalized =
(rows[(uint64_t)r * args.head_dim + tid] - mean[r]) * inv[r] *
norm_weight[tid] + norm_bias[tid];
pooled += (logits[r] / denom) * normalized;
}
const uint64_t dst_index = (uint64_t)pool * args.head_dim + tid;
if (args.cache_f16 != 0u) {
((device half *)pool_cache)[dst_index] = (half)pooled;
} else {
((device float *)pool_cache)[dst_index] = pooled;
}
}
static inline void glm_dense_cache_store_f32_or_f16(
device char *base,
uint64_t index,
@@ -1011,7 +1151,8 @@ kernel void kernel_glm_build_kv_cache(
corr_dims);
}
const float theta_base = (float)pos;
const float inv_ndims = -1.0f / (float)args.qk_rope;
const float inv_ndims = args.qk_rope != 0u ?
-1.0f / (float)args.qk_rope : 0.0f;
for (uint r = tid * 2u; r < args.qk_rope; r += nth * 2u) {
#ifdef DS4_METAL_ROPE_EXP2_LOG2
const float theta = theta_base * exp2(inv_ndims * (float)r * log2(args.freq_base));
@@ -1096,7 +1237,8 @@ kernel void kernel_glm_build_kv_cache_decode_group4(
corr_dims);
}
const float theta_base = (float)pos;
const float inv_ndims = -1.0f / (float)args.qk_rope;
const float inv_ndims = args.qk_rope != 0u ?
-1.0f / (float)args.qk_rope : 0.0f;
for (uint r = tid * 2u; r < args.qk_rope; r += 512u) {
#ifdef DS4_METAL_ROPE_EXP2_LOG2
const float theta = theta_base * exp2(inv_ndims * (float)r * log2(args.freq_base));
@@ -1179,7 +1321,8 @@ kernel void kernel_glm_build_kv_cache_flash(
corr_dims);
}
const float theta_base = (float)pos;
const float inv_ndims = -1.0f / (float)args.qk_rope;
const float inv_ndims = args.qk_rope != 0u ?
-1.0f / (float)args.qk_rope : 0.0f;
for (uint r = tid * 2u; r < args.qk_rope; r += nth * 2u) {
#ifdef DS4_METAL_ROPE_EXP2_LOG2
const float theta = theta_base * exp2(inv_ndims * (float)r * log2(args.freq_base));
@@ -1400,6 +1543,35 @@ kernel void kernel_glm_fill_selected_range_batch(
selected[gid] = slot < visible ? slot : args.pad_row;
}
kernel void kernel_glm53_expand_pool_selection(
constant ds4_metal_args_glm53_expand_pool_selection &args,
device const uint32_t *pool_selected,
device uint32_t *raw_selected,
uint gid [[thread_position_in_grid]]) {
const uint total = args.n_tokens * args.output_width;
if (gid >= total || args.output_width == 0u || args.pool_size == 0u) return;
const uint token = gid / args.output_width;
const uint slot = gid - token * args.output_width;
uint value = 0xffffffffu;
if (slot < args.index_topk) {
const uint pool_slot = slot / args.pool_size;
if (pool_slot < args.selected_pools) {
const uint pool = pool_selected[
(uint64_t)token * args.selected_pools + pool_slot];
value = pool * args.pool_size + slot % args.pool_size;
}
} else {
const uint tail_slot = slot - args.index_topk;
const uint visible = args.pos0 + token + 1u;
const uint tail_count = visible % args.pool_size;
if (tail_slot < tail_count) {
value = visible - tail_count + tail_slot;
}
}
raw_selected[gid] = value;
}
kernel void kernel_glm_indexer_rope_tail_f32(
constant ds4_metal_args_glm_indexer_rope_tail & args,
device char *x,
@@ -1849,7 +2021,7 @@ kernel void kernel_glm_indexer_scores_batch(
device float *dst = (device float *)(scores +
(uint64_t)token * args.score_token_stride) + row;
const uint visible = min(args.pos0 + token + 1u, args.n_rows);
const uint visible = glm_indexer_batch_visible_rows(args, token);
if (row >= visible) {
if (tid == 0) *dst = -INFINITY;
return;
@@ -1910,7 +2082,7 @@ kernel void kernel_glm_indexer_scores_tiled_f32(
const uint last_token = min(token_base + TM, args.n_tokens);
const uint max_visible = last_token > token_base ?
min(args.pos0 + last_token, args.n_rows) : 0u;
glm_indexer_batch_visible_rows(args, last_token - 1u) : 0u;
if (row_base >= max_visible) {
for (uint i = tid; i < TM*TN; i += 128) {
@@ -2005,13 +2177,13 @@ kernel void kernel_glm_indexer_scores_tiled_f32(
}
if (token0 < args.n_tokens && row0 < args.n_rows) {
const uint visible = min(args.pos0 + token0 + 1u, args.n_rows);
const uint visible = glm_indexer_batch_visible_rows(args, token0);
device float *dst = (device float *)(scores +
(uint64_t)token0 * args.score_token_stride) + row0;
*dst = row0 < visible ? acc0 : -INFINITY;
}
if (token1 < args.n_tokens && row1 < args.n_rows) {
const uint visible = min(args.pos0 + token1 + 1u, args.n_rows);
const uint visible = glm_indexer_batch_visible_rows(args, token1);
device float *dst = (device float *)(scores +
(uint64_t)token1 * args.score_token_stride) + row1;
*dst = row1 < visible ? acc1 : -INFINITY;
@@ -2043,7 +2215,7 @@ kernel void kernel_glm_indexer_scores_tiled(
const uint last_token = min(token_base + TM, args.n_tokens);
const uint max_visible = last_token > token_base ?
min(args.pos0 + last_token, args.n_rows) : 0u;
glm_indexer_batch_visible_rows(args, last_token - 1u) : 0u;
if (row_base >= max_visible) {
for (uint i = tid; i < TM*TN; i += 128) {
@@ -2138,13 +2310,13 @@ kernel void kernel_glm_indexer_scores_tiled(
}
if (token0 < args.n_tokens && row0 < args.n_rows) {
const uint visible = min(args.pos0 + token0 + 1u, args.n_rows);
const uint visible = glm_indexer_batch_visible_rows(args, token0);
device float *dst = (device float *)(scores +
(uint64_t)token0 * args.score_token_stride) + row0;
*dst = row0 < visible ? acc0 : -INFINITY;
}
if (token1 < args.n_tokens && row1 < args.n_rows) {
const uint visible = min(args.pos0 + token1 + 1u, args.n_rows);
const uint visible = glm_indexer_batch_visible_rows(args, token1);
device float *dst = (device float *)(scores +
(uint64_t)token1 * args.score_token_stride) + row1;
*dst = row1 < visible ? acc1 : -INFINITY;
@@ -2241,7 +2413,6 @@ kernel void kernel_glm_qk_lowrank_q8_0_glm52_sg(
ushort3 ntg_u [[threads_per_threadgroup]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constexpr uint n_head = 64u;
constexpr uint kv_lora_dim = 512u;
constexpr uint qk_nope = 192u;
constexpr uint qk_dim = 256u;
@@ -2249,8 +2420,8 @@ kernel void kernel_glm_qk_lowrank_q8_0_glm52_sg(
const uint head = tgpig.x;
const uint wt = args.weight_type;
if (head >= n_head ||
args.n_head != n_head ||
if (head >= args.n_head ||
(args.n_head != 32u && args.n_head != 64u) ||
args.kv_lora_dim != kv_lora_dim ||
args.qk_nope != qk_nope ||
args.qk_dim != qk_dim ||
@@ -3624,7 +3795,7 @@ kernel void kernel_glm_attention_indexed_batch_lora_group8_vec_impl(
args.n_selected == 0u ||
args.cache_f16 == 0u ||
args.kv_lora_dim != 512u ||
args.qk_rope != 64u) {
(args.qk_rope != 0u && args.qk_rope != 64u)) {
return;
}
@@ -3668,7 +3839,7 @@ kernel void kernel_glm_attention_indexed_batch_lora_group8_vec_impl(
}
float corr_dims[2] = {0.0f, 0.0f};
if (args.ext_factor != 0.0f) {
if (args.qk_rope != 0u && args.ext_factor != 0.0f) {
glm_rope_yarn_corr_dims((int)args.qk_rope,
(int)args.n_ctx_orig,
args.freq_base,
@@ -3822,9 +3993,8 @@ kernel void kernel_glm_attention_indexed_batch_lora_group8_vec_causal_impl(
const uint head = tgpig.x * group_heads + head_in_group + args.head_base;
if (token >= args.n_tokens ||
args.n_selected == 0u ||
args.cache_f16 == 0u ||
args.kv_lora_dim != 512u ||
args.qk_rope != 64u) {
(args.qk_rope != 0u && args.qk_rope != 64u)) {
return;
}
@@ -3869,7 +4039,7 @@ kernel void kernel_glm_attention_indexed_batch_lora_group8_vec_causal_impl(
}
float corr_dims[2] = {0.0f, 0.0f};
if (args.ext_factor != 0.0f) {
if (args.qk_rope != 0u && args.ext_factor != 0.0f) {
glm_rope_yarn_corr_dims((int)args.qk_rope,
(int)args.n_ctx_orig,
args.freq_base,
@@ -3891,10 +4061,17 @@ kernel void kernel_glm_attention_indexed_batch_lora_group8_vec_causal_impl(
const uint rr = off / kv_vecs;
const uint vv = off - rr * kv_vecs;
const uint row = base + rr;
if (args.cache_f16 != 0u) {
device const half4 *src =
(device const half4 *)((device const half *)kv_lora_cache +
(uint64_t)row * args.kv_lora_dim);
kv_shared[off] = src[vv];
} else {
device const float4 *src =
(device const float4 *)((device const float *)kv_lora_cache +
(uint64_t)row * args.kv_lora_dim);
kv_shared[off] = (half4)src[vv];
}
}
for (uint off = tid; off < rows * rope_vecs; off += 256u) {
const uint rr = off / rope_vecs;
@@ -3903,11 +4080,12 @@ kernel void kernel_glm_attention_indexed_batch_lora_group8_vec_causal_impl(
const uint row = base + rr;
const uint64_t rope_base = (uint64_t)row * args.qk_rope;
const float2 y0 =
glm_cache_load_rotated_rope_pair_f16_only(k_rope_cache,
glm_cache_load_rotated_rope_pair(k_rope_cache,
rope_base,
r,
row,
args.qk_rope,
args.cache_f16,
args.freq_base,
args.freq_scale,
args.ext_factor,
@@ -3915,11 +4093,12 @@ kernel void kernel_glm_attention_indexed_batch_lora_group8_vec_causal_impl(
corr_dims[0],
corr_dims[1]);
const float2 y1 =
glm_cache_load_rotated_rope_pair_f16_only(k_rope_cache,
glm_cache_load_rotated_rope_pair(k_rope_cache,
rope_base,
r + 2u,
row,
args.qk_rope,
args.cache_f16,
args.freq_base,
args.freq_scale,
args.ext_factor,
@@ -4586,14 +4765,15 @@ kernel void kernel_glm_router_select_one(
threadgroup float *scratch [[threadgroup(0)]],
uint token [[threadgroup_position_in_grid]],
uint tid [[thread_position_in_threadgroup]]) {
const uint sort_width = args.n_expert > 256u ? 512u : 256u;
threadgroup float *sel_scores = scratch;
threadgroup int32_t *idx = (threadgroup int32_t *)(scratch + 256);
threadgroup int32_t *idx = (threadgroup int32_t *)(scratch + sort_width);
device const float *token_logits = logits + (uint64_t)token * args.n_expert;
device int32_t *token_selected = selected + (uint64_t)token * args.n_expert_used;
device float *token_weights = weights + (uint64_t)token * args.n_expert_used;
device float *token_probs = probs + (uint64_t)token * args.n_expert;
const uint n_expert = min(args.n_expert, 256u);
const uint n_expert = min(args.n_expert, 512u);
const bool active = tid < n_expert;
const float p = active ? ds4_glm_router_sigmoid(token_logits[tid]) : 0.0f;
if (active) token_probs[tid] = p;
@@ -4601,7 +4781,7 @@ kernel void kernel_glm_router_select_one(
idx[tid] = (int32_t)tid;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint k = 2; k <= 256; k <<= 1) {
for (uint k = 2; k <= sort_width; k <<= 1) {
for (uint j = k >> 1; j > 0; j >>= 1) {
const uint other = tid ^ j;
if (other > tid) {
@@ -4915,6 +5095,270 @@ kernel void kernel_dsv4_router_finalize_weights_one_simd(
}
}
// M3 decode specialization that materializes the probability
// transform in device memory before running the exact SIMD selection and
// weight normalization above. The volatile reload after the device barrier
// pins the same float store/load boundary as the standalone transform dispatch.
kernel void kernel_dsv4_router_transform_finalize_weights_one_simd(
constant ds4_metal_args_dsv4_router_select_one & args,
device const float *logits,
device float *probs,
device const float *bias,
device const int32_t *hash,
device const int32_t *tokens,
device int32_t *selected,
device float *weights,
threadgroup float *scratch [[threadgroup(0)]],
uint tid [[thread_position_in_threadgroup]]) {
if (tid >= 256 || args.hash_mode) return;
if (tid < 64) {
device const float4 *s = (device const float4 *)logits;
device float4 *d = (device float4 *)probs;
const float4 x = s[tid];
const float4 sp = select(log(1.0f + exp(x)), x, x > 20.0f);
d[tid] = sqrt(sp);
}
threadgroup_barrier(mem_flags::mem_device);
device volatile const float *reloaded_probs =
(device volatile const float *)probs;
(void)hash;
(void)tokens;
threadgroup float *score0_tg = scratch;
threadgroup int32_t *idx0_tg =
(threadgroup int32_t *)(scratch + 256);
threadgroup float *score1_tg = scratch + 512;
threadgroup int32_t *idx1_tg =
(threadgroup int32_t *)(scratch + 768);
const float p = reloaded_probs[tid];
float score = args.has_bias ? p + bias[tid] : p;
int32_t idx = (int32_t)tid;
uint cross_stage = 0;
for (uint k = 2; k <= 256; k <<= 1) {
for (uint j = k >> 1; j > 0; j >>= 1) {
float peer_score;
int32_t peer_idx;
bool take_peer;
const bool lower = (tid & j) == 0;
const bool descending = (tid & k) == 0;
if (j < 32) {
peer_score = simd_shuffle_xor(score, (ushort)j);
peer_idx = simd_shuffle_xor(idx, (ushort)j);
take_peer = descending
? (lower ? score < peer_score : score > peer_score)
: (lower ? score > peer_score : score < peer_score);
if (take_peer) {
score = peer_score;
idx = peer_idx;
}
} else {
threadgroup float *score_tg =
(cross_stage & 1u) != 0u ? score1_tg : score0_tg;
threadgroup int32_t *idx_tg =
(cross_stage & 1u) != 0u ? idx1_tg : idx0_tg;
score_tg[tid] = score;
idx_tg[tid] = idx;
threadgroup_barrier(mem_flags::mem_threadgroup);
const uint other = tid ^ j;
peer_score = score_tg[other];
peer_idx = idx_tg[other];
take_peer = descending
? (lower ? score < peer_score : score > peer_score)
: (lower ? score > peer_score : score < peer_score);
if (take_peer) {
score = peer_score;
idx = peer_idx;
}
cross_stage++;
}
}
}
if (tid < 6) {
selected[tid] = idx;
}
threadgroup_barrier(mem_flags::mem_device);
threadgroup volatile float *norm_scratch =
(threadgroup volatile float *)scratch;
if (tid == 0) {
device const int32_t *s = selected;
norm_scratch[0] = 0.0f;
for (uint i = 0; i < 6; i++) {
norm_scratch[0] =
norm_scratch[0] + reloaded_probs[s[i]];
}
norm_scratch[0] = max(norm_scratch[0], 6.103515625e-5f);
norm_scratch[1] = 1.5f / norm_scratch[0];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
if (tid < 6) {
device const int32_t *s = selected;
weights[tid] = reloaded_probs[s[tid]] * norm_scratch[1];
}
}
kernel void kernel_dsv4_router_project_select_fused(
constant ds4_metal_args_mul_mv & args,
constant ds4_metal_args_dsv4_router_select_one & select_args,
device const char * src0_router,
device const char * src1,
device float * logits,
device float * probs,
device const float * bias,
device int32_t * selected,
device float * weights,
device atomic_uint * completion,
threadgroup char * shmem_raw [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
uint3 tpitg [[thread_position_in_threadgroup]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constexpr short NSG = 8;
constexpr short NR0 = 2;
constexpr short NB = 32;
constexpr short NF = 16;
constexpr short NF4 = NF/4;
constexpr short NW = N_SIMDWIDTH;
const uint tid = tpitg.x;
const int nb = args.ne00/NB;
const int r0 = tgpig.x*NR0;
device const float4 *y4 = (device const float4 *)src1;
device const half4 *ax4[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
ax4[row] = (device const half4 *)
(src0_router + (uint64_t)(r0 + row)*args.nb01);
}
float sumf[NR0] = {0.f};
const short ix = tiisg/(NW/NF);
const short il = tiisg%(NW/NF);
const int ib0 = sgitg*NF + ix;
device const float4 *yb4 = y4 + (ib0*NB + il*NF)/4;
for (int ib = ib0; ib < nb; ib += NSG*NF) {
float4 yl4[NF4];
FOR_UNROLL (short i = 0; i < NF4; ++i) {
yl4[i] = yb4[i];
}
FOR_UNROLL (short row = 0; row < NR0; ++row) {
device const half4 *xb4 = ax4[row] + (ib*NB + il*NF)/4;
float sumq = 0.f;
FOR_UNROLL (short i = 0; i < NF4; ++i) {
sumq += dot(float4(xb4[i]), yl4[i]);
}
sumf[row] += sumq;
}
yb4 += NSG*NF*NW/4;
}
helper_mv_reduce_and_write<NR0>(logits, sumf, r0, args.ne01,
tiisg, sgitg, shmem_raw);
threadgroup float *scratch = (threadgroup float *)shmem_raw;
threadgroup_barrier(mem_flags::mem_threadgroup);
atomic_thread_fence(mem_flags::mem_device,
memory_order_seq_cst,
thread_scope_device);
if (tid == 0) {
const uint old = atomic_fetch_add_explicit(
completion, 1u, memory_order_relaxed);
scratch[0] = old == 127u ? 1.0f : 0.0f;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
if (scratch[0] == 0.0f) return;
atomic_thread_fence(mem_flags::mem_device,
memory_order_seq_cst,
thread_scope_device);
if (tid < 64) {
device volatile const float4 *s =
(device volatile const float4 *)logits;
device float4 *d = (device float4 *)probs;
const float4 xv = s[tid];
const float4 sp = select(log(1.0f + exp(xv)), xv, xv > 20.0f);
d[tid] = sqrt(sp);
}
threadgroup_barrier(mem_flags::mem_device);
device volatile const float *reloaded_probs =
(device volatile const float *)probs;
threadgroup float *score0_tg = scratch;
threadgroup int32_t *idx0_tg =
(threadgroup int32_t *)(scratch + 256);
threadgroup float *score1_tg = scratch + 512;
threadgroup int32_t *idx1_tg =
(threadgroup int32_t *)(scratch + 768);
const float p = reloaded_probs[tid];
float score = select_args.has_bias ? p + bias[tid] : p;
int32_t idx = (int32_t)tid;
uint cross_stage = 0;
for (uint k = 2; k <= 256; k <<= 1) {
for (uint j = k >> 1; j > 0; j >>= 1) {
float peer_score;
int32_t peer_idx;
bool take_peer;
const bool lower = (tid & j) == 0;
const bool descending = (tid & k) == 0;
if (j < 32) {
peer_score = simd_shuffle_xor(score, (ushort)j);
peer_idx = simd_shuffle_xor(idx, (ushort)j);
take_peer = descending
? (lower ? score < peer_score : score > peer_score)
: (lower ? score > peer_score : score < peer_score);
if (take_peer) {
score = peer_score;
idx = peer_idx;
}
} else {
threadgroup float *score_tg =
(cross_stage & 1u) != 0u ? score1_tg : score0_tg;
threadgroup int32_t *idx_tg =
(cross_stage & 1u) != 0u ? idx1_tg : idx0_tg;
score_tg[tid] = score;
idx_tg[tid] = idx;
threadgroup_barrier(mem_flags::mem_threadgroup);
const uint other = tid ^ j;
peer_score = score_tg[other];
peer_idx = idx_tg[other];
take_peer = descending
? (lower ? score < peer_score : score > peer_score)
: (lower ? score > peer_score : score < peer_score);
if (take_peer) {
score = peer_score;
idx = peer_idx;
}
cross_stage++;
}
}
}
if (tid < 6) selected[tid] = idx;
threadgroup_barrier(mem_flags::mem_device);
threadgroup volatile float *norm_scratch =
(threadgroup volatile float *)scratch;
if (tid == 0) {
norm_scratch[0] = 0.0f;
for (uint i = 0; i < 6; ++i) {
norm_scratch[0] = norm_scratch[0] + reloaded_probs[selected[i]];
}
norm_scratch[0] = max(norm_scratch[0], 6.103515625e-5f);
norm_scratch[1] = 1.5f / norm_scratch[0];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
if (tid < 6) {
weights[tid] = reloaded_probs[selected[tid]] * norm_scratch[1];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
atomic_thread_fence(mem_flags::mem_device,
memory_order_seq_cst,
thread_scope_device);
if (tid == 0) {
atomic_store_explicit(completion, 0u, memory_order_relaxed);
}
}
// Fills the dense compressed-attention mask with -inf. The selected top-k rows
// are enabled by kernel_dsv4_topk_mask_scatter in a second ordered dispatch.
kernel void kernel_dsv4_topk_mask(
@@ -5319,6 +5763,117 @@ kernel void kernel_dsv4_indexed_mixed_attention_heads8(
dst4[lane + 96] = o3 * inv_s;
}
// Each simdgroup owns two heads and updates both from one staged K/V row.
// This doubles row reuse without increasing the 256-thread workgroup.
kernel void kernel_dsv4_indexed_mixed_attention_heads16_dual(
constant ds4_metal_args_dsv4_indexed_attention &args,
device const char *q,
device const char *raw_kv,
device const char *comp_kv,
device const char *topk,
device const char *sinks,
device char *dst,
threadgroup half4 *kv_shared [[threadgroup(0)]],
uint2 tgpig [[threadgroup_position_in_grid]],
ushort tid [[thread_index_in_threadgroup]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]]) {
const uint token = tgpig.x;
const uint head0 = tgpig.y*16u + (uint)sg;
const uint head1 = head0 + 8u;
if (token >= args.n_tokens || head0 >= args.n_head) return;
device const float4 *qa = (device const float4 *)(q +
(uint64_t)token*args.q_token_stride +
(uint64_t)head0*args.q_head_stride);
half4 qa0 = (half4)qa[lane + 0];
half4 qa1 = (half4)qa[lane + 32];
half4 qa2 = (half4)qa[lane + 64];
half4 qa3 = (half4)qa[lane + 96];
half4 qb0 = half4(0.0h), qb1 = half4(0.0h);
half4 qb2 = half4(0.0h), qb3 = half4(0.0h);
if (head1 < args.n_head) {
device const float4 *qb = (device const float4 *)(q +
(uint64_t)token*args.q_token_stride +
(uint64_t)head1*args.q_head_stride);
qb0 = (half4)qb[lane + 0];
qb1 = (half4)qb[lane + 32];
qb2 = (half4)qb[lane + 64];
qb3 = (half4)qb[lane + 96];
}
float Ma = -FLT_MAX/2.0f, Sa = 0.0f;
float Mb = -FLT_MAX/2.0f, Sb = 0.0f;
float4 ao0 = 0.0f, ao1 = 0.0f, ao2 = 0.0f, ao3 = 0.0f;
float4 bo0 = 0.0f, bo1 = 0.0f, bo2 = 0.0f, bo3 = 0.0f;
const uint qpos = args.pos0 + token;
const uint last_pos = args.pos0 + args.n_tokens - 1u;
const uint first_raw_pos = last_pos + 1u - args.n_raw;
const uint raw_last_pos = first_raw_pos + args.n_raw - 1u;
const uint window_first = (args.window != 0u && qpos + 1u > args.window) ?
qpos + 1u - args.window : 0u;
const uint first = max(first_raw_pos, window_first);
const uint last = min(qpos, raw_last_pos);
if (first <= last) {
for (uint pos = first; pos <= last; pos++) {
const uint logical = pos - first_raw_pos;
const uint row = (args.raw_start + logical)%args.raw_cap;
device const float4 *src = (device const float4 *)(raw_kv +
(uint64_t)row*args.raw_row_stride);
if (tid < 128) kv_shared[tid] = (half4)src[tid];
threadgroup_barrier(mem_flags::mem_threadgroup);
dsv4_attend_shared_h4_row(kv_shared, qa0, qa1, qa2, qa3,
args.scale, lane, Ma, Sa, ao0, ao1, ao2, ao3);
if (head1 < args.n_head) {
dsv4_attend_shared_h4_row(kv_shared, qb0, qb1, qb2, qb3,
args.scale, lane, Mb, Sb, bo0, bo1, bo2, bo3);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
}
const uint visible = min((qpos + 1u)/args.ratio, args.n_comp);
device const int32_t *row_topk = (device const int32_t *)(topk +
(uint64_t)token*args.topk_token_stride);
for (uint i = 0; i < args.top_k; i++) {
const int32_t idx = row_topk[i];
if (idx < 0) continue;
if ((uint)idx >= visible) break;
if (tid < 128) {
kv_shared[tid] = dsv4_load_cache_h4(comp_kv,
args.comp_row_stride, (uint)idx, tid, args.comp_kv_f16 != 0u);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
dsv4_attend_shared_h4_row(kv_shared, qa0, qa1, qa2, qa3,
args.scale, lane, Ma, Sa, ao0, ao1, ao2, ao3);
if (head1 < args.n_head) {
dsv4_attend_shared_h4_row(kv_shared, qb0, qb1, qb2, qb3,
args.scale, lane, Mb, Sb, bo0, bo1, bo2, bo3);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
dsv4_attend_sink(((device const float *)sinks)[head0],
Ma, Sa, ao0, ao1, ao2, ao3);
const float ia = Sa == 0.0f ? 0.0f : 1.0f/Sa;
device float4 *da = (device float4 *)(dst +
(uint64_t)token*args.dst_token_stride +
(uint64_t)head0*args.dst_head_stride);
da[lane + 0] = ao0*ia; da[lane + 32] = ao1*ia;
da[lane + 64] = ao2*ia; da[lane + 96] = ao3*ia;
if (head1 < args.n_head) {
dsv4_attend_sink(((device const float *)sinks)[head1],
Mb, Sb, bo0, bo1, bo2, bo3);
const float ib = Sb == 0.0f ? 0.0f : 1.0f/Sb;
device float4 *db = (device float4 *)(dst +
(uint64_t)token*args.dst_token_stride +
(uint64_t)head1*args.dst_head_stride);
db[lane + 0] = bo0*ib; db[lane + 32] = bo1*ib;
db[lane + 64] = bo2*ib; db[lane + 96] = bo3*ib;
}
}
// Decode specialization of kernel_dsv4_indexed_mixed_attention_heads8.
// Generation attends one token at a time, so the ratio-4 indexed path spends a
// visible amount of time repeatedly staging the same K/V row for the eight
@@ -5450,6 +6005,190 @@ kernel void kernel_dsv4_indexed_mixed_attention_heads8_rb16(
dst4[lane + 96] = o3 * inv_s;
}
// Long-context decode specialization of the indexed mixed-attention path.
//
// The ordinary heads8 kernel reuses each K/V row across eight heads, but only
// launches one threadgroup per head group. Long-context decode therefore has
// too little parallel work while each group scans its raw and selected rows.
// This kernel retains the same eight-head reuse while splitting that row
// sequence across args.n_splits workgroups. A second kernel merges the online
// softmax partials and applies the attention sink.
kernel void kernel_dsv4_indexed_mixed_attention_heads8_split(
constant ds4_metal_args_dsv4_indexed_attention & args,
device const char *q,
device const char *raw_kv,
device const char *comp_kv,
device const char *topk,
device char *tmp,
threadgroup half4 *kv_shared [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tid [[thread_index_in_threadgroup]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]]) {
constexpr uint rows_per_block = 16u;
constexpr uint vecs_per_row = 128u;
const uint token = tgpig.x;
const uint head = tgpig.y * 8u + (uint)sg;
const uint split = tgpig.z;
const uint n_splits = args.n_splits;
if (token >= args.n_tokens || head >= args.n_head ||
n_splits < 2u || n_splits > 31u || split >= n_splits) {
return;
}
device const float4 *q4 = (device const float4 *)(q +
(uint64_t)token * args.q_token_stride +
(uint64_t)head * args.q_head_stride);
const half4 q0 = (half4)q4[lane + 0];
const half4 q1 = (half4)q4[lane + 32];
const half4 q2 = (half4)q4[lane + 64];
const half4 q3 = (half4)q4[lane + 96];
float M = -FLT_MAX/2.0f;
float S = 0.0f;
float4 o0 = 0.0f;
float4 o1 = 0.0f;
float4 o2 = 0.0f;
float4 o3 = 0.0f;
const uint qpos = args.pos0 + token;
const uint last_pos = args.pos0 + args.n_tokens - 1u;
const uint first_raw_pos = last_pos + 1u - args.n_raw;
const uint raw_last_pos = first_raw_pos + args.n_raw - 1u;
const uint window_first = (args.window != 0u && qpos + 1u > args.window) ?
qpos + 1u - args.window : 0u;
const uint raw_first = max(first_raw_pos, window_first);
const uint raw_last = min(qpos, raw_last_pos);
const uint raw_count = raw_first <= raw_last ?
raw_last - raw_first + 1u : 0u;
const uint total_rows = raw_count + args.top_k;
const uint rows_per_split =
(total_rows + n_splits - 1u) / n_splits;
const uint split_first = min(split * rows_per_split, total_rows);
const uint split_last = min(split_first + rows_per_split, total_rows);
const uint visible = min((qpos + 1u) / args.ratio, args.n_comp);
device const int32_t *row_topk = (device const int32_t *)(topk +
(uint64_t)token * args.topk_token_stride);
for (uint seq0 = split_first; seq0 < split_last;
seq0 += rows_per_block) {
const uint n_rows = min(rows_per_block, split_last - seq0);
for (uint off = (uint)tid;
off < n_rows * vecs_per_row;
off += 256u) {
const uint r = off / vecs_per_row;
const uint c = off - r * vecs_per_row;
const uint seq = seq0 + r;
half4 value = half4(0.0h);
if (seq < raw_count) {
const uint pos = raw_first + seq;
const uint logical = pos - first_raw_pos;
const uint row = (args.raw_start + logical) % args.raw_cap;
device const float4 *src = (device const float4 *)(raw_kv +
(uint64_t)row * args.raw_row_stride);
value = (half4)src[c];
} else {
const int32_t idx = row_topk[seq - raw_count];
if (idx >= 0 && (uint)idx < visible) {
value = dsv4_load_cache_h4(comp_kv,
args.comp_row_stride,
(uint)idx,
c,
args.comp_kv_f16 != 0u);
}
}
kv_shared[off] = value;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint r = 0; r < n_rows; r++) {
const uint seq = seq0 + r;
bool valid = true;
if (seq >= raw_count) {
const int32_t idx = row_topk[seq - raw_count];
valid = idx >= 0 && (uint)idx < visible;
}
if (valid) {
dsv4_attend_shared_h4_row_at(kv_shared,
r,
q0, q1, q2, q3,
args.scale,
lane,
M, S,
o0, o1, o2, o3);
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
const uint64_t n_rows = (uint64_t)args.n_tokens * args.n_head;
const uint64_t row = (uint64_t)token * args.n_head + head;
device float4 *partials = (device float4 *)tmp;
partials[(row * vecs_per_row + lane + 0u) * n_splits + split] = o0;
partials[(row * vecs_per_row + lane + 32u) * n_splits + split] = o1;
partials[(row * vecs_per_row + lane + 64u) * n_splits + split] = o2;
partials[(row * vecs_per_row + lane + 96u) * n_splits + split] = o3;
if (lane == 0u) {
device float *stats = (device float *)(partials +
n_rows * vecs_per_row * n_splits);
const uint64_t stat = (row * n_splits + split) * 2u;
stats[stat + 0u] = S;
stats[stat + 1u] = M;
}
}
kernel void kernel_dsv4_indexed_mixed_attention_heads8_split_reduce(
constant ds4_metal_args_dsv4_indexed_attention & args,
device const char *tmp,
device const char *sinks,
device char *dst,
uint tgpig [[threadgroup_position_in_grid]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]]) {
constexpr uint vecs_per_row = 128u;
const uint n_splits = args.n_splits;
const uint64_t n_rows = (uint64_t)args.n_tokens * args.n_head;
const uint64_t row = tgpig;
if (row >= n_rows || n_splits < 2u || n_splits > 31u) {
return;
}
device const float4 *partials = (device const float4 *)tmp;
device const float *stats = (device const float *)(partials +
n_rows * vecs_per_row * n_splits);
float part_sum = 0.0f;
float part_max = -FLT_MAX/2.0f;
if ((uint)lane < n_splits) {
const uint64_t stat = (row * n_splits + (uint)lane) * 2u;
part_sum = stats[stat + 0u];
part_max = stats[stat + 1u];
} else if ((uint)lane == n_splits) {
const uint head = (uint)(row % args.n_head);
part_sum = 1.0f;
part_max = ((device const float *)sinks)[head];
}
const float global_max = simd_max(part_max);
const float part_scale = part_sum > 0.0f ?
exp(part_max - global_max) : 0.0f;
const float total_sum = simd_sum(part_sum * part_scale);
const float inv_sum = total_sum > 0.0f ? 1.0f / total_sum : 0.0f;
device float4 *out = (device float4 *)dst + row * vecs_per_row;
for (uint i = (uint)sg; i < vecs_per_row; i += 4u) {
float4 value = float4(0.0f);
if ((uint)lane < n_splits) {
value = partials[(row * vecs_per_row + i) * n_splits +
(uint)lane] * part_scale;
}
value = simd_sum(value);
if (lane == 0u) {
out[i] = value * inv_sum;
}
}
}
static inline float dsv4_indexer_dot128_shared_q(
float4 c0,
float4 c1,

View File

@@ -41,6 +41,23 @@ struct ds4_metal_args_dsv4_rope_affine_pair {
float beta_slow;
};
struct ds4_metal_args_dsv4_head_norm_rope {
int32_t n_head;
int32_t head_dim;
int32_t head_dim4;
int32_t n_dims;
int32_t n_ctx_orig;
int32_t pos0;
int32_t inverse;
float eps;
float freq_base;
float freq_scale;
float ext_factor;
float attn_factor;
float beta_fast;
float beta_slow;
};
static float rope_yarn_ramp(const float low, const float high, const int i0) {
const float y = (i0 / 2 - low) / max(0.001f, high - low);
return 1.0f - min(1.0f, max(0.0f, y));
@@ -327,36 +344,109 @@ kernel void kernel_dsv4_rope_tail_f32_inplace_pair_shared4(
*((device float *) (dst_base + j1*args.nb0)) = x0*sin_theta + x1*cos_theta;
}
// DS4 positions are always affine within one RoPE dispatch. This variant
// reconstructs the same wrapped int32 position in-kernel, avoiding the host
// position array and its buffer binding while preserving the pair lane mapping
// and all floating-point operations of the specialization above.
kernel void kernel_dsv4_rope_tail_f32_inplace_pair_affine(
constant ds4_metal_args_dsv4_rope_affine_pair & args [[buffer(0)]],
device const char * src0 [[buffer(1)]],
device char * dst [[buffer(4)]],
uint tid [[thread_index_in_threadgroup]],
ushort3 ntg [[threads_per_threadgroup]],
uint3 tgpig [[threadgroup_position_in_grid]]) {
const int i1 = tgpig[0];
const int i2 = tgpig[1];
// Fuses the per-head RMSNorm and partial Q RoPE while retaining the standalone
// norm reduction tree and the mode-0 RoPE lane mapping.
kernel void kernel_dsv4_head_rms_norm_rope_tail_f32(
constant ds4_metal_args_dsv4_head_norm_rope & args,
device char * xraw,
threadgroup float * shmem_f32 [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort3 tpitg [[thread_position_in_threadgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort3 ntg [[threads_per_threadgroup]]) {
if (sgitg == 0) {
shmem_f32[tiisg] = 0.0f;
}
const uint head = tgpig.x;
const uint tok = tgpig.y;
device float4 * x4 = (device float4 *)xraw +
((uint64_t)tok * (uint64_t)args.n_head + head) *
(uint64_t)args.head_dim4;
float sumf = 0.0f;
for (int i00 = tpitg.x; i00 < args.head_dim4; i00 += ntg.x) {
sumf += dot(x4[i00], x4[i00]);
}
sumf = simd_sum(sumf);
threadgroup_barrier(mem_flags::mem_threadgroup);
if (tiisg == 0) {
shmem_f32[sgitg] = sumf;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
sumf = simd_sum(shmem_f32[tiisg]);
const float scale = 1.0f / sqrt(sumf / args.head_dim + args.eps);
const int n_nope = args.head_dim - args.n_dims;
if (n_nope < 0) {
return;
}
float corr_dims[2];
rope_yarn_corr_dims(args.n_dims, args.n_ctx_orig, args.freq_base,
args.beta_fast, args.beta_slow, corr_dims);
const float theta_base = (float)(args.pos0 + (int)tok);
const float inv_ndims = -1.0f / args.n_dims;
device float * xs = (device float *)x4;
for (int i0 = tpitg.x; i0 < args.head_dim; i0 += ntg.x) {
if (i0 < n_nope) {
xs[i0] = xs[i0] * scale;
continue;
}
const int r = i0 - n_nope;
if ((r & 1) != 0) {
continue;
}
#ifdef DS4_METAL_ROPE_EXP2_LOG2
const float theta =
theta_base * exp2(inv_ndims * (float)r * log2(args.freq_base));
#else
const float theta =
theta_base * pow(args.freq_base, inv_ndims * r);
#endif
float cos_theta;
float sin_theta;
rope_yarn(theta, args.freq_scale, corr_dims, r,
args.ext_factor, args.attn_factor,
&cos_theta, &sin_theta);
if (args.inverse) {
sin_theta = -sin_theta;
}
const float x0 = xs[i0] * scale;
const float x1 = xs[i0 + 1] * scale;
xs[i0] = x0 * cos_theta - x1 * sin_theta;
xs[i0 + 1] = x0 * sin_theta + x1 * cos_theta;
}
}
// DS4 positions are always affine within one RoPE dispatch. This variant
// reconstructs the same wrapped int32 position in-kernel, avoiding the host
// position array and its buffer binding while preserving the pair lane mapping
// and all floating-point operations of the specialization above.
/* Shared, deliberately noinline so that every caller gets bit-identical
* trigonometric codegen. The header note about tiny trig codegen changes
* flipping sampled tokens is exactly why this body must be compiled once and
* shared rather than inlined separately into each kernel. */
static __attribute__((noinline)) void ds4_rope_tail_pair_affine_row(
constant ds4_metal_args_dsv4_rope_affine_pair & args,
device const char * src_base,
device char * dst_base,
int n_nope,
uint raw_pos,
uint tid,
uint nthreads) {
float corr_dims[2];
rope_yarn_corr_dims(args.n_dims, args.n_ctx_orig, args.freq_base, args.beta_fast, args.beta_slow, corr_dims);
const uint raw_pos = args.pos0 + (uint)i2 * args.pos_step;
const float theta_base = (float)as_type<int>(raw_pos);
const float inv_ndims = -1.f/args.n_dims;
device const char * src_base =
src0 + (uint64_t)i2*args.token_bytes + (uint64_t)i1*args.row_bytes;
device char * dst_base =
dst + (uint64_t)i2*args.token_bytes + (uint64_t)i1*args.row_bytes;
for (int r = tid; r < args.n_dims; r += ntg.x) {
for (int r = tid; r < args.n_dims; r += nthreads) {
if ((r & 1) != 0) {
continue;
}
@@ -381,5 +471,413 @@ kernel void kernel_dsv4_rope_tail_f32_inplace_pair_affine(
*((device float *) (dst_base + j0*sizeof(float))) = x0*cos_theta - x1*sin_theta;
*((device float *) (dst_base + j1*sizeof(float))) = x0*sin_theta + x1*cos_theta;
}}
kernel void kernel_dsv4_rope_tail_f32_inplace_pair_affine(
constant ds4_metal_args_dsv4_rope_affine_pair & args [[buffer(0)]],
device const char * src0 [[buffer(1)]],
device char * dst [[buffer(4)]],
uint tid [[thread_index_in_threadgroup]],
ushort3 ntg [[threads_per_threadgroup]],
uint3 tgpig [[threadgroup_position_in_grid]]) {
const int i1 = tgpig[0];
const int i2 = tgpig[1];
const int n_nope = args.head_dim - args.n_dims;
if (n_nope < 0) {
return;
}
const uint raw_pos = args.pos0 + (uint)i2 * args.pos_step;
device const char * src_base =
src0 + (uint64_t)i2*args.token_bytes + (uint64_t)i1*args.row_bytes;
device char * dst_base =
dst + (uint64_t)i2*args.token_bytes + (uint64_t)i1*args.row_bytes;
ds4_rope_tail_pair_affine_row(args, src_base, dst_base, n_nope, raw_pos, tid, ntg.x);
}
// Decode-only fusion of the KV RoPE tail with the FP8/raw finalizer. Both were
// already single 64-thread threadgroups on the same row, back to back, so the
// pair cost two dispatches (~12.4 us) to touch 2 KB. The RoPE body below is a
// verbatim copy of kernel_dsv4_rope_tail_f32_inplace_pair_affine specialised to
// the decode grid (one head, one token, so i1 = i2 = 0) and the finalizer body
// is a verbatim copy of kernel_dsv4_kv_fp8_store_f32. The barrier between them
// is required because RoPE writes element pairs across lanes while the raw copy
// reads them per lane. Arithmetic, order and rounding are unchanged; the header
// warning above about trigonometric codegen still applies, so this kernel is
// gated and verified against full-vocabulary logits before promotion.
kernel void kernel_dsv4_kv_rope_fp8_store_f32(
constant ds4_metal_args_dsv4_kv_fp8_store & args,
constant ds4_metal_args_dsv4_rope_affine_pair & rope,
device float * kv,
device float * raw_cache,
threadgroup float * scratch [[threadgroup(0)]],
uint tid [[thread_index_in_threadgroup]]) {
{
const int rope_n_nope = rope.head_dim - rope.n_dims;
if (rope_n_nope < 0) {
return;
}
ds4_rope_tail_pair_affine_row(rope,
(device const char *)kv,
(device char *)kv,
rope_n_nope,
rope.pos0,
tid,
64u);
}
/* The RoPE helper writes device-memory pairs that different lanes read
* below. A threadgroup-only fence does not make those cross-lane device
* writes visible. */
threadgroup_barrier(mem_flags::mem_device_and_threadgroup);
{
const int head_dim = args.head_dim;
const int n_rot = args.n_rot;
const int n_nope = head_dim - n_rot;
if (head_dim <= 0 || n_rot < 0 || n_nope < 0 || tid >= 64) {
return;
}
device float * raw = raw_cache + (int64_t)args.raw_row * head_dim;
for (int off = 0; off < n_nope; off += 64) {
float v = 0.0f;
if (off + (int)tid < n_nope) {
v = kv[off + tid];
scratch[tid] = abs(v);
} else {
scratch[tid] = 0.0f;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = 32; stride > 0; stride >>= 1) {
if (tid < stride) {
scratch[tid] = max(scratch[tid], scratch[tid + stride]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
const float amax = max(scratch[0], 1.0e-4f);
const float fp8_scale = exp2(ceil(log2(amax / 448.0f)));
if (off + (int)tid < n_nope) {
const float q = dsv4_e4m3fn_dequant(clamp(v / fp8_scale, -448.0f, 448.0f)) * fp8_scale;
kv[off + tid] = q;
// Diagnostic only: skip the FP16 round-trip that normally matches the
// half-typed FlashAttention KV buffer's precision. With this enabled the
// indexer will see higher-precision raw values than FlashAttention does,
// which is informative but not a production-ready setting.
#ifdef DS4_METAL_KV_RAW_F32
raw[off + tid] = q;
#else
raw[off + tid] = (float)((half)q);
#endif
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
for (int i = n_nope + tid; i < head_dim; i += 64) {
#ifdef DS4_METAL_KV_RAW_F32
raw[i] = kv[i];
#else
raw[i] = (float)((half)kv[i]);
#endif
}
}
}
/* Decode-only sibling of kernel_flash_attn_ext_vec_reduce that also applies the
* inverse RoPE tail to the row it just produced, removing a whole dispatch per
* layer. Each threadgroup owns one head's entire 512-float row, so the RoPE is
* an intra-threadgroup dependency: reduce, barrier, rotate. Both halves call the
* same shared noinline helpers the standalone kernels use, so the arithmetic and
* its codegen are identical to running the two dispatches back to back. */
kernel void kernel_flash_attn_ext_vec_reduce_rope(
constant ds4_metal_args_flash_attn_ext_vec_reduce & args,
device const char * htmp,
device char * dst,
constant ds4_metal_args_dsv4_rope_affine_pair & rope,
uint tgpig[[threadgroup_position_in_grid]],
ushort tiitg[[thread_index_in_threadgroup]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
ds4_flash_attn_vec_reduce_row(args, htmp, dst, tgpig, tiisg, sgitg,
(short)FC_flash_attn_ext_vec_reduce_NWG,
(short)FC_flash_attn_ext_vec_reduce_DV);
threadgroup_barrier(mem_flags::mem_device);
const int n_nope = rope.head_dim - rope.n_dims;
if (n_nope < 0) {
return;
}
device char * row = dst + (uint64_t)tgpig * rope.row_bytes;
ds4_rope_tail_pair_affine_row(rope,
(device const char *)row,
row,
n_nope,
rope.pos0,
tiitg,
(uint)(32 * FC_flash_attn_ext_vec_reduce_NWG));
}
struct ds4_metal_args_dsv4_comp_finalize {
ds4_metal_args_dsv4_rope_affine_pair rope;
float rms_eps;
uint32_t pad0;
};
/* Decode-only emit-path fusion. Every ratio-th token, each layer finalizes
* one freshly pooled compressor row per compressor: RMS norm, RoPE tail, and
* then the FP8 round-trip + F16 commit copy (attention, 512 floats) or the
* Hadamard+FP4 QAT (indexer, 128 floats). Those were seven single-row
* dispatches; this kernel is one dispatch with two threadgroups.
*
* Each phase reproduces its standalone kernel bit-exactly:
* - norm: kernel_rms_norm_mul_f32_4's tree (float4 lanes, simd_sum, zero-
* padded 32-slot cross-simdgroup reduce); 512 uses 128 virtual threads on
* simdgroups 0-3, 128 uses 32 virtual threads on simdgroup 0.
* - rope: ds4_rope_tail_pair_affine_row verbatim (lanes 0-63, nthreads=64).
* - fp8: kernel_dsv4_fp8_kv_quantize_f32's 64-lane shmem max tree and
* round-trip, src==dst so the verbatim tail copy is a no-op and dropped.
* - commit: per-element f32->f16 conversion (value-wise exact).
* - qat: kernel_dsv4_indexer_hadamard_fp4_f32's butterfly and per-32 amax
* tree on lanes 0-127.
* Threads outside a phase's virtual width still execute every barrier, so
* threadgroup barriers stay uniform across the 256-thread threadgroup. */
kernel void kernel_dsv4_comp_row_finalize_f32(
constant ds4_metal_args_dsv4_comp_finalize & args [[buffer(0)]],
device float * attn_row [[buffer(1)]],
device const float * attn_norm_w [[buffer(2)]],
device char * attn_cache [[buffer(3)]],
device float * index_row [[buffer(4)]],
device const float * index_norm_w [[buffer(5)]],
device float * attn_state_kv [[buffer(6)]],
device float * attn_state_score [[buffer(7)]],
device float * index_state_kv [[buffer(8)]],
device float * index_state_score [[buffer(9)]],
threadgroup float * shmem [[threadgroup(0)]],
uint tgpig [[threadgroup_position_in_grid]],
ushort tiitg [[thread_index_in_threadgroup]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constant ds4_metal_args_dsv4_rope_affine_pair & rope_args = args.rope;
if (tgpig == 0) {
/* -------- attention compressor row (512 floats) -------- */
{
device float4 * y4 = (device float4 *)attn_row;
device const float4 * x4 = (device const float4 *)attn_row;
device const float4 * w4 = (device const float4 *)attn_norm_w;
if (sgitg == 0) {
shmem[tiisg] = 0.0f;
}
float sumf = 0.0f;
if (tiitg < 128) {
sumf = dot(x4[tiitg], x4[tiitg]);
}
sumf = simd_sum(sumf);
threadgroup_barrier(mem_flags::mem_threadgroup);
if (tiitg < 128 && tiisg == 0) {
shmem[sgitg] = sumf;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float total = 0.0f;
if (tiitg < 128) {
total = simd_sum(shmem[tiisg]);
}
const float mean = total / 512.0f;
const float scale = 1.0f/sqrt(mean + args.rms_eps);
if (tiitg < 128) {
y4[tiitg] = (x4[tiitg]*scale)*w4[tiitg];
}
}
threadgroup_barrier(mem_flags::mem_device);
ds4_rope_tail_pair_affine_row(rope_args,
(device const char *)attn_row,
(device char *)attn_row,
512 - rope_args.n_dims,
rope_args.pos0,
tiitg,
64u);
threadgroup_barrier(mem_flags::mem_device);
for (int off = 0; off < 512 - rope_args.n_dims; off += 64) {
float v = 0.0f;
if (tiitg < 64) {
v = attn_row[off + tiitg];
shmem[tiitg] = abs(v);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = 32; stride > 0; stride >>= 1) {
if (tiitg < stride) {
shmem[tiitg] = max(shmem[tiitg], shmem[tiitg + stride]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
const float amax = max(shmem[0], 1.0e-4f);
const float scale = exp2(ceil(log2(amax / 448.0f)));
if (tiitg < 64) {
const float q = dsv4_e4m3fn_dequant(clamp(v / scale, -448.0f, 448.0f)) * scale;
attn_row[off + tiitg] = q;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
threadgroup_barrier(mem_flags::mem_device);
if (tiitg < 128) {
device const float4 * x4 = (device const float4 *)attn_row;
device half4 * o4 = (device half4 *)attn_cache;
const float4 v = x4[tiitg];
o4[tiitg] = half4(v);
}
return;
}
if (tgpig >= 2u) {
/* Ratio-4 state shifts for both compressors (elementwise row move,
* so the flat gid mapping is bit-exact): 4*1024 attention elements
* then 4*256 indexer elements. */
const uint gid = (tgpig - 2u) * 256u + tiitg;
const uint n0 = 4u * 1024u;
if (gid < n0) {
attn_state_kv[gid] = attn_state_kv[n0 + gid];
attn_state_score[gid] = attn_state_score[n0 + gid];
return;
}
const uint gid1 = gid - n0;
const uint n1 = 4u * 256u;
if (gid1 >= n1) return;
index_state_kv[gid1] = index_state_kv[n1 + gid1];
index_state_score[gid1] = index_state_score[n1 + gid1];
return;
}
/* -------- indexer compressor row (128 floats) -------- */
{
device float4 * y4 = (device float4 *)index_row;
device const float4 * x4 = (device const float4 *)index_row;
device const float4 * w4 = (device const float4 *)index_norm_w;
if (sgitg == 0) {
shmem[tiisg] = 0.0f;
}
float sumf = 0.0f;
if (tiitg < 32) {
sumf = dot(x4[tiitg], x4[tiitg]);
}
sumf = simd_sum(sumf);
threadgroup_barrier(mem_flags::mem_threadgroup);
if (tiisg == 0) {
shmem[sgitg] = sumf;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float total = 0.0f;
if (tiitg < 32) {
total = simd_sum(shmem[tiisg]);
}
const float mean = total / 128.0f;
const float scale = 1.0f/sqrt(mean + args.rms_eps);
if (tiitg < 32) {
y4[tiitg] = (x4[tiitg]*scale)*w4[tiitg];
}
}
threadgroup_barrier(mem_flags::mem_device);
ds4_rope_tail_pair_affine_row(rope_args,
(device const char *)index_row,
(device char *)index_row,
128 - rope_args.n_dims,
rope_args.pos0,
tiitg,
64u);
threadgroup_barrier(mem_flags::mem_device);
{
threadgroup float *vals = shmem;
threadgroup float *absbuf = shmem + 128;
if (tiitg < 128) {
vals[tiitg] = index_row[tiitg];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = 1u; stride < 128u; stride <<= 1u) {
if (tiitg < 128 && (tiitg & stride) == 0u) {
const uint base = (tiitg & ~(2u * stride - 1u)) + (tiitg & (stride - 1u));
const float a = vals[base];
const float b = vals[base + stride];
vals[base] = a + b;
vals[base + stride] = a - b;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
float v = 0.0f;
if (tiitg < 128) {
v = vals[tiitg] * 0.08838834764831845f;
absbuf[tiitg] = abs(v);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
const uint block = tiitg >> 5u;
const uint lane = tiitg & 31u;
const uint block_base = block * 32u;
for (uint stride = 16u; stride > 0u; stride >>= 1u) {
if (tiitg < 128 && lane < stride) {
absbuf[block_base + lane] = max(absbuf[block_base + lane],
absbuf[block_base + lane + stride]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (tiitg < 128) {
const float amax = max(absbuf[block_base], 7.052966104933725e-38f);
const float scale = exp2(ceil(log2(amax / 6.0f)));
index_row[tiitg] = dsv4_e2m1fn_dequant(clamp(v / scale, -6.0f, 6.0f)) * scale;
}
}
}
// Host-visible packed FlashAttention + exact inverse-RoPE decode kernel.
kernel void kernel_dsv4_flash_attn_vec_packed32_reduce_rope_f16_dk512_dv512(
constant ds4_metal_args_flash_attn_ext_vec & args [[buffer(0)]],
device const char * q [[buffer(1)]],
device const char * k [[buffer(2)]],
device const char * v [[buffer(3)]],
device const char * mask [[buffer(4)]],
device const char * sinks [[buffer(5)]],
device const char * pad [[buffer(6)]],
device char * dst [[buffer(7)]],
constant ds4_metal_args_dsv4_rope_affine_pair & rope
[[buffer(8)]],
threadgroup char * shmem [[threadgroup(0)]],
uint head [[threadgroup_position_in_grid]],
ushort tiitg [[thread_index_in_threadgroup]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
/* Uniform specialization guard; host applies the same eligibility gate. */
if (!FC_flash_attn_ext_vec_has_mask ||
!FC_flash_attn_ext_vec_has_sinks ||
FC_flash_attn_ext_vec_has_bias ||
FC_flash_attn_ext_vec_has_scap ||
FC_flash_attn_ext_vec_nsg != 1 ||
FC_flash_attn_ext_vec_nwg != 32 ||
FC_flash_attn_ext_vec_ns10 != 512 ||
FC_flash_attn_ext_vec_ns20 != 512 ||
args.ne01 != 1 || args.ne02 != 64 || args.ne03 != 1 ||
args.ne_12_2 != 1 || args.ne_12_3 != 1 ||
args.ne31 != 1 || args.ne32 != 1 || args.ne33 != 1 ||
args.ne11 <= 0 || args.ne11 > 1024 || head >= (uint)args.ne02 ||
args.nb02 != 2048 || args.nb11 != 1024 || args.nb21 != 1024 ||
rope.head_dim != 512 || rope.n_dims != 64 ||
rope.row_bytes != 2048 || rope.inverse == 0) {
return;
}
ds4_flash_attn_vec_packed8_reduce_f16_512(
args, q, k, v, mask, sinks, pad, dst, shmem,
head, tiisg, sgitg);
/* Same producer/consumer boundary as the current reduce+RoPE kernel. */
threadgroup_barrier(mem_flags::mem_device);
const int n_nope = rope.head_dim - rope.n_dims;
device char * row = dst + (uint64_t)head * rope.row_bytes;
ds4_rope_tail_pair_affine_row(rope,
(device const char *)row,
row,
n_nope,
rope.pos0,
tiitg,
32u * 32u);
}

View File

@@ -1398,24 +1398,26 @@ constant int32_t FC_flash_attn_ext_vec_reduce_NWG [[function_constant(FC_FLASH_A
// Reduces split-K decode FlashAttention partials. It combines each workgroup's
// output vector and softmax (sum,max) pair into the final attention result.
kernel void kernel_flash_attn_ext_vec_reduce(
/* Shared and deliberately noinline so the split-K reduction is compiled once and
* every caller gets identical codegen. The RoPE-fused sibling in dsv4_rope.metal
* calls this same body, which is what keeps the fusion bit-exact. */
static __attribute__((noinline)) void ds4_flash_attn_vec_reduce_row(
constant ds4_metal_args_flash_attn_ext_vec_reduce & args,
device const char * htmp,
device char * dst,
uint tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
#define NWG (FC_flash_attn_ext_vec_reduce_NWG)
#define DV (FC_flash_attn_ext_vec_reduce_DV)
uint tgpig,
ushort tiisg,
ushort sgitg,
short NWG_,
short DV_) {
const uint64_t rid = tgpig;
const short iwg = tiisg;
device const float * ss = (device const float *) htmp + (uint64_t)args.nrows*DV*NWG;
device const float * ss = (device const float *) htmp + (uint64_t)args.nrows*DV_*NWG_;
float S = ss[rid*(2*NWG) + 2*iwg + 0];
float M = ss[rid*(2*NWG) + 2*iwg + 1];
float S = ss[rid*(2*NWG_) + 2*iwg + 0];
float M = ss[rid*(2*NWG_) + 2*iwg + 1];
const float m = simd_max(M);
const float ms = exp(M - m);
@@ -1423,19 +1425,268 @@ kernel void kernel_flash_attn_ext_vec_reduce(
S = simd_sum(S*ms);
S = S == 0.0f ? 0.0f : 1.0f/S;
const short DV4 = DV/4;
const short DV4 = DV_/4;
device const float4 * htmp4 = (device const float4 *) htmp + rid*DV4*NWG;
device const float4 * htmp4 = (device const float4 *) htmp + rid*DV4*NWG_;
device float4 * dst4 = (device float4 *) dst + rid*DV4;
for (short i = sgitg; i < DV4; i += NWG) {
const float4 v = simd_sum(htmp4[i*NWG + iwg]*ms);
for (short i = sgitg; i < DV4; i += NWG_) {
const float4 v = simd_sum(htmp4[i*NWG_ + iwg]*ms);
if (iwg == 0) {
dst4[i] = v*S;
}
}
#undef NWG
#undef DV
}
kernel void kernel_flash_attn_ext_vec_reduce(
constant ds4_metal_args_flash_attn_ext_vec_reduce & args,
device const char * htmp,
device char * dst,
uint tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
ds4_flash_attn_vec_reduce_row(args, htmp, dst, tgpig, tiisg, sgitg,
(short)FC_flash_attn_ext_vec_reduce_NWG,
(short)FC_flash_attn_ext_vec_reduce_DV);
}
// M5 decode specialization: time-slice all 32 split-K workgroups through eight
// physical simdgroups, then reduce through the same 32-lane topology without a
// device partial buffer. The host gate fixes the exact F16 512-wide geometry.
static inline void ds4_flash_attn_vec_packed8_reduce_f16_512(
constant ds4_metal_args_flash_attn_ext_vec & args,
device const char * q,
device const char * k,
device const char * v,
device const char * mask,
device const char * sinks,
device const char * pad,
device char * dst,
threadgroup char * shmem,
uint head,
ushort tiisg,
ushort sgitg) {
constexpr short NW = 32;
constexpr short C = 32;
constexpr short NSG = 8;
constexpr short NWG = 32;
constexpr short DK4 = 128;
constexpr short DV4 = 128;
constexpr short SH = 128;
/* 24,448 dynamic bytes: shared Q, eight score/mask banks, all 32
* split-local weights and stats, sink scales, and a padded 32x33 F32
* float4 partial plane. */
threadgroup half4 *q_shared = (threadgroup half4 *)shmem;
threadgroup half *score_banks =
(threadgroup half *)(q_shared + DK4);
threadgroup volatile float *weights =
(threadgroup volatile float *)(score_banks + NSG * SH);
threadgroup volatile float *stats = weights + NWG * C;
threadgroup volatile float *sink_scale = stats + 2 * NWG;
threadgroup volatile float4 *partial_plane =
(threadgroup volatile float4 *)(sink_scale + NWG);
const short lane = (short)tiisg;
threadgroup half *bank = score_banks + (short)sgitg * SH;
threadgroup float *ss = (threadgroup float *)bank;
threadgroup half *sm = bank + 2 * C;
device const float4 *q4 =
(device const float4 *)(q + (uint64_t)head * args.nb02);
if (sgitg == 0) {
for (short i = lane; i < DK4; i += NW) {
q_shared[i] = (half4)q4[i];
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
/* Eight physical simdgroups time-slice the exact 32 legacy split-K
* workgroups. The official <=1024-key gate gives each virtual split at
* most one 32-row block, so its value partial can be formed later from
* these materialized weights without changing online-softmax order. */
for (short iwg = (short)sgitg; iwg < NWG; iwg += NSG) {
float S = 0.0f;
float M = -FLT_MAX / 2;
float out_scale = 1.0f;
const int ic_original = (int)iwg * C;
weights[(uint)iwg * C + (uint)lane] = 0.0f;
ss[lane] = 0.0f;
sm[lane] = (half)0.0h;
simdgroup_barrier(mem_flags::mem_threadgroup);
if (ic_original < args.ne11) {
device const char *k_block = k;
device const char *v_block = v;
device const half *pm = (device const half *)mask;
int ic = ic_original;
if (FC_flash_attn_ext_vec_has_kvpad && ic + C > args.ne11) {
k_block = pad;
const uint64_t k_pad_bytes =
args.nb11 * (uint64_t)C *
(uint64_t)args.ne_12_2 * (uint64_t)args.ne_12_3;
const uint64_t v_pad_bytes =
args.nb21 * (uint64_t)C *
(uint64_t)args.ne_12_2 * (uint64_t)args.ne_12_3;
if (FC_flash_attn_ext_vec_shared_kvpad) {
v_block = k_block;
pm = (device const half *)(k_block +
k_pad_bytes + v_pad_bytes);
} else {
v_block = k_block + k_pad_bytes;
pm = (device const half *)(v_block + v_pad_bytes);
}
ic = 0;
}
sm[lane] = pm[ic + lane];
if (simd_max(sm[lane]) > -MAXHALF) {
device const half4 *pk4 =
(device const half4 *)(k_block +
(uint64_t)ic * args.nb11);
threadgroup const half4 *pq4 = q_shared;
pk4 += lane;
pq4 += lane;
float lane_mqk = 0.0f;
FOR_UNROLL (short cc = 0; cc < C; ++cc) {
float mqk = 0.0f;
FOR_UNROLL (short ii = 0; ii < DK4 / NW; ++ii) {
mqk += dot((float4)pk4[cc * DK4 + ii * NW],
(float4)pq4[ii * NW]);
}
mqk = simd_sum(mqk);
if (lane == cc) {
lane_mqk = mqk;
}
}
ss[lane] = fma(lane_mqk, args.scale,
(float)sm[lane]);
simdgroup_barrier(mem_flags::mem_threadgroup);
const float old_m = M;
const float score = ss[lane];
M = simd_max(max(M, score));
const float ms = exp(old_m - M);
const float vs = exp(score - M);
S = S * ms + simd_sum(vs);
ss[lane] = vs;
simdgroup_barrier(mem_flags::mem_threadgroup);
weights[(uint)iwg * C + (uint)lane] = ss[lane];
}
if (FC_flash_attn_ext_vec_has_sinks && iwg == 0) {
const float old_m = M;
const float sink = lane == 0
? ((device const float *)sinks)[head]
: -FLT_MAX / 2;
M = simd_max(max(M, sink));
const float ms = exp(old_m - M);
const float vs = exp(sink - M);
S = S * ms + simd_sum(vs);
out_scale = ms;
}
} else if (FC_flash_attn_ext_vec_has_sinks && iwg == 0) {
const float old_m = M;
const float sink = lane == 0
? ((device const float *)sinks)[head]
: -FLT_MAX / 2;
M = simd_max(max(M, sink));
const float ms = exp(old_m - M);
const float vs = exp(sink - M);
S = S * ms + simd_sum(vs);
out_scale = ms;
}
if (lane == 0) {
stats[2 * (uint)iwg + 0] = S;
stats[2 * (uint)iwg + 1] = M;
sink_scale[(uint)iwg] = out_scale;
}
simdgroup_barrier(mem_flags::mem_threadgroup);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
/* Recreate the legacy 32-lane reducer exactly: lane is the virtual split
* index, including neutral idle splits in their original tree positions. */
const short split = lane;
float reduce_S = stats[2 * (uint)split + 0];
float reduce_M = stats[2 * (uint)split + 1];
const float reduce_max = simd_max(reduce_M);
const float reduce_ms = exp(reduce_M - reduce_max);
reduce_S = simd_sum(reduce_S * reduce_ms);
const float reduce_inv =
reduce_S == 0.0f ? 0.0f : 1.0f / reduce_S;
device float4 *dst4 =
(device float4 *)(dst +
(uint64_t)head * 512u * sizeof(float));
/* Form one 32-float4 output quadrant at a time. During production each
* physical simdgroup time-slices four virtual splits while SIMD lanes are
* contiguous output columns, exactly matching the legacy V loads and
* cc-major accumulation. A padded 33-column plane avoids a 32-way TG-bank
* conflict when the reducer transposes lanes back to virtual splits. */
for (short quadrant = 0; quadrant < 4; ++quadrant) {
for (short iwg = (short)sgitg; iwg < NWG; iwg += NSG) {
float4 lo = float4(0.0f);
const int ic_original = (int)iwg * C;
if (ic_original < args.ne11) {
device const char *v_block = v;
int ic = ic_original;
if (FC_flash_attn_ext_vec_has_kvpad && ic + C > args.ne11) {
device const char *k_block = pad;
const uint64_t k_pad_bytes =
args.nb11 * (uint64_t)C *
(uint64_t)args.ne_12_2 * (uint64_t)args.ne_12_3;
if (FC_flash_attn_ext_vec_shared_kvpad) {
v_block = k_block;
} else {
v_block = k_block + k_pad_bytes;
}
ic = 0;
}
device const half4 *pv4 =
(device const half4 *)(v_block +
(uint64_t)ic * args.nb21);
threadgroup volatile float *split_weights =
weights + (uint)iwg * C;
const short oc = quadrant * NW + lane;
FOR_UNROLL (short cc = 0; cc < C; ++cc) {
lo += float4(pv4[cc * DV4 + oc]) *
float4(split_weights[cc]);
}
float4 acc = float4(0.0f);
acc += lo;
if (iwg == 0) {
acc *= sink_scale[0];
}
lo = acc;
}
partial_plane[(uint)iwg * 33u + (uint)lane] = lo;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
/* lane is now the legacy split index and each physical simdgroup
* reduces four output columns through the identical simd_sum tree. */
for (short out_lane = (short)sgitg; out_lane < NW; out_lane += NSG) {
const float4 materialized =
(float4)partial_plane[(uint)lane * 33u + (uint)out_lane];
const float4 reduced = simd_sum(materialized * reduce_ms);
if (lane == 0) {
dst4[quadrant * NW + out_lane] = reduced * reduce_inv;
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
}

143
metal/glm53_bf16.metal Normal file
View File

@@ -0,0 +1,143 @@
// BF16 model-weight kernels used by GLM-5.3 Flash.
static inline float glm53_bf16_to_f32(ushort value) {
return as_type<float>((uint)value << 16);
}
struct glm53_bf16_matmul_args {
uint in_dim;
uint out_dim;
uint n_rows;
};
kernel void kernel_glm53_embedding_bf16(
constant glm53_bf16_matmul_args &args,
device const ushort *weights,
device const int *tokens,
device float *out,
uint2 gid [[thread_position_in_grid]]) {
const uint d = gid.x;
const uint row = gid.y;
if (d >= args.in_dim || row >= args.n_rows) return;
const int token = tokens[row];
out[(ulong)row * args.in_dim + d] =
token >= 0 && (uint)token < args.out_dim
? glm53_bf16_to_f32(weights[(ulong)(uint)token * args.in_dim + d])
: 0.0f;
}
static inline void glm53_mul_mv_bf16_f32_row(
constant glm53_bf16_matmul_args &args,
device const ushort *weights,
device const float *x,
device float *out,
uint2 tgpig,
ushort lane,
ushort sg,
ushort nsg) {
const uint out_row = tgpig.x * (uint)nsg + sg;
const uint token = tgpig.y;
if (out_row >= args.out_dim || token >= args.n_rows) return;
device const ushort *w = weights + (ulong)out_row * args.in_dim;
device const float *xr = x + (ulong)token * args.in_dim;
float sum = 0.0f;
uint k = lane;
for (; k + 224u < args.in_dim; k += 256u) {
const ushort w0 = w[k];
const ushort w1 = w[k + 32u];
const ushort w2 = w[k + 64u];
const ushort w3 = w[k + 96u];
const ushort w4 = w[k + 128u];
const ushort w5 = w[k + 160u];
const ushort w6 = w[k + 192u];
const ushort w7 = w[k + 224u];
const float x0 = xr[k];
const float x1 = xr[k + 32u];
const float x2 = xr[k + 64u];
const float x3 = xr[k + 96u];
const float x4 = xr[k + 128u];
const float x5 = xr[k + 160u];
const float x6 = xr[k + 192u];
const float x7 = xr[k + 224u];
sum = fma(glm53_bf16_to_f32(w0), x0, sum);
sum = fma(glm53_bf16_to_f32(w1), x1, sum);
sum = fma(glm53_bf16_to_f32(w2), x2, sum);
sum = fma(glm53_bf16_to_f32(w3), x3, sum);
sum = fma(glm53_bf16_to_f32(w4), x4, sum);
sum = fma(glm53_bf16_to_f32(w5), x5, sum);
sum = fma(glm53_bf16_to_f32(w6), x6, sum);
sum = fma(glm53_bf16_to_f32(w7), x7, sum);
}
for (; k < args.in_dim; k += 32u) {
sum = fma(glm53_bf16_to_f32(w[k]), xr[k], sum);
}
sum = simd_sum(sum);
if (lane == 0u) out[(ulong)token * args.out_dim + out_row] = sum;
}
/* One simdgroup owns one output row. Eight independent loads expose enough
* memory-level parallelism for decode without changing the reduction tree. */
kernel void kernel_glm53_mul_mv_bf16_f32(
constant glm53_bf16_matmul_args &args,
device const ushort *weights,
device const float *x,
device float *out,
uint2 tgpig [[threadgroup_position_in_grid]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]],
ushort nsg [[simdgroups_per_threadgroup]]) {
glm53_mul_mv_bf16_f32_row(args, weights, x, out,
tgpig, lane, sg, nsg);
}
kernel void kernel_glm53_mul_mv_bf16_f32_qkv(
constant glm53_bf16_matmul_args &args,
device const ushort *weights_q,
device const ushort *weights_k,
device const ushort *weights_v,
device const float *x,
device float *out_q,
device float *out_k,
device float *out_v,
uint3 tgpig [[threadgroup_position_in_grid]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]],
ushort nsg [[simdgroups_per_threadgroup]]) {
device const ushort *weights = tgpig.z == 0u ? weights_q :
(tgpig.z == 1u ? weights_k : weights_v);
device float *out = tgpig.z == 0u ? out_q :
(tgpig.z == 1u ? out_k : out_v);
glm53_mul_mv_bf16_f32_row(args, weights, x, out,
tgpig.xy, lane, sg, nsg);
}
struct glm53_bf16_block16 {
ushort v[16];
};
template <typename type4x4>
void glm53_dequantize_bf16(
device const glm53_bf16_block16 *src,
short il,
thread type4x4 &reg) {
(void)il;
float4x4 values;
for (short i = 0; i < 16; i++) {
values[i / 4][i % 4] = glm53_bf16_to_f32(src->v[i]);
}
reg = (type4x4)values;
}
typedef decltype(kernel_mul_mm<
half, half4x4, simdgroup_half8x8,
half, half2x4, simdgroup_half8x8,
glm53_bf16_block16, 1, glm53_dequantize_bf16,
float, float4x4, float, float2x4>) glm53_mul_mm_bf16_t;
template [[host_name("kernel_glm53_mul_mm_bf16_f32")]]
kernel glm53_mul_mm_bf16_t kernel_mul_mm<
half, half4x4, simdgroup_half8x8,
half, half2x4, simdgroup_half8x8,
glm53_bf16_block16, 1, glm53_dequantize_bf16,
half, half4x4, float, float2x4>;

314
metal/glm53_kda.metal Normal file
View File

@@ -0,0 +1,314 @@
// Kimi Delta Attention kernels, adapted from the kimi-k3 branch.
struct glm53_kda_args {
uint n_heads;
uint n_rows;
float lower_bound;
float norm_eps;
};
/*
* One threadgroup owns one (sequence, head). Four simdgroups update four
* value rows concurrently; every lane owns four adjacent key columns.
*/
kernel void kernel_glm53_kda_decode(
constant glm53_kda_args &args,
device const float *q_in,
device const float *k_in,
device const float *v_in,
device const float *raw_gate,
device const float *raw_beta,
device const float *output_gate,
device const float *q_conv,
device const float *k_conv,
device const float *v_conv,
device const float *a_log,
device const float *dt_bias,
device const float *output_norm,
device float *conv_state,
device float *state,
device float *out,
threadgroup float *scratch [[threadgroup(0)]],
uint2 tgpig [[threadgroup_position_in_grid]],
ushort tid [[thread_index_in_threadgroup]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]]) {
constexpr uint D = 128u;
constexpr uint HISTORY = 3u;
const uint row = tgpig.x;
const uint head = tgpig.y;
if (row >= args.n_rows || head >= args.n_heads) return;
threadgroup float *sq = scratch;
threadgroup float *sk = sq + D;
threadgroup float *sd = sk + D;
threadgroup float *sv = sd + D;
threadgroup float *so = sv + D;
threadgroup float *reduce_q = so + D;
threadgroup float *reduce_k = reduce_q + 4u;
threadgroup float *reduce_o = reduce_k + 4u;
threadgroup float *beta_shared = reduce_o + 4u;
const uint projection = args.n_heads * D;
const uint channel = head * D + tid;
const ulong input_base = (ulong)row * projection + head * D;
const ulong conv_row_stride = 3ul * HISTORY * projection;
if (tid < D) {
float q_acc = 0.0f;
float k_acc = 0.0f;
float v_acc = 0.0f;
device float *q_state = conv_state +
(ulong)row * conv_row_stride;
device float *k_state = q_state + HISTORY * projection;
device float *v_state = k_state + HISTORY * projection;
for (uint w = 0; w < HISTORY; w++) {
q_acc = fma(q_state[(ulong)w * projection + channel],
q_conv[(ulong)channel * 4u + w], q_acc);
k_acc = fma(k_state[(ulong)w * projection + channel],
k_conv[(ulong)channel * 4u + w], k_acc);
v_acc = fma(v_state[(ulong)w * projection + channel],
v_conv[(ulong)channel * 4u + w], v_acc);
}
const float q_new = q_in[input_base + tid];
const float k_new = k_in[input_base + tid];
const float v_new = v_in[input_base + tid];
q_acc = fma(q_new, q_conv[(ulong)channel * 4u + 3u], q_acc);
k_acc = fma(k_new, k_conv[(ulong)channel * 4u + 3u], k_acc);
v_acc = fma(v_new, v_conv[(ulong)channel * 4u + 3u], v_acc);
q_state[channel] = q_state[projection + channel];
q_state[projection + channel] = q_state[2ul * projection + channel];
q_state[2ul * projection + channel] = q_new;
k_state[channel] = k_state[projection + channel];
k_state[projection + channel] = k_state[2ul * projection + channel];
k_state[2ul * projection + channel] = k_new;
v_state[channel] = v_state[projection + channel];
v_state[projection + channel] = v_state[2ul * projection + channel];
v_state[2ul * projection + channel] = v_new;
sq[tid] = q_acc / (1.0f + exp(-q_acc));
sk[tid] = k_acc / (1.0f + exp(-k_acc));
sv[tid] = v_acc / (1.0f + exp(-v_acc));
const float gate = raw_gate[input_base + tid] + dt_bias[channel];
sd[tid] = exp(args.lower_bound *
(1.0f / (1.0f + exp(-exp(a_log[head]) * gate))));
}
if (tid == 0u) {
beta_shared[0] =
1.0f / (1.0f + exp(-raw_beta[(ulong)row * args.n_heads + head]));
}
threadgroup_barrier(mem_flags::mem_threadgroup |
mem_flags::mem_device);
float q_sumsq = sq[tid] * sq[tid];
float k_sumsq = sk[tid] * sk[tid];
q_sumsq = simd_sum(q_sumsq);
k_sumsq = simd_sum(k_sumsq);
if (lane == 0u) {
reduce_q[sg] = q_sumsq;
reduce_k[sg] = k_sumsq;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float q_total = lane < 4u ? reduce_q[lane] : 0.0f;
float k_total = lane < 4u ? reduce_k[lane] : 0.0f;
q_total = simd_sum(q_total);
k_total = simd_sum(k_total);
const float q_scale = rsqrt(q_total + 1.0e-6f) * 0x1.6a09e6p-4f;
const float k_scale = rsqrt(k_total + 1.0e-6f);
if (tid < D) {
sq[tid] *= q_scale;
sk[tid] *= k_scale;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
const uint k0 = lane * 4u;
const float4 q4 = *((threadgroup float4 *)(sq + k0));
const float4 k4 = *((threadgroup float4 *)(sk + k0));
const float4 decay4 = *((threadgroup float4 *)(sd + k0));
const ulong state_head =
((ulong)row * args.n_heads + head) * D * D;
for (uint value = sg; value < D; value += 4u) {
device float4 *hptr =
(device float4 *)(state + state_head + (ulong)value * D + k0);
float4 h = *hptr * decay4;
float hk = dot(h, k4);
hk = simd_sum(hk);
const float delta_v = (sv[value] - hk) * beta_shared[0];
h = fma(k4, float4(delta_v), h);
*hptr = h;
float hq = simd_sum(dot(h, q4));
if (lane == 0u) so[value] = hq;
}
threadgroup_barrier(mem_flags::mem_threadgroup |
mem_flags::mem_device);
float o_sumsq = so[tid] * so[tid];
o_sumsq = simd_sum(o_sumsq);
if (lane == 0u) reduce_o[sg] = o_sumsq;
threadgroup_barrier(mem_flags::mem_threadgroup);
float o_total = lane < 4u ? reduce_o[lane] : 0.0f;
o_total = simd_sum(o_total);
const float o_scale = rsqrt(o_total / (float)D + args.norm_eps);
if (tid < D) {
const ulong index = input_base + tid;
const float gate =
1.0f / (1.0f + exp(-output_gate[index]));
out[index] = so[tid] * o_scale * output_norm[tid] * gate;
}
}
kernel void kernel_glm53_kda_prefill_prepare(
constant glm53_kda_args &args,
device float *q,
device float *k,
device float *v,
device float *raw_gate,
device const float *q_conv,
device const float *k_conv,
device const float *v_conv,
device const float *a_log,
device const float *dt_bias,
device float *conv_state,
threadgroup float *scratch [[threadgroup(0)]],
uint head [[threadgroup_position_in_grid]],
ushort tid [[thread_index_in_threadgroup]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]]) {
constexpr uint D = 128u;
constexpr uint HISTORY = 3u;
if (head >= args.n_heads) return;
threadgroup float *sq = scratch;
threadgroup float *sk = sq + D;
threadgroup float *reduce_q = sk + D;
threadgroup float *reduce_k = reduce_q + 4u;
const uint projection = args.n_heads * D;
const uint channel = head * D + tid;
device float *q_state = conv_state;
device float *k_state = q_state + HISTORY * projection;
device float *v_state = k_state + HISTORY * projection;
for (uint token = 0; token < args.n_rows; token++) {
const ulong index = (ulong)token * projection + channel;
float q_acc = 0.0f;
float k_acc = 0.0f;
float v_acc = 0.0f;
for (uint w = 0; w < HISTORY; w++) {
q_acc = fma(q_state[(ulong)w * projection + channel],
q_conv[(ulong)channel * 4u + w], q_acc);
k_acc = fma(k_state[(ulong)w * projection + channel],
k_conv[(ulong)channel * 4u + w], k_acc);
v_acc = fma(v_state[(ulong)w * projection + channel],
v_conv[(ulong)channel * 4u + w], v_acc);
}
const float q_new = q[index];
const float k_new = k[index];
const float v_new = v[index];
q_acc = fma(q_new, q_conv[(ulong)channel * 4u + 3u], q_acc);
k_acc = fma(k_new, k_conv[(ulong)channel * 4u + 3u], k_acc);
v_acc = fma(v_new, v_conv[(ulong)channel * 4u + 3u], v_acc);
q_state[channel] = q_state[projection + channel];
q_state[projection + channel] = q_state[2ul * projection + channel];
q_state[2ul * projection + channel] = q_new;
k_state[channel] = k_state[projection + channel];
k_state[projection + channel] = k_state[2ul * projection + channel];
k_state[2ul * projection + channel] = k_new;
v_state[channel] = v_state[projection + channel];
v_state[projection + channel] = v_state[2ul * projection + channel];
v_state[2ul * projection + channel] = v_new;
sq[tid] = q_acc / (1.0f + exp(-q_acc));
sk[tid] = k_acc / (1.0f + exp(-k_acc));
v[index] = v_acc / (1.0f + exp(-v_acc));
const float gate = raw_gate[index] + dt_bias[channel];
raw_gate[index] = exp(args.lower_bound *
(1.0f / (1.0f + exp(-exp(a_log[head]) * gate))));
threadgroup_barrier(mem_flags::mem_threadgroup |
mem_flags::mem_device);
float q_sumsq = simd_sum(sq[tid] * sq[tid]);
float k_sumsq = simd_sum(sk[tid] * sk[tid]);
if (lane == 0u) {
reduce_q[sg] = q_sumsq;
reduce_k[sg] = k_sumsq;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float q_total = lane < 4u ? reduce_q[lane] : 0.0f;
float k_total = lane < 4u ? reduce_k[lane] : 0.0f;
q_total = simd_sum(q_total);
k_total = simd_sum(k_total);
q[index] = sq[tid] * rsqrt(q_total + 1.0e-6f) *
0x1.6a09e6p-4f;
k[index] = sk[tid] * rsqrt(k_total + 1.0e-6f);
threadgroup_barrier(mem_flags::mem_threadgroup |
mem_flags::mem_device);
}
}
kernel void kernel_glm53_kda_prefill_recurrence(
constant glm53_kda_args &args,
device const float *q,
device const float *k,
device const float *v,
device const float *decay,
device const float *raw_beta,
device float *state,
device float *out,
uint2 tgpig [[threadgroup_position_in_grid]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]]) {
constexpr uint D = 128u;
const uint head = tgpig.x;
const uint value = tgpig.y * 4u + sg;
if (head >= args.n_heads || value >= D) return;
const uint projection = args.n_heads * D;
const uint k0 = lane * 4u;
device float4 *state_ptr = (device float4 *)(
state + ((ulong)head * D + value) * D + k0);
float4 h = *state_ptr;
for (uint token = 0; token < args.n_rows; token++) {
const ulong base = (ulong)token * projection + head * D;
const float4 q4 = *((device const float4 *)(q + base + k0));
const float4 k4 = *((device const float4 *)(k + base + k0));
const float4 decay4 =
*((device const float4 *)(decay + base + k0));
h *= decay4;
const float hk = simd_sum(dot(h, k4));
const float beta = 1.0f /
(1.0f + exp(-raw_beta[(ulong)token * args.n_heads + head]));
const float delta_v = (v[base + value] - hk) * beta;
h = fma(k4, float4(delta_v), h);
const float result = simd_sum(dot(h, q4));
if (lane == 0u) out[base + value] = result;
}
*state_ptr = h;
}
kernel void kernel_glm53_kda_prefill_output(
constant glm53_kda_args &args,
device float *out,
device const float *output_gate,
device const float *output_norm,
threadgroup float *partial [[threadgroup(0)]],
uint2 tgpig [[threadgroup_position_in_grid]],
ushort tid [[thread_index_in_threadgroup]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]]) {
constexpr uint D = 128u;
const uint token = tgpig.x;
const uint head = tgpig.y;
if (token >= args.n_rows || head >= args.n_heads) return;
const uint projection = args.n_heads * D;
const ulong base = (ulong)token * projection + head * D;
const float raw = out[base + tid];
float sumsq = simd_sum(raw * raw);
if (lane == 0u) partial[sg] = sumsq;
threadgroup_barrier(mem_flags::mem_threadgroup);
float total = lane < 4u ? partial[lane] : 0.0f;
total = simd_sum(total);
const float scale = rsqrt(total / (float)D + args.norm_eps);
out[base + tid] = raw * scale * output_norm[tid] /
(1.0f + exp(-output_gate[base + tid]));
}

271
metal/glm53_vision.metal Normal file
View File

@@ -0,0 +1,271 @@
// GLM-5.3 Flash vision operations not covered by the shared BF16 matmuls.
struct glm53_vision_rows_args {
uint width;
uint rows;
float eps;
};
struct glm53_vision_qkv_args {
uint rows;
uint grid_h;
uint grid_w;
float eps;
};
struct glm53_vision_attention_args {
uint rows;
float scale;
};
struct glm53_vision_scatter_args {
uint dst_row;
uint image_row;
uint rows;
uint total_rows;
uint width;
uint hc;
};
static inline float glm53_vision_erf(float x) {
const float sign = x < 0.0f ? -1.0f : 1.0f;
const float a = abs(x);
const float t = 1.0f / (1.0f + 0.3275911f * a);
const float p = (((((1.061405429f * t - 1.453152027f) * t) +
1.421413741f) * t - 0.284496736f) * t +
0.254829592f) * t;
return sign * (1.0f - p * exp(-a * a));
}
kernel void kernel_glm53_vision_add_bias(
constant glm53_vision_rows_args &args,
device float *x,
device const ushort *bias,
uint2 gid [[thread_position_in_grid]]) {
if (gid.x >= args.width || gid.y >= args.rows) return;
x[(ulong)gid.y * args.width + gid.x] += glm53_bf16_to_f32(bias[gid.x]);
}
kernel void kernel_glm53_vision_rms_bf16(
constant glm53_vision_rows_args &args,
device const float *x,
device const ushort *weight,
device float *out,
threadgroup float *partial,
uint row [[threadgroup_position_in_grid]],
uint tid [[thread_index_in_threadgroup]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]],
ushort nsg [[simdgroups_per_threadgroup]]) {
if (row >= args.rows) return;
device const float *xr = x + (ulong)row * args.width;
device float *yr = out + (ulong)row * args.width;
float sum = 0.0f;
for (uint d = tid; d < args.width; d += 256u) sum = fma(xr[d], xr[d], sum);
sum = simd_sum(sum);
if (lane == 0u) partial[sg] = sum;
threadgroup_barrier(mem_flags::mem_threadgroup);
if (sg == 0u) {
float v = lane < nsg ? partial[lane] : 0.0f;
v = simd_sum(v);
if (lane == 0u) partial[0] = rsqrt(v / (float)args.width + args.eps);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
const float inv = partial[0];
for (uint d = tid; d < args.width; d += 256u) {
yr[d] = xr[d] * inv * glm53_bf16_to_f32(weight[d]);
}
}
kernel void kernel_glm53_vision_qkv_rope(
constant glm53_vision_qkv_args &args,
device const float *qkv,
device const ushort *bias,
device const ushort *q_weight,
device const ushort *k_weight,
device float *q,
device float *k,
device float *v,
uint2 group [[threadgroup_position_in_grid]],
ushort lane [[thread_index_in_simdgroup]]) {
const uint row = group.x;
const uint head = group.y;
if (row >= args.rows || head >= 16u) return;
const ulong qkv_base = (ulong)row * 3072u + (ulong)head * 64u;
const ulong out_base = (ulong)row * 1024u + (ulong)head * 64u;
float q0 = qkv[qkv_base + lane] + glm53_bf16_to_f32(bias[(ulong)head * 64u + lane]);
float q1 = qkv[qkv_base + lane + 32u] +
glm53_bf16_to_f32(bias[(ulong)head * 64u + lane + 32u]);
float k0 = qkv[qkv_base + 1024u + lane] +
glm53_bf16_to_f32(bias[1024u + (ulong)head * 64u + lane]);
float k1 = qkv[qkv_base + 1024u + lane + 32u] +
glm53_bf16_to_f32(bias[1024u + (ulong)head * 64u + lane + 32u]);
const float qsum = simd_sum(fma(q0, q0, q1 * q1));
const float ksum = simd_sum(fma(k0, k0, k1 * k1));
const float qinv = rsqrt(qsum / 64.0f + args.eps);
const float kinv = rsqrt(ksum / 64.0f + args.eps);
q0 *= qinv * glm53_bf16_to_f32(q_weight[lane]);
q1 *= qinv * glm53_bf16_to_f32(q_weight[lane + 32u]);
k0 *= kinv * glm53_bf16_to_f32(k_weight[lane]);
k1 *= kinv * glm53_bf16_to_f32(k_weight[lane + 32u]);
const uint merge_w = args.grid_w / 2u;
const uint group_index = row / 4u;
const uint within = row & 3u;
const uint py = (group_index / merge_w) * 2u + within / 2u;
const uint px = (group_index % merge_w) * 2u + within % 2u;
const uint freq_index = lane & 15u;
const uint pos = lane < 16u ? py : px;
const float inv_freq = powr(10000.0f, -(float)freq_index / 16.0f);
const float angle = (float)pos * inv_freq;
const float cs = cos(angle);
const float sn = sin(angle);
q[out_base + lane] = q0 * cs - q1 * sn;
q[out_base + lane + 32u] = q1 * cs + q0 * sn;
k[out_base + lane] = k0 * cs - k1 * sn;
k[out_base + lane + 32u] = k1 * cs + k0 * sn;
v[out_base + lane] = qkv[qkv_base + 2048u + lane] +
glm53_bf16_to_f32(bias[2048u + (ulong)head * 64u + lane]);
v[out_base + lane + 32u] = qkv[qkv_base + 2048u + lane + 32u] +
glm53_bf16_to_f32(bias[2048u + (ulong)head * 64u + lane + 32u]);
}
/* A simdgroup owns one query/head and keeps its 64 output values in registers.
* This is quadratic in compute, as the model graph requires, but linear in
* memory and never materializes the attention matrix. */
kernel void kernel_glm53_vision_attention(
constant glm53_vision_attention_args &args,
device const float *q,
device const float *k,
device const float *v,
device float *out,
uint2 group [[threadgroup_position_in_grid]],
ushort lane [[thread_index_in_simdgroup]]) {
const uint row = group.x;
const uint head = group.y;
if (row >= args.rows || head >= 16u) return;
const ulong base = (ulong)row * 1024u + (ulong)head * 64u;
const float q0 = q[base + lane];
const float q1 = q[base + lane + 32u];
float acc0 = 0.0f, acc1 = 0.0f;
float max_score = -INFINITY;
float denom = 0.0f;
for (uint key_row = 0; key_row < args.rows; key_row++) {
const ulong kb = (ulong)key_row * 1024u + (ulong)head * 64u;
float score = simd_sum(q0 * k[kb + lane] + q1 * k[kb + lane + 32u]);
score *= args.scale;
const float next_max = max(max_score, score);
const float old_scale = max_score == -INFINITY ? 0.0f : exp(max_score - next_max);
const float new_scale = exp(score - next_max);
denom = denom * old_scale + new_scale;
acc0 = acc0 * old_scale + new_scale * v[kb + lane];
acc1 = acc1 * old_scale + new_scale * v[kb + lane + 32u];
max_score = next_max;
}
out[base + lane] = acc0 / denom;
out[base + lane + 32u] = acc1 / denom;
}
kernel void kernel_glm53_vision_bias_residual(
constant glm53_vision_rows_args &args,
device float *x,
device const ushort *bias,
device const float *residual,
uint2 gid [[thread_position_in_grid]]) {
if (gid.x >= args.width || gid.y >= args.rows) return;
const ulong off = (ulong)gid.y * args.width + gid.x;
x[off] += glm53_bf16_to_f32(bias[gid.x]) + residual[off];
}
kernel void kernel_glm53_vision_swiglu_bias(
constant glm53_vision_rows_args &args,
device const float *gate,
device const ushort *gate_bias,
device const float *up,
device const ushort *up_bias,
device float *out,
uint2 gid [[thread_position_in_grid]]) {
if (gid.x >= args.width || gid.y >= args.rows) return;
const ulong off = (ulong)gid.y * args.width + gid.x;
const float g = min(gate[off] + glm53_bf16_to_f32(gate_bias[gid.x]), 10.0f);
const float u = clamp(up[off] + glm53_bf16_to_f32(up_bias[gid.x]), -10.0f, 10.0f);
out[off] = (g / (1.0f + exp(-g))) * u;
}
kernel void kernel_glm53_vision_downsample_reorder(
device const float *x,
device float *out,
uint2 gid [[thread_position_in_grid]]) {
const uint d = gid.x;
const uint row = gid.y;
if (d >= 4096u) return;
const uint channel = d / 4u;
const uint within = d & 3u;
out[(ulong)row * 4096u + d] = x[((ulong)row * 4u + within) * 1024u + channel];
}
kernel void kernel_glm53_vision_layernorm_gelu(
constant glm53_vision_rows_args &args,
device const float *x,
device const ushort *weight,
device const ushort *bias,
device float *out,
threadgroup float *partial,
uint row [[threadgroup_position_in_grid]],
uint tid [[thread_index_in_threadgroup]],
ushort lane [[thread_index_in_simdgroup]],
ushort sg [[simdgroup_index_in_threadgroup]],
ushort nsg [[simdgroups_per_threadgroup]]) {
if (row >= args.rows) return;
device const float *xr = x + (ulong)row * args.width;
device float *yr = out + (ulong)row * args.width;
float sum = 0.0f;
for (uint d = tid; d < args.width; d += 256u) sum += xr[d];
sum = simd_sum(sum);
if (lane == 0u) partial[sg] = sum;
threadgroup_barrier(mem_flags::mem_threadgroup);
if (sg == 0u) {
float v = lane < nsg ? partial[lane] : 0.0f;
v = simd_sum(v);
if (lane == 0u) partial[0] = v / (float)args.width;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
const float mean = partial[0];
float var = 0.0f;
for (uint d = tid; d < args.width; d += 256u) {
const float centered = xr[d] - mean;
var = fma(centered, centered, var);
}
var = simd_sum(var);
if (lane == 0u) partial[sg] = var;
threadgroup_barrier(mem_flags::mem_threadgroup);
if (sg == 0u) {
float v = lane < nsg ? partial[lane] : 0.0f;
v = simd_sum(v);
if (lane == 0u) partial[0] = rsqrt(v / (float)args.width + args.eps);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
const float inv = partial[0];
const float inv_sqrt2 = 0.7071067811865475f;
for (uint d = tid; d < args.width; d += 256u) {
float v = (xr[d] - mean) * inv * glm53_bf16_to_f32(weight[d]) +
glm53_bf16_to_f32(bias[d]);
yr[d] = 0.5f * v * (1.0f + glm53_vision_erf(v * inv_sqrt2));
}
}
kernel void kernel_glm53_vision_scatter_hc(
constant glm53_vision_scatter_args &args,
device float *hc,
device const float *image,
uint2 gid [[thread_position_in_grid]]) {
const uint d = gid.x;
const uint linear_row = gid.y;
if (d >= args.width || linear_row >= args.rows * args.hc) return;
const uint image_delta = linear_row / args.hc;
const uint hc_index = linear_row % args.hc;
const ulong dst = ((ulong)(args.dst_row + image_delta) * args.hc + hc_index) *
args.width + d;
const ulong src = (ulong)(args.image_row + image_delta) * args.width + d;
if (args.dst_row + image_delta < args.total_rows) hc[dst] = image[src];
}

File diff suppressed because it is too large Load Diff

View File

@@ -241,3 +241,145 @@ kernel void kernel_dsv4_qkv_rms_norm_f32_4(
y[i] = (x[i] * scale) * w[i];
}
}
// Decode-only triple fusion: the q/kv RMS norm, the KV RoPE tail, and the
// FP8/raw finalizer were three back-to-back dispatches on the same rows.
// The q threadgroup is byte-identical to kernel_dsv4_qkv_rms_norm_f32_4.
// The kv threadgroup continues with the shared affine-row RoPE helper (lane
// mapping preserved: r == lane on the first 64 lanes) and a verbatim copy of
// kernel_dsv4_kv_fp8_store_f32 with its work predicated to the first 64
// lanes (barriers stay uniform across the whole threadgroup). Arithmetic,
// order and rounding are unchanged; gated and verified against
// full-vocabulary logits before promotion.
kernel void kernel_dsv4_qkv_rms_norm_kv_rope_fp8_store_f32(
constant ds4_metal_args_qkv_rms_norm & args,
constant ds4_metal_args_dsv4_rope_affine_pair & rope,
constant ds4_metal_args_dsv4_kv_fp8_store & store,
device const float4 * q_src,
device const float4 * q_weight,
device float4 * q_dst,
device const float4 * kv_src,
device const float4 * kv_weight,
device float4 * kv_dst,
device float * raw_cache,
threadgroup float * shmem_f32 [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort3 tpitg[[thread_position_in_threadgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort3 ntg[[threads_per_threadgroup]]) {
if (sgitg == 0) {
shmem_f32[tiisg] = 0.0f;
}
const uint row = tgpig.x;
const bool kv_task = tgpig.y != 0;
const int n = kv_task ? args.kv_n : args.q_n;
const int n4 = kv_task ? args.kv_n4 : args.q_n4;
const uint64_t row_stride4 = (kv_task ? args.kv_row_stride : args.q_row_stride) / sizeof(float4);
device const float4 * x = kv_task ? kv_src + row * row_stride4 : q_src + row * row_stride4;
device const float4 * w = kv_task ? kv_weight : q_weight;
device float4 * y = kv_task ? kv_dst + row * row_stride4 : q_dst + row * row_stride4;
float sumf = 0.0f;
for (int i = tpitg.x; i < n4; i += ntg.x) {
const float4 v = x[i];
sumf += dot(v, v);
}
sumf = simd_sum(sumf);
threadgroup_barrier(mem_flags::mem_threadgroup);
if (tiisg == 0) {
shmem_f32[sgitg] = sumf;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
sumf = shmem_f32[tiisg];
sumf = simd_sum(sumf);
#ifdef DS4_METAL_NORM_RSQRT_DISABLE
const float scale = 1.0f / sqrt(sumf / float(n) + args.eps);
#else
const float scale = rsqrt(sumf / float(n) + args.eps);
#endif
for (int i = tpitg.x; i < n4; i += ntg.x) {
y[i] = (x[i] * scale) * w[i];
}
if (!kv_task) {
return;
}
// KV RoPE tail in place, then the FP8/raw finalizer (verbatim bodies).
threadgroup_barrier(mem_flags::mem_device_and_threadgroup);
device char *kv_row = (device char *)(kv_dst + row * row_stride4);
const int rope_n_nope = rope.head_dim - rope.n_dims;
if (rope_n_nope < 0) {
return;
}
ds4_rope_tail_pair_affine_row(rope,
(device const char *)kv_row,
kv_row,
rope_n_nope,
rope.pos0,
tpitg.x,
ntg.x);
threadgroup_barrier(mem_flags::mem_device_and_threadgroup);
const int head_dim = store.head_dim;
const int n_rot = store.n_rot;
const int n_nope = head_dim - n_rot;
if (head_dim <= 0 || n_rot < 0 || n_nope < 0) {
return;
}
const uint tid = tpitg.x;
device float *kv = (device float *)kv_row;
device float *raw = raw_cache + (int64_t)store.raw_row * head_dim;
threadgroup float *scratch = shmem_f32 + 32;
for (int off = 0; off < n_nope; off += 64) {
float v = 0.0f;
if (tid < 64u && off + (int)tid < n_nope) {
v = kv[off + tid];
scratch[tid] = abs(v);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = 32; stride > 0; stride >>= 1) {
if (tid < stride) {
scratch[tid] = max(scratch[tid], scratch[tid + stride]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
const float amax = max(scratch[0], 1.0e-4f);
const float fp8_scale = exp2(ceil(log2(amax / 448.0f)));
if (tid < 64u && off + (int)tid < n_nope) {
const float q = dsv4_e4m3fn_dequant(clamp(v / fp8_scale, -448.0f, 448.0f)) * fp8_scale;
kv[off + tid] = q;
#ifdef DS4_METAL_KV_RAW_F32
raw[off + tid] = q;
#else
raw[off + tid] = (float)((half)q);
#endif
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (tid < 64u) {
for (int i = n_nope + tid; i < head_dim; i += 64) {
#ifdef DS4_METAL_KV_RAW_F32
raw[i] = kv[i];
#else
raw[i] = (float)((half)kv[i]);
#endif
}
}
}

3552
metal/qwen38.metal Normal file

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1 @@
ALTER TABLE preferences DROP COLUMN last_project_id;

View File

@@ -0,0 +1 @@
ALTER TABLE preferences ADD COLUMN last_project_id INTEGER;

View File

@@ -0,0 +1 @@
ALTER TABLE projects DROP COLUMN collapsed;

View File

@@ -0,0 +1,2 @@
ALTER TABLE projects ADD COLUMN collapsed BOOLEAN NOT NULL DEFAULT FALSE
CHECK (collapsed IN (FALSE, TRUE));

View File

@@ -0,0 +1 @@
ALTER TABLE preferences DROP COLUMN sidebar_width;

View File

@@ -0,0 +1 @@
ALTER TABLE preferences ADD COLUMN sidebar_width INTEGER NOT NULL DEFAULT 276;

View File

@@ -0,0 +1,4 @@
ALTER TABLE preferences DROP COLUMN kv_continued_interval_tokens;
ALTER TABLE preferences DROP COLUMN kv_cold_max_tokens;
ALTER TABLE preferences DROP COLUMN kv_min_tokens;
ALTER TABLE preferences DROP COLUMN kv_budget_gib;

View File

@@ -0,0 +1,4 @@
ALTER TABLE preferences ADD COLUMN kv_budget_gib BIGINT;
ALTER TABLE preferences ADD COLUMN kv_min_tokens INTEGER;
ALTER TABLE preferences ADD COLUMN kv_cold_max_tokens INTEGER;
ALTER TABLE preferences ADD COLUMN kv_continued_interval_tokens INTEGER;

View File

@@ -0,0 +1 @@
-- Preferences live in config.yaml; the table is not recreated.

View File

@@ -0,0 +1 @@
DROP TABLE preferences;

View File

@@ -0,0 +1 @@
ALTER TABLE sessions DROP COLUMN compacted_summary;

View File

@@ -0,0 +1 @@
ALTER TABLE sessions ADD COLUMN compacted_summary TEXT;

View File

@@ -0,0 +1 @@
ALTER TABLE messages DROP COLUMN system;

View File

@@ -0,0 +1 @@
ALTER TABLE messages ADD COLUMN system BOOLEAN NOT NULL DEFAULT 0 CHECK (system IN (0, 1));

View File

@@ -0,0 +1,2 @@
ALTER TABLE messages DROP COLUMN compaction_tail_start;
ALTER TABLE messages DROP COLUMN compaction;

View File

@@ -0,0 +1,2 @@
ALTER TABLE messages ADD COLUMN compaction BOOLEAN NOT NULL DEFAULT FALSE;
ALTER TABLE messages ADD COLUMN compaction_tail_start INTEGER;

View File

@@ -0,0 +1 @@
DROP TABLE a2ui_messages;

View File

@@ -0,0 +1,10 @@
CREATE TABLE a2ui_messages (
id INTEGER NOT NULL PRIMARY KEY AUTOINCREMENT,
session_id INTEGER NOT NULL,
message_id INTEGER NOT NULL,
json TEXT NOT NULL,
FOREIGN KEY (session_id) REFERENCES sessions(id),
FOREIGN KEY (message_id) REFERENCES messages(id)
);
CREATE INDEX a2ui_messages_session_id ON a2ui_messages(session_id, id);

View File

@@ -0,0 +1 @@
ALTER TABLE a2ui_messages DROP COLUMN dismissed;

View File

@@ -0,0 +1,2 @@
ALTER TABLE a2ui_messages
ADD COLUMN dismissed BOOLEAN NOT NULL DEFAULT 0;

View File

@@ -0,0 +1,4 @@
ALTER TABLE messages DROP COLUMN output_tokens;
ALTER TABLE messages DROP COLUMN cached_tokens;
ALTER TABLE messages DROP COLUMN input_tokens;
ALTER TABLE messages DROP COLUMN generation_duration_ms;

View File

@@ -0,0 +1,4 @@
ALTER TABLE messages ADD COLUMN generation_duration_ms INTEGER;
ALTER TABLE messages ADD COLUMN input_tokens INTEGER;
ALTER TABLE messages ADD COLUMN cached_tokens INTEGER;
ALTER TABLE messages ADD COLUMN output_tokens INTEGER;

View File

@@ -0,0 +1 @@
ALTER TABLE messages DROP COLUMN model_content;

View File

@@ -0,0 +1 @@
ALTER TABLE messages ADD COLUMN model_content TEXT;

View File

@@ -0,0 +1 @@
ALTER TABLE sessions DROP COLUMN last_used;

View File

@@ -0,0 +1,4 @@
ALTER TABLE sessions ADD COLUMN last_used BIGINT NOT NULL DEFAULT 0 CHECK (last_used >= 0);
UPDATE sessions
SET last_used = unixepoch();

View File

@@ -0,0 +1 @@
ALTER TABLE sessions DROP COLUMN permission_mode;

View File

@@ -0,0 +1,2 @@
ALTER TABLE sessions ADD COLUMN permission_mode TEXT NOT NULL DEFAULT 'heuristic'
CHECK (permission_mode IN ('heuristic', 'ai'));

View File

@@ -0,0 +1 @@
ALTER TABLE messages DROP COLUMN tool_approval_reasons;

View File

@@ -0,0 +1 @@
ALTER TABLE messages ADD COLUMN tool_approval_reasons TEXT;

View File

@@ -0,0 +1 @@
ALTER TABLE messages DROP COLUMN instruction_metadata;

View File

@@ -0,0 +1 @@
ALTER TABLE messages ADD COLUMN instruction_metadata TEXT;

View File

@@ -0,0 +1 @@
ALTER TABLE sessions DROP COLUMN model;

View File

@@ -0,0 +1,7 @@
ALTER TABLE sessions ADD COLUMN model TEXT
CHECK (model IS NULL OR model IN (
'deepseek-v4-flash-0731',
'deepseek-v4-pro',
'glm-5.2',
'glm-5.3-flash'
));

View File

@@ -0,0 +1,13 @@
ALTER TABLE sessions ADD COLUMN model_without_qwen TEXT
CHECK (model_without_qwen IS NULL OR model_without_qwen IN (
'deepseek-v4-flash-0731',
'deepseek-v4-pro',
'glm-5.2',
'glm-5.3-flash'
));
UPDATE sessions SET model_without_qwen = CASE
WHEN model = 'qwen3.8-flash-next' THEN NULL
ELSE model
END;
ALTER TABLE sessions DROP COLUMN model;
ALTER TABLE sessions RENAME COLUMN model_without_qwen TO model;

View File

@@ -0,0 +1,11 @@
ALTER TABLE sessions ADD COLUMN model_with_qwen TEXT
CHECK (model_with_qwen IS NULL OR model_with_qwen IN (
'deepseek-v4-flash-0731',
'deepseek-v4-pro',
'glm-5.2',
'glm-5.3-flash',
'qwen3.8-flash-next'
));
UPDATE sessions SET model_with_qwen = model;
ALTER TABLE sessions DROP COLUMN model;
ALTER TABLE sessions RENAME COLUMN model_with_qwen TO model;

66
native/media/ds4_media.m Normal file
View File

@@ -0,0 +1,66 @@
#import <AppKit/AppKit.h>
#import <AVKit/AVKit.h>
#include <stdbool.h>
@interface DS4MediaWindowController : NSWindowController <NSWindowDelegate>
@end
static NSMutableSet<DS4MediaWindowController *> *DS4MediaWindows(void) {
static NSMutableSet<DS4MediaWindowController *> *windows;
static dispatch_once_t once;
dispatch_once(&once, ^{
windows = [NSMutableSet set];
});
return windows;
}
@implementation DS4MediaWindowController
- (void)windowWillClose:(NSNotification *)notification {
AVPlayerView *view = (AVPlayerView *)self.window.contentView;
[view.player pause];
[DS4MediaWindows() removeObject:self];
}
@end
bool ds4_media_open(const char *url_bytes, const char *title_bytes, bool video) {
if (url_bytes == NULL || title_bytes == NULL) {
return false;
}
NSString *url_string = [NSString stringWithUTF8String:url_bytes];
NSString *title = [NSString stringWithUTF8String:title_bytes];
NSURL *url = [NSURL URLWithString:url_string];
if (url == nil || title == nil) {
return false;
}
dispatch_async(dispatch_get_main_queue(), ^{
NSRect frame = NSMakeRect(0, 0, video ? 800 : 560, video ? 500 : 180);
NSWindow *window = [[NSWindow alloc]
initWithContentRect:frame
styleMask:NSWindowStyleMaskTitled |
NSWindowStyleMaskClosable |
NSWindowStyleMaskMiniaturizable |
NSWindowStyleMaskResizable
backing:NSBackingStoreBuffered
defer:NO];
window.title = title.length > 0 ? title : url.lastPathComponent;
window.minSize = video ? NSMakeSize(480, 300) : NSMakeSize(420, 150);
AVPlayerView *view = [[AVPlayerView alloc] initWithFrame:frame];
view.autoresizingMask = NSViewWidthSizable | NSViewHeightSizable;
view.player = [AVPlayer playerWithURL:url];
window.contentView = view;
DS4MediaWindowController *controller =
[[DS4MediaWindowController alloc] initWithWindow:window];
window.delegate = controller;
[DS4MediaWindows() addObject:controller];
[controller showWindow:nil];
[window center];
[window makeKeyAndOrderFront:nil];
[view.player play];
});
return true;
}

View File

@@ -1,15 +1,14 @@
# Vendored DS4 Metal boundary
These files are a one-time snapshot of the DS4 Metal boundary from commit
`efdadd41e20134af4f3381e1ed90e96fe4faef6f`:
The Metal bridge and GPU declarations track the DS4 boundary audited at commit
`ec7642cdd9ec81d01ad4b1fd8f8a3d1511533748` (2026-08-31):
- `ds4_metal.m`
- `ds4.h`
- `ds4_gpu.h`
- `ds4_ssd.h`
- `LICENSE`
The matching Metal kernels live in the repository-level `metal/` directory.
The matching Metal kernels, including `glm53_*.metal`, live in the
repository-level `metal/` directory. The other headers remain DS4Server-owned
interfaces and are not copies of DS4 application code.
DS4Server builds and bundles this local snapshot; it does not read a sibling
DS4 checkout. Rust owns the model, graph, session, sampling, and lifecycle.
Objective-C remains only at the platform Metal boundary.

View File

@@ -126,13 +126,10 @@ typedef struct {
typedef struct {
const char *model_path;
const char *mtp_path;
ds4_backend backend;
int n_threads;
int context_size;
uint32_t prefill_chunk;
int mtp_draft_tokens;
float mtp_margin;
float dspark_confidence_threshold;
const char *directional_steering_file;
const char *expert_profile_path;
@@ -150,6 +147,7 @@ typedef struct {
bool glm_mtp_timing;
bool dspark;
bool dspark_strict;
bool dspark_exact_sampling;
bool dspark_confidence_threshold_set;
bool cuda_tensor_parallel;
bool ssd_streaming;
@@ -158,6 +156,8 @@ typedef struct {
bool inspect_only;
/* Multi-GPU placement uses this to price per-layer KV storage. */
int placement_ctx_hint;
/* Number of independently allocated session graphs/caches to reserve. */
int placement_session_count_hint;
/* Server batch mode serializes execution and can share prefill scratch. */
bool share_session_prefill_workspace;
bool first_token_test;
@@ -371,6 +371,30 @@ int ds4_test_sample_logits(const float *logits, uint32_t n_vocab,
float temperature, int top_k,
float top_p, float min_p, uint64_t *rng,
float *prob_scratch);
int ds4_test_sampling_probabilities(const float *logits, uint32_t n_vocab,
float temperature, int top_k,
float top_p, float min_p, float *probs);
int ds4_test_speculative_sample(const float *target_logits,
const float *draft_logits,
uint32_t n_vocab,
float temperature,
int top_k,
float top_p,
float min_p,
uint64_t *rng,
float *target_probs,
float *draft_probs);
int ds4_test_speculative_delta_sample(const float *target_logits,
uint32_t n_vocab,
int draft_token,
float temperature,
int top_k,
float top_p,
float min_p,
uint64_t *rng,
float *target_probs);
int ds4_test_argmax_excluding_logits(const float *logits, uint32_t n_vocab,
int excluded_id);
uint64_t ds4_test_mixed_native_count(void);
#endif
int ds4_session_top_logprobs(ds4_session *s, ds4_token_score *out, int k);
@@ -403,6 +427,15 @@ int ds4_session_eval_speculative_argmax(ds4_session *s, int first_token,
int max_tokens, int eos_token,
int *accepted, int accepted_cap,
char *err, size_t errlen);
/* Evaluate one already-sampled target token and speculatively extend it.
* Positive-temperature DSpark normally commits greedily verified draft
* tokens; dspark_exact_sampling selects exact stochastic p/q acceptance. */
int ds4_session_eval_speculative(ds4_session *s, int first_token,
int max_tokens, int eos_token,
float temperature, int top_k,
float top_p, float min_p, uint64_t *rng,
int *accepted, int accepted_cap,
char *err, size_t errlen);
/* TP worker side of a mirrored speculative-verify block: run its half of the
* batch verify for KV side effects, then obey the leader's commit frame
* (keep, or roll back and replay). Only called from ds4_tp_worker_run. */
@@ -415,8 +448,6 @@ int ds4_session_ctx(ds4_session *s);
int ds4_session_prefill_cap(ds4_session *s);
int ds4_engine_routed_quant_bits(ds4_engine *e);
bool ds4_engine_has_output_head(ds4_engine *e);
bool ds4_engine_has_mtp(ds4_engine *e);
int ds4_engine_mtp_draft_tokens(ds4_engine *e);
const ds4_tokens *ds4_session_tokens(ds4_session *s);
/* Low-level graph slice entry points used by distributed inference. The

View File

@@ -43,8 +43,10 @@ typedef struct {
int ds4_gpu_init(void);
void ds4_gpu_cleanup(void);
int ds4_gpu_metal4_tensor_api_enabled(void);
ds4_gpu_tensor *ds4_gpu_tensor_alloc(uint64_t bytes);
ds4_gpu_tensor *ds4_gpu_tensor_alloc_untracked(uint64_t bytes);
ds4_gpu_tensor *ds4_gpu_tensor_alloc_managed(uint64_t bytes);
ds4_gpu_tensor *ds4_gpu_tensor_view(const ds4_gpu_tensor *base, uint64_t offset, uint64_t bytes);
void ds4_gpu_tensor_free(ds4_gpu_tensor *tensor);
@@ -78,6 +80,24 @@ int ds4_gpu_begin_commands(void);
int ds4_gpu_flush_encoder(void);
int ds4_gpu_flush_commands(void);
int ds4_gpu_commands_active(void);
#ifdef __APPLE__
int ds4_gpu_parallel_ffn_finish(void);
void ds4_gpu_parallel_ffn_abort(void);
int ds4_gpu_parallel_ffn_start(
ds4_gpu_tensor *gate,
ds4_gpu_tensor *up,
ds4_gpu_tensor *mid,
ds4_gpu_tensor *shared_out,
const void *model_map,
uint64_t model_size,
uint64_t gate_offset,
uint64_t up_offset,
uint64_t down_offset,
uint32_t model_dim,
uint32_t shared_dim,
const ds4_gpu_tensor *x,
float clamp);
#endif
int ds4_gpu_signal_selected_readback_ready(uint64_t *event_value);
int ds4_gpu_commit_and_wait_selected_readback(uint64_t event_value, const char *label);
int ds4_gpu_wait_selected_readback_ready(uint64_t event_value, const char *label);
@@ -95,8 +115,47 @@ int ds4_gpu_synchronize(void);
int ds4_gpu_set_model_map(const void *model_map, uint64_t model_size);
int ds4_gpu_set_model_fd(int fd);
int ds4_gpu_set_model_fd_for_map(int fd, const void *model_map);
int ds4_gpu_build_derived_artifacts(const void *model_map, uint64_t model_size,
const char *model_path);
int ds4_gpu_model_range_replaced(const void *model_map, uint64_t offset,
uint64_t bytes);
int ds4_gpu_set_model_map_range(const void *model_map, uint64_t model_size, uint64_t map_offset, uint64_t map_size, uint64_t max_tensor_bytes);
int ds4_gpu_set_transient_model_map_range(const void *model_map, uint64_t model_size, uint64_t map_offset, uint64_t map_size, uint64_t max_tensor_bytes);
/* Caller must finish every command that references this mapping first. */
int ds4_gpu_release_transient_model_map(const void *model_map, uint64_t model_size);
int ds4_gpu_model_map_active(const void *model_map, uint64_t model_size);
/* Add a secondary GGUF mapping without replacing the primary model mapping. */
int ds4_gpu_set_aux_model_map_range(const void *model_map,
uint64_t model_size,
uint64_t map_offset,
uint64_t map_size);
int ds4_gpu_set_model_map_spans(const void *model_map, uint64_t model_size, const uint64_t *offsets, const uint64_t *sizes, uint32_t count, uint64_t max_tensor_bytes);
typedef struct {
const ds4_gpu_tensor *tensor;
const void *map;
uint64_t size;
uint64_t offset;
uint64_t bytes;
} ds4_gpu_qwen_weight_view;
typedef struct {
uint32_t u[16];
float f[8];
} ds4_gpu_qwen_kernel_args;
int ds4_gpu_qwen_dispatch(
const char *kernel,
ds4_gpu_tensor *out,
const ds4_gpu_tensor *a,
const ds4_gpu_tensor *b,
const ds4_gpu_tensor *c,
const ds4_gpu_tensor *d,
const ds4_gpu_qwen_weight_view *weights,
uint32_t weight_count,
const ds4_gpu_qwen_kernel_args *args,
uint32_t grid_x,
uint32_t grid_y);
int ds4_gpu_cache_model_range(const void *model_map, uint64_t model_size, uint64_t offset, uint64_t bytes, const char *label);
int ds4_gpu_cache_q8_f16_range(const void *model_map, uint64_t model_size, uint64_t offset, uint64_t bytes, uint64_t in_dim, uint64_t out_dim, const char *label);
int ds4_gpu_q8_cache_suppressed(void);
@@ -138,16 +197,56 @@ int ds4_gpu_preload_q4_expert_tables(const void *model_map, uint64_t model_size,
int ds4_gpu_should_use_managed_kv_cache(uint64_t kv_cache_bytes, uint64_t context_bytes);
void ds4_gpu_set_quality(bool quality);
void ds4_gpu_set_glm_model(bool enabled);
void ds4_gpu_set_qwen_model(bool enabled);
void ds4_gpu_set_ssd_streaming(bool enabled);
void ds4_gpu_set_glm_streaming_prefill_full_layer(bool enabled);
#ifdef __APPLE__
int ds4_gpu_device_is_pre_m5_apple_silicon(void);
int ds4_gpu_device_is_m5_apple_silicon(void);
int ds4_gpu_set_decode_pipeline_fast_lookup(int enabled);
/* Strict test oracle for the fixed decode mul_mv pipeline lookup cache. */
int ds4_gpu_test_decode_pipeline_fast_lookup(void);
/* Strict test oracle for the extended decode mul_mv_ext (nsg + nxpsg) cache. */
int ds4_gpu_test_decode_pipeline_fast_lookup_ext(void);
/* Strict test oracle for the generated resident-prefill MXFP4 half LUT. */
int ds4_gpu_test_mxfp4_down_half_lut(uint16_t *legacy_bits,
uint16_t *lut_bits);
enum {
DS4_GPU_TEST_MXFP4_PAIR_TAIL_CULL = 1u << 0,
DS4_GPU_TEST_MXFP4_PAIR_COMPACT_TILE = 1u << 1,
DS4_GPU_TEST_MXFP4_MAP_SCATTER = 1u << 2,
DS4_GPU_TEST_MXFP4_DOWN_TAIL_CULL = 1u << 3,
DS4_GPU_TEST_MXFP4_DOWN_HALF_LUT = 1u << 4,
DS4_GPU_TEST_OUTPUT_HC_WEIGHTS4 = 1u << 5,
DS4_GPU_TEST_HC_RMS_SCALE_PROJ = 1u << 6,
};
void ds4_gpu_test_set_flags(uint32_t flags);
void ds4_gpu_release_zero_prefix_prefill_mask_cache(void);
#else
static inline int ds4_gpu_device_is_pre_m5_apple_silicon(void) { return 0; }
static inline int ds4_gpu_device_is_m5_apple_silicon(void) { return 0; }
#endif
void ds4_gpu_set_streaming_expert_cache_budget(uint32_t experts);
void ds4_gpu_set_streaming_expert_cache_expert_bytes(uint64_t bytes);
uint64_t ds4_gpu_recommended_working_set_size(void);
uint32_t ds4_gpu_stream_expert_cache_configured_count(void);
uint32_t ds4_gpu_stream_expert_cache_current_count(void);
typedef struct ds4_gpu_stream_expert_cache_stats {
uint32_t configured_count;
uint32_t current_count;
uint64_t hits;
uint64_t misses;
uint64_t evictions;
uint64_t wraps;
uint64_t buffer_allocs;
uint64_t buffer_reuses;
uint64_t evict_advise_bytes;
uint64_t willneed_advise_bytes;
uint64_t pread_bytes;
double pread_ms;
} ds4_gpu_stream_expert_cache_stats;
void ds4_gpu_stream_expert_cache_get_stats(
ds4_gpu_stream_expert_cache_stats *stats);
typedef struct ds4_gpu_stream_expert_table {
const void *model_map;
uint64_t model_size;
@@ -200,6 +299,7 @@ int ds4_gpu_stream_expert_cache_seed_from_layer_selected(
uint32_t n_tokens,
uint32_t n_seed_tokens,
uint32_t n_selected);
int ds4_gpu_stream_expert_cache_finish_pending_batch(void);
int ds4_gpu_stream_expert_cache_release_layer_cache(void);
#endif
int ds4_gpu_stream_expert_cache_seed_experts(
@@ -207,6 +307,14 @@ int ds4_gpu_stream_expert_cache_seed_experts(
const int32_t *expert_ids,
const uint32_t *expert_priorities,
uint32_t n_experts);
#ifdef __APPLE__
/* Seed from mapped weights with blits appended to the active command buffer. */
int ds4_gpu_stream_expert_cache_seed_experts_gpu_copy(
const ds4_gpu_stream_expert_table *table,
const int32_t *expert_ids,
const uint32_t *expert_priorities,
uint32_t n_experts);
#endif
void ds4_gpu_print_memory_report(const char *label);
/* Tensor-parallel per-layer gates (Metal only). The encoder calls
@@ -598,6 +706,17 @@ int ds4_gpu_matmul_q8_0_pair_tensor(
const ds4_gpu_tensor *x,
uint64_t n_tok);
int ds4_gpu_matmul_q4_K_pair_decode_tensor(
ds4_gpu_tensor *out0,
ds4_gpu_tensor *out1,
const void *model_map,
uint64_t model_size,
uint64_t weight0_offset,
uint64_t weight1_offset,
uint64_t in_dim,
uint64_t out_dim,
const ds4_gpu_tensor *x);
/* Multi-row decode projections that preserve the one-row reduction order. */
int ds4_gpu_matmul_q8_0_decode_rows_exact_tensor(
ds4_gpu_tensor *out,
@@ -643,6 +762,36 @@ int ds4_gpu_shared_gate_up_swiglu_q8_0_tensor(
uint64_t out_dim,
const ds4_gpu_tensor *x,
float clamp);
int ds4_gpu_router_shared_gate_up_q8_0_tensor(
ds4_gpu_tensor *router_logits,
ds4_gpu_tensor *gate,
ds4_gpu_tensor *up,
ds4_gpu_tensor *mid,
const void *model_map,
uint64_t model_size,
uint64_t router_weight_offset,
uint64_t gate_offset,
uint64_t up_offset,
uint64_t in_dim,
uint64_t router_out_dim,
uint64_t out_dim,
const ds4_gpu_tensor *x,
float clamp,
bool router_only);
#ifdef __APPLE__
int ds4_gpu_router_project_select_fused_tensor(
ds4_gpu_tensor *router_logits,
ds4_gpu_tensor *probs,
ds4_gpu_tensor *selected,
ds4_gpu_tensor *weights,
const void *model_map,
uint64_t model_size,
uint64_t router_weight_offset,
uint64_t bias_offset,
bool has_bias,
const ds4_gpu_tensor *x);
#endif
int ds4_gpu_shared_mid_swiglu_q8_0_decode_exact_tensor(
ds4_gpu_tensor *mid,
const void *model_map,
@@ -720,6 +869,20 @@ int ds4_gpu_matmul_f16_tensor(
const ds4_gpu_tensor *x,
uint64_t n_tok);
/* CUDA batch path: fold an input RMS normalization into the FP16 activation
* conversion used by the following projection. Returns 0 without touching
* out when the optimized path is unavailable. */
int ds4_gpu_matmul_f16_rms_fold_tensor(
ds4_gpu_tensor *out,
const void *model_map,
uint64_t model_size,
uint64_t weight_offset,
uint64_t in_dim,
uint64_t out_dim,
const ds4_gpu_tensor *x,
uint64_t n_tok,
float norm_eps);
/* Exact multi-row form of the DeepSeek 4096x256 F16 router projection. */
int ds4_gpu_matmul_f16_router_rows_exact_tensor(
ds4_gpu_tensor *out,
@@ -761,6 +924,95 @@ int ds4_gpu_matmul_f16_pair_compressor_store_tensor(
uint32_t ratio,
uint32_t pos);
int ds4_gpu_matmul_f16_quad_compressor_store_tensor(
ds4_gpu_tensor *out0_kv,
ds4_gpu_tensor *out0_score,
ds4_gpu_tensor *out1_kv,
ds4_gpu_tensor *out1_score,
ds4_gpu_tensor *state0_kv,
ds4_gpu_tensor *state0_score,
ds4_gpu_tensor *state1_kv,
ds4_gpu_tensor *state1_score,
const void *model_map,
uint64_t model_size,
uint64_t weight0_kv_offset,
uint64_t weight0_score_offset,
uint64_t weight1_kv_offset,
uint64_t weight1_score_offset,
uint64_t ape0_offset,
uint32_t ape0_type,
uint64_t ape1_offset,
uint32_t ape1_type,
uint64_t in_dim,
uint32_t width0,
uint32_t width1,
const ds4_gpu_tensor *x,
uint32_t ratio,
uint32_t pos);
/* Decode-only M5 fusion: emit-path compressor row finalize (norm + rope +
* fp8/commit + indexer qat) in one dispatch. Bit-exact vs the separate
* dispatches. Returns 1 when fused, 0 to fall back. */
int ds4_gpu_dsv4_comp_row_finalize_tensor(
ds4_gpu_tensor *attn_stage,
ds4_gpu_tensor *attn_cache,
uint32_t attn_comp_row,
uint64_t attn_norm_offset,
ds4_gpu_tensor *index_cache,
uint32_t index_comp_row,
uint64_t index_norm_offset,
ds4_gpu_tensor *attn_state_kv,
ds4_gpu_tensor *attn_state_score,
ds4_gpu_tensor *index_state_kv,
ds4_gpu_tensor *index_state_score,
const void *model_map,
uint64_t model_size,
uint32_t pos,
uint32_t n_rot,
uint32_t n_ctx_orig,
float freq_base,
float freq_scale,
float ext_factor,
float attn_factor,
float beta_fast,
float beta_slow,
float rms_eps);
/* Decode-only M5 fusion: q_a/kv Q8 pair projection + F16 quad compressor
* projection/store in one dispatch. Bit-exact vs the separate dispatches.
* Returns 1 when fused, 0 to fall back, -1 on error. */
int ds4_gpu_qkv_pair_quad_compressor_store_tensor(
ds4_gpu_tensor *qr,
ds4_gpu_tensor *kv_raw,
ds4_gpu_tensor *out0_kv,
ds4_gpu_tensor *out0_score,
ds4_gpu_tensor *out1_kv,
ds4_gpu_tensor *out1_score,
ds4_gpu_tensor *state0_kv,
ds4_gpu_tensor *state0_score,
ds4_gpu_tensor *state1_kv,
ds4_gpu_tensor *state1_score,
const void *model_map,
uint64_t model_size,
uint64_t q_a_offset,
uint64_t kv_offset,
uint64_t weight0_kv_offset,
uint64_t weight0_score_offset,
uint64_t weight1_kv_offset,
uint64_t weight1_score_offset,
uint64_t ape0_offset,
uint32_t ape0_type,
uint64_t ape1_offset,
uint32_t ape1_type,
uint32_t in_dim,
uint32_t q_rank,
uint32_t kv_dim,
uint32_t width0,
uint32_t width1,
const ds4_gpu_tensor *x,
uint32_t ratio,
uint32_t pos);
int ds4_gpu_matmul_f32_tensor(
ds4_gpu_tensor *out,
const void *model_map,
@@ -841,6 +1093,31 @@ int ds4_gpu_dsv4_qkv_rms_norm_rows_tensor(
uint32_t rows,
float eps);
int ds4_gpu_dsv4_qkv_rms_norm_kv_rope_fp8_store_tensor(
ds4_gpu_tensor *q_out,
const ds4_gpu_tensor *q,
const void *model_map,
uint64_t model_size,
uint64_t q_weight_offset,
uint32_t q_n,
ds4_gpu_tensor *kv_out,
const ds4_gpu_tensor *kv,
uint64_t kv_weight_offset,
uint32_t kv_n,
ds4_gpu_tensor *raw_cache,
uint64_t raw_cap,
uint32_t raw_row,
uint32_t n_rot,
uint32_t pos0,
uint32_t n_ctx_orig,
float freq_base,
float freq_scale,
float ext_factor,
float attn_factor,
float beta_fast,
float beta_slow,
float eps);
int ds4_gpu_dsv4_qkv_rms_norm_rows_kv_rope_tensor(
ds4_gpu_tensor *q_out,
const ds4_gpu_tensor *q,
@@ -926,6 +1203,8 @@ int ds4_gpu_dsv4_indexer_qat_tensor(
uint32_t n_rows,
uint32_t head_dim);
int ds4_gpu_rope_tail_tensor(
ds4_gpu_tensor *x,
uint32_t n_tok,
@@ -1046,6 +1325,37 @@ int ds4_gpu_glm_store_indexer_k_tensor(
float beta_slow,
bool cache_f16);
/* GLM-5.3 pools four normalized indexer keys with a learned, per-channel
* softmax. Partial pools are retained in tail_k/tail_gate across calls. */
int ds4_gpu_glm53_indexer_pool_update_tensor(
ds4_gpu_tensor *pool_cache,
ds4_gpu_tensor *tail_k,
ds4_gpu_tensor *tail_gate,
const ds4_gpu_tensor *raw_k,
const ds4_gpu_tensor *gate,
const void *model_map,
uint64_t model_size,
uint64_t norm_weight_offset,
uint64_t norm_bias_offset,
uint64_t ape_offset,
uint32_t pos0,
uint32_t n_tokens,
uint32_t cache_cap,
uint32_t head_dim,
uint32_t pool_size,
float eps,
bool cache_f16);
int ds4_gpu_glm53_expand_pool_selection_tensor(
ds4_gpu_tensor *raw_selected,
const ds4_gpu_tensor *pool_selected,
uint32_t n_tokens,
uint32_t pos0,
uint32_t selected_pools,
uint32_t index_topk,
uint32_t pool_size,
uint32_t output_width);
int ds4_gpu_glm_build_kv_cache_tensor(
ds4_gpu_tensor *key_cache,
ds4_gpu_tensor *value_cache,
@@ -1158,6 +1468,20 @@ int ds4_gpu_glm_indexer_scores_batch_tensor(
float scale,
bool cache_f16);
int ds4_gpu_glm53_indexer_scores_batch_tensor(
ds4_gpu_tensor *scores,
const ds4_gpu_tensor *q,
const ds4_gpu_tensor *weights,
const ds4_gpu_tensor *indexer_key_cache,
uint32_t n_rows,
uint32_t n_tokens,
uint32_t pos0,
uint32_t pool_size,
uint32_t n_head,
uint32_t head_dim,
float scale,
bool cache_f16);
int ds4_gpu_glm_qk_lowrank_q8_0_tensor(
ds4_gpu_tensor *qk_low,
const ds4_gpu_tensor *q,
@@ -1468,6 +1792,21 @@ int ds4_gpu_glm_attention_indexed_batch_lora_causal_tensor(
float beta_fast,
float beta_slow);
/* Dense causal MLA over the shared compact latent cache. qk_low and lora_out
* are [token, head, kv_lora_dim]; the F16 cache is shared by all heads. */
int ds4_gpu_glm_attention_dense_compact_lora_causal_tensor(
ds4_gpu_tensor *lora_out,
const ds4_gpu_tensor *qk_low,
const ds4_gpu_tensor *kv_lora_cache,
uint32_t q_row0,
uint32_t n_q,
uint32_t n_kv,
uint32_t cache_cap,
bool cache_f16,
uint32_t n_head,
uint32_t kv_lora_dim,
uint32_t qk_dim);
int ds4_gpu_glm_attention_indexed_batch_lora_valid_tensor(
ds4_gpu_tensor *lora_out,
const ds4_gpu_tensor *q,
@@ -1593,7 +1932,9 @@ int ds4_gpu_compressor_update_tensor(
float beta_fast,
float beta_slow,
float rms_eps,
bool state_already_stored);
bool state_already_stored,
bool decode_one_token,
bool defer_finalize);
int ds4_gpu_compressor_store_batch_tensor(
const ds4_gpu_tensor *kv,
@@ -2134,6 +2475,7 @@ int ds4_gpu_glm_routed_moe_one_tensor(
const ds4_gpu_tensor *weights,
uint32_t n_total_expert,
uint32_t n_expert,
float swiglu_clamp,
uint32_t layer_index,
const ds4_gpu_tensor *x,
bool force_resident);
@@ -2162,6 +2504,7 @@ int ds4_gpu_glm_routed_moe_batch_tensor(
const ds4_gpu_tensor *weights,
uint32_t n_total_expert,
uint32_t n_expert,
float swiglu_clamp,
uint32_t layer_index,
const ds4_gpu_tensor *x,
uint32_t n_tokens,
@@ -2192,6 +2535,7 @@ int ds4_gpu_glm_routed_moe_batch_direct_scalar_q4_tensor(
const ds4_gpu_tensor *weights,
uint32_t n_total_expert,
uint32_t n_expert,
float swiglu_clamp,
uint32_t layer_index,
const ds4_gpu_tensor *x,
uint32_t n_tokens,
@@ -2199,6 +2543,10 @@ int ds4_gpu_glm_routed_moe_batch_direct_scalar_q4_tensor(
int ds4_gpu_routed_moe_set_selected_override(const int32_t *selected, uint32_t n_selected);
void ds4_gpu_set_glm_mtp_verify_mode(bool enabled);
#ifdef DS4_ROCM_BUILD
int ds4_gpu_dspark_gfx1151_fast_path(void);
void ds4_gpu_set_dspark_verify_mode(bool enabled);
#endif
int ds4_gpu_matmul_q8_0_kslice_hc_expand_add_tensor(
ds4_gpu_tensor *out_hc,
@@ -2393,18 +2741,6 @@ int ds4_gpu_hc_weighted_sum_tensor(
uint32_t n_embd,
uint32_t n_hc);
int ds4_gpu_hc_weighted_sum_norm_tensor(
ds4_gpu_tensor *out,
ds4_gpu_tensor *norm_out,
const ds4_gpu_tensor *residual_hc,
const ds4_gpu_tensor *weights,
const void *model_map,
uint64_t model_size,
uint64_t norm_weight_offset,
uint32_t n_embd,
uint32_t n_hc,
float norm_eps);
int ds4_gpu_hc_weighted_sum_split_tensor(
ds4_gpu_tensor *out,
const ds4_gpu_tensor *residual_hc,
@@ -2445,6 +2781,17 @@ int ds4_gpu_hc_split_weighted_sum_norm_tensor(
float eps,
float norm_eps);
int ds4_gpu_hc_rms_norm_mix_f16_available(void);
int ds4_gpu_hc_rms_norm_mix_f16_tensor(
ds4_gpu_tensor *out,
const ds4_gpu_tensor *x,
const void *model_map,
uint64_t model_size,
uint64_t weight_offset,
uint32_t n,
uint32_t out_dim,
float eps);
/* Batched HC RMSNorm followed by its narrow F16 mixer projection. On the
* tuned Metal path, scale_scratch stores one float per row instead of the
* full normalized HC tensor; other shapes retain the established fallback. */
@@ -2460,6 +2807,29 @@ int ds4_gpu_hc_rms_scale_project_f16_tensor(
uint32_t n_rows,
float eps);
#ifdef __APPLE__
int ds4_gpu_hc_rms_norm_mix_split_norm_f16_tensor(
ds4_gpu_tensor *mix,
ds4_gpu_tensor *out,
ds4_gpu_tensor *norm_out,
ds4_gpu_tensor *split,
const ds4_gpu_tensor *residual_hc,
const void *model_map,
uint64_t model_size,
uint64_t mix_weight_offset,
uint64_t scale_offset,
uint64_t base_offset,
uint64_t norm_weight_offset,
uint32_t n,
uint32_t mix_dim,
uint32_t n_embd,
uint32_t n_hc,
uint32_t sinkhorn_iters,
float eps,
float hc_eps,
float norm_eps);
#endif
int ds4_gpu_output_hc_weights_tensor(
ds4_gpu_tensor *out,
const ds4_gpu_tensor *pre,
@@ -2596,6 +2966,174 @@ int ds4_gpu_matmul_q8_0_hc_expand_tensor(
uint32_t n_embd,
uint32_t n_hc);
int ds4_gpu_glm53_embedding_bf16(
ds4_gpu_tensor *out,
const void *model_map,
uint64_t model_size,
uint64_t weight_offset,
const ds4_gpu_tensor *token_ids,
uint32_t n_tokens,
uint32_t n_embd,
uint32_t n_vocab);
int ds4_gpu_glm53_matmul_bf16(
ds4_gpu_tensor *out,
const void *model_map,
uint64_t model_size,
uint64_t weight_offset,
uint32_t in_dim,
uint32_t out_dim,
const ds4_gpu_tensor *x,
uint32_t n_rows);
int ds4_gpu_glm53_matmul_bf16_qkv(
ds4_gpu_tensor *out_q,
ds4_gpu_tensor *out_k,
ds4_gpu_tensor *out_v,
const void *model_map,
uint64_t model_size,
uint64_t weight_q_offset,
uint64_t weight_k_offset,
uint64_t weight_v_offset,
uint32_t in_dim,
uint32_t out_dim,
const ds4_gpu_tensor *x);
#ifndef DS4_GLM53_VISION_TYPES_DEFINED
#define DS4_GLM53_VISION_TYPES_DEFINED
#define DS4_GLM53_VISION_LAYERS 24u
typedef struct {
uint64_t norm1;
uint64_t qkv_weight;
uint64_t qkv_bias;
uint64_t q_norm;
uint64_t k_norm;
uint64_t attn_proj_weight;
uint64_t attn_proj_bias;
uint64_t norm2;
uint64_t gate_weight;
uint64_t gate_bias;
uint64_t up_weight;
uint64_t up_bias;
uint64_t down_weight;
uint64_t down_bias;
} ds4_glm53_vision_layer_weights;
typedef struct {
uint64_t patch_weight;
uint64_t patch_bias;
uint64_t post_norm;
uint64_t downsample_weight;
uint64_t downsample_bias;
uint64_t merger_proj;
uint64_t merger_norm;
uint64_t merger_norm_bias;
uint64_t merger_gate;
uint64_t merger_up;
uint64_t merger_down;
ds4_glm53_vision_layer_weights layer[DS4_GLM53_VISION_LAYERS];
} ds4_glm53_vision_weights;
#endif
/* Encode normalized, block-major image patches into 4096-wide language-model
* embeddings. GPU implementations keep every intermediate on device. */
int ds4_gpu_glm53_vision_encode(
float *out,
const float *patches,
uint32_t grid_h,
uint32_t grid_w,
const void *model_map,
uint64_t model_size,
const ds4_glm53_vision_weights *weights);
/* Replace token rows with projected image embeddings and repeat each row into
* every GLM hyperconnection stream. Must be called in an active command batch. */
int ds4_gpu_glm53_scatter_image_hc(
ds4_gpu_tensor *hc,
const ds4_gpu_tensor *image,
uint32_t dst_row,
uint32_t image_row,
uint32_t rows,
uint32_t total_rows,
uint32_t n_embd,
uint32_t n_hc);
/* GLM-5.3 Kimi Delta Attention. Recurrent and convolution state stay FP32. */
int ds4_gpu_glm53_kda_decode(
ds4_gpu_tensor *out,
ds4_gpu_tensor *conv_state,
ds4_gpu_tensor *recurrent_state,
const ds4_gpu_tensor *q,
const ds4_gpu_tensor *k,
const ds4_gpu_tensor *v,
const ds4_gpu_tensor *raw_gate,
const ds4_gpu_tensor *raw_beta,
const ds4_gpu_tensor *output_gate,
const void *model_map,
uint64_t model_size,
uint64_t q_conv_offset,
uint64_t k_conv_offset,
uint64_t v_conv_offset,
uint64_t a_log_offset,
uint64_t dt_bias_offset,
uint64_t output_norm_offset,
uint32_t n_heads,
uint32_t n_rows,
float gate_lower_bound,
float norm_eps);
int ds4_gpu_glm53_kda_prefill(
ds4_gpu_tensor *out,
ds4_gpu_tensor *conv_state,
ds4_gpu_tensor *recurrent_state,
ds4_gpu_tensor *q,
ds4_gpu_tensor *k,
ds4_gpu_tensor *v,
ds4_gpu_tensor *raw_gate,
const ds4_gpu_tensor *raw_beta,
const ds4_gpu_tensor *output_gate,
const void *model_map,
uint64_t model_size,
uint64_t q_conv_offset,
uint64_t k_conv_offset,
uint64_t v_conv_offset,
uint64_t a_log_offset,
uint64_t dt_bias_offset,
uint64_t output_norm_offset,
uint32_t n_heads,
uint32_t n_tokens,
float gate_lower_bound,
float norm_eps);
/* Decode-island CUDA graph capture (CUDA backend; Metal/ROCm/CPU stub it
* out and stay eager). Design ported from the Entrpi/ds4 batched-serving
* fork's per-layer decode graph capture. The key identifies a captured
* island: layer, island index, and the activation buffers whose addresses
* the captured kernels bake in. ds4_cuda.cu mirrors this struct
* byte-for-byte (it does not include this header); keep both in sync. */
typedef struct ds4_decode_graph_key {
uint32_t il;
uint32_t island; /* 0: layer top to pre-rope; 1: attn-out to layer end */
uint32_t variant;
uint32_t _pad;
void *cur_hc;
void *after_attn_hc;
void *after_ffn_hc;
void *attn_norm;
} ds4_decode_graph_key;
int ds4_gpu_decode_graphs_supported(void);
/* 1: replayed (island already executed; skip encoding it)
* 0: capturing (encode the island, then call _end)
* -1: run eagerly */
int ds4_gpu_decode_graph_begin(const ds4_decode_graph_key *key);
/* 0: capture committed and launched; -1: capture failed (entry retired;
* the caller must re-encode the island eagerly -- no work was executed). */
int ds4_gpu_decode_graph_end(const ds4_decode_graph_key *key);
void ds4_gpu_decode_graph_abort(const ds4_decode_graph_key *key);
void ds4_gpu_decode_graphs_invalidate(void);
#ifdef __cplusplus
}
#endif

File diff suppressed because it is too large Load Diff

View File

@@ -1,22 +0,0 @@
MIT License
Copyright (c) 2026 The ds4.c authors
Copyright (c) 2023-2026 The ggml authors
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.

File diff suppressed because it is too large Load Diff

View File

@@ -1,33 +0,0 @@
#ifndef DS4_WEB_H
#define DS4_WEB_H
#include <stddef.h>
#include <stdbool.h>
typedef int (*ds4_web_confirm_fn)(void *privdata, const char *message,
char *err, size_t err_len);
typedef void (*ds4_web_log_fn)(void *privdata, const char *message);
typedef bool (*ds4_web_cancel_fn)(void *privdata);
typedef struct {
const char *home_dir;
int port;
ds4_web_confirm_fn confirm;
void *confirm_privdata;
ds4_web_log_fn log;
void *log_privdata;
ds4_web_cancel_fn cancel;
void *cancel_privdata;
} ds4_web_config;
typedef struct ds4_web ds4_web;
ds4_web *ds4_web_create(const ds4_web_config *cfg);
void ds4_web_free(ds4_web *web);
char *ds4_web_google_search(ds4_web *web, const char *query,
char *err, size_t err_len);
char *ds4_web_visit_page(ds4_web *web, const char *url,
char *err, size_t err_len);
#endif

27
osv-scanner.toml Normal file
View File

@@ -0,0 +1,27 @@
[[IgnoredVulns]]
id = "RUSTSEC-2024-0436"
reason = "INFO Unmaintained with no patched version; transitive through Iced 0.14's wgpu Metal backend."
[[IgnoredVulns]]
id = "RUSTSEC-2026-0206"
reason = "INFO Unmaintained with no patched version; transitive through Iced 0.14's SVG renderer."
[[IgnoredVulns]]
id = "RUSTSEC-2026-0192"
reason = "INFO Unmaintained with no patched version; transitive through Iced 0.14's text and SVG renderers."
[[IgnoredVulns]]
id = "RUSTSEC-2026-0253"
reason = "No compatible Iced 0.14 upgrade; Cryoglyph does not call the affected LruCache::pop API."
[[IgnoredVulns]]
id = "RUSTSEC-2023-0071"
reason = "No patched rsa release; the release packager only creates ad-hoc signatures without RSA keys."
[[IgnoredVulns]]
id = "RUSTSEC-2020-0095"
reason = "INFO Unmaintained with no patched version; release-only transitive dependency through apple-codesign."
[[IgnoredVulns]]
id = "RUSTSEC-2025-0134"
reason = "INFO Unmaintained with no patched version; transitive through apple-codesign's unused remote-signing stack."

View File

@@ -45,7 +45,7 @@ def post(path, payload):
first = post(
"/v1/chat/completions",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"messages": [{"role": "user", "content": prompt}],
"tools": [chat_tool],
"reasoning_effort": "none",
@@ -58,7 +58,7 @@ print("chat-first", json.dumps(first["usage"], separators=(",", ":")))
second = post(
"/v1/chat/completions",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"messages": [
{"role": "user", "content": prompt},
{"role": "assistant", "content": "", "tool_calls": [call]},
@@ -79,7 +79,7 @@ print("chat", json.dumps(second, separators=(",", ":")))
first = post(
"/v1/messages",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"messages": [{"role": "user", "content": prompt}],
"tools": [anthropic_tool],
"thinking": {"type": "disabled"},
@@ -92,7 +92,7 @@ print("anthropic-first", json.dumps(first["usage"], separators=(",", ":")))
second = post(
"/v1/messages",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"messages": [
{"role": "user", "content": prompt},
{"role": "assistant", "content": [call]},
@@ -118,7 +118,7 @@ print("anthropic", json.dumps(second, separators=(",", ":")))
first = post(
"/v1/responses",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"input": prompt,
"tools": [responses_tool],
"reasoning": {"effort": "none"},
@@ -131,7 +131,7 @@ print("responses-first", json.dumps(first["usage"], separators=(",", ":")))
second = post(
"/v1/responses",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"input": [
{
"type": "message",

View File

@@ -16,7 +16,7 @@ cases = [
"chat",
"/v1/chat/completions",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"messages": [
{
"role": "user",
@@ -42,7 +42,7 @@ cases = [
"anthropic",
"/v1/messages",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"messages": [
{
"role": "user",
@@ -65,7 +65,7 @@ cases = [
"responses",
"/v1/responses",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"input": "Call echo with text hi. Do not answer normally.",
"tools": [
{

View File

@@ -11,7 +11,7 @@ cases = [
"chat",
"/v1/chat/completions",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"messages": [{"role": "user", "content": prompt}],
"reasoning_effort": "low",
"temperature": 0,
@@ -22,7 +22,7 @@ cases = [
"anthropic",
"/v1/messages",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"messages": [{"role": "user", "content": prompt}],
"thinking": {"type": "enabled", "budget_tokens": 64},
"output_config": {"effort": "low"},
@@ -34,7 +34,7 @@ cases = [
"responses",
"/v1/responses",
{
"model": "deepseek-v4-flash",
"model": "deepseek-v4-flash-0731",
"input": prompt,
"reasoning": {"effort": "low", "summary": "auto"},
"temperature": 0,

67
scripts/execution_parity.py Executable file
View File

@@ -0,0 +1,67 @@
#!/usr/bin/env python3
"""Run the DS4 token oracles, then the existing endpoint smoke corpus."""
import argparse
import pathlib
import subprocess
import sys
ROOT = pathlib.Path(__file__).resolve().parent.parent
HARDWARE_TESTS = (
"flash_resident_and_ssd_streaming_choose_the_same_tokens",
"flash_ssd_streaming_maps_batched_prefill_layers",
"resident_multi_session_switching_preserves_each_kv_frontier",
"dspark_runs_a_target_owned_greedy_cycle",
"ssd_streaming_supports_dspark",
"directional_steering_matches_the_ds4_token_oracle",
"resident_and_streamed_glm_match_ds4_decode_oracles",
"streamed_glm_uses_ds4_indexed_prefill_for_long_prompts",
"glm_mtp_preserves_target_tokens_and_drafts",
"pro_ssd_streaming_full_layer_and_selected_batch_agree",
)
ENDPOINT_SCRIPTS = (
"endpoint_parity.py",
"endpoint_reasoning.py",
"endpoint_continuation.py",
)
def run(command):
print("+", " ".join(map(str, command)), flush=True)
subprocess.run(command, cwd=ROOT, check=True)
def main():
parser = argparse.ArgumentParser()
parser.add_argument(
"--endpoint",
help="also run the existing endpoint parity scripts against this URL",
)
parser.add_argument(
"--skip-hardware",
action="store_true",
help="skip tests requiring the local Flash/MTP/DSpark GGUF fixtures",
)
args = parser.parse_args()
if not args.skip_hardware:
for test in HARDWARE_TESTS:
run(
[
"cargo",
"test",
"--all-features",
test,
"--",
"--ignored",
"--nocapture",
]
)
if args.endpoint:
for script in ENDPOINT_SCRIPTS:
run([sys.executable, ROOT / "scripts" / script, args.endpoint])
if __name__ == "__main__":
main()

View File

@@ -0,0 +1,23 @@
Reproduce the following C code EXACTLY, character for character, inside a single code block and output nothing else:
```c
static uint32_t clamp_u32(uint32_t v, uint32_t lo, uint32_t hi) {
if (v < lo) return lo;
if (v > hi) return hi;
return v;
}
static uint32_t ring_advance(uint32_t pos, uint32_t cap) {
uint32_t next = pos + 1u;
return next >= cap ? 0u : next;
}
static int scratch_init(scratch *s, uint32_t ctx_size) {
if (ctx_size == 0u) ctx_size = 1u;
s->ctx_size = ctx_size;
s->comp_cap = ctx_size / 4u + 2u;
s->rows = clamp_u32(s->comp_cap, 1u, 4096u);
s->head = 0u;
return s->rows > 0u ? 0 : -1;
}
```

40
scripts/import_hotlists.py Executable file
View File

@@ -0,0 +1,40 @@
#!/usr/bin/env python3
"""Regenerate the Rust expert hotlists from a DS4 source checkout."""
import argparse
import pathlib
import re
ARRAY = re.compile(
r"static const uint16_t ds4_default_streaming_hotlist_(\w+)\[\]\[2\] = \{(.*?)\n\};",
re.DOTALL,
)
PAIR = re.compile(r"\{(\d+),\s*(\d+)\}")
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("source", type=pathlib.Path)
parser.add_argument("output", type=pathlib.Path)
args = parser.parse_args()
texts = [
(args.source / "ds4_streaming_hotlist.inc").read_text(),
(args.source / "ds4_streaming_hotlist_glm52.inc").read_text(),
]
arrays = {
name: PAIR.findall(body)
for text in texts
for name, body in ARRAY.findall(text)
}
names = (("PRO", "pro"), ("FLASH", "flash"), ("GLM52", "glm52"))
lines = ["// Generated mechanically by scripts/import_hotlists.py.\n"]
for constant, source_name in names:
lines.append(f"pub(super) const {constant}: &[(u16, u16)] = &[\n")
lines.extend(f" ({layer}, {expert}),\n" for layer, expert in arrays[source_name])
lines.append("];\n")
args.output.write_text("".join(lines))
if __name__ == "__main__":
main()

131
scripts/speculative_parity.py Executable file
View File

@@ -0,0 +1,131 @@
#!/usr/bin/env python3
"""Compare DS4 and Rust endpoints configured for the same decoding mode."""
import argparse
import json
import urllib.request
CASES = [
"hi",
"Reply with exactly three words describing a calm sea.",
"Write the first eight positive odd numbers separated by commas.",
]
def post(base_url, path, payload):
request = urllib.request.Request(
base_url.rstrip("/") + path,
data=json.dumps(payload).encode(),
headers={"Content-Type": "application/json"},
)
with urllib.request.urlopen(request, timeout=600) as response:
return json.load(response)
def chat(base_url, prompt):
result = post(
base_url,
"/v1/chat/completions",
{
"model": "deepseek-v4-flash-0731",
"messages": [{"role": "user", "content": prompt}],
"reasoning_effort": "none",
"temperature": 0,
"max_tokens": 64,
},
)
choice = result["choices"][0]
message = choice["message"]
return {
"content": message.get("content"),
"reasoning": message.get("reasoning_content"),
"finish_reason": choice.get("finish_reason"),
"usage": result.get("usage"),
}
def anthropic(base_url, prompt):
result = post(
base_url,
"/v1/messages",
{
"model": "deepseek-v4-flash-0731",
"messages": [{"role": "user", "content": prompt}],
"thinking": {"type": "disabled"},
"temperature": 0,
"max_tokens": 64,
},
)
return {
"content": result.get("content"),
"stop_reason": result.get("stop_reason"),
"usage": result.get("usage"),
}
def responses(base_url, prompt):
result = post(
base_url,
"/v1/responses",
{
"model": "deepseek-v4-flash-0731",
"input": prompt,
"reasoning": {"effort": "none"},
"temperature": 0,
"max_output_tokens": 64,
},
)
return {
"output": normalize(result.get("output")),
"status": result.get("status"),
"usage": result.get("usage"),
}
def normalize(value):
if isinstance(value, list):
return [normalize(item) for item in value]
if isinstance(value, dict):
return {
key: normalize(item)
for key, item in value.items()
if key not in {"id", "created_at"}
}
return value
def main():
parser = argparse.ArgumentParser()
parser.add_argument("reference_url", help="DS4 reference server URL")
parser.add_argument("rust_url", help="Rust DS4Server URL in the same mode")
args = parser.parse_args()
failures = []
for prompt in CASES:
for name, request in (
("chat", chat),
("anthropic", anthropic),
("responses", responses),
):
reference = request(args.reference_url, prompt)
rust = request(args.rust_url, prompt)
if reference != rust:
failures.append(
{
"case": name,
"prompt": prompt,
"reference": reference,
"rust": rust,
}
)
else:
print(f"ok {name}: {prompt}")
if failures:
print(json.dumps(failures, indent=2, ensure_ascii=False))
raise SystemExit(1)
print("all deterministic DS4/Rust endpoint outputs and usage records match")
if __name__ == "__main__":
main()

1217
src/a2ui.rs Normal file

File diff suppressed because it is too large Load Diff

451
src/a2ui/evaluation.rs Normal file
View File

@@ -0,0 +1,451 @@
use super::*;
pub(super) fn resolve(value: &Value, data: &Value) -> Result<Value, String> {
resolve_at(value, data, data)
}
pub(super) fn resolve_at(value: &Value, data: &Value, context: &Value) -> Result<Value, String> {
if let Some(path) = value
.as_object()
.and_then(|value| value.get("path"))
.and_then(Value::as_str)
{
let source = if path.starts_with('/') { data } else { context };
let pointer = if path.starts_with('/') {
normalize_pointer(path).to_owned()
} else {
format!("/{path}")
};
return Ok(source.pointer(&pointer).cloned().unwrap_or(Value::Null));
}
if let Some(call) = value
.as_object()
.and_then(|value| value.get("call"))
.and_then(Value::as_str)
{
return evaluate(
call,
value.get("args").unwrap_or(&Value::Null),
data,
context,
);
}
match value {
Value::Array(values) => values
.iter()
.map(|value| resolve_at(value, data, context))
.collect::<Result<Vec<_>, _>>()
.map(Value::Array),
Value::Object(values) => values
.iter()
.map(|(key, value)| Ok((key.clone(), resolve_at(value, data, context)?)))
.collect::<Result<Map<_, _>, String>>()
.map(Value::Object),
_ => Ok(value.clone()),
}
}
pub(super) fn evaluate(
call: &str,
args: &Value,
data: &Value,
context: &Value,
) -> Result<Value, String> {
if !FUNCTIONS.contains(&call) {
return Err(format!("function `{call}` is not declared by the catalog"));
}
let args = args
.as_object()
.ok_or_else(|| format!("function `{call}` requires object args"))?;
let value = || resolve_at(args.get("value").unwrap_or(&Value::Null), data, context);
match call {
"required" => Ok(Value::Bool(match value()? {
Value::Null => false,
Value::String(value) => !value.is_empty(),
Value::Array(value) => !value.is_empty(),
Value::Object(value) => !value.is_empty(),
_ => true,
})),
"regex" => {
let pattern = args
.get("pattern")
.and_then(Value::as_str)
.ok_or_else(|| "regex requires a pattern".to_owned())?;
let regex = Regex::new(pattern).map_err(|error| format!("invalid regex: {error}"))?;
Ok(Value::Bool(regex.is_match(&display_value(&value()?))))
}
"length" => {
let length = display_value(&value()?).chars().count() as u64;
let min = args.get("min").and_then(Value::as_u64).unwrap_or(0);
let max = args.get("max").and_then(Value::as_u64).unwrap_or(u64::MAX);
Ok(Value::Bool((min..=max).contains(&length)))
}
"numeric" => {
let resolved = value()?;
let number = resolved
.as_f64()
.or_else(|| resolved.as_str().and_then(|value| value.parse().ok()));
let min = args
.get("min")
.and_then(Value::as_f64)
.unwrap_or(f64::NEG_INFINITY);
let max = args
.get("max")
.and_then(Value::as_f64)
.unwrap_or(f64::INFINITY);
Ok(Value::Bool(
number.is_some_and(|number| (min..=max).contains(&number)),
))
}
"email" => {
let email = display_value(&value()?);
let valid = Regex::new(r"^[^\s@]+@[^\s@]+\.[^\s@]+$")
.unwrap()
.is_match(&email);
Ok(Value::Bool(valid))
}
"formatString" => Ok(Value::String(interpolate(
&display_value(&value()?),
data,
context,
)?)),
"formatNumber" => {
let number = value()?.as_f64().unwrap_or(0.0);
let digits = resolve_at(args.get("decimals").unwrap_or(&json!(2)), data, context)?
.as_f64()
.map(|value| value.max(0.0) as u64)
.unwrap_or(2)
.min(12) as usize;
let formatted = format!("{number:.digits$}");
Ok(Value::String(
if resolve_at(
args.get("grouping").unwrap_or(&Value::Bool(true)),
data,
context,
)?
.as_bool()
.unwrap_or(true)
{
group_number(&formatted)
} else {
formatted
},
))
}
"formatCurrency" => {
let number = value()?.as_f64().unwrap_or(0.0);
let currency = display_value(&resolve_at(
args.get("currency").unwrap_or(&Value::Null),
data,
context,
)?);
let digits = resolve_at(args.get("decimals").unwrap_or(&json!(2)), data, context)?
.as_f64()
.map(|value| value.max(0.0) as u64)
.unwrap_or(2)
.min(12) as usize;
let formatted = format!("{number:.digits$}");
let number = if resolve_at(
args.get("grouping").unwrap_or(&Value::Bool(true)),
data,
context,
)?
.as_bool()
.unwrap_or(true)
{
group_number(&formatted)
} else {
formatted
};
Ok(Value::String(format!("{currency} {number}")))
}
"formatDate" => {
let input = display_value(&value()?);
let date =
time::OffsetDateTime::parse(&input, &time::format_description::well_known::Rfc3339)
.map_err(|error| format!("formatDate requires an RFC 3339 value: {error}"))?;
let pattern = display_value(&resolve_at(
args.get("format").unwrap_or(&Value::Null),
data,
context,
)?);
Ok(Value::String(format_date(date, &pattern)))
}
"pluralize" => {
let number = value()?.as_f64().unwrap_or(0.0);
let key = if number == 0.0 && args.contains_key("zero") {
"zero"
} else if number == 1.0 && args.contains_key("one") {
"one"
} else if number == 2.0 && args.contains_key("two") {
"two"
} else {
"other"
};
resolve_at(args.get(key).unwrap_or(&Value::Null), data, context)
}
"and" | "or" => {
let values = args
.get("values")
.and_then(Value::as_array)
.ok_or_else(|| format!("{call} requires array `values`"))?;
let values = values
.iter()
.map(|value| {
resolve_at(value, data, context)?
.as_bool()
.ok_or_else(|| format!("{call} requires boolean values"))
})
.collect::<Result<Vec<_>, _>>()?;
Ok(Value::Bool(if call == "and" {
values.into_iter().all(|value| value)
} else {
values.into_iter().any(|value| value)
}))
}
"not" => {
Ok(Value::Bool(!value()?.as_bool().ok_or_else(|| {
"not requires a boolean value".to_owned()
})?))
}
"openUrl" => Ok(Value::Null),
"@index" => {
let index = TEMPLATE_INDEX
.with(Cell::get)
.ok_or_else(|| "@index is only available in a list template".to_owned())?;
let offset = args
.get("offset")
.map(|value| resolve_at(value, data, context))
.transpose()?
.and_then(|value| value.as_i64())
.unwrap_or(0);
Ok(json!(index as i64 + offset))
}
_ => unreachable!(),
}
}
fn interpolate(template: &str, data: &Value, context: &Value) -> Result<String, String> {
let mut output = String::new();
let mut offset = 0;
while let Some(relative) = template[offset..].find("${") {
let start = offset + relative;
if start > offset && template.as_bytes()[start - 1] == b'\\' {
output.push_str(&template[offset..start - 1]);
output.push_str("${");
offset = start + 2;
continue;
}
output.push_str(&template[offset..start]);
let end = expression_end(template, start + 2)
.ok_or_else(|| "formatString contains an unclosed expression".to_owned())?;
output.push_str(&display_value(&evaluate_expression(
&template[start + 2..end],
data,
context,
)?));
offset = end + 1;
}
output.push_str(&template[offset..]);
Ok(output)
}
fn expression_end(text: &str, mut offset: usize) -> Option<usize> {
let mut depth = 1;
let mut quote = None;
while offset < text.len() {
let character = text[offset..].chars().next()?;
if let Some(current) = quote {
if character == current && text.as_bytes().get(offset.wrapping_sub(1)) != Some(&b'\\') {
quote = None;
}
} else if matches!(character, '\'' | '"') {
quote = Some(character);
} else if text[offset..].starts_with("${") {
depth += 1;
offset += 2;
continue;
} else if character == '}' {
depth -= 1;
if depth == 0 {
return Some(offset);
}
}
offset += character.len_utf8();
}
None
}
fn evaluate_expression(expression: &str, data: &Value, context: &Value) -> Result<Value, String> {
let expression = expression.trim();
if let Some(open) = expression.find('(')
&& expression.ends_with(')')
{
let call = expression[..open].trim();
let mut args = Map::new();
for argument in split_expression_args(&expression[open + 1..expression.len() - 1]) {
let colon = top_level_separator(argument, ':')
.ok_or_else(|| format!("formatString argument `{argument}` must be named"))?;
let name = argument[..colon].trim();
if name.is_empty() {
return Err("formatString contains an empty argument name".into());
}
args.insert(
name.to_owned(),
expression_value(argument[colon + 1..].trim(), data, context)?,
);
}
return evaluate(call, &Value::Object(args), data, context);
}
let (source, pointer) = if expression.starts_with('/') {
(data, normalize_pointer(expression).to_owned())
} else {
(context, format!("/{expression}"))
};
Ok(source.pointer(&pointer).cloned().unwrap_or(Value::Null))
}
fn split_expression_args(input: &str) -> Vec<&str> {
let mut arguments = Vec::new();
let mut start = 0;
while let Some(relative) = top_level_separator(&input[start..], ',') {
arguments.push(input[start..start + relative].trim());
start += relative + 1;
}
if !input[start..].trim().is_empty() {
arguments.push(input[start..].trim());
}
arguments
}
fn top_level_separator(input: &str, separator: char) -> Option<usize> {
let mut round = 0;
let mut braces = 0;
let mut quote = None;
for (index, character) in input.char_indices() {
if let Some(current) = quote {
if character == current && input.as_bytes().get(index.wrapping_sub(1)) != Some(&b'\\') {
quote = None;
}
continue;
}
match character {
'\'' | '"' => quote = Some(character),
'(' => round += 1,
')' => round -= 1,
'{' => braces += 1,
'}' => braces -= 1,
_ if character == separator && round == 0 && braces == 0 => return Some(index),
_ => {}
}
}
None
}
fn expression_value(expression: &str, data: &Value, context: &Value) -> Result<Value, String> {
if expression.starts_with("${") && expression.ends_with('}') {
return evaluate_expression(&expression[2..expression.len() - 1], data, context);
}
if expression.starts_with('\'') && expression.ends_with('\'') && expression.len() >= 2 {
return Ok(Value::String(
expression[1..expression.len() - 1].to_owned(),
));
}
if let Ok(value) = serde_json::from_str(expression) {
return Ok(value);
}
if expression.starts_with('/') {
return Ok(data
.pointer(normalize_pointer(expression))
.cloned()
.unwrap_or(Value::Null));
}
Ok(Value::String(expression.to_owned()))
}
fn group_number(number: &str) -> String {
let (whole, fraction) = number.split_once('.').unwrap_or((number, ""));
let (sign, digits) = whole
.strip_prefix('-')
.map_or(("", whole), |digits| ("-", digits));
let mut grouped = String::with_capacity(number.len() + number.len() / 3);
grouped.push_str(sign);
for (index, digit) in digits.chars().enumerate() {
if index > 0 && (digits.len() - index).is_multiple_of(3) {
grouped.push(',');
}
grouped.push(digit);
}
if !fraction.is_empty() {
grouped.push('.');
grouped.push_str(fraction);
}
grouped
}
pub(super) fn format_date(date: time::OffsetDateTime, pattern: &str) -> String {
const MONTHS: [&str; 12] = [
"January",
"February",
"March",
"April",
"May",
"June",
"July",
"August",
"September",
"October",
"November",
"December",
];
const DAYS: [&str; 7] = [
"Monday",
"Tuesday",
"Wednesday",
"Thursday",
"Friday",
"Saturday",
"Sunday",
];
let month = MONTHS[date.month() as usize - 1];
let day = DAYS[date.weekday().number_days_from_monday() as usize];
let hour_12 = match date.hour() % 12 {
0 => 12,
hour => hour,
};
let replacements = [
("EEEE", day.to_owned()),
("MMMM", month.to_owned()),
("yyyy", format!("{:04}", date.year())),
("MMM", month[..3].to_owned()),
("yy", format!("{:02}", date.year().rem_euclid(100))),
("MM", format!("{:02}", date.month() as u8)),
("dd", format!("{:02}", date.day())),
("HH", format!("{:02}", date.hour())),
("hh", format!("{hour_12:02}")),
("mm", format!("{:02}", date.minute())),
("ss", format!("{:02}", date.second())),
("E", day[..3].to_owned()),
("M", (date.month() as u8).to_string()),
("d", date.day().to_string()),
("H", date.hour().to_string()),
("h", hour_12.to_string()),
("a", if date.hour() < 12 { "AM" } else { "PM" }.to_owned()),
];
let mut output = String::new();
let mut remaining = pattern;
while !remaining.is_empty() {
if let Some((token, replacement)) = replacements
.iter()
.find(|(token, _)| remaining.starts_with(token))
{
output.push_str(replacement);
remaining = &remaining[token.len()..];
} else {
let character = remaining.chars().next().unwrap();
output.push(character);
remaining = &remaining[character.len_utf8()..];
}
}
output
}

881
src/a2ui/export.rs Normal file
View File

@@ -0,0 +1,881 @@
use super::{Surface, bound_value_at, display_value, validate_surface_composition};
use pulldown_cmark::{Event, Options, Parser, Tag, TagEnd, html};
use serde_json::{Map, Value};
use std::collections::{BTreeSet, HashMap};
use std::fmt::Write;
const WIDTH: usize = 900;
const HEIGHT: usize = 720;
const COLORS: [&str; 6] = [
"#75beff", "#89d185", "#d18616", "#f14c4c", "#b180d7", "#e2c440",
];
pub(crate) fn surface_svg(
surface: &Surface,
images: &HashMap<String, String>,
) -> Result<String, String> {
validate_surface_composition(surface)?;
let renderer = Renderer { surface, images };
let content = renderer.component("root", &surface.data, BTreeSet::new());
let agent = surface
.surface_properties
.get("agentDisplayName")
.and_then(Value::as_str)
.unwrap_or("A2UI");
let label = format!("{agent} surface {}", surface.id);
let mut svg = format!(
r#"<svg xmlns="http://www.w3.org/2000/svg" class="a2ui-snapshot" viewBox="0 0 {WIDTH} {HEIGHT}" role="img" aria-label="{}"><foreignObject x="0" y="0" width="{WIDTH}" height="{HEIGHT}"><div xmlns="http://www.w3.org/1999/xhtml" class="a2ui-document"><style>{}</style><header><strong>{}</strong><span>{}</span></header><main>"#,
escaped(&label),
STYLE,
escaped(agent),
escaped(&surface.id),
);
svg.push_str(&content);
svg.push_str("</main></div></foreignObject></svg>");
Ok(svg)
}
const STYLE: &str = r#"
*{box-sizing:border-box} .a2ui-document{height:100%;overflow:auto;background:#17181b;color:#e8e8ea;font:14px/1.45 -apple-system,BlinkMacSystemFont,"Segoe UI",sans-serif;border:1px solid #34363d;border-radius:12px}.a2ui-document>header{position:sticky;top:0;z-index:2;display:flex;justify-content:space-between;padding:11px 14px;background:#202125;border-bottom:1px solid #34363d;color:#aeb0b6;font-size:12px}.a2ui-document>main{display:flex;flex-direction:column;gap:10px;padding:16px}.row,.column,.list{display:flex;gap:10px}.row,.list.horizontal{flex-direction:row;align-items:flex-start}.column,.list{flex-direction:column}.card,.map,.table{padding:13px;background:#202125;border:1px solid #34363d;border-radius:9px}.button,.chip{display:inline-block;padding:7px 11px;background:#30323a;border:1px solid #464951;border-radius:7px}.button.primary{background:#416a91}.button.borderless{background:transparent;border-color:transparent}.field{display:flex;flex-direction:column;gap:5px}.field-label,.muted,.chart-kind,.caption{color:#999ba3;font-size:12px}.input{min-height:34px;padding:7px 9px;background:#101114;border:1px solid #393b42;border-radius:6px;white-space:pre-wrap}.check{display:flex;gap:7px;align-items:center}.check-mark{width:18px;height:18px;text-align:center;border:1px solid #555862;border-radius:5px}.divider{height:1px;background:#3a3c43}.divider.vertical{width:1px;height:auto;min-height:32px}.metric-value{font-size:28px}.error{color:#ff8585}.bar{display:grid;grid-template-columns:minmax(80px,150px) 1fr auto;gap:8px;align-items:center}.bar-track{height:9px;overflow:hidden;background:#30323a;border-radius:5px}.bar-fill{height:100%;background:#75beff}.pie{width:150px;height:150px;border-radius:50%}.donut{position:relative}.donut:after{content:"";position:absolute;inset:34px;background:#17181b;border-radius:50%}.legend{display:flex;flex-wrap:wrap;gap:8px 14px}.swatch{display:inline-block;width:10px;height:10px;margin-right:5px;border-radius:2px}.table{width:100%;border-collapse:collapse}.table th,.table td{padding:7px 9px;border:1px solid #3a3c43;text-align:left}.timeline-item,.location,.mind-node{display:flex;gap:8px}.tabs>details{margin-top:7px;padding:7px 9px;border:1px solid #34363d;border-radius:7px}.tabs summary{cursor:pointer}.media{padding:10px;border:1px solid #34363d;border-radius:7px}.surface-image{display:block;max-width:100%;max-height:260px;margin:auto;border-radius:7px}.chart,.timeline,.mindmap{display:flex;flex-direction:column;gap:8px}h1,h2,h3,p{margin-top:0}pre{overflow:auto}a{color:#8bbceb}
"#;
struct Renderer<'a> {
surface: &'a Surface,
images: &'a HashMap<String, String>,
}
impl Renderer<'_> {
fn component(&self, id: &str, context: &Value, mut ancestors: BTreeSet<String>) -> String {
if !ancestors.insert(id.to_owned()) {
return format!(
"<div class=\"error\">Cyclic component {}</div>",
escaped(id)
);
}
let Some(component) = self.surface.components.get(id).and_then(Value::as_object) else {
return format!(
"<div class=\"error\">Missing component {}</div>",
escaped(id)
);
};
let kind = component
.get("component")
.and_then(Value::as_str)
.unwrap_or("Unknown");
match kind {
"Text" => {
let content = safe_markdown(&self.value(component, "text", context));
if component.get("variant").and_then(Value::as_str) == Some("caption") {
format!("<div class=\"caption\">{content}</div>")
} else {
content
}
}
"Image" => self.image(component, context),
"Video" | "AudioPlayer" => self.media(component, context, kind),
"Icon" => format!(
"<span aria-label=\"{}\">{}</span>",
escaped(&self.value(component, "name", context)),
icon(&self.value(component, "name", context))
),
"Divider" => format!(
"<div class=\"divider{}\"></div>",
if component.get("axis").and_then(Value::as_str) == Some("vertical") {
" vertical"
} else {
""
}
),
"Row" | "Column" | "List" => {
let horizontal = kind == "Row"
|| (kind == "List"
&& component.get("direction").and_then(Value::as_str)
== Some("horizontal"));
let class = if kind == "List" {
if horizontal {
"list horizontal"
} else {
"list"
}
} else if kind == "Row" {
"row"
} else {
"column"
};
format!(
"<div class=\"{class}\"{}>{}</div>",
layout_style(component),
self.children(component.get("children"), context, &ancestors)
)
}
"Card" => format!(
"<div class=\"card\">{}</div>",
self.child(component, "child", context, &ancestors)
),
"Tabs" => self.tabs(component, context, &ancestors),
"Modal" => format!(
"<details><summary>{}</summary><div class=\"card\">{}</div></details>",
self.child(component, "trigger", context, &ancestors),
self.child(component, "content", context, &ancestors)
),
"Button" => self.with_error(
component,
format!(
"<span class=\"button {}\">{}</span>",
component
.get("variant")
.and_then(Value::as_str)
.unwrap_or("default"),
self.child(component, "child", context, &ancestors)
),
),
"TextField" => self.text_field(component, context),
"CheckBox" => self.checkbox(component, context),
"Slider" => self.slider(component, context),
"DateTimeInput" => self.with_error(
component,
self.field(
&self.value(component, "label", context),
&self.value(component, "value", context),
),
),
"ChoicePicker" => self.choices(component, context),
"Chart" => self.chart(component, context),
"Table" => self.table(component, context),
"Metric" => self.metric(component, context),
"Timeline" => self.timeline(component, context),
"Map" => self.map(component, context),
"MindMap" => self.mind_map(component, context),
"Form" => format!(
"<section class=\"card column\"><strong>{}</strong>{}<span class=\"button\">{}</span></section>",
escaped(&self.value(component, "title", context)),
self.children(component.get("children"), context, &ancestors),
escaped(
component
.get("submitLabel")
.and_then(Value::as_str)
.unwrap_or("Submit")
)
),
_ => format!(
"<div class=\"error\">Unsupported component {}</div>",
escaped(kind)
),
}
}
fn value(&self, component: &Map<String, Value>, key: &str, context: &Value) -> String {
display_value(&bound_value_at(
component.get(key),
&self.surface.data,
context,
))
}
fn with_error(&self, component: &Map<String, Value>, mut content: String) -> String {
if let Some(error) = super::first_failed_check(component, &self.surface.data) {
let _ = write!(content, "<div class=\"error\">{}</div>", escaped(&error));
}
content
}
fn child(
&self,
component: &Map<String, Value>,
key: &str,
context: &Value,
ancestors: &BTreeSet<String>,
) -> String {
component
.get(key)
.and_then(Value::as_str)
.map(|id| self.component(id, context, ancestors.clone()))
.unwrap_or_default()
}
fn children(
&self,
children: Option<&Value>,
context: &Value,
ancestors: &BTreeSet<String>,
) -> String {
if let Some(children) = children.and_then(Value::as_array) {
return children
.iter()
.filter_map(Value::as_str)
.map(|id| self.weighted_component(id, context, ancestors))
.collect();
}
let Some(template) = children.and_then(Value::as_object) else {
return String::new();
};
let id = template
.get("componentId")
.and_then(Value::as_str)
.unwrap_or("");
let path = template.get("path").and_then(Value::as_str).unwrap_or("/");
self.surface
.data
.pointer(path)
.and_then(Value::as_array)
.into_iter()
.flatten()
.enumerate()
.map(|(index, item)| {
super::with_template_index(index, || self.weighted_component(id, item, ancestors))
})
.collect()
}
fn weighted_component(
&self,
id: &str,
context: &Value,
ancestors: &BTreeSet<String>,
) -> String {
let content = self.component(id, context, ancestors.clone());
self.surface
.components
.get(id)
.and_then(|component| component.get("weight"))
.and_then(Value::as_f64)
.filter(|weight| weight.is_finite() && *weight > 0.0)
.map_or(content.clone(), |weight| {
format!("<div style=\"flex:{weight} 1 0\">{content}</div>")
})
}
fn tabs(
&self,
component: &Map<String, Value>,
context: &Value,
ancestors: &BTreeSet<String>,
) -> String {
let mut output = String::from("<div class=\"tabs\">");
for (index, tab) in component
.get("tabs")
.and_then(Value::as_array)
.into_iter()
.flatten()
.enumerate()
{
let title = display_value(&bound_value_at(
tab.get("title"),
&self.surface.data,
context,
));
let child = tab.get("child").and_then(Value::as_str).unwrap_or("");
let _ = write!(
output,
"<details{}><summary>{}</summary>{}</details>",
if index == 0 { " open" } else { "" },
escaped(&title),
self.component(child, context, ancestors.clone())
);
}
output.push_str("</div>");
output
}
fn image(&self, component: &Map<String, Value>, context: &Value) -> String {
let url = self.value(component, "url", context);
let description = self.value(component, "description", context);
let source = self
.images
.get(&url)
.map(String::as_str)
.or_else(|| url.starts_with("data:image/").then_some(url.as_str()));
source.map_or_else(
|| {
format!(
"<div class=\"media\">▧ {}<div class=\"muted\">{}</div></div>",
escaped(&description),
escaped(&url)
)
},
|source| {
format!(
"<img class=\"surface-image\" src=\"{}\" alt=\"{}\"/>",
escaped(source),
escaped(&description)
)
},
)
}
fn media(&self, component: &Map<String, Value>, context: &Value, kind: &str) -> String {
format!(
"<div class=\"media\">{} {}<div class=\"muted\">{}</div></div>",
if kind == "Video" { "" } else { "" },
escaped(&self.value(component, "description", context)),
escaped(&self.value(component, "url", context))
)
}
fn text_field(&self, component: &Map<String, Value>, context: &Value) -> String {
let mut value = self.value(component, "value", context);
if component.get("variant").and_then(Value::as_str) == Some("obscured") {
value = "••••••••".to_owned();
} else if value.is_empty() {
value = self.value(component, "placeholder", context);
}
self.with_error(
component,
self.field(&self.value(component, "label", context), &value),
)
}
fn field(&self, label: &str, value: &str) -> String {
format!(
"<div class=\"field\"><span class=\"field-label\">{}</span><div class=\"input\">{}</div></div>",
escaped(label),
escaped(value)
)
}
fn checkbox(&self, component: &Map<String, Value>, context: &Value) -> String {
let checked = bound_value_at(component.get("value"), &self.surface.data, context)
.as_bool()
.unwrap_or(false);
self.with_error(
component,
format!(
"<div class=\"check\"><span class=\"check-mark\">{}</span>{}</div>",
if checked { "" } else { "" },
escaped(&self.value(component, "label", context))
),
)
}
fn slider(&self, component: &Map<String, Value>, context: &Value) -> String {
let value = bound_value_at(component.get("value"), &self.surface.data, context)
.as_f64()
.unwrap_or(0.0);
let min = component.get("min").and_then(Value::as_f64).unwrap_or(0.0);
let max = component
.get("max")
.and_then(Value::as_f64)
.unwrap_or(100.0);
let percent = if max > min {
((value - min) / (max - min) * 100.0).clamp(0.0, 100.0)
} else {
0.0
};
self.with_error(component, format!(
"<div class=\"field\"><span class=\"field-label\">{}</span><div class=\"bar-track\"><div class=\"bar-fill\" style=\"width:{percent:.2}%\"></div></div><span>{}</span></div>",
escaped(&self.value(component, "label", context)),
escaped(&display_value(&Value::from(value)))
))
}
fn choices(&self, component: &Map<String, Value>, context: &Value) -> String {
let selected = bound_value_at(component.get("value"), &self.surface.data, context);
let multiple =
component.get("variant").and_then(Value::as_str) == Some("multipleSelection");
let mut output = format!(
"<div class=\"field\"><span class=\"field-label\">{}</span><div class=\"row\">",
escaped(&self.value(component, "label", context))
);
for option in component
.get("options")
.and_then(Value::as_array)
.into_iter()
.flatten()
{
let value = option.get("value").unwrap_or(&Value::Null);
let active = selected
.as_array()
.is_some_and(|values| values.contains(value))
|| (!multiple && &selected == value);
let _ = write!(
output,
"<span class=\"chip\">{} {}</span>",
if active { "" } else { "" },
escaped(&option.get("label").map(display_value).unwrap_or_default())
);
}
output.push_str("</div></div>");
self.with_error(component, output)
}
fn chart(&self, component: &Map<String, Value>, context: &Value) -> String {
let series = bound_value_at(component.get("series"), &self.surface.data, context);
let series = series.as_array().map(Vec::as_slice).unwrap_or(&[]);
let chart_type = component
.get("chartType")
.and_then(Value::as_str)
.unwrap_or("bar");
let mut output = format!(
"<section class=\"chart\"><strong>{}</strong><span class=\"chart-kind\">{}</span>",
escaped(&self.value(component, "title", context)),
escaped(&chart_type.to_uppercase())
);
if matches!(chart_type, "pie" | "donut") {
let total = series
.iter()
.filter_map(|point| point.get("value")?.as_f64())
.filter(|value| *value > 0.0)
.sum::<f64>();
let mut position = 0.0;
let mut stops = Vec::new();
for (index, point) in series.iter().enumerate() {
let value = point
.get("value")
.and_then(Value::as_f64)
.unwrap_or(0.0)
.max(0.0);
let end = if total > 0.0 {
position + value / total * 100.0
} else {
position
};
stops.push(format!(
"{} {position:.3}% {end:.3}%",
COLORS[index % COLORS.len()]
));
position = end;
}
let _ = write!(
output,
"<div class=\"row\"><div class=\"pie {}\" style=\"background:conic-gradient({})\"></div><div class=\"legend\">",
if chart_type == "donut" { "donut" } else { "" },
stops.join(",")
);
for (index, point) in series.iter().enumerate() {
let _ = write!(
output,
"<span><i class=\"swatch\" style=\"background:{}\"></i>{} {}</span>",
COLORS[index % COLORS.len()],
escaped(&point.get("label").map(display_value).unwrap_or_default()),
escaped(&point.get("value").map(display_value).unwrap_or_default())
);
}
output.push_str("</div></div>");
} else if chart_type == "heatmap" {
let columns = series
.iter()
.flat_map(|row| {
row.get("segments")
.and_then(Value::as_array)
.into_iter()
.flatten()
})
.filter_map(|segment| segment.get("label").and_then(Value::as_str))
.fold(Vec::<String>::new(), |mut columns, label| {
if !columns.iter().any(|column| column == label) {
columns.push(label.to_owned());
}
columns
});
let max = series
.iter()
.flat_map(|row| {
row.get("segments")
.and_then(Value::as_array)
.into_iter()
.flatten()
})
.filter_map(|segment| segment.get("value").and_then(Value::as_f64))
.fold(0.0_f64, f64::max)
.max(1.0);
output.push_str("<table class=\"table\"><thead><tr><th></th>");
for column in &columns {
let _ = write!(output, "<th>{}</th>", escaped(column));
}
output.push_str("</tr></thead><tbody>");
for row in series {
let segments = row
.get("segments")
.and_then(Value::as_array)
.map(Vec::as_slice)
.unwrap_or(&[]);
let _ = write!(
output,
"<tr><th>{}</th>",
escaped(&row.get("label").map(display_value).unwrap_or_default())
);
for column in &columns {
let value = segments
.iter()
.find(|segment| {
segment.get("label").and_then(Value::as_str) == Some(column)
})
.and_then(|segment| segment.get("value").and_then(Value::as_f64))
.unwrap_or(0.0);
let alpha = (value / max).clamp(0.0, 1.0);
let _ = write!(
output,
"<td style=\"background:rgba(183,72,72,{alpha:.3})\">{}</td>",
escaped(&display_value(&Value::from(value)))
);
}
output.push_str("</tr>");
}
output.push_str("</tbody></table>");
} else {
let max = series
.iter()
.filter_map(|point| point.get("value")?.as_f64())
.fold(0.0_f64, f64::max)
.max(1.0);
for point in series {
let value = point.get("value").and_then(Value::as_f64).unwrap_or(0.0);
let percent = (value / max * 100.0).clamp(0.0, 100.0);
let _ = write!(
output,
"<div class=\"bar\"><span>{}</span><div class=\"bar-track\"><div class=\"bar-fill\" style=\"width:{percent:.2}%\"></div></div><span>{}</span></div>",
escaped(&point.get("label").map(display_value).unwrap_or_default()),
escaped(&point.get("value").map(display_value).unwrap_or_default())
);
if let Some(segments) = point.get("segments").and_then(Value::as_array) {
let details = segments
.iter()
.map(|segment| {
format!(
"{} {}",
segment.get("label").map(display_value).unwrap_or_default(),
segment.get("value").map(display_value).unwrap_or_default()
)
})
.collect::<Vec<_>>()
.join(" · ");
let _ = write!(output, "<div class=\"muted\">{}</div>", escaped(&details));
}
}
}
output.push_str("</section>");
output
}
fn table(&self, component: &Map<String, Value>, context: &Value) -> String {
let columns = bound_value_at(component.get("columns"), &self.surface.data, context);
let rows = bound_value_at(component.get("rows"), &self.surface.data, context);
let mut output = format!(
"<section><strong>{}</strong><table class=\"table\"><thead><tr>",
escaped(&self.value(component, "title", context))
);
for column in columns.as_array().into_iter().flatten() {
let _ = write!(output, "<th>{}</th>", escaped(&display_value(column)));
}
output.push_str("</tr></thead><tbody>");
for row in rows
.as_array()
.into_iter()
.flatten()
.filter_map(Value::as_array)
{
output.push_str("<tr>");
for cell in row {
let _ = write!(output, "<td>{}</td>", escaped(&display_value(cell)));
}
output.push_str("</tr>");
}
output.push_str("</tbody></table></section>");
output
}
fn metric(&self, component: &Map<String, Value>, context: &Value) -> String {
format!(
"<section><div class=\"muted\">{}</div><div class=\"metric-value\">{}</div><div class=\"muted\">{} {}</div></section>",
escaped(&self.value(component, "label", context)),
escaped(&self.value(component, "value", context)),
escaped(&self.value(component, "detail", context)),
escaped(&self.value(component, "trend", context))
)
}
fn timeline(&self, component: &Map<String, Value>, context: &Value) -> String {
let events = bound_value_at(component.get("events"), &self.surface.data, context);
let mut output = format!(
"<section class=\"timeline\"><strong>{}</strong>",
escaped(&self.value(component, "title", context))
);
for event in events.as_array().into_iter().flatten() {
let _ = write!(
output,
"<div class=\"timeline-item\"><span>●</span><div><strong>{}</strong><div class=\"muted\">{} · {} · {}</div></div></div>",
escaped(&event.get("title").map(display_value).unwrap_or_default()),
escaped(&event.get("time").map(display_value).unwrap_or_default()),
escaped(
&event
.get("description")
.map(display_value)
.unwrap_or_default()
),
escaped(&event.get("status").map(display_value).unwrap_or_default())
);
}
output.push_str("</section>");
output
}
fn map(&self, component: &Map<String, Value>, context: &Value) -> String {
let locations = bound_value_at(component.get("locations"), &self.surface.data, context);
let mut output = format!(
"<section class=\"map\"><strong>{}</strong>",
escaped(&self.value(component, "title", context))
);
for location in locations.as_array().into_iter().flatten() {
let _ = write!(
output,
"<div class=\"location\"><span>⌖</span><div><strong>{}</strong><div class=\"muted\">{}, {} · {}</div></div></div>",
escaped(&location.get("label").map(display_value).unwrap_or_default()),
escaped(
&location
.get("latitude")
.map(display_value)
.unwrap_or_default()
),
escaped(
&location
.get("longitude")
.map(display_value)
.unwrap_or_default()
),
escaped(
&location
.get("detail")
.map(display_value)
.unwrap_or_default()
)
);
}
output.push_str("</section>");
output
}
fn mind_map(&self, component: &Map<String, Value>, context: &Value) -> String {
let nodes = bound_value_at(component.get("nodes"), &self.surface.data, context);
let nodes = nodes.as_array().map(Vec::as_slice).unwrap_or(&[]);
let mut output = format!(
"<section class=\"mindmap\"><strong>{}</strong>",
escaped(&self.value(component, "title", context))
);
if let Some(root) = nodes.first() {
output.push_str(&mind_node(nodes, root, 0, BTreeSet::new()));
}
output.push_str("</section>");
output
}
}
fn mind_node(nodes: &[Value], node: &Value, depth: usize, mut seen: BTreeSet<String>) -> String {
let id = node.get("id").map(display_value).unwrap_or_default();
if !seen.insert(id) {
return "<div class=\"muted\">Cycle</div>".to_owned();
}
let mut output = format!(
"<div class=\"mind-node\" style=\"padding-left:{}px\"><span>{}</span><span>{}</span></div>",
depth * 18,
if depth == 0 { "" } else { "" },
escaped(&node.get("label").map(display_value).unwrap_or_default())
);
for child in node
.get("children")
.and_then(Value::as_array)
.into_iter()
.flatten()
{
if let Some(child) = nodes
.iter()
.find(|candidate| candidate.get("id") == Some(child))
{
output.push_str(&mind_node(nodes, child, depth + 1, seen.clone()));
}
}
output
}
fn layout_style(component: &Map<String, Value>) -> String {
let justify = match component.get("justify").and_then(Value::as_str) {
Some("center") => "center",
Some("end") => "flex-end",
Some("spaceBetween") => "space-between",
Some("spaceAround") => "space-around",
Some("spaceEvenly") => "space-evenly",
_ => "flex-start",
};
let align = match component.get("align").and_then(Value::as_str) {
Some("center") => "center",
Some("end") => "flex-end",
Some("stretch") => "stretch",
_ => "flex-start",
};
format!(" style=\"justify-content:{justify};align-items:{align}\"")
}
fn safe_markdown(markdown: &str) -> String {
let options = Options::ENABLE_TABLES
| Options::ENABLE_FOOTNOTES
| Options::ENABLE_STRIKETHROUGH
| Options::ENABLE_TASKLISTS;
let parser = Parser::new_ext(markdown, options).scan(Vec::new(), |images, event| {
Some(match event {
Event::Html(html) | Event::InlineHtml(html) => Event::Text(html),
Event::Start(tag @ Tag::Image { .. }) => {
let embedded = matches!(
&tag,
Tag::Image { dest_url, .. } if dest_url.starts_with("data:image/")
);
images.push(embedded);
if embedded {
Event::Start(tag)
} else {
Event::Text("[Image: ".into())
}
}
Event::End(TagEnd::Image) => {
if images.pop().unwrap_or(false) {
Event::End(TagEnd::Image)
} else {
Event::Text("]".into())
}
}
event => event,
})
});
let mut output = String::new();
html::push_html(&mut output, parser);
output
}
fn escaped(text: &str) -> String {
let mut output = String::with_capacity(text.len());
for character in text.chars() {
match character {
'&' => output.push_str("&amp;"),
'<' => output.push_str("&lt;"),
'>' => output.push_str("&gt;"),
'"' => output.push_str("&quot;"),
'\'' => output.push_str("&#39;"),
_ => output.push(character),
}
}
output
}
fn icon(name: &str) -> &'static str {
match name {
"accountCircle" | "person" => "",
"add" => "+",
"arrowBack" => "",
"arrowForward" => "",
"attachFile" => "",
"calendarToday" | "event" => "",
"call" | "phone" => "",
"camera" => "",
"check" => "",
"close" => "×",
"delete" => "",
"download" => "",
"edit" => "",
"error" => "",
"fastForward" => "",
"favorite" => "",
"favoriteOff" => "",
"folder" => "",
"help" => "?",
"home" => "",
"info" => "",
"locationOn" => "",
"lock" => "",
"lockOpen" => "",
"mail" => "",
"menu" => "",
"moreVert" => "",
"moreHoriz" => "",
"notifications" => "",
"notificationsOff" => "",
"pause" => "",
"payment" => "¤",
"photo" => "",
"play" => "",
"print" => "",
"refresh" => "",
"rewind" => "",
"search" => "",
"send" => "",
"settings" => "",
"share" => "",
"shoppingCart" => "",
"skipNext" => "▸|",
"skipPrevious" => "|◂",
"star" => "",
"starHalf" | "starOff" => "",
"stop" => "",
"upload" => "",
"visibility" => "",
"visibilityOff" => "",
"volumeDown" => "",
"volumeMute" => "",
"volumeOff" => "×",
"volumeUp" => "",
"warning" => "",
_ => "",
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use std::collections::BTreeMap;
#[test]
fn static_svg_renderer_covers_the_complete_catalog() {
let entries = [
json!({"id":"text","component":"Text","text":"Hello"}),
json!({"id":"image","component":"Image","url":"https://example.com/a.png","description":"Preview"}),
json!({"id":"icon","component":"Icon","name":"settings"}),
json!({"id":"video","component":"Video","url":"https://example.com/a.mp4"}),
json!({"id":"audio","component":"AudioPlayer","url":"https://example.com/a.mp3"}),
json!({"id":"divider","component":"Divider"}),
json!({"id":"row","component":"Row","children":[]}),
json!({"id":"list","component":"List","children":[]}),
json!({"id":"card","component":"Card","child":"text"}),
json!({"id":"modal","component":"Modal","trigger":"text","content":"text"}),
json!({"id":"tabs","component":"Tabs","tabs":[{"title":"One","child":"text"}]}),
json!({"id":"button","component":"Button","child":"text"}),
json!({"id":"field","component":"TextField","label":"Name","value":"Ada"}),
json!({"id":"check","component":"CheckBox","label":"Ready","value":true}),
json!({"id":"slider","component":"Slider","label":"Amount","value":4,"max":10}),
json!({"id":"date","component":"DateTimeInput","label":"When","value":"2026-08-31"}),
json!({"id":"choice","component":"ChoicePicker","label":"Pick","options":[{"label":"One","value":"one"}],"value":["one"]}),
json!({"id":"chart","component":"Chart","title":"Chart","chartType":"heatmap","series":[{"label":"A","segments":[{"label":"B","value":2}]}]}),
json!({"id":"table","component":"Table","title":"Table","columns":["A"],"rows":[["B"]]}),
json!({"id":"metric","component":"Metric","label":"Count","value":"2"}),
json!({"id":"timeline","component":"Timeline","title":"Timeline","events":[{"time":"Now","title":"Done"}]}),
json!({"id":"map","component":"Map","title":"Map","locations":[{"label":"Here","latitude":1,"longitude":2}]}),
json!({"id":"mind","component":"MindMap","title":"Mind","nodes":[{"id":"n","label":"Root","children":[]}]}),
json!({"id":"form","component":"Form","title":"Form","children":[]}),
];
let ids = entries
.iter()
.filter_map(|entry| entry.get("id").and_then(Value::as_str))
.map(str::to_owned)
.collect::<Vec<_>>();
let mut components = entries
.into_iter()
.map(|entry| (entry["id"].as_str().unwrap().to_owned(), entry))
.collect::<BTreeMap<_, _>>();
components.insert(
"root".into(),
json!({"id":"root","component":"Column","children":ids}),
);
let surface = Surface {
id: "catalog".into(),
catalog_id: super::super::CATALOG_ID.into(),
surface_properties: json!({}),
send_data_model: false,
components,
data: json!({}),
owner_message_id: 1,
};
let svg = surface_svg(&surface, &HashMap::new()).unwrap();
assert!(svg.starts_with("<svg"));
assert!(!svg.contains("Unsupported component"));
assert!(svg.contains("HEATMAP"));
assert!(svg.contains("Submit"));
}
}

659
src/a2ui/validation.rs Normal file
View File

@@ -0,0 +1,659 @@
use super::*;
pub(super) fn validate_component(
component: &Map<String, Value>,
catalog_id: &str,
) -> Result<(), String> {
let id = required_string(component, "id")?;
let kind = required_string(component, "component")?;
let allowed = if catalog_id == CATALOG_ID {
ds4_catalog()
.pointer("/components")
.and_then(Value::as_object)
.is_some_and(|components| components.contains_key(kind))
} else {
BASIC_COMPONENTS.contains(&kind)
};
if !allowed {
return Err(format!(
"component `{id}` uses `{kind}`, which is not in catalog `{catalog_id}`"
));
}
let required: &[&str] = match kind {
"Text" => &["text"],
"Image" => &["url"],
"Icon" => &["name"],
"Video" | "AudioPlayer" => &["url"],
"Row" | "Column" | "List" => &["children"],
"Card" => &["child"],
"Modal" => &["trigger", "content"],
"Tabs" => &["tabs"],
"Button" => &["child", "action"],
"TextField" => &["label"],
"CheckBox" => &["label", "value"],
"Slider" => &["value", "max"],
"DateTimeInput" => &["value"],
"ChoicePicker" => &["options", "value"],
"Chart" => &["chartType", "series"],
"Table" => &["columns", "rows"],
"Metric" => &["label", "value"],
"Timeline" => &["events"],
"Map" => &["locations"],
"MindMap" => &["nodes"],
"Form" => &["children"],
_ => &[],
};
if let Some(field) = required
.iter()
.find(|field| !component.contains_key(**field))
{
return Err(format!("component `{id}` ({kind}) requires `{field}`"));
}
let common = ["id", "component", "accessibility", "weight", "checks"];
let specific: &[&str] = match kind {
"Text" => &["text", "variant"],
"Image" => &["url", "description", "fit", "variant"],
"Icon" => &["name"],
"Video" => &["url", "posterUrl"],
"AudioPlayer" => &["url", "description"],
"Divider" => &["axis"],
"Row" | "Column" => &["children", "justify", "align"],
"List" => &["children", "direction", "align"],
"Card" => &["child"],
"Modal" => &["trigger", "content"],
"Tabs" => &["tabs"],
"Button" => &["child", "variant", "action"],
"TextField" => &["label", "value", "placeholder", "variant"],
"CheckBox" => &["label", "value"],
"Slider" => &["label", "min", "max", "value", "steps"],
"DateTimeInput" => &["label", "value", "enableDate", "enableTime", "min", "max"],
"ChoicePicker" => &[
"label",
"variant",
"options",
"value",
"displayStyle",
"filterable",
],
"Chart" => &["title", "chartType", "series"],
"Table" => &["title", "columns", "rows"],
"Metric" => &["label", "value", "detail", "trend"],
"Timeline" => &["title", "events"],
"Map" => &["title", "locations"],
"MindMap" => &["title", "nodes"],
"Form" => &["title", "children", "submitLabel", "action"],
_ => &[],
};
if let Some(field) = component
.keys()
.find(|field| !common.contains(&field.as_str()) && !specific.contains(&field.as_str()))
{
return Err(format!(
"component `{id}` ({kind}) contains unknown property `{field}`"
));
}
if let Some(accessibility) = component.get("accessibility") {
let accessibility = object(Some(accessibility), "accessibility")?;
for field in ["label", "description"] {
validate_optional_dynamic(accessibility, field, Value::is_string, id)?;
}
}
if component
.get("weight")
.is_some_and(|value| !value.is_number())
{
return Err(format!("component `{id}` weight must be a number"));
}
if component
.get("title")
.is_some_and(|value| !value.is_string())
{
return Err(format!("component `{id}` title must be a string"));
}
if catalog_id != CATALOG_ID
&& component.contains_key("checks")
&& !matches!(
kind,
"Button" | "TextField" | "CheckBox" | "ChoicePicker" | "Slider" | "DateTimeInput"
)
{
return Err(format!("component `{id}` ({kind}) does not support checks"));
}
match kind {
"Row" | "Column" | "List" | "Form" => {
validate_children(component.get("children").unwrap(), id)?
}
"Card" | "Button" => expect_string(component, "child", id)?,
"Modal" => {
expect_string(component, "trigger", id)?;
expect_string(component, "content", id)?;
}
"Tabs" => validate_tabs(component.get("tabs").unwrap(), id)?,
"ChoicePicker" => validate_options(component.get("options").unwrap(), id)?,
"Slider" => {
expect_number(component, "max", id)?;
if let Some(min) = component.get("min")
&& !min.is_number()
{
return Err(format!("component `{id}` min must be a number"));
}
if let Some(steps) = component.get("steps")
&& !steps.as_u64().is_some_and(|steps| steps > 0)
{
return Err(format!("component `{id}` steps must be a positive integer"));
}
}
_ => {}
}
match kind {
"Text" => validate_dynamic_field(component, "text", Value::is_string, id)?,
"Image" | "Video" | "AudioPlayer" => {
validate_dynamic_field(component, "url", Value::is_string, id)?;
validate_optional_dynamic(component, "description", Value::is_string, id)?;
validate_optional_dynamic(component, "posterUrl", Value::is_string, id)?;
}
"Icon" => validate_dynamic_field(component, "name", Value::is_string, id)?,
"TextField" | "DateTimeInput" => {
validate_optional_dynamic(component, "label", Value::is_string, id)?;
validate_optional_dynamic(component, "value", Value::is_string, id)?;
}
"CheckBox" => {
validate_dynamic_field(component, "label", Value::is_string, id)?;
validate_dynamic_field(component, "value", Value::is_boolean, id)?;
}
"Slider" => validate_dynamic_field(component, "value", Value::is_number, id)?,
"ChoicePicker" => {
validate_optional_dynamic(component, "label", Value::is_string, id)?;
validate_dynamic_field(
component,
"value",
|value| {
value
.as_array()
.is_some_and(|values| values.iter().all(Value::is_string))
},
id,
)?;
}
"Chart" | "Table" | "Timeline" | "Map" | "MindMap" => {
for field in required {
if *field != "chartType" {
validate_dynamic_field(component, field, Value::is_array, id)?;
}
}
}
_ => {}
}
match kind {
"Text" => validate_enum(component, "variant", &["caption", "body"], id)?,
"Image" => {
validate_enum(
component,
"fit",
&["contain", "cover", "fill", "none", "scaleDown"],
id,
)?;
validate_enum(
component,
"variant",
&[
"icon",
"avatar",
"smallFeature",
"mediumFeature",
"largeFeature",
"header",
],
id,
)?;
}
"Icon" if component.get("name").is_some_and(Value::is_string) => validate_enum(
component,
"name",
&[
"accountCircle",
"add",
"arrowBack",
"arrowForward",
"attachFile",
"calendarToday",
"call",
"camera",
"check",
"close",
"delete",
"download",
"edit",
"event",
"error",
"fastForward",
"favorite",
"favoriteOff",
"folder",
"help",
"home",
"info",
"locationOn",
"lock",
"lockOpen",
"mail",
"menu",
"moreVert",
"moreHoriz",
"notificationsOff",
"notifications",
"pause",
"payment",
"person",
"phone",
"photo",
"play",
"print",
"refresh",
"rewind",
"search",
"send",
"settings",
"share",
"shoppingCart",
"skipNext",
"skipPrevious",
"star",
"starHalf",
"starOff",
"stop",
"upload",
"visibility",
"visibilityOff",
"volumeDown",
"volumeMute",
"volumeOff",
"volumeUp",
"warning",
],
id,
)?,
"Divider" => validate_enum(component, "axis", &["horizontal", "vertical"], id)?,
"Row" | "Column" => {
validate_enum(
component,
"justify",
&[
"start",
"center",
"end",
"spaceBetween",
"spaceAround",
"spaceEvenly",
"stretch",
],
id,
)?;
validate_enum(
component,
"align",
&["start", "center", "end", "stretch"],
id,
)?;
}
"List" => {
validate_enum(component, "direction", &["vertical", "horizontal"], id)?;
validate_enum(
component,
"align",
&["start", "center", "end", "stretch"],
id,
)?;
}
"Button" => validate_enum(
component,
"variant",
&["default", "primary", "borderless"],
id,
)?,
"TextField" => validate_enum(
component,
"variant",
&["longText", "number", "shortText", "obscured"],
id,
)?,
"ChoicePicker" => {
validate_enum(
component,
"variant",
&["multipleSelection", "mutuallyExclusive"],
id,
)?;
validate_enum(component, "displayStyle", &["checkbox", "chips"], id)?;
optional_bool(component, "filterable")?;
}
"DateTimeInput" => {
optional_bool(component, "enableDate")?;
optional_bool(component, "enableTime")?;
validate_optional_dynamic(component, "min", Value::is_string, id)?;
validate_optional_dynamic(component, "max", Value::is_string, id)?;
}
"Chart" => validate_enum(
component,
"chartType",
&[
"bar",
"line",
"area",
"stackedBar",
"pie",
"donut",
"heatmap",
],
id,
)?,
_ => {}
}
if let Some(action) = component.get("action") {
validate_action(action, id)?;
}
if let Some(checks) = component.get("checks") {
let checks = checks
.as_array()
.ok_or_else(|| format!("component `{id}` checks must be an array"))?;
for check in checks {
let check = object(Some(check), "check")?;
reject_unknown(check, &["condition", "message"], "check")?;
let condition = check
.get("condition")
.ok_or_else(|| format!("component `{id}` check requires condition"))?;
required_string(check, "message")?;
validate_function(condition)?;
}
}
validate_dynamic_values(component)?;
Ok(())
}
fn validate_dynamic_field(
object: &Map<String, Value>,
field: &str,
literal: impl Fn(&Value) -> bool,
id: &str,
) -> Result<(), String> {
let value = object
.get(field)
.ok_or_else(|| format!("component `{id}` requires `{field}`"))?;
if literal(value) {
return Ok(());
}
let dynamic = value
.as_object()
.ok_or_else(|| format!("component `{id}` {field} has the wrong literal or dynamic type"))?;
if let Some(path) = dynamic.get("path") {
if dynamic.len() == 1 && path.is_string() {
return Ok(());
}
return Err(format!(
"component `{id}` {field} has an invalid data binding"
));
}
validate_function(value).map_err(|error| format!("component `{id}` {field}: {error}"))
}
fn validate_optional_dynamic(
object: &Map<String, Value>,
field: &str,
literal: impl Fn(&Value) -> bool,
id: &str,
) -> Result<(), String> {
if object.contains_key(field) {
validate_dynamic_field(object, field, literal, id)
} else {
Ok(())
}
}
pub(super) fn ds4_catalog() -> &'static Value {
static CATALOG: OnceLock<Value> = OnceLock::new();
CATALOG.get_or_init(|| {
serde_json::from_str(CATALOG_JSON).expect("embedded DS4Server A2UI catalog must be valid")
})
}
fn expect_string(object: &Map<String, Value>, field: &str, id: &str) -> Result<(), String> {
if !object.get(field).is_some_and(Value::is_string) {
return Err(format!("component `{id}` {field} must be a string"));
}
Ok(())
}
fn expect_number(object: &Map<String, Value>, field: &str, id: &str) -> Result<(), String> {
if !object.get(field).is_some_and(Value::is_number) {
return Err(format!("component `{id}` {field} must be a number"));
}
Ok(())
}
fn validate_enum(
object: &Map<String, Value>,
field: &str,
allowed: &[&str],
id: &str,
) -> Result<(), String> {
if let Some(value) = object.get(field) {
let value = value
.as_str()
.ok_or_else(|| format!("component `{id}` {field} must be a string"))?;
if !allowed.contains(&value) {
return Err(format!(
"component `{id}` has invalid {field} `{value}`; expected one of: {}",
allowed.join(", ")
));
}
}
Ok(())
}
fn validate_children(value: &Value, id: &str) -> Result<(), String> {
if value
.as_array()
.is_some_and(|children| children.iter().all(Value::is_string))
{
return Ok(());
}
let template = object(Some(value), "children")?;
reject_unknown(template, &["componentId", "path"], "children template")?;
required_string(template, "componentId")?;
required_string(template, "path")?;
if !required_string(template, "path")?.starts_with('/') {
return Err(format!(
"component `{id}` child template path must be a JSON Pointer"
));
}
Ok(())
}
fn validate_tabs(value: &Value, id: &str) -> Result<(), String> {
let tabs = value
.as_array()
.filter(|tabs| !tabs.is_empty())
.ok_or_else(|| format!("component `{id}` tabs must be a non-empty array"))?;
for tab in tabs {
let tab = object(Some(tab), "tab")?;
reject_unknown(tab, &["title", "child"], "tab")?;
validate_dynamic_field(tab, "title", Value::is_string, id)?;
required_string(tab, "child")?;
}
Ok(())
}
fn validate_options(value: &Value, id: &str) -> Result<(), String> {
let options = value
.as_array()
.ok_or_else(|| format!("component `{id}` options must be an array"))?;
for option in options {
let option = object(Some(option), "choice option")?;
reject_unknown(option, &["label", "value"], "choice option")?;
validate_dynamic_field(option, "label", Value::is_string, id)?;
required_string(option, "value")?;
}
Ok(())
}
fn validate_action(value: &Value, id: &str) -> Result<(), String> {
let action = object(Some(value), "action")?;
if let Some(event) = action.get("event") {
reject_unknown(action, &["event"], "action")?;
let event = object(Some(event), "action.event")?;
reject_unknown(
event,
&["name", "context", "wantResponse", "responsePath"],
"action.event",
)?;
required_string(event, "name")?;
if let Some(context) = event.get("context")
&& !context.is_object()
{
return Err(format!("component `{id}` action context must be an object"));
}
optional_bool(event, "wantResponse")?;
if let Some(path) = event.get("responsePath")
&& !path.as_str().is_some_and(|path| path.starts_with('/'))
{
return Err(format!(
"component `{id}` responsePath must be a JSON Pointer"
));
}
return Ok(());
}
reject_unknown(action, &["functionCall"], "action")?;
validate_function(
action
.get("functionCall")
.ok_or_else(|| format!("component `{id}` action requires event or functionCall"))?,
)
}
fn validate_dynamic_values(value: &Map<String, Value>) -> Result<(), String> {
for value in value.values() {
if value.get("call").is_some() {
validate_function(value)?;
}
match value {
Value::Object(object) => validate_dynamic_values(object)?,
Value::Array(values) => {
for value in values {
if let Value::Object(object) = value {
validate_dynamic_values(object)?;
}
}
}
_ => {}
}
}
Ok(())
}
pub(super) fn validate_function(value: &Value) -> Result<(), String> {
let function = object(Some(value), "function call")?;
reject_unknown(function, &["call", "args"], "function call")?;
let call = value
.get("call")
.and_then(Value::as_str)
.ok_or_else(|| "function call requires `call`".to_owned())?;
if !FUNCTIONS.contains(&call) {
return Err(format!("function `{call}` is not declared by the catalog"));
}
if value.get("args").is_some_and(|args| !args.is_object()) {
return Err(format!("function `{call}` requires object `args`"));
}
let args = value
.get("args")
.and_then(Value::as_object)
.cloned()
.unwrap_or_default();
let allowed: &[&str] = match call {
"required" | "email" | "formatString" | "not" => &["value"],
"regex" => &["value", "pattern"],
"length" | "numeric" => &["value", "min", "max"],
"formatNumber" => &["value", "decimals", "grouping"],
"formatCurrency" => &["value", "currency", "decimals", "grouping"],
"formatDate" => &["value", "format"],
"pluralize" => &["value", "zero", "one", "two", "few", "many", "other"],
"openUrl" => &["url"],
"and" | "or" => &["values"],
"@index" => &["offset"],
_ => &[],
};
reject_unknown(&args, allowed, &format!("function `{call}` args"))?;
let required: &[&str] = match call {
"required" | "regex" | "length" | "numeric" | "email" | "formatString" | "formatNumber"
| "not" | "pluralize" => &["value"],
"formatCurrency" => &["value", "currency"],
"formatDate" => &["value", "format"],
"openUrl" => &["url"],
"and" | "or" => &["values"],
"@index" => &[],
_ => &[],
};
if let Some(field) = required.iter().find(|field| !args.contains_key(**field)) {
return Err(format!("function `{call}` requires argument `{field}`"));
}
if call == "regex" && !args.get("pattern").is_some_and(Value::is_string) {
return Err("function `regex` requires string argument `pattern`".into());
}
if matches!(call, "length" | "numeric")
&& !args.contains_key("min")
&& !args.contains_key("max")
{
return Err(format!("function `{call}` requires `min` or `max`"));
}
if call == "pluralize" && !args.contains_key("other") {
return Err("function `pluralize` requires argument `other`".into());
}
if matches!(call, "and" | "or")
&& !args
.get("values")
.and_then(Value::as_array)
.is_some_and(|values| values.len() >= 2)
{
return Err(format!("function `{call}` requires at least two values"));
}
match call {
"regex" | "length" | "email" | "formatString" => {
validate_dynamic_field(&args, "value", Value::is_string, call)?;
}
"numeric" | "formatNumber" | "formatCurrency" | "pluralize" => {
validate_dynamic_field(&args, "value", Value::is_number, call)?;
}
"not" => validate_dynamic_field(&args, "value", Value::is_boolean, call)?,
"formatDate" => validate_dynamic_field(&args, "format", Value::is_string, call)?,
"openUrl" => {
let url = required_string(&args, "url")?;
url::Url::parse(url).map_err(|error| format!("invalid openUrl URL: {error}"))?;
}
"and" | "or" => {
for item in args["values"].as_array().unwrap() {
if !item.is_boolean() {
let item = item.as_object().ok_or_else(|| {
format!("function `{call}` values must be dynamic booleans")
})?;
if !item.contains_key("path") && !item.contains_key("call") {
return Err(format!("function `{call}` values must be dynamic booleans"));
}
}
}
}
_ => {}
}
for field in ["min", "max", "decimals", "offset"] {
if args.contains_key(field) {
validate_dynamic_field(&args, field, Value::is_number, call)?;
}
}
if args.contains_key("grouping") {
validate_dynamic_field(&args, "grouping", Value::is_boolean, call)?;
}
for field in ["currency", "zero", "one", "two", "few", "many", "other"] {
if args.contains_key(field) {
validate_dynamic_field(&args, field, Value::is_string, call)?;
}
}
Ok(())
}

264
src/a2ui_validation.rs Normal file
View File

@@ -0,0 +1,264 @@
use crate::a2ui::{BASIC_NON_MEDIA_COMPONENTS, Store, extract_lines, validate_surface_composition};
use crate::config::Config;
use crate::model::ModelChoice;
use serde_json::{Value, json};
use std::time::Duration;
const CASES: &[Case] = &[
Case {
name: "pie-natural",
prompt: "Do not call tools. The file extension counts are rs 41, md 7, and toml 2. Give me a pie chart with the numbers.",
components: &["Chart"],
chart_type: Some("pie"),
},
Case {
name: "pie-single-slice",
prompt: "Do not call tools. There are 41 rs files and no other file extensions. Show that distribution as a one-slice pie chart.",
components: &["Chart"],
chart_type: Some("pie"),
},
Case {
name: "donut-natural",
prompt: "Do not call tools. Show a donut chart for 62 completed, 23 active, and 15 blocked tasks. Include the total.",
components: &["Chart"],
chart_type: Some("donut"),
},
Case {
name: "heatmap-natural",
prompt: "Do not call tools. Show a heatmap of pull requests reviewed in 2025 and 2026, with each year as a row and January, February, and March as columns. Use values 2, 5, 3 for 2025 and 4, 1, 6 for 2026.",
components: &["Chart"],
chart_type: Some("heatmap"),
},
Case {
name: "form-controls",
prompt: "Do not call tools. Build an A2UI Form for a name, multiline notes, due date, and multiple checkbox priorities, with a submit button.",
components: &["Form"],
chart_type: None,
},
Case {
name: "filterable-choices",
prompt: "Do not call tools. Build an A2UI surface with a filterable mutually-exclusive ChoicePicker for Rust, Python, and TypeScript.",
components: &["ChoicePicker"],
chart_type: None,
},
Case {
name: "composed-basics",
prompt: "Do not call tools. Build one complete A2UI project dashboard. Compose a Card and Column containing Markdown text, an avatar Image, an Icon, a Divider, a weighted Row, Tabs, a dynamic List, a Modal, a Button, TextField, CheckBox, Slider, DateTimeInput, and a filterable mutually-exclusive ChoicePicker. Bind the controls to the data model and give the button an event.",
components: BASIC_NON_MEDIA_COMPONENTS,
chart_type: None,
},
Case {
name: "media-players",
prompt: "Do not call tools. Build one A2UI Card containing a Video with a poster URL and an AudioPlayer with a description. Use valid HTTPS media URLs and compose both into the root tree.",
components: &["Video", "AudioPlayer"],
chart_type: None,
},
];
struct Case {
name: &'static str,
prompt: &'static str,
components: &'static [&'static str],
chart_type: Option<&'static str>,
}
struct Options {
endpoint: Option<String>,
model: Option<String>,
case: Option<String>,
attempts: u32,
}
pub(crate) fn run(args: impl Iterator<Item = String>) -> Result<(), String> {
let options = parse_options(args)?;
let config = Config::load(&crate::app::config_path())?;
let endpoint = options
.endpoint
.unwrap_or_else(|| format!("http://127.0.0.1:{}", config.endpoint.port));
let model_id = options
.model
.unwrap_or_else(|| config.model.id().to_owned());
let model =
ModelChoice::from_id(&model_id).ok_or_else(|| format!("unknown model `{model_id}`"))?;
let cases = CASES
.iter()
.filter(|case| options.case.as_deref().is_none_or(|name| name == case.name))
.collect::<Vec<_>>();
if cases.is_empty() {
return Err(format!(
"unknown case; choose one of: {}",
CASES
.iter()
.map(|case| case.name)
.collect::<Vec<_>>()
.join(", ")
));
}
let agent: ureq::Agent = ureq::Agent::config_builder()
.timeout_connect(Some(Duration::from_secs(5)))
.timeout_recv_response(Some(Duration::from_secs(30 * 60)))
.timeout_recv_body(Some(Duration::from_secs(30)))
.build()
.into();
let mut system = crate::agent::system_prompt(model, &config.system_prompt, false);
system.push_str("\n\n");
system.push_str(crate::a2ui::SYSTEM_PROMPT);
let metadata = Store::default().client_metadata();
let mut passed = 0;
let total = cases.len() as u32 * options.attempts;
println!("A2UI live validation: {model_id} at {endpoint}");
for case in cases {
for attempt in 1..=options.attempts {
let user = format!("{}\n\nA2UI client metadata:\n{}", case.prompt, metadata);
let payload = json!({
"model": model_id,
"messages": [
{"role": "system", "content": system},
{"role": "user", "content": user}
],
"reasoning_effort": "none",
"temperature": 0,
"max_tokens": 4096
});
let content = request(&agent, &endpoint, &payload)?;
match validate(case, &content) {
Ok(summary) => {
passed += 1;
println!("PASS {}#{attempt}: {summary}", case.name);
}
Err(error) => {
println!("FAIL {}#{attempt}: {error}", case.name);
println!(
"--- response ---\n{}\n--- end response ---",
bounded(&content)
);
}
}
}
}
println!("A2UI validation summary: {passed}/{total} passed");
if passed == total {
Ok(())
} else {
Err(format!("{} live validation case(s) failed", total - passed))
}
}
fn parse_options(mut args: impl Iterator<Item = String>) -> Result<Options, String> {
let mut options = Options {
endpoint: None,
model: None,
case: None,
attempts: 1,
};
while let Some(argument) = args.next() {
let value = match argument.as_str() {
"--endpoint" | "--model" | "--case" | "--attempts" => args
.next()
.ok_or_else(|| format!("{argument} requires a value"))?,
"--help" | "-h" => {
println!(
"Usage: ds4-server validate-a2ui [--endpoint URL] [--model ID] [--case NAME] [--attempts N]"
);
std::process::exit(0);
}
_ => return Err(format!("unknown argument `{argument}`")),
};
match argument.as_str() {
"--endpoint" => options.endpoint = Some(value.trim_end_matches('/').to_owned()),
"--model" => options.model = Some(value),
"--case" => options.case = Some(value),
"--attempts" => {
options.attempts = value
.parse::<u32>()
.ok()
.filter(|attempts| *attempts > 0)
.ok_or_else(|| "--attempts must be a positive integer".to_owned())?;
}
_ => unreachable!(),
}
}
Ok(options)
}
fn request(agent: &ureq::Agent, endpoint: &str, payload: &Value) -> Result<String, String> {
let mut response = agent
.post(&format!("{endpoint}/v1/chat/completions"))
.header("Content-Type", "application/json")
.send(payload.to_string().as_bytes())
.map_err(|error| format!("could not call the local endpoint: {error}"))?;
let body = response
.body_mut()
.read_to_string()
.map_err(|error| format!("could not read the local endpoint response: {error}"))?;
let value: Value = serde_json::from_str(&body)
.map_err(|error| format!("local endpoint returned invalid JSON: {error}"))?;
value
.pointer("/choices/0/message/content")
.and_then(Value::as_str)
.map(str::to_owned)
.ok_or_else(|| format!("local endpoint response has no assistant content: {body}"))
}
fn validate(case: &Case, content: &str) -> Result<String, String> {
let lines = extract_lines(content);
if lines.is_empty() {
return Err("the model emitted no complete fenced A2UI JSONL messages".into());
}
let mut store = Store::default();
let mut errors = Vec::new();
for (index, line) in lines.iter().enumerate() {
if let Err(error) = &line.value {
errors.push(format!("line {} is invalid JSON: {error}", index + 1));
} else if let Err(error) = store.apply_raw(&line.raw, 1) {
errors.push(format!(
"line {} failed catalog validation: {error}",
index + 1
));
}
}
if !errors.is_empty() {
return Err(errors.join("; "));
}
let surfaces = store.surfaces().collect::<Vec<_>>();
let mut reachable_components = std::collections::BTreeSet::new();
for surface in &surfaces {
reachable_components.extend(
validate_surface_composition(surface)
.map_err(|error| format!("surface `{}` is not composable: {error}", surface.id))?,
);
}
let components = surfaces
.iter()
.flat_map(|surface| surface.components.values())
.filter_map(Value::as_object)
.collect::<Vec<_>>();
for expected in case.components {
if !reachable_components.contains(*expected) {
return Err(format!(
"valid surface did not compose expected {expected} component into its root tree"
));
}
}
if let Some(expected) = case.chart_type
&& !components.iter().any(|component| {
component.get("component").and_then(Value::as_str) == Some("Chart")
&& component.get("chartType").and_then(Value::as_str) == Some(expected)
})
{
return Err(format!("valid surface omitted expected {expected} chart"));
}
Ok(format!(
"{} protocol message(s), {} surface(s), {} component(s)",
lines.len(),
surfaces.len(),
components.len()
))
}
fn bounded(content: &str) -> &str {
content
.char_indices()
.nth(8_000)
.map_or(content, |(end, _)| &content[..end])
}

File diff suppressed because it is too large Load Diff

480
src/agent/web.rs Normal file
View File

@@ -0,0 +1,480 @@
use headless_chrome::browser::{default_executable, tab::Tab};
use headless_chrome::protocol::cdp::{Emulation, Runtime};
use headless_chrome::{Browser as Chrome, LaunchOptions};
use serde_json::Value;
use std::ffi::OsStr;
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, mpsc};
use std::thread;
use std::time::Duration;
const PORT: u16 = 9333;
const CONNECT_TIMEOUT: Duration = Duration::from_secs(3);
const CDP_TIMEOUT: Duration = Duration::from_secs(20);
const POLL_INTERVAL: Duration = Duration::from_millis(50);
const CLICK_GOOGLE_CONSENT_JS: &str = r#"(() => {
const clean=s=>(s||'').replace(/\s+/g,' ').trim();
const pats=[/accept all/i,/i agree/i,/agree/i,/accetta tutto/i,/tout accepter/i,/aceptar todo/i,/alle akzeptieren/i];
const els=[...document.querySelectorAll('button,[role=button],input[type=submit],a')];
for (const el of els){const t=clean(el.innerText||el.value||el.textContent);if(!t)continue;if(pats.some(p=>p.test(t))){el.click();return 'clicked '+t;}}
return '';
})()"#;
const EXTRACT_SEARCH_JS: &str = r#"(() => {
const clean=s=>(s||'').replace(/\s+/g,' ').trim();
const esc=s=>clean(s).replace(/\\/g,'\\\\').replace(/\[/g,'\\[').replace(/\]/g,'\\]').replace(/\n/g,' ');
const visible=el=>{const r=el.getBoundingClientRect();const st=getComputedStyle(el);return r.width>0&&r.height>0&&st.display!=='none'&&st.visibility!=='hidden'&&st.opacity!=='0';};
const bad=h=>(/(^|\.)google\./.test(h)||/(^|\.)gstatic\./.test(h)||/(^|\.)googleusercontent\./.test(h));
const lines=['# Google search results','',`URL: ${location.href}`,'','## Visible links'];
const seen=new Set();
for(const a of document.querySelectorAll('a[href]')){if(!visible(a))continue;let href=a.href||'';try{const u=new URL(href);if(u.pathname==='/url'&&u.searchParams.get('q'))href=u.searchParams.get('q');}catch{}let u;try{u=new URL(href);}catch{continue;}if(!/^https?:$/.test(u.protocol))continue;if(bad(u.hostname))continue;const text=esc(a.innerText||a.textContent);if(text.length<3)continue;if(seen.has(u.href))continue;seen.add(u.href);lines.push(`- [${text.slice(0,180)}](${u.href})`);if(seen.size>=20)break;}
lines.push('','## Text snapshot',clean(document.body.innerText).slice(0,1200));
return lines.join('\n');
})()"#;
const EXTRACT_PAGE_JS: &str = r#"(() => {
const clean=s=>(s||'').replace(/\s+/g,' ').trim();
const esc=s=>clean(s).replace(/\\/g,'\\\\').replace(/\[/g,'\\[').replace(/\]/g,'\\]').replace(/\n/g,' ');
const visible=el=>{const r=el.getBoundingClientRect();const st=getComputedStyle(el);return r.width>0&&r.height>0&&st.display!=='none'&&st.visibility!=='hidden'&&st.opacity!=='0';};
const inline=n=>{if(!n)return'';if(n.nodeType===3)return n.nodeValue;if(n.nodeType!==1)return'';const el=n;if(el.tagName==='SCRIPT'||el.tagName==='STYLE'||el.tagName==='NOSCRIPT')return'';if(el.tagName==='A'){const t=esc(el.innerText||el.textContent);const h=el.href||'';return t&&h?`[${t}](${h})`:t;}if(el.tagName==='CODE')return '`'+clean(el.innerText||el.textContent).replace(/`/g,'\\`')+'`';return [...el.childNodes].map(inline).join('');};
const lines=[`# ${clean(document.title)||location.href}`,'',`URL: ${location.href}`,'','## Content'];
const blocks=[...document.body.querySelectorAll('h1,h2,h3,h4,h5,h6,p,li,pre,blockquote,td,th,[id="content-text"],[class*="comment-body"],[class*="comment-content"],[data-testid*="comment-text"]')];
const seen=new Set();
for(const el of blocks){if(!visible(el))continue;let s='';const tag=el.tagName;if(/^H[1-6]$/.test(tag)){s='#'.repeat(Number(tag[1]))+' '+inline(el);}else if(tag==='LI'){s='- '+inline(el);}else if(tag==='PRE'){s='```\n'+(el.innerText||el.textContent||'').trimEnd()+'\n```';}else if(tag==='BLOCKQUOTE'){s='> '+clean(el.innerText||el.textContent);}else{s=inline(el);}s=s.trim();if(!s||seen.has(s))continue;seen.add(s);lines.push('',s);if(lines.join('\n').length>900000){lines.push('','[Content truncated by browser extractor.]');break;}}
lines.push('','## Visible links');let n=0;const linkSeen=new Set();
for(const a of document.querySelectorAll('a[href]')){if(!visible(a))continue;const t=esc(a.innerText||a.textContent);if(t.length<3)continue;let u;try{u=new URL(a.href);}catch{continue;}if(!/^https?:$/.test(u.protocol)||linkSeen.has(u.href))continue;linkSeen.add(u.href);lines.push(`- [${t.slice(0,160)}](${u.href})`);if(++n>=80)break;}
return lines.join('\n');
})()"#;
const SCROLL_DYNAMIC_PAGE_JS: &str = r#"(() => new Promise(resolve => {
const root=()=>document.scrollingElement||document.documentElement||document.body;
const blockSel='h1,h2,h3,h4,h5,h6,p,li,pre,blockquote,td,th,[id="content-text"],[class*="comment-body"],[class*="comment-content"],[data-testid*="comment-text"]';
const lazySel='[onscroll],[loading="lazy"],[data-src],[data-lazy],[class*="lazy"],[class*="infinite"],[class*="virtual"],[role="feed"],[id*="comment"],[class*="comment"],[data-testid*="comment"]';
const hookCount=()=>{let n=0;try{if(window.onscroll)n++;if(document.onscroll)n++;if(document.body&&document.body.onscroll)n++;}catch(e){}try{if(typeof getEventListeners==='function'){for(const o of [window,document,document.body]){if(!o)continue;const ev=getEventListeners(o);if(ev&&ev.scroll)n+=ev.scroll.length;}}catch(e){}try{n+=document.querySelectorAll(lazySel).length;}catch(e){}return n;};
const metrics=()=>{const r=root();return {height:r?r.scrollHeight:0,view:innerHeight||900,y:scrollY||(r&&r.scrollTop)||0,text:((document.body&&document.body.innerText)||'').length,links:document.links?document.links.length:0,blocks:document.body?document.body.querySelectorAll(blockSel).length:0,hooks:hookCount()};};
const sig=m=>[m.height,m.text,m.links,m.blocks].join('|');
const grew=(a,b)=>b.height>a.height+20||b.text>a.text+200||b.links>a.links+2||b.blocks>a.blocks+2;
const scrollOnce=()=>{const r=root();if(!r)return;const h=Math.max(700,Math.floor((innerHeight||900)*0.85));window.scrollTo(0,Math.min(r.scrollHeight,(scrollY||r.scrollTop||0)+h));};
let last=metrics(),lastSig=sig(last),same=0,steps=0;
const scrollable=last.height>last.view*1.35;
if(!scrollable||last.hooks===0){resolve('scroll skipped hooks='+last.hooks+' text='+last.text);return;}
const tick=()=>{if(steps>=28){resolve('scrolled '+steps+' text='+last.text);return;}const before=last;scrollOnce();steps++;setTimeout(()=>{const now=metrics(),nowSig=sig(now);if(nowSig===lastSig)same++;else same=0;const loaded=grew(before,now);last=now;lastSig=nowSig;if(steps===1&&!loaded){resolve('scroll probe unchanged text='+now.text);return;}const atBottom=now.y+now.view+20>=now.height;if(same>=4||(atBottom&&same>=1)){resolve('scrolled '+steps+' text='+now.text);return;}tick();},900);};
tick();
}))()"#;
pub(super) struct Browser {
chrome: Option<Chrome>,
profile_dir: PathBuf,
}
impl Browser {
pub(super) fn new() -> Self {
Self {
chrome: None,
profile_dir: crate::app::browser_profile_path(),
}
}
pub(super) fn google_search(
&mut self,
query: &str,
cancel: &AtomicBool,
) -> Result<String, String> {
if query.is_empty() {
return Err("google_search requires query".into());
}
self.run_page(&google_search_url(query), EXTRACT_SEARCH_JS, false, cancel)
}
pub(super) fn visit_page(&mut self, url: &str, cancel: &AtomicBool) -> Result<String, String> {
if url.is_empty() {
return Err("visit_page requires url".into());
}
self.run_page(url, EXTRACT_PAGE_JS, true, cancel)
}
fn run_page(
&mut self,
url: &str,
extract: &str,
dynamic_scroll: bool,
cancel: &AtomicBool,
) -> Result<String, String> {
let chrome = self.ensure_browser(cancel)?.clone();
let tab = new_tab(chrome, cancel)?;
tab.set_default_timeout(CDP_TIMEOUT);
let result = (|| {
prepare_page(&tab, cancel)?;
check_cancel(cancel)?;
tab.navigate_to(url)
.map_err(|error| format!("page navigation failed: {error}"))?;
wait_navigated_ready(&tab, cancel)?;
if let Ok(clicked) = evaluate_string(&tab, CLICK_GOOGLE_CONSENT_JS, cancel)
&& !clicked.is_empty()
{
sleep(cancel, Duration::from_millis(1500))?;
wait_navigated_ready(&tab, cancel)?;
}
if dynamic_scroll {
scroll_dynamic_page(&tab, cancel)?;
}
evaluate_string(&tab, extract, cancel)
})();
close_tab(tab, cancel.load(Ordering::Relaxed));
result
}
fn ensure_browser(&mut self, cancel: &AtomicBool) -> Result<&Chrome, String> {
if self.chrome.is_none() {
let profile_dir = self.profile_dir.clone();
self.chrome = Some(background(cancel, move || start_browser(profile_dir))?);
}
Ok(self.chrome.as_ref().unwrap())
}
}
fn start_browser(profile_dir: PathBuf) -> Result<Chrome, String> {
if let Some(browser) = connect_existing() {
return Ok(browser);
}
if profile_dir.to_str().is_none() {
return Err("The browser profile path is not valid UTF-8.".into());
}
std::fs::create_dir_all(&profile_dir)
.map_err(|error| format!("failed to create Chrome profile dir: {error}"))?;
let executable = std::env::var_os("DS4_CHROME")
.filter(|path| !path.is_empty())
.map(PathBuf::from)
.map_or_else(default_executable, Ok)
.map_err(|error| format!("could not find Chrome: {error}"))?;
let args = [
OsStr::new("--remote-allow-origins=*"),
OsStr::new("--no-default-browser-check"),
OsStr::new("--disable-sync"),
OsStr::new("--use-mock-keychain"),
OsStr::new("--password-store=basic"),
OsStr::new("--mute-audio"),
OsStr::new("about:blank"),
];
let options = LaunchOptions::default_builder()
.headless(true)
.path(Some(executable))
.user_data_dir(Some(profile_dir))
.port(Some(PORT))
.idle_browser_timeout(CDP_TIMEOUT)
.ignore_certificate_errors(false)
.enable_gpu(true)
.disable_default_args(true)
.args(args.to_vec())
.build()
.map_err(|error| format!("could not configure Chrome: {error}"))?;
Chrome::new(options).map_err(|error| format!("could not start Chrome: {error}"))
}
fn connect_existing() -> Option<Chrome> {
let agent: ureq::Agent = ureq::Agent::config_builder()
.timeout_connect(Some(CONNECT_TIMEOUT))
.timeout_recv_response(Some(CONNECT_TIMEOUT))
.timeout_recv_body(Some(CONNECT_TIMEOUT))
.build()
.into();
let mut response = agent
.get(&format!("http://127.0.0.1:{PORT}/json/version"))
.call()
.ok()?;
let body = response.body_mut().read_to_string().ok()?;
let websocket = serde_json::from_str::<Value>(&body)
.ok()?
.get("webSocketDebuggerUrl")?
.as_str()?
.to_owned();
Chrome::connect_with_timeout(websocket, CDP_TIMEOUT).ok()
}
fn new_tab(chrome: Chrome, cancel: &AtomicBool) -> Result<Arc<Tab>, String> {
let (sender, receiver) = mpsc::sync_channel(1);
thread::spawn(move || {
let result = chrome
.new_tab()
.map_err(|error| format!("could not open browser tab: {error}"));
if let Err(mpsc::SendError(Ok(tab))) = sender.send(result) {
let _ = tab.close_target();
}
});
wait(receiver, cancel)
}
fn prepare_page(tab: &Arc<Tab>, cancel: &AtomicBool) -> Result<(), String> {
check_cancel(cancel)?;
tab.call_method(Runtime::Enable(None))
.map_err(|error| format!("could not enable page runtime: {error}"))?;
let _ = tab.call_method(Emulation::SetFocusEmulationEnabled { enabled: true });
let _ = tab.call_method(Emulation::SetDeviceMetricsOverride {
width: 1365,
height: 900,
device_scale_factor: 1.0,
mobile: false,
scale: None,
screen_width: None,
screen_height: None,
position_x: None,
position_y: None,
dont_set_visible_size: None,
screen_orientation: None,
viewport: None,
display_feature: None,
device_posture: None,
});
wait_ready(tab, cancel)
}
fn wait_ready(tab: &Arc<Tab>, cancel: &AtomicBool) -> Result<(), String> {
for _ in 0..80 {
check_cancel(cancel)?;
if let Ok(state) = evaluate_string(tab, "document.readyState", cancel)
&& matches!(state.as_str(), "complete" | "interactive")
{
return sleep(cancel, Duration::from_millis(800));
}
sleep(cancel, Duration::from_millis(250))?;
}
Ok(())
}
fn wait_navigated_ready(tab: &Arc<Tab>, cancel: &AtomicBool) -> Result<(), String> {
let mut last_len = -1;
let mut stable = 0;
let mut saw_real_url = false;
for iteration in 0..100 {
check_cancel(cancel)?;
let Ok(probe) = evaluate_string(
tab,
"location.href+'\\n'+document.readyState+'\\n'+((document.body&&document.body.innerText)||'').length",
cancel,
) else {
sleep(cancel, Duration::from_millis(250))?;
continue;
};
let Ok((href, ready, text_len)) = parse_probe(&probe) else {
sleep(cancel, Duration::from_millis(250))?;
continue;
};
let real_url = !href.is_empty() && href != "about:blank" && !href.starts_with("chrome://");
let ready = matches!(ready, "complete" | "interactive");
saw_real_url |= real_url;
if text_len > 0 && text_len == last_len {
stable += 1;
} else {
stable = 0;
}
last_len = text_len;
if saw_real_url && ready && text_len > 0 && stable >= 2 {
return sleep(cancel, Duration::from_millis(500));
}
if saw_real_url && ready && iteration >= 24 {
return Ok(());
}
sleep(cancel, Duration::from_millis(250))?;
}
Ok(())
}
fn scroll_dynamic_page(tab: &Arc<Tab>, cancel: &AtomicBool) -> Result<(), String> {
check_cancel(cancel)?;
let result = evaluate_string(tab, SCROLL_DYNAMIC_PAGE_JS, cancel);
check_cancel(cancel)?;
match result {
Err(error) if error == "interrupted" => Err(error),
_ => Ok(()),
}
}
fn evaluate_string(
tab: &Arc<Tab>,
expression: &str,
cancel: &AtomicBool,
) -> Result<String, String> {
let tab = Arc::clone(tab);
let expression = expression.to_owned();
background(cancel, move || {
let response = tab
.call_method(Runtime::Evaluate {
expression,
object_group: None,
include_command_line_api: Some(true),
silent: None,
context_id: None,
return_by_value: Some(true),
generate_preview: None,
user_gesture: None,
await_promise: Some(true),
throw_on_side_effect: None,
timeout: None,
disable_breaks: None,
repl_mode: None,
allow_unsafe_eval_blocked_by_csp: None,
unique_context_id: None,
serialization_options: None,
})
.map_err(|error| format!("JavaScript evaluation failed: {error}"))?;
if response.exception_details.is_some() {
return Err("JavaScript evaluation failed".into());
}
response
.result
.value
.and_then(|value| value.as_str().map(str::to_owned))
.ok_or_else(|| "Runtime.evaluate did not return a string".into())
})
}
fn background<T, F>(cancel: &AtomicBool, operation: F) -> Result<T, String>
where
T: Send + 'static,
F: FnOnce() -> Result<T, String> + Send + 'static,
{
let (sender, receiver) = mpsc::sync_channel(1);
thread::spawn(move || {
let _ = sender.send(operation());
});
wait(receiver, cancel)
}
fn wait<T>(receiver: mpsc::Receiver<Result<T, String>>, cancel: &AtomicBool) -> Result<T, String> {
loop {
match receiver.recv_timeout(POLL_INTERVAL) {
Ok(result) => return result,
Err(mpsc::RecvTimeoutError::Timeout) => check_cancel(cancel)?,
Err(mpsc::RecvTimeoutError::Disconnected) => {
return Err("browser worker stopped unexpectedly".into());
}
}
}
}
fn close_tab(tab: Arc<Tab>, background: bool) {
if background {
thread::spawn(move || {
let _ = tab.close_target();
});
} else {
let _ = tab.close_target();
}
}
fn check_cancel(cancel: &AtomicBool) -> Result<(), String> {
if cancel.load(Ordering::Relaxed) {
Err("interrupted".into())
} else {
Ok(())
}
}
fn sleep(cancel: &AtomicBool, duration: Duration) -> Result<(), String> {
let mut remaining = duration;
while !remaining.is_zero() {
check_cancel(cancel)?;
let step = remaining.min(POLL_INTERVAL);
thread::sleep(step);
remaining = remaining.saturating_sub(step);
}
check_cancel(cancel)
}
fn parse_probe(probe: &str) -> Result<(&str, &str, i64), String> {
let mut parts = probe.splitn(3, '\n');
let href = parts.next().unwrap_or_default();
let ready = parts
.next()
.ok_or_else(|| "page readiness probe returned malformed data".to_owned())?;
let text_len = parts
.next()
.ok_or_else(|| "page readiness probe returned malformed data".to_owned())?
.parse()
.unwrap_or(0);
Ok((href, ready, text_len))
}
fn google_search_url(query: &str) -> String {
const HEX: &[u8; 16] = b"0123456789ABCDEF";
let mut url = String::from("https://www.google.com/search?q=");
for byte in query.bytes() {
if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b'~') {
url.push(char::from(byte));
} else {
url.push('%');
url.push(char::from(HEX[(byte >> 4) as usize]));
url.push(char::from(HEX[(byte & 15) as usize]));
}
}
url
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::{Read, Write};
use std::net::TcpListener;
#[test]
fn browser_helpers_preserve_ds4_wire_values() {
assert_eq!(
google_search_url("café & Rust~"),
"https://www.google.com/search?q=caf%C3%A9%20%26%20Rust~"
);
assert_eq!(
parse_probe("https://example.com\ncomplete\n42").unwrap(),
("https://example.com", "complete", 42)
);
assert!(parse_probe("incomplete").is_err());
}
#[test]
#[ignore = "requires an installed Chrome browser"]
fn headless_browser_extracts_a_local_page() {
let listener = TcpListener::bind(("127.0.0.1", 0)).unwrap();
let address = listener.local_addr().unwrap();
let server = thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
let mut request = [0; 4096];
let _ = stream.read(&mut request);
let body = "<title>Fixture</title><h1>Hello</h1><p>Browser works.</p>";
write!(
stream,
"HTTP/1.1 200 OK\r\nContent-Type: text/html\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
)
.unwrap();
});
let profile_dir = std::env::temp_dir().join(format!(
"ds4-browser-test-{}",
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let mut browser = Browser {
chrome: None,
profile_dir: profile_dir.clone(),
};
let output = browser
.run_page(
&format!("http://{address}"),
EXTRACT_PAGE_JS,
false,
&AtomicBool::new(false),
)
.unwrap();
assert!(output.contains("# Fixture"));
assert!(output.contains("# Hello"));
assert!(output.contains("Browser works."));
drop(browser);
server.join().unwrap();
std::fs::remove_dir_all(profile_dir).unwrap();
}
}

3192
src/app.rs

File diff suppressed because it is too large Load Diff

159
src/app/extensions.rs Normal file
View File

@@ -0,0 +1,159 @@
use super::*;
impl App {
pub(super) fn install_extension(&mut self) {
let source = self.extension_source.trim().to_owned();
let requested_ref = self.extension_ref.trim().to_owned();
if source.is_empty() {
self.extension_error = Some("Enter an HTTPS Git repository URL.".into());
return;
}
self.cancel_running_extension_hooks();
let root = extensions_path();
self.start_extension_operation("Installing extension", move || {
crate::extensions::ExtensionRegistry::install(
&root,
&source,
(!requested_ref.is_empty()).then_some(requested_ref.as_str()),
)
});
}
pub(super) fn update_extension(&mut self, id: String) {
self.cancel_running_extension_hooks();
let root = extensions_path();
self.start_extension_operation("Updating extension", move || {
crate::extensions::ExtensionRegistry::update(&root, &id)
});
}
pub(super) fn toggle_extension(&mut self, id: String, enabled: bool) {
let Some(extension) = self
.extensions
.extensions
.iter()
.find(|extension| extension.id == id)
else {
return;
};
if enabled && extension.has_commands() && !extension.trusted {
self.pending_extension_trust = Some(id);
return;
}
self.cancel_running_extension_hooks();
let root = extensions_path();
self.start_extension_operation(
if enabled {
"Enabling extension"
} else {
"Disabling extension"
},
move || crate::extensions::ExtensionRegistry::set_enabled(&root, &id, enabled),
);
}
pub(super) fn confirm_extension_trust(&mut self) {
let Some(id) = self.pending_extension_trust.take() else {
return;
};
self.cancel_running_extension_hooks();
let root = extensions_path();
self.start_extension_operation("Enabling trusted extension", move || {
crate::extensions::ExtensionRegistry::trust_and_enable(&root, &id)
});
}
pub(super) fn confirm_uninstall_extension(&mut self) {
let Some(id) = self.pending_extension_uninstall.take() else {
return;
};
self.cancel_running_extension_hooks();
let root = extensions_path();
self.start_extension_operation("Uninstalling extension", move || {
crate::extensions::ExtensionRegistry::uninstall(&root, &id)
});
}
fn start_extension_operation(
&mut self,
label: &str,
operation: impl FnOnce() -> Result<crate::extensions::ExtensionRegistry, String>
+ Send
+ 'static,
) {
if self.extension_operation.is_some() {
self.extension_error = Some("Another extension operation is already running.".into());
return;
}
let (sender, receiver) = mpsc::channel();
thread::spawn(move || {
let _ = sender.send(operation());
});
self.extension_operation = Some(ActiveExtensionOperation {
label: label.into(),
receiver,
});
self.extension_error = None;
}
pub(super) fn poll_extension_operation(&mut self) {
let Some(operation) = &self.extension_operation else {
return;
};
match operation.receiver.try_recv() {
Ok(Ok(extensions)) => {
self.extensions = extensions;
self.extension_operation = None;
self.extension_source.clear();
self.extension_ref.clear();
self.extension_error = None;
#[cfg(target_os = "macos")]
{
self.agent_tools = None;
}
}
Ok(Err(error)) => {
self.extension_operation = None;
self.extension_error = Some(error);
}
Err(TryRecvError::Empty) => {}
Err(TryRecvError::Disconnected) => {
self.extension_operation = None;
self.extension_error = Some("The extension operation stopped unexpectedly.".into());
}
}
}
pub(super) fn extension_operation_label(&self) -> Option<&str> {
self.extension_operation
.as_ref()
.map(|operation| operation.label.as_str())
}
pub(super) fn cancel_running_extension_hooks(&mut self) {
#[cfg(target_os = "macos")]
{
let requests = self
.active_extension_hooks
.iter()
.chain(
self.background_chats
.values()
.filter_map(|chat| chat.active_extension_hooks.as_ref()),
)
.map(|request| (Arc::clone(&request.cancel), Arc::clone(&request.finished)))
.collect::<Vec<_>>();
for (cancel, _) in &requests {
cancel.store(true, Ordering::Relaxed);
}
let deadline = Instant::now() + Duration::from_secs(1);
while requests
.iter()
.any(|(_, finished)| !finished.load(Ordering::Acquire))
&& Instant::now() < deadline
{
thread::sleep(Duration::from_millis(10));
}
}
}
}

Some files were not shown because too many files have changed in this diff Show More