Skip to content
Closed
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
36 changes: 34 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -89,8 +89,40 @@ the right experts get. It works because routing has measurable structure (see
the [expert atlas](https://github.com/JustVugg/colibri/issues/175)) — and
structure is cacheable.

The engine is a single C file (`c/glm.c`) plus small headers. No BLAS, no Python
at runtime, no GPU required.
The engine is a single C file (`c/colibri.c`) plus small headers. No BLAS, no
Python at runtime, no GPU required.

### Local cluster mode

The coordinator keeps token generation, routing, and KV state local while
disk-backed expert workers execute routed FFNs on other Macs. A layer's routed
batch-union is sent as one persistent TCP request, so a token does not incur one
round trip per expert.

Start the optional registration service:

```bash
./coli cluster coordinator --host 0.0.0.0 --port 8765
```

On each worker, with the same converted model available locally:

```bash
./coli cluster worker --model /nvme/glm52_i4 --port 9100 \
--coordinator http://COORDINATOR:8765 --advertise-host WORKER_IP
```

Run the coordinator with discovery, or provide `--cluster-workers
HOST:PORT,...` for a static setup:

```bash
./coli serve --model /nvme/glm52_i4 \
--cluster-coordinator http://127.0.0.1:8765
```

The transport is disabled unless workers are configured, so the existing
single-machine path remains unchanged. Dense-layer sharding and browser/WebGPU
workers are separate follow-up seams.

## How it works

Expand Down
148 changes: 148 additions & 0 deletions c/cluster.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,148 @@
#!/usr/bin/env python3
"""Registration and discovery control plane for local expert workers."""

import argparse
import json
import threading
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from urllib.request import Request, urlopen


PROTOCOL_VERSION = 1


class ClusterRegistry:
def __init__(self, stale_after=30.0):
self.stale_after = float(stale_after)
self._nodes = {}
self._lock = threading.Lock()

def register(self, node):
required = {"node_id", "host", "port", "role"}
missing = sorted(required - set(node))
if missing:
raise ValueError("missing node fields: " + ", ".join(missing))
port = int(node["port"])
if not 1 <= port <= 65535:
raise ValueError("port must be between 1 and 65535")
role = str(node["role"])
if role not in ("expert", "dense", "coordinator"):
raise ValueError("role must be expert, dense, or coordinator")
record = dict(node)
record.update(protocol_version=PROTOCOL_VERSION, port=port, last_seen=time.time())
with self._lock:
self._nodes[str(node["node_id"])] = record
return record

def heartbeat(self, node_id):
with self._lock:
node = self._nodes.get(str(node_id))
if node is None:
raise KeyError(node_id)
node["last_seen"] = time.time()
return dict(node)

def snapshot(self):
now = time.time()
with self._lock:
nodes = [dict(node) for node in self._nodes.values()
if now - node["last_seen"] <= self.stale_after]
nodes.sort(key=lambda node: (node["role"], node["node_id"]))
return {"protocol_version": PROTOCOL_VERSION, "nodes": nodes}

def expert_endpoints(self):
return [f"{node['host']}:{node['port']}"
for node in self.snapshot()["nodes"] if node["role"] == "expert"]


class _Handler(BaseHTTPRequestHandler):
server_version = "colibri-cluster/1"

def log_message(self, *_args):
return

def _json(self, status, value):
payload = json.dumps(value, separators=(",", ":")).encode()
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)

def _body(self):
length = int(self.headers.get("Content-Length", "0"))
if length > 1 << 20:
raise ValueError("request body is too large")
return json.loads(self.rfile.read(length) or b"{}")

def do_GET(self): # noqa: N802 - stdlib handler API
if self.path in ("/health", "/v1/cluster/topology"):
self._json(200, self.server.registry.snapshot())
return
self._json(404, {"error": "not found"})

def do_POST(self): # noqa: N802 - stdlib handler API
try:
body = self._body()
if self.path == "/v1/cluster/register":
self._json(200, self.server.registry.register(body))
elif self.path == "/v1/cluster/heartbeat":
self._json(200, self.server.registry.heartbeat(body["node_id"]))
else:
self._json(404, {"error": "not found"})
except (KeyError, ValueError, TypeError, json.JSONDecodeError) as error:
self._json(400, {"error": str(error)})


class ClusterServer(ThreadingHTTPServer):
daemon_threads = True

def __init__(self, address, registry):
super().__init__(address, _Handler)
self.registry = registry


def serve(host="127.0.0.1", port=8765, stale_after=30.0):
server = ClusterServer((host, port), ClusterRegistry(stale_after))
print(f"colibri cluster coordinator listening on http://{host}:{port}", flush=True)
try:
server.serve_forever()
finally:
server.server_close()


def register(coordinator, node):
request = Request(coordinator.rstrip("/") + "/v1/cluster/register",
data=json.dumps(node).encode(),
headers={"Content-Type": "application/json"}, method="POST")
with urlopen(request, timeout=5) as response:
return json.load(response)


def heartbeat(coordinator, node_id):
request = Request(coordinator.rstrip("/") + "/v1/cluster/heartbeat",
data=json.dumps({"node_id": node_id}).encode(),
headers={"Content-Type": "application/json"}, method="POST")
with urlopen(request, timeout=5) as response:
return json.load(response)


def discover_workers(coordinator):
with urlopen(coordinator.rstrip("/") + "/v1/cluster/topology", timeout=5) as response:
topology = json.load(response)
return [f"{node['host']}:{int(node['port'])}"
for node in topology.get("nodes", []) if node.get("role") == "expert"]


def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, default=8765)
parser.add_argument("--stale-after", type=float, default=30.0)
args = parser.parse_args()
serve(args.host, args.port, args.stale_after)


if __name__ == "__main__":
main()
69 changes: 68 additions & 1 deletion c/coli
Original file line number Diff line number Diff line change
Expand Up @@ -149,6 +149,14 @@ def need_model(model):
if not os.path.exists(GLM):
sys.exit(f"{C.yel}engine is not built.{C.r} Run: coli build")

def need_worker_model(model):
if not os.path.isdir(model):
sys.exit(f"{C.yel}model not found:{C.r} {model}\n set COLI_MODEL or use --model")
if not os.path.exists(os.path.join(model,"config.json")):
sys.exit(f"{C.yel}config.json is missing from:{C.r} {model}")
if not os.path.exists(GLM):
sys.exit(f"{C.yel}engine is not built.{C.r} Run: coli build")

def cuda_binary():
if not os.path.exists(GLM): return False
if sys.platform == "linux":
Expand Down Expand Up @@ -211,6 +219,8 @@ def env_for(a):
e.setdefault("PIPE", "1")
e.setdefault("PILOT_REAL", "1")
e["COLI_POLICY"]=a.policy
if getattr(a, "cluster_workers", None): e["CLUSTER_WORKERS"] = a.cluster_workers
if getattr(a, "cluster_coordinator", None): e["CLUSTER_COORDINATOR"] = a.cluster_coordinator
if a.ram: e["RAM_GB"]=str(a.ram)
if a.ngen: e["NGEN"]=str(a.ngen)
if a.topp: e["TOPP"]=str(a.topp)
Expand Down Expand Up @@ -828,14 +838,53 @@ def cmd_serve(a):
with open(serve_pidfile(a.port),"w") as f: f.write(f"{os.getpid()} {a.model}\n")
except OSError: pass
from openai_server import serve
e=env_for(a)
if a.cluster_coordinator and not a.cluster_workers:
from cluster import discover_workers
try:
workers=discover_workers(a.cluster_coordinator)
except OSError as error:
sys.exit(f"{C.yel}cannot discover cluster workers:{C.r} {error}")
if not workers:
sys.exit(f"{C.yel}cluster coordinator has no live expert workers{C.r}")
e["CLUSTER_WORKERS"] = ",".join(workers)
print(f" {C.dim}[CLUSTER] discovered {len(workers)} expert worker(s){C.r}", file=sys.stderr)
try:
serve(a.model, a.host, a.port, a.model_id, a.api_key,
a.cap,a.ngen,GLM,env_for(a),a.cors_origin,
a.cap,a.ngen,GLM,e,a.cors_origin,
a.max_queue,a.queue_timeout,a.kv_slots)
finally:
try: os.unlink(serve_pidfile(a.port))
except OSError: pass

def cmd_cluster_coordinator(a):
from cluster import serve
serve(a.host, a.port, a.stale_after)

def cmd_cluster_worker(a):
need_worker_model(a.model)
coordinator=a.coordinator or os.environ.get("CLUSTER_COORDINATOR")
host=a.advertise_host or os.environ.get("CLUSTER_ADVERTISE_HOST", "127.0.0.1")
node_id=a.node_id or f"{host}:{a.port}"
stop_heartbeat=threading.Event()
if coordinator:
from cluster import heartbeat, register
register(coordinator, {"node_id":node_id, "host":host, "port":a.port,
"role":"expert", "layers":a.layers})
def keep_registered():
while not stop_heartbeat.wait(10):
try: heartbeat(coordinator, node_id)
except OSError: pass
threading.Thread(target=keep_registered, name="colibri-cluster-heartbeat", daemon=True).start()
e=env_for(a)
e.update({"EXPERT_WORKER":"1", "CLUSTER_WORKER_PORT":str(a.port),
"COLI_MMAP":os.environ.get("COLI_MMAP", "1")})
print(f" {C.dim}[CLUSTER] expert worker · {host}:{a.port} · layers {a.layers}{C.r}")
try:
return subprocess.call([GLM,str(a.cap),str(a.ebits),str(a.dbits)],env=e)
finally:
stop_heartbeat.set()

def cmd_stop(a):
"""Shut down a running `coli serve` AND its engine — one command, no pkill.
The engine re-execs itself for OMP tuning, so its process is named `exe`,
Expand Down Expand Up @@ -963,6 +1012,10 @@ def main():
common.add_argument("--cap", type=int, default=8); common.add_argument("--ngen", type=int, default=1024) # rete di sicurezza: la fine vera la decidono gli stop token
common.add_argument("--topp", type=float, default=0); common.add_argument("--topk", type=int, default=0)
common.add_argument("--temp", type=float, default=None) # temperatura token (0=greedy, default 1.0+nucleus .95)
common.add_argument("--cluster-workers", default=os.environ.get("CLUSTER_WORKERS"),
help="comma-separated expert workers, host:port,...")
common.add_argument("--cluster-coordinator", default=os.environ.get("CLUSTER_COORDINATOR"),
help="control-plane URL used to discover expert workers")
ap=argparse.ArgumentParser(prog="coli", parents=[common], description="colibri — run GLM-5.2 locally")
ap.add_argument("--version", action="version", version=f"colibri {_version}")
sub=ap.add_subparsers(dest="cmd")
Expand All @@ -988,6 +1041,17 @@ def main():
ps.add_argument("--max-queue",type=int,default=int(os.environ.get("COLI_MAX_QUEUE","8")))
ps.add_argument("--queue-timeout",type=float,default=float(os.environ.get("COLI_QUEUE_TIMEOUT","300")))
ps.add_argument("--kv-slots",type=int,default=int(os.environ.get("COLI_KV_SLOTS","1")))
pcluster=sub.add_parser("cluster", help="run the local-cluster control plane or an expert worker")
cluster_sub=pcluster.add_subparsers(dest="cluster_cmd")
pcoord=cluster_sub.add_parser("coordinator", help="serve worker registration and discovery")
pcoord.add_argument("--host",default="127.0.0.1"); pcoord.add_argument("--port",type=int,default=8765)
pcoord.add_argument("--stale-after",type=float,default=30.0)
pworker=cluster_sub.add_parser("worker", parents=[common], help="serve disk-backed expert compute")
pworker.add_argument("--port",type=int,default=int(os.environ.get("CLUSTER_WORKER_PORT","9100")))
pworker.add_argument("--coordinator",default=os.environ.get("CLUSTER_COORDINATOR"))
pworker.add_argument("--node-id",default=None); pworker.add_argument("--advertise-host",default=None)
pworker.add_argument("--layers",default="all",help="topology label, e.g. 0-37")
pworker.add_argument("--ebits",type=int,default=8); pworker.add_argument("--dbits",type=int,default=8)
pst=sub.add_parser("stop", parents=[common], help="shut down a running coli serve and its engine")
pst.add_argument("--port",type=int,default=8000); pst.add_argument("--dry-run",action="store_true")
pw=sub.add_parser("web", parents=[common], help="serve + open the dashboard in a browser")
Expand Down Expand Up @@ -1015,6 +1079,9 @@ def main():
handler={"build":cmd_build,"info":cmd_info,"plan":cmd_plan,"doctor":cmd_doctor,
"run":cmd_run,"chat":cmd_chat,"serve":cmd_serve,"stop":cmd_stop,"bench":cmd_bench,
"convert":cmd_convert,"web":cmd_web}.get(a.cmd)
if a.cmd=="cluster":
if a.cluster_cmd=="coordinator": handler=cmd_cluster_coordinator
elif a.cluster_cmd=="worker": handler=cmd_cluster_worker
if handler: sys.exit(handler(a) or 0)
banner(); print(__doc__)

Expand Down
Loading
Loading