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
136 changes: 136 additions & 0 deletions src/frontend/cua_s1.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,136 @@
"""Small loopback HTTP worker; the Rust frontend remains the public serving layer."""

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

do we need to push server under the model folder?


from __future__ import annotations

import argparse
import json
import logging
import socket
import threading
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer

from models.cua_s1.multimodal.protocol import (
MAX_BODY,
InvalidRequest,
MalformedJSON,
decode_request,
parse_request,
)

LOG = logging.getLogger(__name__)


class WorkerServer(ThreadingHTTPServer):
daemon_threads = True

def __init__(self, address, engine):
self.engine = engine
self.inference_lock = threading.Lock()
super().__init__(address, Handler)


class Handler(BaseHTTPRequestHandler):
def setup(self):
super().setup()
self.connection.settimeout(15)

def log_message(self, format, *args):
# Do not log paths, input images, instructions or arbitrary request headers.
pass

def send_json(self, status, value):
raw = json.dumps(value, ensure_ascii=False, allow_nan=False).encode("utf-8")
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(raw)))
self.end_headers()
try:
self.wfile.write(raw)
except (BrokenPipeError, ConnectionResetError):
pass

def do_GET(self):
if self.path == "/health":
self.send_json(200, {"status": "ready", "modality": "multimodal"})
else:
self.send_json(404, {"detail": "unknown route"})

def do_POST(self):
if self.path != "/v1/systemone":
self.send_json(404, {"detail": "unknown route"})
return
if self.headers.get("Transfer-Encoding"):
self.send_json(
411,
{
"detail": "Content-Length is required; chunked requests are unsupported"
},
)
return
lengths = self.headers.get_all("Content-Length", [])
if len(lengths) != 1:
self.send_json(411, {"detail": "one Content-Length is required"})
return
try:
length = int(lengths[0])
except ValueError:
self.send_json(400, {"detail": "invalid Content-Length"})
return
if length < 0 or length > MAX_BODY:
self.send_json(413, {"detail": "request exceeds body limit"})
return
if self.headers.get_content_type() != "application/json":
self.send_json(415, {"detail": "Content-Type must be application/json"})
return
if not self.server.inference_lock.acquire(blocking=False):
self.send_json(503, {"detail": "worker busy"})
return
try:
raw = self.rfile.read(length)
if len(raw) != length:
self.send_json(400, {"detail": "incomplete body"})
return
parsed = parse_request(decode_request(raw))
result = self.server.engine.predict(parsed)
self.send_json(200, result)
except MalformedJSON as exc:
self.send_json(400, {"detail": str(exc)})
except InvalidRequest as exc:
self.send_json(422, {"detail": str(exc)})
except (TimeoutError, socket.timeout):
self.send_json(408, {"detail": "request body timed out"})
except Exception as exc:
LOG.error("inference failed: %s", type(exc).__name__)
self.send_json(500, {"detail": "inference failed"})
finally:
self.server.inference_lock.release()


def main():
from models.cua_s1.multimodal.model import MultimodalEngine

p = argparse.ArgumentParser(description=__doc__)
p.add_argument(
"--base", required=True, help="verified local base checkpoint directory"
)
p.add_argument(
"--adapter", required=True, help="verified local multimodal adapter directory"
)
p.add_argument("--port", type=int, default=8000)
args = p.parse_args()
logging.basicConfig(level=logging.INFO)
engine = MultimodalEngine(args.base, args.adapter)
engine.warmup()
# Bind only after model loading and a representative inference succeed.
server = WorkerServer(("127.0.0.1", args.port), engine)
LOG.info("multimodal worker ready on 127.0.0.1:%s", args.port)
try:
server.serve_forever()
except KeyboardInterrupt:
pass
finally:
server.server_close()


if __name__ == "__main__":
main()
174 changes: 174 additions & 0 deletions src/models/cua_s1/multimodal/model.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,174 @@
"""Direct Transformers/PEFT execution. No production dependency on cua_s1."""

from __future__ import annotations

import hashlib
import json
from pathlib import Path

from .protocol import InvalidRequest, Question, Request, answer, build_messages

REFERENCE_REVISION = "0e75660ce4c2edda519e0c795fa3ad98abf4e76f"
BASE_REVISION = "851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a"
ADAPTER_REVISION = "16818868b0cc7813808aae4e87b417657046ab79"
IDENTITY = f"cua-ai/cua-s1-4b-0.2@{ADAPTER_REVISION}:multimodal"
WEIGHTS_MANIFEST_SHA256 = (
"9820bd232c5762f114e19680c0f8203d7e1faaf8a60c196cfe01964d6d8a6c09"
)
MAX_TOKENS = 4096


def letter_ids(tokenizer, count: int) -> list[int]:
ids = []
for index in range(count):
encoded = tokenizer.encode(chr(65 + index), add_special_tokens=False)
if len(encoded) != 1:
raise ValueError("each candidate letter must be a single token")
ids.append(encoded[0])
return ids


def validate_adapter_config(config: dict):
targets = {
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
"linear_fc1",
"linear_fc2",
}
if (
config.get("peft_type") != "LORA"
or config.get("r") != 16
or config.get("lora_alpha") != 32
or set(config.get("target_modules", [])) != targets
or config.get("base_model_name_or_path") != "Qwen/Qwen3.5-4B"
):
raise ValueError("expected the pinned 0.2 multimodal LoRA adapter")


def parse_weights_manifest(raw: bytes) -> dict:
"""Accept only the manifest from the pinned upstream reference commit."""
if hashlib.sha256(raw).hexdigest() != WEIGHTS_MANIFEST_SHA256:
raise ValueError("upstream weights manifest checksum mismatch")
return json.loads(raw)


def verify_weights(base: Path, adapter: Path):
"""Check local artifacts before assigning the pinned identity to responses."""
lock = parse_weights_manifest((base.parent / "weights.lock.json").read_bytes())
allowed = {base: set(), adapter: set()}
for artifact in lock["artifacts"]:
for name, expected in artifact["files"].items():
if artifact["role"] == "adapter":
if not name.startswith("multimodal/"):
continue
path = adapter / name.removeprefix("multimodal/")
else:
path = base / name
root = adapter if artifact["role"] == "adapter" else base
allowed[root].add(path.relative_to(root).as_posix())
if not path.is_file() or path.stat().st_size != expected["size"]:
raise ValueError(f"missing or wrong-size pinned artifact: {path.name}")
with path.open("rb") as handle:
digest = hashlib.file_digest(handle, "sha256").hexdigest()
if digest != expected["sha256"]:
raise ValueError(f"checksum mismatch: {path.name}")
for root, names in allowed.items():
for path in root.rglob("*"):
relative = path.relative_to(root)
if path.is_file() and relative.parts[0] != ".cache":
if relative.as_posix() not in names:
raise ValueError(
f"unlisted artifact may override pinned files: {relative}"
)


class MultimodalEngine:
def __init__(
self, base: str, adapter: str, device: str = "cuda", dtype: str = "bfloat16"
):
import torch
from peft import PeftModel
from peft.tuners.lora import LoraLayer
from transformers import (
AutoModelForImageTextToText,
AutoProcessor,
AutoTokenizer,
)

base_path, adapter_path = Path(base), Path(adapter)
verify_weights(base_path, adapter_path)
validate_adapter_config(
json.loads((adapter_path / "adapter_config.json").read_text())
)
self.tokenizer = AutoTokenizer.from_pretrained(base, local_files_only=True)
self.processor = AutoProcessor.from_pretrained(base, local_files_only=True)
model = AutoModelForImageTextToText.from_pretrained(
base,
torch_dtype=getattr(torch, dtype),
device_map=device,
local_files_only=True,
)
self.model = PeftModel.from_pretrained(model, adapter, local_files_only=True)
modules = [
name
for name, module in self.model.named_modules()
if isinstance(module, LoraLayer)
]
if len(modules) != 178 or not any(".visual." in name for name in modules):
raise RuntimeError(
"multimodal adapter did not attach to all 178 expected modules"
)
self.adapter_modules = len(modules)
self.model.eval()
self.dtype = dtype

def prepare(self, image, question: Question):
messages = build_messages(question)
text = self.processor.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
inputs = self.processor(text=[text], images=[image], return_tensors="pt")
if inputs["input_ids"].shape[-1] > MAX_TOKENS:
raise InvalidRequest(f"processed prompt exceeds {MAX_TOKENS} tokens")
return inputs

def score(self, inputs, question: Question) -> list[float]:
import torch

ids = letter_ids(self.tokenizer, len(question.keys))
inputs = inputs.to(self.model.device)
with torch.no_grad():
output = self.model(**inputs)
logits = output.logits[0, -1, :]
return torch.softmax(
logits[torch.tensor(ids, device=logits.device)].float(), dim=-1
).tolist()

def predict(self, request: Request) -> dict:
# Validate all processed lengths before executing any question.
prepared = [self.prepare(request.image, q) for q in request.questions]
answers = {
q.name: answer(q, self.score(inputs, q))
for q, inputs in zip(request.questions, prepared)
}
return {
"model": IDENTITY,
"answers": answers,
"usage": {
"input_tokens": sum(x["input_ids"].shape[-1] for x in prepared),
"output_tokens": 0,
},
}

def warmup(self):
from PIL import Image

q = Question(
"warmup", ("continue", "cancel"), ("Continue", "Cancel"), "Continue"
)
self.predict(Request(Image.new("RGB", (224, 224), "white"), (q,)))
Loading