Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 24 additions & 0 deletions deploy/install_prefill_worker_launchd.sh
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,16 @@ MODEL_REVISION="${KAKEYA_MODEL_REVISION:-}"
TOKENIZER_REVISION="${KAKEYA_TOKENIZER_REVISION:-}"
QUANTIZATION="${KAKEYA_CACHE_QUANTIZATION:-4bit-mlx}"
ROPE_HASH="${KAKEYA_ROPE_HASH:-}"
SINK="${KAKEYA_WORKER_SINK:-4}"
WINDOW="${KAKEYA_WORKER_WINDOW:-64}"
BLOCK_TOKENS="${KAKEYA_CACHE_BLOCK_TOKENS:-64}"
PREFILL_TPS="${KAKEYA_WORKER_PREFILL_TPS:-20}"
NETWORK="${KAKEYA_WORKER_NETWORK:-lan}"
PRIORITY="${KAKEYA_WORKER_PRIORITY:-50}"
RTT_MS="${KAKEYA_WORKER_RTT_MS:-1.0}"
PEER="${KAKEYA_WORKER_PEER:-}"
MAX_CONCURRENT_JOBS="${KAKEYA_WORKER_MAX_CONCURRENT_JOBS:-1}"
MAX_PROMPT_TOKENS="${KAKEYA_WORKER_MAX_PROMPT_TOKENS:-131072}"
LABEL="ai.kakeya.prefill-worker"
PLIST="$HOME/Library/LaunchAgents/$LABEL.plist"
LOG_DIR="$HOME/.kakeya"
Expand All @@ -27,6 +37,10 @@ psk_xml=""
if [[ -n "$PSK_FILE" ]]; then
psk_xml="<string>--fleet-psk-file</string><string>$PSK_FILE</string>"
fi
peer_xml=""
if [[ -n "$PEER" ]]; then
peer_xml="<string>--peer</string><string>$PEER</string>"
fi

cat > "$PLIST" <<EOF
<?xml version="1.0" encoding="UTF-8"?>
Expand All @@ -49,6 +63,16 @@ cat > "$PLIST" <<EOF
<string>--layer-geometry-hash</string><string>$KAKEYA_LAYER_GEOMETRY_HASH</string>
<string>--tenant-id</string><string>$TENANT</string>
<string>--cache-gb</string><string>$CACHE_GB</string>
<string>--sink</string><string>$SINK</string>
<string>--window</string><string>$WINDOW</string>
<string>--block-size-tokens</string><string>$BLOCK_TOKENS</string>
<string>--prefill-tps</string><string>$PREFILL_TPS</string>
<string>--network</string><string>$NETWORK</string>
<string>--priority</string><string>$PRIORITY</string>
<string>--rtt-ms</string><string>$RTT_MS</string>
<string>--max-concurrent-jobs</string><string>$MAX_CONCURRENT_JOBS</string>
<string>--max-prompt-tokens</string><string>$MAX_PROMPT_TOKENS</string>
$peer_xml
$psk_xml
</array>
<key>WorkingDirectory</key><string>$KAKEYA_WORKER_REPO</string>
Expand Down
14 changes: 13 additions & 1 deletion deploy/launchd/ai.kakeya.grpc-runtime-prefill.plist
Original file line number Diff line number Diff line change
Expand Up @@ -17,19 +17,31 @@
<string>--skip-cache-check</string>
<string>--enable-prefill-cache</string>
<string>--prefill-cache-gb</string><string>1</string>
<string>--cache-peer</string><string>169.254.27.104:52051</string>
<string>--peer</string><string>169.254.27.104:53051</string>
<string>--cache-peer</string><string>169.254.27.104:53051</string>
<string>--cache-model-id</string><string>gemma-4-26B-A4B-it-mlx-4bit</string>
<string>--model-revision</string><string>local-4bit-v1</string>
<string>--tokenizer-revision</string><string>gemma4-v1</string>
<string>--cache-quantization</string><string>4bit-mlx</string>
<string>--cache-kv-dtype</string><string>bfloat16</string>
<string>--cache-block-tokens</string><string>64</string>
<string>--cache-tenant-id</string><string>private-fleet</string>
<string>--fleet-psk-file</string><string>/Users/fluffy314/.kakeya/fleet.psk</string>
<string>--node-id</string><string>head-runtime</string>
<string>--advertise</string><string>169.254.187.239:51051</string>
<string>--cache-advertise</string><string>169.254.187.239:51051</string>
<string>--network-label</string><string>thunderbolt</string>
<string>--network-priority</string><string>100</string>
<string>--measured-rtt-ms</string><string>0.55</string>
<string>--remote-prefill-min-tokens</string><string>128</string>
<string>--cache-link-mbps</string><string>10000</string>
<string>--cache-default-rtt-ms</string><string>0.55</string>
<string>--prefill-min-savings-ratio</string><string>0</string>
<string>--primary-prefill-penalty-ms</string><string>1000</string>
<string>--network-http-host</string><string>127.0.0.1</string>
<string>--network-http-port</string><string>8090</string>
<string>--network-api-key</string><string>__NETWORK_KEY__</string>
<string>--network-state</string><string>/Users/fluffy314/.kakeya/inference_network.json</string>
<string>--network-telemetry-url</string><string>http://127.0.0.1:8090/v1/network/telemetry/tokens</string>
<string>--network-telemetry-api-key</string><string>__NETWORK_KEY__</string>
<string>--log-level</string><string>INFO</string>
Expand Down
22 changes: 21 additions & 1 deletion docs/ops/distributed-prefill-kv-network.md
Original file line number Diff line number Diff line change
Expand Up @@ -126,8 +126,16 @@ export KAKEYA_CACHE_MODEL_ID="gemma-4-26B-A4B-it-mlx-4bit"
export KAKEYA_MODEL_REVISION="local-4bit-v1"
export KAKEYA_TOKENIZER_REVISION="gemma4-v1"
export KAKEYA_WORKER_NODE_ID="prefill-mini-1"
export KAKEYA_WORKER_BIND="169.254.27.104:53051"
export KAKEYA_WORKER_ADVERTISE="169.254.27.104:53051"
export KAKEYA_LAYER_GEOMETRY_HASH="<same-value-as-primary>"
export KAKEYA_WORKER_SINK="4"
export KAKEYA_WORKER_WINDOW="2048"
export KAKEYA_CACHE_BLOCK_TOKENS="64"
export KAKEYA_WORKER_NETWORK="thunderbolt"
export KAKEYA_WORKER_PRIORITY="100"
export KAKEYA_WORKER_RTT_MS="0.55"
export KAKEYA_WORKER_PEER="169.254.187.239:51051"
export KAKEYA_FLEET_PSK_FILE="$HOME/.kakeya/fleet.psk"
export KAKEYA_TENANT_ID="private-fleet"
bash deploy/install_prefill_worker_launchd.sh
Expand Down Expand Up @@ -158,7 +166,7 @@ KEY="$(cat ~/.kakeya/network_api_key)"
curl -fsS -X POST http://127.0.0.1:8090/v1/network/nodes/register \
-H "Content-Type: application/json" \
-H "X-API-Key: $KEY" \
-d '{"alias":"peer-mini","address":"169.254.27.104:52051","region":"Private","role":"cache"}'
-d '{"alias":"peer-mini","address":"169.254.27.104:53051","region":"Private","role":"hybrid"}'
```

Create a paired group:
Expand Down Expand Up @@ -192,8 +200,20 @@ Minimal acceptance:
curl -fsS https://kakeya.ai/healthz
curl -fsS https://kakeya.ai/v1/network/summary
curl -fsS https://kakeya.ai/v1/network/tokens
curl -fsS https://kakeya.ai/v1/network/prefill

PYTHONPATH=.:sdks/python python scripts/verify_remote_prefill_e2e.py \
--address 127.0.0.1:51051 \
--dashboard http://127.0.0.1:8090 \
--tokenizer-id ~/kakeya-models/gemma-4-26B-A4B-it-mlx-4bit
```

The verifier exits non-zero unless a live worker capability is present and one
cold unique prefix increments `remote_jobs`, `remote_hits`, and
`tokens_reused`. Decode throughput is reported separately; remote prefill is
accepted on lower TTFT/prefill time and higher request throughput, not a change
to single-stream decode tokens/s.

## Rollback

The cache is an optimization; inference correctness does not depend on it.
Expand Down
4 changes: 4 additions & 0 deletions inference_engine/network/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,10 @@ def tokens():
"hit_rate": summary["kv_hit_rate"],
}

@app.get("/v1/network/prefill")
def prefill():
return state.prefill_stats()

@app.post(
"/v1/network/telemetry/tokens",
dependencies=[Depends(require_key)],
Expand Down
8 changes: 4 additions & 4 deletions inference_engine/network/dashboard.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@ def dashboard_html() -> str:
h1{font-size:24px;margin:0}h2{font-size:17px;margin:22px 0 10px}.sub,.muted{color:var(--muted)}
button{background:transparent;border:1px solid var(--line);color:var(--text);border-radius:7px;padding:8px 12px;cursor:pointer}
button.active,.primary{background:var(--accent);color:#0b1020;border-color:var(--accent);font-weight:650}
.tabs{display:flex;gap:7px}.stats,.grid2,.grid3{display:grid;gap:12px}.stats{grid-template-columns:repeat(5,1fr);margin-top:22px}.grid2{grid-template-columns:1.2fr .8fr}.grid3{grid-template-columns:repeat(3,1fr)}
.tabs{display:flex;gap:7px}.stats,.grid2,.grid3{display:grid;gap:12px}.stats{grid-template-columns:repeat(auto-fit,minmax(140px,1fr));margin-top:22px}.grid2{grid-template-columns:1.2fr .8fr}.grid3{grid-template-columns:repeat(3,1fr)}
.card{background:var(--panel);border:1px solid var(--line);border-radius:9px;padding:15px}.stat b{font-size:23px;display:block}.stat span{color:var(--muted);font-size:12px}
.map{height:280px;position:relative;overflow:hidden}.map:before{content:"";position:absolute;inset:34px;border:1px dashed var(--line);border-radius:45%}
.dot{position:absolute;transform:translate(-50%,-50%)}.dot i{display:block;width:12px;height:12px;background:var(--ok);border:2px solid var(--bg);border-radius:50%}.dot.head i{width:16px;height:16px;background:var(--accent)}.dot small{white-space:nowrap}
Expand All @@ -31,8 +31,8 @@ def dashboard_html() -> str:
<body><main class="wrap">
<div class="top"><div><h1>Kakeya Inference Network</h1><div class="sub">P2P Prefill KV sharing across trusted inference nodes</div></div>
<div class="tabs"><button class="active" data-tab="overview">Overview</button><button data-tab="nodes">Nodes</button><button data-tab="groups">Groups</button><button class="primary" id="registerBtn">Register node</button></div></div>
<section id="register" class="card register hidden"><h3>Register inference node</h3><div class="grid3"><label>Alias<input id="alias" placeholder="cache-peer-tb"></label><label>Address<input id="address" placeholder="169.254.27.104:52051"></label><label>Region<input id="region" placeholder="Hong Kong"></label></div><label>Admin API key<input id="adminKey" type="password" placeholder="Required for network changes"></label><button class="primary" id="createRegistration">Create pairing token</button><code id="pairing" class="hidden"></code></section>
<section class="stats"><div class="card stat"><b id="online">0</b><span>Online nodes</span></div><div class="card stat"><b id="groupCount">0</b><span>Inference groups</span></div><div class="card stat"><b id="tokens">0</b><span>Completed tokens</span></div><div class="card stat"><b id="hitRate">0%</b><span>KV-assisted tokens</span></div><div class="card stat"><b id="cache">0 GB</b><span>Shared cache online</span></div></section>
<section id="register" class="card register hidden"><h3>Register inference node</h3><div class="grid3"><label>Alias<input id="alias" placeholder="prefill-worker-tb"></label><label>Address<input id="address" placeholder="169.254.27.104:53051"></label><label>Region<input id="region" placeholder="Hong Kong"></label></div><label>Admin API key<input id="adminKey" type="password" placeholder="Required for network changes"></label><button class="primary" id="createRegistration">Create pairing token</button><code id="pairing" class="hidden"></code></section>
<section class="stats"><div class="card stat"><b id="online">0</b><span>Online nodes</span></div><div class="card stat"><b id="groupCount">0</b><span>Inference groups</span></div><div class="card stat"><b id="tokens">0</b><span>Completed tokens</span></div><div class="card stat"><b id="hitRate">0%</b><span>KV-assisted tokens</span></div><div class="card stat"><b id="cache">0 GB</b><span>Shared cache online</span></div><div class="card stat"><b id="remoteJobs">0</b><span>Remote prefill jobs</span></div><div class="card stat"><b id="remoteHits">0</b><span>Remote KV imports</span></div><div class="card stat"><b id="reusedTokens">0</b><span>Tokens reused</span></div></section>
<section id="overview" class="tab">
<div class="grid2"><div><h2>Online node distribution</h2><div class="card map" id="map"></div></div><div><h2>Live KV discovery</h2><div class="card" id="events"><div class="event"><time>live</time><div><b>Waiting for node telemetry</b><div class="muted">Capability gossip and prefix lookups appear here.</div></div></div></div><h2>Cache capacity</h2><div class="card"><span id="capacityLabel">0 / 0 GB</span><div class="bar"><i id="capacityBar" style="width:0%"></i></div></div></div></div>
</section>
Expand All @@ -49,7 +49,7 @@ def dashboard_html() -> str:
$('createGroup').onclick=async()=>{await fetch('/v1/network/groups',{method:'POST',headers:writeHeaders(),body:JSON.stringify({name:$('groupName').value,node_ids:$('groupNodes').value.split(',').map(x=>x.trim()).filter(Boolean)})});load()};
function nodePosition(i,total){let a=(i/Math.max(total,1))*Math.PI*2;return {x:50+38*Math.cos(a),y:53+35*Math.sin(a)}}
async function load(){let [s,n,g]=await Promise.all([fetch('/v1/network/summary').then(r=>r.json()),fetch('/v1/network/nodes').then(r=>r.json()),fetch('/v1/network/groups').then(r=>r.json())]);
$('online').textContent=s.online_nodes;$('groupCount').textContent=s.groups;$('tokens').textContent=fmt(s.completed_tokens);$('hitRate').textContent=(s.kv_hit_rate*100).toFixed(0)+'%';$('cache').textContent=gb(s.cache_bytes_used+s.cache_bytes_free)+' GB';
$('online').textContent=s.online_nodes;$('groupCount').textContent=s.groups;$('tokens').textContent=fmt(s.completed_tokens);$('hitRate').textContent=(s.kv_hit_rate*100).toFixed(0)+'%';$('cache').textContent=gb(s.cache_bytes_used+s.cache_bytes_free)+' GB';let p=s.prefill||{};$('remoteJobs').textContent=fmt(p.remote_jobs);$('remoteHits').textContent=fmt(p.remote_hits);$('reusedTokens').textContent=fmt(p.tokens_reused);
let total=s.cache_bytes_used+s.cache_bytes_free,pct=total?s.cache_bytes_used/total*100:0;$('capacityLabel').textContent=`${gb(s.cache_bytes_used)} / ${gb(total)} GB`;$('capacityBar').style.width=pct+'%';
$('map').innerHTML=n.map((x,i)=>{let p=nodePosition(i,n.length);return `<div class="dot ${x.role.includes('head')?'head':''}" style="left:${p.x}%;top:${p.y}%"><i></i><small>${x.region}<br>${x.alias}</small></div>`}).join('');
$('nodesBody').innerHTML=n.map(x=>`<tr><td>${x.alias}</td><td>${x.role}</td><td>${x.region}</td><td>${x.cache?(x.cache.model_id+' / '+x.cache.format):'—'}</td><td>${x.endpoint.network} ${x.endpoint.rtt_ms?x.endpoint.rtt_ms+'ms':''}</td><td class="${x.status}">${x.status}</td></tr>`).join('');
Expand Down
14 changes: 13 additions & 1 deletion inference_engine/network/state.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,9 @@
import secrets
import threading
import time
from dataclasses import asdict, is_dataclass
from pathlib import Path
from typing import Any
from typing import Any, Callable

from inference_engine.distributed.capability import CapabilityRegistry
from inference_engine.distributed.prefill_cache import PrefixCacheStore
Expand All @@ -20,10 +21,12 @@ def __init__(
cache_store: PrefixCacheStore,
*,
state_path: str | Path,
prefill_stats_provider: Callable[[], Any] | None = None,
) -> None:
self.registry = registry
self.cache_store = cache_store
self.state_path = Path(state_path).expanduser()
self.prefill_stats_provider = prefill_stats_provider
self._lock = threading.RLock()
self._data = self._load()

Expand Down Expand Up @@ -218,8 +221,17 @@ def summary(self) -> dict[str, Any]:
"local_lookup_hits": cache_stats.lookup_hits,
"local_lookup_misses": cache_stats.lookup_misses,
"local_tokens_served": cache_stats.tokens_served,
"prefill": self.prefill_stats(),
}

def prefill_stats(self) -> dict[str, Any]:
if self.prefill_stats_provider is None:
return {}
stats = self.prefill_stats_provider()
if is_dataclass(stats) and not isinstance(stats, type):
return asdict(stats)
return dict(stats)

def topology(self) -> dict[str, Any]:
nodes = self.nodes()
edges = []
Expand Down
4 changes: 4 additions & 0 deletions scripts/start_grpc_runtime_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -557,6 +557,10 @@ async def _serve(args: argparse.Namespace) -> int:
registry,
prefill_store,
state_path=args.network_state,
prefill_stats_provider=(
(lambda: prefill_hook.stats)
if prefill_hook is not None else None
),
)
http_server = uvicorn.Server(uvicorn.Config(
create_network_app(
Expand Down
101 changes: 101 additions & 0 deletions scripts/verify_remote_prefill_e2e.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,101 @@
#!/usr/bin/env python3
"""Run one cold-prefix request and prove that a remote prefill worker served it."""
from __future__ import annotations

import argparse
import json
import secrets
import time
import urllib.request


def _get_json(url: str):
with urllib.request.urlopen(url, timeout=5) as response:
return json.load(response)


def _delta(before: dict, after: dict, key: str) -> int:
return int(after.get(key, 0)) - int(before.get(key, 0))


def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--address", default="127.0.0.1:51051")
parser.add_argument("--dashboard", default="http://127.0.0.1:8090")
parser.add_argument("--tokenizer-id", required=True)
parser.add_argument("--minimum-prefix-tokens", type=int, default=128)
args = parser.parse_args()

from kakeya import Client
from transformers import AutoTokenizer

from scripts.chat_grpc import _resolve_eos_token_ids

tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_id)
nonce = secrets.token_hex(8)
sentence = (
"Kakeya remote prefill verification context. "
f"Unique run {nonce}. "
)
text = sentence
token_ids = tokenizer.encode(text, add_special_tokens=True)
while len(token_ids) < args.minimum_prefix_tokens:
text += sentence
token_ids = tokenizer.encode(text, add_special_tokens=True)

nodes = _get_json(f"{args.dashboard}/v1/network/nodes")
workers = [
node for node in nodes
if node.get("prefill_worker") and node.get("status") == "online"
]
if not workers:
print(json.dumps({
"ok": False,
"reason": "no online prefill worker capability",
"nodes": nodes,
}, indent=2))
return 2

before = _get_json(f"{args.dashboard}/v1/network/prefill")
started = time.perf_counter()
with Client(args.address) as client:
with client.create_session(
eos_token_ids=_resolve_eos_token_ids(tokenizer),
client_label="remote-prefill-e2e",
) as session:
session.append(token_ids)
list(session.generate(max_tokens=1))
elapsed = time.perf_counter() - started
after = _get_json(f"{args.dashboard}/v1/network/prefill")

result = {
"ok": (
_delta(before, after, "remote_jobs") >= 1
and _delta(before, after, "remote_hits") >= 1
and _delta(before, after, "tokens_reused")
>= args.minimum_prefix_tokens
),
"worker_nodes": [node["id"] for node in workers],
"prefix_tokens": len(token_ids),
"wall_seconds": elapsed,
"delta": {
key: _delta(before, after, key)
for key in (
"remote_jobs",
"remote_hits",
"tokens_reused",
"tokens_computed",
"bytes_received",
"remote_job_failures",
"fallbacks",
)
},
"before": before,
"after": after,
}
print(json.dumps(result, indent=2))
return 0 if result["ok"] else 1


if __name__ == "__main__":
raise SystemExit(main())
42 changes: 42 additions & 0 deletions tests/inference_engine/bridge/test_prefill_worker_launchd.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
from pathlib import Path


ROOT = Path(__file__).resolve().parents[3]
INSTALLER = ROOT / "deploy" / "install_prefill_worker_launchd.sh"
HEAD_PLIST = ROOT / "deploy" / "launchd" / "ai.kakeya.grpc-runtime-prefill.plist"


def test_worker_installer_emits_full_cache_compatibility_contract():
source = INSTALLER.read_text()
for flag in (
"--sink",
"--window",
"--block-size-tokens",
"--prefill-tps",
"--network",
"--priority",
"--rtt-ms",
"--max-concurrent-jobs",
"--max-prompt-tokens",
):
assert f"<string>{flag}</string>" in source
assert 'PEER="${KAKEYA_WORKER_PEER:-}"' in source
assert "<string>--peer</string>" in source


def test_head_runtime_discovers_and_uses_worker_cache_port():
plist = HEAD_PLIST.read_text()
assert (
"<string>--peer</string><string>169.254.27.104:53051</string>"
in plist
)
assert (
"<string>--cache-peer</string><string>169.254.27.104:53051</string>"
in plist
)
assert "<string>--primary-prefill-penalty-ms</string>" in plist
assert (
"<string>--cache-tenant-id</string><string>private-fleet</string>"
in plist
)
assert "<string>--fleet-psk-file</string>" in plist
Loading
Loading