diff --git a/src/frontend/cua_s1.py b/src/frontend/cua_s1.py new file mode 100644 index 0000000..ab96ad4 --- /dev/null +++ b/src/frontend/cua_s1.py @@ -0,0 +1,157 @@ +"""Small loopback HTTP worker; the Rust frontend remains the public serving layer.""" + +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.graph_runtime import GraphConfig + 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) + p.add_argument( + "--graph", action="store_true", help="enable segmented CUDA Graph replay" + ) + p.add_argument("--graph-max-shapes", type=int, default=8) + p.add_argument("--graph-max-memory-mib", type=int, default=1024) + p.add_argument("--graph-min-uses", type=int, default=2) + p.add_argument("--graph-max-tokens", type=int, default=2048) + args = p.parse_args() + try: + graph_config = ( + GraphConfig( + max_shapes=args.graph_max_shapes, + max_bytes=args.graph_max_memory_mib * 1024 * 1024, + min_uses=args.graph_min_uses, + max_tokens=args.graph_max_tokens, + ) + if args.graph + else None + ) + except ValueError as exc: + p.error(str(exc)) + logging.basicConfig(level=logging.INFO) + engine = MultimodalEngine(args.base, args.adapter, graph_config=graph_config) + 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() diff --git a/src/models/cua_s1/multimodal/graph_runtime.py b/src/models/cua_s1/multimodal/graph_runtime.py new file mode 100644 index 0000000..13a122a --- /dev/null +++ b/src/models/cua_s1/multimodal/graph_runtime.py @@ -0,0 +1,483 @@ +"""Bounded CUDA Graph replay for Qwen3.5 multimodal linear-attention runs. + +Full-attention SDPA changes BF16 results under whole-model Graph capture on the +pinned RTX 4090 stack. Keep those layers eager and capture the intervening +Gated DeltaNet runs, whose replay matches the original forward bit for bit. +""" + +from __future__ import annotations + +import ctypes +import gc +import threading +import time +from collections import OrderedDict +from contextlib import suppress +from dataclasses import dataclass, field + + +def tensor_signature(value): + if value is None: + return None + return ( + tuple(value.shape), + tuple(value.stride()), + str(value.dtype), + str(value.device), + ) + + +@dataclass(frozen=True) +class GraphConfig: + max_shapes: int = 8 + max_bytes: int = 1024 * 1024 * 1024 + min_uses: int = 2 + max_tokens: int = 2048 + + def __post_init__(self): + if min(self.max_shapes, self.max_bytes, self.min_uses, self.max_tokens) <= 0: + raise ValueError("graph limits must be positive") + + +class GraphCache: + """LRU cache whose entries own all segments for one tensor layout.""" + + def __init__(self, max_shapes: int, max_bytes: int): + self.max_shapes = max_shapes + self.max_bytes = max_bytes + self.entries = OrderedDict() + self.bytes = 0 + + def __len__(self): + return len(self.entries) + + def get(self, key): + value = self.entries.get(key) + if value is not None: + self.entries.move_to_end(key) + return value + + def put(self, key, value): + if value.bytes > self.max_bytes: + return None + removed = [] + if key in self.entries: + previous = self.entries.pop(key) + self.bytes -= previous.bytes + removed.append(previous) + self.entries[key] = value + self.bytes += value.bytes + while len(self.entries) > self.max_shapes or self.bytes > self.max_bytes: + _, previous = self.entries.popitem(last=False) + self.bytes -= previous.bytes + removed.append(previous) + return removed + + def clear(self): + removed = list(self.entries.values()) + self.entries.clear() + self.bytes = 0 + return removed + + +class _GraphPool: + """One shape's private allocator pool and exclusively owned capture stream. + + Segments share scratch space only within a shape, in capture/replay order. + All graphs must be reset before this stream is destroyed. Other shapes use + separate streams because the pinned PyTorch clears cuBLAS workspaces by + capture stream when destroying a graph. + """ + + def __init__(self, device): + import torch + + self._torch = torch + self.device = device + self._stream_ptr = None + if torch.cuda.get_allocator_backend() != "native": + raise RuntimeError( + "graph memory accounting requires the native CUDA allocator" + ) + self.pool_id = torch.cuda.graph_pool_handle() + with torch.cuda.device(device): + pointer = ctypes.c_void_p() + torch.cuda.check_error( + torch.cuda.cudart().cudaStreamCreate(ctypes.addressof(pointer)) + ) + self._stream_ptr = pointer.value + self.capture_stream = torch.cuda.ExternalStream( + self._stream_ptr, device=device + ) + + def reserved_bytes(self): + # Allocated tensor deltas omit inactive graph scratch space, which stays + # reserved for replay. Query this pool directly, not process-wide deltas. + return sum( + segment["total_size"] + for segment in self._torch.cuda.memory_snapshot( + mempool_id=self.pool_id, include_traces=False + ) + ) + + def close(self): + if getattr(self, "_stream_ptr", None) is None: + return + torch = self._torch + with torch.cuda.device(self.device): + torch.cuda.check_error( + torch.cuda.cudart().cudaStreamDestroy(self._stream_ptr) + ) + self._stream_ptr = None + self.capture_stream = None + + def __del__(self): + try: + self.close() + except Exception: # noqa: BLE001, S110 -- best-effort destructor cleanup + pass + + +class _GraphSegment: + """One fixed-layout capture of adjacent linear-attention decoder layers.""" + + def __init__(self, layers, start, end, hidden, mask, pool=None): + import torch + + self._torch = torch + self._device = hidden.device + self._pool = None + self._owns_pool = pool is None + self.graph = None + self.capture_stream = None + self.layers = layers + self.start = start + self.end = end + self.input_signature = tensor_signature(hidden) + self.mask_signature = tensor_signature(mask) + self.static_hidden = hidden.clone() + self.static_mask = mask.clone() if mask is not None else None + self.external_bytes = self.static_hidden.untyped_storage().nbytes() + if self.static_mask is not None: + self.external_bytes += self.static_mask.untyped_storage().nbytes() + + stream = torch.cuda.Stream(device=hidden.device) + stream.wait_stream(torch.cuda.current_stream(hidden.device)) + with torch.cuda.stream(stream), torch.no_grad(): + for _ in range(3): + self._forward() + torch.cuda.current_stream(hidden.device).wait_stream(stream) + + try: + with torch.cuda.device(self._device): + self._pool = pool if pool is not None else _GraphPool(self._device) + self.capture_stream = self._pool.capture_stream + self.graph = torch.cuda.CUDAGraph() + with ( + torch.cuda.graph( + self.graph, + stream=self.capture_stream, + pool=self._pool.pool_id, + ), + torch.no_grad(), + ): + self.static_output = self._forward() + except BaseException: + # Preserve the actual capture failure if a poisoned CUDA context + # also prevents cleanup. The destructor makes a best-effort retry. + with suppress(Exception): + self.close() + raise + + def close(self): + """Retire a graph; shared-pool segments must all retire together.""" + if getattr(self, "_pool", None) is None: + return + torch = self._torch + with torch.cuda.device(self._device): + # Replays run on the caller's stream, not the capture stream. + torch.cuda.synchronize(self._device) + if self.graph is not None: + self.graph.reset() + self.graph = None + self.static_output = None + self.static_hidden = None + self.static_mask = None + if self._owns_pool: + self._pool.close() + self._pool = None + self.capture_stream = None + + def __del__(self): + # Interpreter teardown and failed CUDA contexts cannot safely raise. + try: + self.close() + except Exception: # noqa: BLE001, S110 -- best-effort destructor cleanup + pass + + def _forward(self): + hidden = self.static_hidden + for index in range(self.start, self.end): + hidden = self.layers[index]( + hidden, + position_embeddings=None, + attention_mask=self.static_mask, + position_ids=None, + past_key_values=None, + use_cache=False, + ) + return hidden + + def replay(self, hidden, mask): + if ( + tensor_signature(hidden) != self.input_signature + or tensor_signature(mask) != self.mask_signature + ): + raise ValueError("CUDA Graph segment layout changed") + self.static_hidden.copy_(hidden) + if self.static_mask is not None: + self.static_mask.copy_(mask) + self.graph.replay() + return self.static_output + + +@dataclass +class _ShapeEntry: + blocks: dict = field(default_factory=dict) + bytes: int = 0 + pool: _GraphPool | None = None + + def update_bytes(self): + self.bytes = self.pool.reserved_bytes() + sum( + block.external_bytes for block in self.blocks.values() + ) + + def close(self): + # No surviving segment may replay after the first reset clears the + # shared stream's workspace. The runtime retires an entire shape at once. + for block in self.blocks.values(): + block.close() + self.blocks.clear() + if self.pool is not None: + self.pool.close() + self.pool = None + + def __del__(self): + try: + self.close() + except Exception: # noqa: BLE001, S110 -- best-effort destructor cleanup + pass + + +class GraphRuntime: + """Exact-length graph cache for a loaded, immutable multimodal model.""" + + def __init__(self, model, config: GraphConfig | None = None): + self.model = model + self.config = config or GraphConfig() + self.cache = GraphCache(self.config.max_shapes, self.config.max_bytes) + self.uses = OrderedDict() + self.disabled = OrderedDict() + self.lock = threading.Lock() + self.stats = { + "eager": 0, + "unsupported": 0, + "warmup": 0, + "disabled": 0, + "replays": 0, + "captures": 0, + "capture_ms": 0.0, + "evictions": 0, + "rejected": 0, + "memory_budget": 0, + "capture_oom": 0, + "capture_error": 0, + "numerical_mismatch": 0, + } + + def invalidate(self): + """Call after replacing model weights or adapters.""" + with self.lock: + retired = self.cache.clear() + self.uses.clear() + self.disabled.clear() + for entry in retired: + entry.close() + del retired + gc.collect() + + def _eager(self, values): + self.stats["eager"] += 1 + return self.model(**values, logits_to_keep=1, use_cache=False).logits[0, -1, :] + + def _supported(self, values): + hidden = values["inputs_embeds"] + if ( + self.model.training + or hidden.device.type != "cuda" + or hidden.ndim != 3 + or hidden.shape[0] != 1 + or hidden.shape[1] > self.config.max_tokens + or values["position_ids"].shape != (3, 1, hidden.shape[1]) + ): + return False + core = self.model.get_base_model().model + text = core.language_model + types = text.config.layer_types + return len(types) == text.config.num_hidden_layers == len(text.layers) and set( + types + ) == {"linear_attention", "full_attention"} + + def _run_segments(self, values, entry): + from transformers.masking_utils import ( + create_causal_mask, + create_recurrent_attention_mask, + ) + + core = self.model.get_base_model().model + text = core.language_model + hidden = values["inputs_embeds"] + kwargs = { + "config": text.config, + "inputs_embeds": hidden, + "attention_mask": values["attention_mask"], + "past_key_values": None, + "position_ids": None, + } + masks = { + "full_attention": create_causal_mask(**kwargs), + "linear_attention": create_recurrent_attention_mask(**kwargs), + } + rope = text.rotary_emb(hidden, values["position_ids"]) + index = 0 + while index < text.config.num_hidden_layers: + kind = text.config.layer_types[index] + if kind == "full_attention": + hidden = text.layers[index]( + hidden, + position_embeddings=rope, + attention_mask=masks[kind], + position_ids=None, + past_key_values=None, + use_cache=False, + ) + index += 1 + continue + end = index + 1 + while ( + end < text.config.num_hidden_layers + and text.config.layer_types[end] == "linear_attention" + ): + end += 1 + block = entry.blocks.get((index, end)) + if block is None: + if entry.pool is None: + entry.pool = _GraphPool(hidden.device) + block = _GraphSegment( + text.layers, + index, + end, + hidden, + masks["linear_attention"], + pool=entry.pool, + ) + entry.blocks[(index, end)] = block + entry.update_bytes() + hidden = block.replay(hidden, masks["linear_attention"]) + index = end + hidden = text.norm(hidden) + return self.model.get_base_model().lm_head(hidden[:, -1:, :])[0, -1, :] + + def forward(self, values): + """Return fresh logits; no static graph output escapes the runtime lock.""" + import torch + + with self.lock, torch.no_grad(): + if not self._supported(values): + self.stats["unsupported"] += 1 + return self._eager(values) + key = ( + id(self.model), + self.model.get_base_model().model.language_model.config._attn_implementation, + *( + tensor_signature(values[name]) + for name in ("inputs_embeds", "position_ids", "attention_mask") + ), + ) + if key in self.disabled: + self.stats["disabled"] += 1 + return self._eager(values) + entry = self.cache.get(key) + if entry is not None: + output = self._run_segments(values, entry) + self.stats["replays"] += 1 + return output + + uses = self.uses.get(key, 0) + 1 + self.uses[key] = uses + self.uses.move_to_end(key) + while len(self.uses) > 128: + self.uses.popitem(last=False) + if uses < self.config.min_uses: + self.stats["warmup"] += 1 + return self._eager(values) + + reference = self._eager(values) + candidate = _ShapeEntry() + try: + torch.cuda.synchronize(values["inputs_embeds"].device) + capture_start = time.perf_counter() + output = self._run_segments(values, candidate) + torch.cuda.synchronize(values["inputs_embeds"].device) + capture_ms = (time.perf_counter() - capture_start) * 1000 + except torch.cuda.OutOfMemoryError: + self.stats["capture_oom"] += 1 + candidate.close() + del candidate + gc.collect() + torch.cuda.empty_cache() + self._disable(key) + return reference + except RuntimeError: + # Only a healthy CUDA context can continue serving eagerly. + torch.cuda.synchronize(values["inputs_embeds"].device) + self.stats["capture_error"] += 1 + candidate.close() + del candidate + self._disable(key) + gc.collect() + return reference + + if not torch.equal(reference, output): + self.stats["numerical_mismatch"] += 1 + candidate.close() + del candidate, output + self._disable(key) + gc.collect() + return reference + retired = self.cache.put(key, candidate) + if retired is None: + self.stats["memory_budget"] += 1 + candidate.close() + del candidate + self._disable(key) + gc.collect() + return reference + self.stats["captures"] += 1 + self.stats["capture_ms"] += capture_ms + self.stats["evictions"] += len(retired) + if retired: + for entry in retired: + entry.close() + del entry + del retired + gc.collect() + else: + del retired + return output + + def _disable(self, key): + self.stats["rejected"] += 1 + self.disabled[key] = None + while len(self.disabled) > 128: + self.disabled.popitem(last=False) diff --git a/src/models/cua_s1/multimodal/model.py b/src/models/cua_s1/multimodal/model.py new file mode 100644 index 0000000..ceb6405 --- /dev/null +++ b/src/models/cua_s1/multimodal/model.py @@ -0,0 +1,299 @@ +"""Direct Transformers/PEFT execution. No production dependency on cua_s1.""" + +from __future__ import annotations + +import copy +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" + and relative.as_posix() not in names + ): + raise ValueError( + f"unlisted artifact may override pinned files: {relative}" + ) + + +class _RequestImageProcessor: + """Reuse one image result; each processor call owns its mutable mapping.""" + + def __init__(self, processor): + self.processor = processor + self.result = None + + def __getattr__(self, name): + return getattr(self.processor, name) + + def __call__(self, *args, **kwargs): + if self.result is None: + self.result = self.processor(*args, **kwargs) + # BatchFeature.to and processor token expansion can mutate mappings. + # Share only tensor values, which the pinned processor does not modify. + return type(self.result)(dict(self.result)) + + +class MultimodalEngine: + def __init__( + self, + base: str, + adapter: str, + device: str = "cuda", + dtype: str = "bfloat16", + graph_config=None, + ): + 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 + self.graph_runtime = None + if graph_config is not None: + from .graph_runtime import GraphRuntime + + self.graph_runtime = GraphRuntime(self.model, graph_config) + + def prepare(self, image, question: Question): + return self._prepare(self.processor, image, question) + + @staticmethod + def _prepare(processor, image, question: Question): + messages = build_messages(question) + text = processor.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) + inputs = 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 prepare_reused(self, image, questions): + """Tokenize each question normally, preprocessing this request's image once.""" + processor = copy.copy(self.processor) + processor.image_processor = _RequestImageProcessor( + self.processor.image_processor + ) + return [self._prepare(processor, image, question) for question in questions] + + def encode_image(self, inputs): + """Encode once through the vision modules with their active LoRA adapters.""" + import torch + + core = self.model.get_base_model().model + with torch.no_grad(): + output = core.get_image_features( + inputs["pixel_values"].to(self.model.device), + inputs["image_grid_thw"].to(self.model.device), + return_dict=True, + ) + return torch.cat(output.pooler_output, dim=0) + + def score_reused(self, inputs, question: Question, features) -> list[float]: + """Build fresh text embeddings and 3D positions around shared image features.""" + import torch + + ids = letter_ids(self.tokenizer, len(question.keys)) + # Keep prepared CPU inputs intact and avoid another image transfer. + text_inputs = { + key: value.to(self.model.device) + for key, value in inputs.items() + if key != "pixel_values" + } + core = self.model.get_base_model().model + with torch.no_grad(): + input_ids = text_inputs.pop("input_ids") + embeds = core.get_input_embeddings()(input_ids) + image_embeds = features.to(embeds.device, embeds.dtype) + image_mask, _ = core.get_placeholder_mask( + input_ids, inputs_embeds=embeds, image_features=image_embeds + ) + embeds = embeds.masked_scatter(image_mask, image_embeds) + position_ids, _ = core.get_rope_index( + input_ids=input_ids, + mm_token_type_ids=text_inputs.pop("mm_token_type_ids"), + image_grid_thw=text_inputs.pop("image_grid_thw"), + attention_mask=text_inputs.get("attention_mask"), + ) + # Candidate scoring reads only the final position. + graph_runtime = getattr(self, "graph_runtime", None) + if graph_runtime is None: + output = self.model( + inputs_embeds=embeds, + position_ids=position_ids, + logits_to_keep=1, + **text_inputs, + ) + logits = output.logits[0, -1, :] + else: + logits = graph_runtime.forward( + { + "inputs_embeds": embeds, + "position_ids": position_ids, + **text_inputs, + } + ) + return torch.softmax( + logits[torch.tensor(ids, device=logits.device)].float(), dim=-1 + ).tolist() + + 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_reference(self, request: Request) -> dict: + """Original full-processor/full-model path retained for paired experiments.""" + # 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 predict(self, request: Request) -> dict: + if len(request.questions) == 1: + return self.predict_reference(request) + # No image encoder or language model runs until every prompt is valid. + prepared = self.prepare_reused(request.image, request.questions) + features = self.encode_image(prepared[0]) + answers = { + q.name: answer(q, self.score_reused(inputs, q, features)) + 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,))) diff --git a/src/models/cua_s1/multimodal/protocol.py b/src/models/cua_s1/multimodal/protocol.py new file mode 100644 index 0000000..4346257 --- /dev/null +++ b/src/models/cua_s1/multimodal/protocol.py @@ -0,0 +1,261 @@ +"""Bounded screenshot-only extension of the Cua-S1 choice contract in PR #11.""" + +from __future__ import annotations + +import base64 +import binascii +import io +import json +import math +from dataclasses import dataclass + +from PIL import Image, UnidentifiedImageError + +MODEL = "cua-s1-4b-0.2" +MAX_BODY = 8 * 1024 * 1024 +MAX_IMAGE_BYTES = 4 * 1024 * 1024 +MAX_PIXELS = 1024 * 1024 +MAX_SIDE = 2048 +MAX_ASPECT_RATIO = 200 # Pinned Transformers Qwen image processor's smart_resize. +MAX_QUESTIONS = 8 +MAX_TEXT = 16384 + +# Prompt and letter layout follow trycua/cua at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f. +# MIT License +# +# Copyright (c) 2025 Cua AI, Inc. +# +# 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. +SYSTEM_PROMPT = ( + "You are a one-pass computer-use decision model. You are shown the " + "current state of a screen and a fixed, closed list of candidate " + "(element, action) options, each given a single letter. Choose exactly " + "one option: the single best next action to take. Answer with ONLY that " + "option's letter -- no words, no punctuation, no explanation." +) + + +class InvalidRequest(ValueError): + """Input cannot be evaluated under the supported contract.""" + + +class MalformedJSON(InvalidRequest): + """The body is not a usable UTF-8 JSON object (HTTP 400).""" + + +@dataclass(frozen=True) +class Question: + name: str + keys: tuple[str, ...] + labels: tuple[str, ...] + goal: str + + +@dataclass(frozen=True) +class Request: + image: Image.Image + questions: tuple[Question, ...] + + +def _object(pairs): + result = {} + for key, value in pairs: + if key in result: + raise MalformedJSON("duplicate JSON keys are not supported") + result[key] = value + return result + + +def _nonfinite(value): + raise MalformedJSON("non-finite JSON numbers are not supported") + + +def decode_request(raw: bytes) -> dict: + if len(raw) > MAX_BODY: + raise InvalidRequest("request body exceeds 8 MiB") + try: + value = json.loads( + raw.decode("utf-8"), object_pairs_hook=_object, parse_constant=_nonfinite + ) + # The decoder accepts lone surrogate escapes and overflowing floats. + # Reject both before any text can reach the tokenizer or a response. + json.dumps(value, ensure_ascii=False, allow_nan=False).encode("utf-8") + except MalformedJSON: + raise + except (ValueError, UnicodeError, RecursionError) as exc: + raise MalformedJSON( + "request body must contain valid JSON and UTF-8 text" + ) from exc + if not isinstance(value, dict): + raise MalformedJSON("request must be a JSON object") + return value + + +def _text(value, field): + if not isinstance(value, (str, dict, list)): + raise InvalidRequest(f"{field} must be a string, object or array") + try: + result = ( + value + if isinstance(value, str) + else json.dumps(value, ensure_ascii=False, allow_nan=False) + ) + except (ValueError, TypeError, RecursionError) as exc: + raise InvalidRequest(f"invalid {field}") from exc + if len(result) > MAX_TEXT: + raise InvalidRequest(f"{field} exceeds {MAX_TEXT} characters") + if any( + token in result + for token in ( + "<|image_pad|>", + "<|video_pad|>", + "<|vision_start|>", + "<|vision_end|>", + ) + ): + raise InvalidRequest(f"{field} contains an unsupported media control token") + return result + + +def _image(state): + if not isinstance(state, dict) or set(state) != {"image"}: + raise InvalidRequest("state must contain exactly one image data URL") + url = state["image"] + if not isinstance(url, str): + raise InvalidRequest("state.image must be a PNG/JPEG base64 data URL") + prefix, separator, encoded = url.partition(",") + expected = {"data:image/png;base64": "PNG", "data:image/jpeg;base64": "JPEG"} + if not separator or prefix not in expected: + raise InvalidRequest("only inline PNG/JPEG images are supported") + if len(encoded) > 4 * ((MAX_IMAGE_BYTES + 2) // 3): + raise InvalidRequest("encoded image exceeds 4 MiB") + try: + raw = base64.b64decode(encoded, validate=True) + if len(raw) > MAX_IMAGE_BYTES: + raise InvalidRequest("image exceeds 4 MiB") + with Image.open(io.BytesIO(raw)) as source: + if source.format != expected[prefix]: + raise InvalidRequest("image format does not match its MIME type") + w, h = source.size + if max(w, h) > MAX_ASPECT_RATIO * min(w, h): + raise InvalidRequest( + f"image aspect ratio must not exceed {MAX_ASPECT_RATIO}:1" + ) + if ( + max(w, h) > MAX_SIDE + or w * h > MAX_PIXELS + or getattr(source, "n_frames", 1) != 1 + ): + raise InvalidRequest( + "image must be single-frame, at most 2048 per side and 1048576 pixels" + ) + source.load() + return source.convert("RGB") + except ( + binascii.Error, + UnidentifiedImageError, + OSError, + Image.DecompressionBombError, + ValueError, + ) as exc: + if isinstance(exc, InvalidRequest): + raise + raise InvalidRequest("invalid image data") from exc + + +def parse_request(value: dict) -> Request: + if not isinstance(value, dict) or set(value) != {"model", "state", "questions"}: + raise InvalidRequest("request must contain model, state and questions only") + if value["model"] != MODEL: + raise InvalidRequest(f"model must be {MODEL}") + questions = value["questions"] + if not isinstance(questions, dict) or not 1 <= len(questions) <= MAX_QUESTIONS: + raise InvalidRequest("questions must contain 1 to 8 questions") + parsed = [] + for name, q in questions.items(): + if not isinstance(name, str) or not name or len(name) > 256: + raise InvalidRequest("question names must contain 1 to 256 characters") + if not isinstance(q, dict) or q.get("type") != "choice": + raise InvalidRequest("only choice questions are supported") + if set(q) - {"type", "instructions", "criteria"}: + raise InvalidRequest("unsupported question fields") + if "instructions" not in q: + raise InvalidRequest("instructions is required for every question") + criteria = q.get("criteria") + if not isinstance(criteria, dict) or not 1 <= len(criteria) <= 26: + raise InvalidRequest("choice requires 1 to 26 options") + labels = [] + for key, label in criteria.items(): + if not isinstance(key, str) or not key or len(key) > 256: + raise InvalidRequest("option keys must contain 1 to 256 characters") + text = _text(key if label is None else label, "criteria") + labels.append(json.dumps(text, ensure_ascii=False)[1:-1]) + goal = ( + "" + if q["instructions"] is None + else _text(q["instructions"], "instructions") + ) + if len(goal) + sum(map(len, labels)) > MAX_TEXT: + raise InvalidRequest("combined question text exceeds 16384 characters") + parsed.append(Question(name, tuple(criteria), tuple(labels), goal)) + return Request(_image(value["state"]), tuple(parsed)) + + +def build_messages(q: Question, image_marker: str = "inline.png") -> list[dict]: + lines = "\n".join( + f'{chr(65 + i)}. Decision "{label}" -> select' + for i, label in enumerate(q.labels) + ) + text = ( + (f"Goal: {q.goal}\n\n" if q.goal else "") + + "App: Cua Driver\nTask family: closed-candidate decision\n\n" + + "The current screenshot is attached.\n\n" + + f"Options:\n{lines}\n\nAnswer with a single letter." + ) + return [ + {"role": "system", "content": SYSTEM_PROMPT}, + { + "role": "user", + "content": [ + {"type": "image", "image": image_marker}, + {"type": "text", "text": text}, + ], + }, + ] + + +def answer(q: Question, probabilities: list[float]) -> dict: + if len(probabilities) != len(q.keys) or any( + not math.isfinite(p) or not 0 <= p <= 1 for p in probabilities + ): + raise ValueError("model returned invalid probabilities") + if not math.isclose(sum(probabilities), 1.0, abs_tol=1e-5): + raise ValueError("model probabilities do not sum to one") + n = len(probabilities) + confidence = ( + 1.0 + if n == 1 + else 1 + sum(p * math.log(p) for p in probabilities if p) / math.log(n) + ) + return { + "type": "choice", + "choice": q.keys[max(range(n), key=probabilities.__getitem__)], + "probabilities": dict(zip(q.keys, probabilities)), + "confidence": max(0.0, min(1.0, confidence)), + }