diff --git a/src/frontend/cua_s1.py b/src/frontend/cua_s1.py new file mode 100644 index 0000000..cefa44d --- /dev/null +++ b/src/frontend/cua_s1.py @@ -0,0 +1,205 @@ +"""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 = False # Join accepted handlers before retiring GPU buffers. + + def __init__(self, address, engine): + self.engine = engine + self.inference_lock = threading.Lock() + self._engine_closed = False + super().__init__(address, Handler) + + def server_close(self): + super().server_close() + if not self._engine_closed: + close = getattr(self.engine, "close", None) + if close is not None: + close() + self._engine_closed = True + + +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 parse_args(argv=None): + from models.cua_s1.multimodal.graph_runtime import GraphConfig + + 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-mode", + choices=("exact", "rule-bucket", "auto"), + help="enable the selected Graph execution mode", + ) + p.add_argument("--graph-bucket-width", type=int, default=None) + 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, + help="distinct requests needed before capture", + ) + p.add_argument("--graph-max-tokens", type=int, default=2048) + p.add_argument("--graph-admission-window", type=int, default=8) + p.add_argument("--graph-cooldown-requests", type=int, default=32) + p.add_argument("--graph-capture-window", type=int, default=32) + p.add_argument("--graph-max-captures", type=int, default=4) + p.add_argument("--graph-capture-budget-ms", type=float, default=2000.0) + args = p.parse_args(argv) + mode = args.graph_mode or ("exact" if args.graph else None) + if args.graph_bucket_width is not None and mode not in {"rule-bucket", "auto"}: + p.error("--graph-bucket-width requires --graph-mode rule-bucket or auto") + try: + graph_config = ( + GraphConfig( + mode=mode, + bucket_width=64 + if args.graph_bucket_width is None + else args.graph_bucket_width, + 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, + admission_window=args.graph_admission_window, + cooldown_requests=args.graph_cooldown_requests, + capture_window=args.graph_capture_window, + max_captures=args.graph_max_captures, + capture_budget_ms=args.graph_capture_budget_ms, + ) + if mode is not None + else None + ) + except ValueError as exc: + p.error(str(exc)) + args.graph_config = graph_config + return args + + +def main(): + from models.cua_s1.multimodal.model import MultimodalEngine + + args = parse_args() + logging.basicConfig(level=logging.INFO) + engine = MultimodalEngine(args.base, args.adapter, graph_config=args.graph_config) + server = None + try: + engine.warmup() + # Bind only after model loading and 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) + server.serve_forever() + except KeyboardInterrupt: + pass + finally: + if server is not None: + server.server_close() + else: + engine.close() + + +if __name__ == "__main__": + main() diff --git a/src/models/cua_s1/multimodal/graph_admission.py b/src/models/cua_s1/multimodal/graph_admission.py new file mode 100644 index 0000000..fe8ad60 --- /dev/null +++ b/src/models/cua_s1/multimodal/graph_admission.py @@ -0,0 +1,71 @@ +"""Bounded request-frequency admission and sliding capture-work budget.""" + +from collections import OrderedDict, deque + + +class AdmissionPolicy: + """Called under the runtime lock; counters advance once per predict request. + + A completed attempt can exceed the time budget because capture is synchronous. + Every subsequent attempt waits for enough budget to expire. Cache hits never + spend capture budget. History is intentionally bounded and forgetting is cold. + """ + + def __init__(self, config): + self.config = config + self.history = OrderedDict() + self.cooldowns = OrderedDict() + self.attempts = deque() + self.request_index = 0 + + def reset(self): + self.history.clear() + self.cooldowns.clear() + self.attempts.clear() + self.request_index = 0 + + def begin_request(self): + self.request_index += 1 + while self.attempts and ( + self.request_index - self.attempts[0][0] >= self.config.capture_window + ): + self.attempts.popleft() + + def reason(self, key): + """Observe a cache miss and return its eager fallback reason, if any.""" + if self.request_index <= self.cooldowns.get(key, -1): + return "cooldown" + self.cooldowns.pop(key, None) + last, count = self.history.get(key, (-self.config.admission_window, 0)) + if self.request_index - last > self.config.admission_window: + count = 0 + if last != self.request_index: + count += 1 + self.history[key] = (self.request_index, count) + self.history.move_to_end(key) + while len(self.history) > 128: + self.history.popitem(last=False) + if count < self.config.min_uses: + return "warmup" + if ( + len(self.attempts) >= self.config.max_captures + or sum(item[1] for item in self.attempts) >= self.config.capture_budget_ms + ): + return "capture_budget" + return None + + def start_capture(self): + ticket = [self.request_index, 0.0] + self.attempts.append(ticket) + return ticket + + @staticmethod + def finish_capture(ticket, elapsed_ms): + ticket[1] = elapsed_ms + + def evict(self, key): + self.history.pop(key, None) + self.cooldowns[key] = self.request_index + self.config.cooldown_requests + self.cooldowns.move_to_end(key) + while len(self.cooldowns) > 128: + self.cooldowns.popitem(last=False) diff --git a/src/models/cua_s1/multimodal/graph_auto.py b/src/models/cua_s1/multimodal/graph_auto.py new file mode 100644 index 0000000..590c7a1 --- /dev/null +++ b/src/models/cua_s1/multimodal/graph_auto.py @@ -0,0 +1,149 @@ +"""Opt-in adaptive dispatch with shared Graph budgets and strict numeric gates.""" + +from collections import Counter, deque +from contextlib import contextmanager + +from .graph_buckets import bucket_length +from .graph_runtime import GraphConfig +from .graph_selector import AutoSelector +from .graph_shared import SharedGraphRuntime + + +class AutoGraphRuntime(SharedGraphRuntime): + def __init__(self, model, config=None): + super().__init__(model, config or GraphConfig(mode="auto")) + self.selector = AutoSelector(self.config) + self._pending = deque(maxlen=16) + self._decisions = {} + self.decisions = Counter() + self.selections = Counter(eager=0, exact=0, **{"rule-bucket": 0}) + + @property + def stats(self): + result = super().stats + result.update( + { + "selected_" + mode.replace("-", "_"): count + for mode, count in self.selections.items() + } + ) + return result + + def select_mode(self, mode): + raise RuntimeError("automatic runtime chooses its own execution mode") + + def _collect_timings(self): + while self._pending and self._pending[0][1].query(): + start, end, bucket, mode = self._pending.popleft() + self.selector.record(bucket, mode, latency_ms=start.elapsed_time(end)) + + @contextmanager + def request(self): + with super().request(): + self._collect_timings() + self.selector.begin_request(self.admission.request_index) + self._decisions.clear() + yield + + def forward(self, values): + import torch + + with self.lock, torch.no_grad(): + if self._closed: + raise RuntimeError("shared Graph runtime is closed") + exact = self._runtimes["exact"] + bucket_runtime = self._runtimes["rule-bucket"] + if not self._in_request: + exact.stats["no_request"] += 1 + return exact._eager(values) + if ( + not exact._supported(values) + or bucket_length( + values["inputs_embeds"].shape[1], self.config.bucket_width + ) + > self.config.max_tokens + ): + exact.stats["unsupported"] += 1 + self.selections["eager"] += 1 + return exact._eager(values) + key, bucket = exact._key(values), bucket_runtime._key(values) + if ( + key not in self._decisions + and len(self._decisions) >= self.selector.max_layouts + ): + self.selections["eager"] += 1 + self.decisions["request_layout_limit"] += 1 + return exact._eager(values) + length = values["inputs_embeds"].shape[1] + self._collect_timings() + self.selector.observe(key, bucket, length) + # Observe both candidate histories once/request, including eager + # observations. reason() deduplicates another call by the selected + # runtime; discarded answers do not spend capture budget. + exact.admission.reason(key) + bucket_runtime.admission.reason(bucket) + if key not in self._decisions: + mode = self.selector.choose( + key, + bucket, + length, + exact_entry=self.cache.entries.get(("exact", key)), + bucket_entry=self.cache.entries.get(("rule-bucket", bucket)), + cache_shapes=len(self.cache), + cache_bytes=self.cache.bytes, + exact_allowed=key not in exact.disabled, + bucket_allowed=bucket not in bucket_runtime.disabled, + ) + self._decisions[key] = mode + self.decisions[self.selector.last_reason] += 1 + mode = self._decisions[key] + self.selections[mode] += 1 + runtime = exact if mode == "eager" else self._runtimes[mode] + before = dict(runtime.stats) + samples = sum(self.selections.values()) + # Keep overhead bounded even when a selected mode repeatedly falls + # back and therefore never acquires its own replay-cost sample. + sample = samples <= 8 or samples % 16 == 0 + start = end = None + with torch.cuda.device(values["inputs_embeds"].device): + if sample: + start, end = ( + torch.cuda.Event(enable_timing=True) for _ in range(2) + ) + start.record() + result = ( + runtime._eager(values) + if mode == "eager" + else runtime.forward(values) + ) + if end is not None: + end.record() + attempts = runtime.stats["capture_attempts"] - before["capture_attempts"] + checks = runtime.stats.get("length_checks", 0) - before.get( + "length_checks", 0 + ) + if attempts: + self.selector.record( + bucket, + mode, + capture_ms=runtime.stats["capture_attempt_ms"] + - before["capture_attempt_ms"], + ) + elif not checks and start is not None: + actual_mode = ( + mode if runtime.stats["replays"] > before["replays"] else "eager" + ) + self._pending.append((start, end, bucket, actual_mode)) + entry_key = key if mode == "exact" else bucket + entry = self.cache.entries.get((mode, entry_key)) + if entry is not None: + self.selector.record(bucket, mode, owned_bytes=entry.bytes) + self.selector.selected(key, mode) + return result + + def invalidate(self): + with self.lock: + super().invalidate() + self.selector.reset() + self._pending.clear() + self._decisions.clear() diff --git a/src/models/cua_s1/multimodal/graph_buckets.py b/src/models/cua_s1/multimodal/graph_buckets.py new file mode 100644 index 0000000..a2916d3 --- /dev/null +++ b/src/models/cua_s1/multimodal/graph_buckets.py @@ -0,0 +1,229 @@ +"""Opt-in rule-only CUDA Graph buckets for the pinned multimodal worker.""" + +from __future__ import annotations + +import time + +from .graph_runtime import ( + GraphConfig, + GraphRuntime, + _GraphPool, + _GraphSegment, + tensor_signature, +) +from .rule_prefill import pinned_implementation, rule_prefill + + +def bucket_length(length, width): + return ((length + width - 1) // width) * width + + +class RuleBucketRuntime(GraphRuntime): + """Bounded cache of rule graphs, with explicit model-instance dispatch.""" + + def __init__(self, model, config=None, **resources): + config = config or GraphConfig(mode="rule-bucket") + super().__init__(model, config, **resources) + self.width = config.bucket_width + self.stats.update( + length_checks=0, + length_rejections=0, + length_disabled=0, + length_check_ms=0.0, + real_tokens=0, + padded_tokens=0, + ) + + @staticmethod + def _dense(values): + hidden = values["inputs_embeds"] + mask = values["attention_mask"] + return ( + hidden.is_contiguous() + and mask is not None + and tuple(mask.shape) == tuple(hidden.shape[:2]) + and bool((mask == 1).all()) + ) + + def _key(self, values): + hidden = values["inputs_embeds"] + return ( + id(self.model), + self.model.get_base_model().model.language_model.config._attn_implementation, + bucket_length(hidden.shape[1], self.width), + hidden.shape[2], + str(hidden.dtype), + str(hidden.device), + ) + + def _validate_replay(self, values, entry, output): + import torch + + length = values["inputs_embeds"].shape[1] + if length in entry.verified_lengths: + return output + started = time.perf_counter() + reference = self._eager(values) + self.stats["length_checks"] += 1 + if torch.equal(reference, output): + entry.verified_lengths.add(length) + else: + entry.rejected_lengths.add(length) + self.stats["length_rejections"] += 1 + output = reference + self.stats["length_check_ms"] += (time.perf_counter() - started) * 1000 + return output + + def forward(self, values): + import torch + + with self.lock, torch.no_grad(): + if self._closed: + raise RuntimeError("Graph runtime is closed") + if not self._in_request: + self.stats["no_request"] += 1 + return self._eager(values) + if not self._supported(values) or not self._dense(values): + self.stats["unsupported"] += 1 + return self._eager(values) + length = values["inputs_embeds"].shape[1] + if bucket_length(length, self.width) > self.config.max_tokens: + self.stats["unsupported"] += 1 + return self._eager(values) + key = self._key(values) + if key in self.disabled: + self.stats["disabled"] += 1 + return self._eager(values) + entry = self.cache.get(key) + if entry is not None: + if length in entry.rejected_lengths: + self.stats["length_disabled"] += 1 + return self._eager(values) + output = self._run_segments(values, entry) + self.stats["replays"] += 1 + return self._validate_replay(values, entry, output) + reason = self.admission.reason(key) + if reason is not None: + self.stats[reason] += 1 + return self._eager(values) + ticket = self.admission.start_capture() + self.stats["capture_attempts"] += 1 + attempt_start = time.perf_counter() + try: + output = self._capture(values, key) + entry = self.cache.get(key) + if entry is not None: + entry.verified_lengths = {length} + entry.rejected_lengths = set() + return output + finally: + elapsed = (time.perf_counter() - attempt_start) * 1000 + self.admission.finish_capture(ticket, elapsed) + self.stats["capture_attempt_ms"] += elapsed + + def _run_segments(self, values, entry): + from transformers.masking_utils import create_causal_mask + + implementation = pinned_implementation() + text = self.model.get_base_model().model.language_model + hidden = values["inputs_embeds"] + length = hidden.shape[1] + self.stats["real_tokens"] += length + self.stats["padded_tokens"] += bucket_length(length, self.width) - length + mask = create_causal_mask( + config=text.config, + inputs_embeds=hidden, + attention_mask=values["attention_mask"], + past_key_values=None, + position_ids=None, + ) + rope = text.rotary_emb(hidden, values["position_ids"]) + for index, layer in enumerate(text.layers): + if text.config.layer_types[index] == "full_attention": + hidden = layer( + hidden, + position_embeddings=rope, + attention_mask=mask, + position_ids=None, + past_key_values=None, + use_cache=False, + ) + continue + if type(layer.linear_attn) is not implementation.Qwen3_5GatedDeltaNet: + raise RuntimeError("unsupported linear-attention implementation") + + def dispatch(query, key, value, *, g, beta, **kwargs): + packed = pack_rule_inputs( + dict(query=query, key=key, value=value, g=g, beta=beta), self.width + ) + block = entry.blocks.get(index) + if block is None: + if entry.pool is None: + entry.pool = _GraphPool(query.device) + block = RuleSegment( + implementation.torch_chunk_gated_delta_rule, packed, entry.pool + ) + entry.blocks[index] = block + entry.update_bytes() + output, state = block.replay_values(packed) + return output[:, :length].contiguous(), state + + residual = hidden + hidden = rule_prefill( + layer.linear_attn, layer.input_layernorm(hidden), dispatch + ) + hidden = residual + hidden + residual = hidden + hidden = residual + layer.mlp(layer.post_attention_layernorm(hidden)) + hidden = text.norm(hidden) + return self.model.get_base_model().lm_head(hidden[:, -1:, :])[0, -1, :] + + +def pack_rule_inputs(values, width): + import torch.nn.functional as F + + length = values["query"].shape[1] + padding = bucket_length(length, width) - length + return { + # Zero-padding is a clone that can preserve noncontiguous strides. + # Canonicalize even on an exact boundary so one bucket has one layout. + name: F.pad(value, (0, 0) * (value.ndim - 2) + (0, padding)).contiguous() + for name, value in values.items() + } + + +class RuleSegment(_GraphSegment): + """Capture just the fallback rule, reusing #33 stream/pool ownership.""" + + def __init__(self, function, values, pool): + self.function = function + self.signatures = {k: tensor_signature(v) for k, v in values.items()} + self.static_extra = {k: v.clone() for k, v in values.items() if k != "query"} + super().__init__([], 0, 0, values["query"], None, pool=pool) + self.external_bytes += sum( + v.untyped_storage().nbytes() for v in self.static_extra.values() + ) + + def _forward(self): + return self.function( + self.static_hidden, + self.static_extra["key"], + self.static_extra["value"], + g=self.static_extra["g"], + beta=self.static_extra["beta"], + initial_state=None, + output_final_state=False, + use_qk_l2norm_in_kernel=True, + ) + + def replay_values(self, values): + if {k: tensor_signature(v) for k, v in values.items()} != self.signatures: + raise ValueError("CUDA Graph rule input layout changed") + for name, value in self.static_extra.items(): + value.copy_(values[name]) + return super().replay(values["query"], None) + + def close(self): + super().close() + if hasattr(self, "static_extra"): + self.static_extra.clear() 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..fc15b4c --- /dev/null +++ b/src/models/cua_s1/multimodal/graph_runtime.py @@ -0,0 +1,589 @@ +"""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 math +import threading +import time +from collections import OrderedDict +from contextlib import contextmanager, suppress +from dataclasses import dataclass, field + +from .graph_admission import AdmissionPolicy + + +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 + admission_window: int = 8 + cooldown_requests: int = 32 + capture_window: int = 32 + max_captures: int = 4 + capture_budget_ms: float = 2000.0 + mode: str = "exact" + bucket_width: int = 64 + + def __post_init__(self): + if self.mode not in {"exact", "rule-bucket", "auto"}: + raise ValueError("graph mode must be exact, rule-bucket or auto") + if ( + type(self.bucket_width) is not int + or self.bucket_width <= 0 + or self.bucket_width % 64 + ): + raise ValueError("bucket width must be a positive integer multiple of 64") + if self.mode == "exact" and self.bucket_width != 64: + raise ValueError("bucket width is only configurable in rule-bucket mode") + if self.mode in {"rule-bucket", "auto"} and self.bucket_width > self.max_tokens: + raise ValueError("bucket width cannot exceed graph max tokens") + if not math.isfinite(self.capture_budget_ms): + raise ValueError("capture time budget must be finite") + if ( + min( + self.max_shapes, + self.max_bytes, + self.min_uses, + self.max_tokens, + self.admission_window, + self.cooldown_requests, + self.capture_window, + self.max_captures, + self.capture_budget_ms, + ) + <= 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: + key: tuple | None = None + blocks: dict = field(default_factory=dict) + bytes: int = 0 + pool: _GraphPool | None = None + verified_lengths: set[int] = field(default_factory=set) + rejected_lengths: set[int] = field(default_factory=set) + + 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, + *, + cache=None, + admission=None, + lock=None, + ): + self.model = model + self.config = config or GraphConfig() + self.cache = ( + cache + if cache is not None + else GraphCache(self.config.max_shapes, self.config.max_bytes) + ) + self.admission = ( + admission if admission is not None else AdmissionPolicy(self.config) + ) + self._in_request = False + self._closed = False + self.disabled = OrderedDict() + self.lock = lock if lock is not None else threading.RLock() + self.stats = { + "requests": 0, + "no_request": 0, + "cooldown": 0, + "capture_budget": 0, + "capture_attempts": 0, + "capture_attempt_ms": 0.0, + "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, + } + + @contextmanager + def request(self): + """Serialize a complete prediction and count duplicate layouts only once.""" + with self.lock: + if self._closed: + raise RuntimeError("Graph runtime is closed") + if self._in_request: + raise RuntimeError("nested Graph requests are unsupported") + self._in_request = True + self.admission.begin_request() + self.stats["requests"] += 1 + try: + yield + finally: + self._in_request = False + + def invalidate(self): + """Call after replacing model weights or adapters.""" + with self.lock: + if self._in_request: + raise RuntimeError("cannot invalidate during a Graph request") + retired = self.cache.clear() + self.admission.reset() + self.disabled.clear() + for entry in retired: + entry.close() + del retired + gc.collect() + + def close(self): + """Synchronously retire graph resources; reject subsequent requests.""" + with self.lock: + if self._closed: + return + self.invalidate() + self._closed = True + + 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 _key(self, values): + return ( + 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") + ), + ) + + def forward(self, values): + """Return fresh logits; no static graph output escapes the runtime lock.""" + import torch + + with self.lock, torch.no_grad(): + if self._closed: + raise RuntimeError("Graph runtime is closed") + if not self._in_request: + self.stats["no_request"] += 1 + return self._eager(values) + if not self._supported(values): + self.stats["unsupported"] += 1 + return self._eager(values) + key = self._key(values) + 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 + + reason = self.admission.reason(key) + if reason is not None: + self.stats[reason] += 1 + return self._eager(values) + + ticket = self.admission.start_capture() + self.stats["capture_attempts"] += 1 + attempt_start = time.perf_counter() + try: + return self._capture(values, key) + finally: + elapsed = (time.perf_counter() - attempt_start) * 1000 + self.admission.finish_capture(ticket, elapsed) + self.stats["capture_attempt_ms"] += elapsed + + def _capture(self, values, key): + import torch + + reference = self._eager(values) + candidate = _ShapeEntry(key=key) + 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: + self.admission.evict(entry.key) + 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/graph_selector.py b/src/models/cua_s1/multimodal/graph_selector.py new file mode 100644 index 0000000..1992d6c --- /dev/null +++ b/src/models/cua_s1/multimodal/graph_selector.py @@ -0,0 +1,182 @@ +"""Bounded, causal mode selection with conservative capture amortization. + +The 32-request observation window forecasts at most the next 64 requests. The +500 ms capture and 0.45/0.65 replay priors are heuristics for the pinned fallback +path, replaced by observations when available. No future workload is consulted. +""" + +from __future__ import annotations + +import math +from collections import OrderedDict +from dataclasses import dataclass, field + + +@dataclass +class _Layout: + bucket: tuple + length: int + visits: list = field(default_factory=list) + mode: str | None = None + switched_at: int = 0 + + +class AutoSelector: + window = 32 + max_layouts = 128 + switch_cooldown = 32 + margin = 1.25 + + def __init__(self, config): + self.config = config + self.history = OrderedDict() + self.costs = OrderedDict() + self.index = 0 + self.last_reason = "cold" + + def reset(self): + self.history.clear() + self.costs.clear() + self.index = 0 + + def begin_request(self, index): + self.index = index + for key, layout in list(self.history.items()): + layout.visits[:] = [ + visit for visit in layout.visits if index - visit[0] < self.window + ] + if not layout.visits: + del self.history[key] + + def observe(self, key, bucket, length): + layout = self.history.setdefault(key, _Layout(bucket, length)) + if layout.visits and layout.visits[-1][0] == self.index: + layout.visits[-1][1] += 1 + else: + layout.visits.append([self.index, 1]) + self.history.move_to_end(key) + while len(self.history) > self.max_layouts: + self.history.popitem(last=False) + + def record( + self, bucket, mode, *, latency_ms=None, capture_ms=None, owned_bytes=None + ): + costs = self.costs.setdefault(bucket, {}) + for key, value in [(mode, latency_ms), (mode + "_capture", capture_ms)]: + if value is not None and math.isfinite(value) and value > 0: + costs[key] = ( + value if key not in costs else 0.8 * costs[key] + 0.2 * value + ) + if owned_bytes is not None and owned_bytes > 0: + costs[mode + "_bytes"] = max(costs.get(mode + "_bytes", 0), owned_bytes) + self.costs.move_to_end(bucket) + while len(self.costs) > self.max_layouts: + self.costs.popitem(last=False) + + def selected(self, key, mode): + layout = self.history.get(key) + if layout is not None and layout.mode != mode: + layout.mode = mode + layout.switched_at = self.index + + def _result(self, mode, reason): + self.last_reason = reason + return mode + + def choose( + self, + key, + bucket, + length, + *, + exact_entry=None, + bucket_entry=None, + cache_shapes=0, + cache_bytes=0, + exact_allowed=True, + bucket_allowed=True, + ): + costs = self.costs.get(bucket, {}) + eager = costs.get("eager", 100.0) + exact = costs.get("exact", eager * 0.45) + rule = costs.get("rule-bucket", eager * 0.65) + if exact_entry is not None: + return self._result("exact", "resident_exact") + layout = self.history[key] + required = max(2, self.config.min_uses) + if bucket_entry is not None and length in bucket_entry.rejected_lengths: + bucket_allowed = False + if len(layout.visits) < required: + # A previously verified resident bucket length needs no extra gate. + if ( + bucket_allowed + and bucket_entry is not None + and length in bucket_entry.verified_lengths + and rule < eager + ): + return self._result("rule-bucket", "resident_bucket") + return self._result("eager", "cold") + gap = (layout.visits[-1][0] - layout.visits[0][0]) / (len(layout.visits) - 1) + copies = sum(v[1] for v in layout.visits) / len(layout.visits) + future_exact = 2 * self.window / gap * copies + repeated = sum(len(item.visits) >= required for item in self.history.values()) + pressure = repeated > max(1, self.config.max_shapes // 2) + exact_fits = ( + cache_shapes < self.config.max_shapes + and cache_bytes + costs.get("exact_bytes", 0) <= self.config.max_bytes + ) + switching = layout.mode not in (None, "exact") + cooled = ( + not switching or self.index - layout.switched_at >= self.switch_cooldown + ) + current_cost = ( + min(eager, rule) if bucket_allowed and bucket_entry is not None else eager + ) + if ( + exact_allowed + and gap <= self.config.admission_window + and not pressure + and exact_fits + and cooled + and future_exact * max(0, current_cost - exact) + > self.margin * costs.get("exact_capture", 500.0) + ): + return self._result("exact", "amortized_exact") + + if not bucket_allowed: + return self._result("eager", "bucket_rejected") + if bucket_entry is not None: + gate = 0 if length in bucket_entry.verified_lengths else eager + if future_exact * max(0, eager - rule) > self.margin * gate: + return self._result("rule-bucket", "resident_bucket") + return self._result("eager", "length_gate_cost") + if ( + layout.mode == "exact" + and self.index - layout.switched_at < self.switch_cooldown + ): + return self._result("eager", "switch_cooldown") + + visits = {} + lengths = set() + for item in self.history.values(): + if item.bucket == bucket: + lengths.add(item.length) + for request, count in item.visits: + visits[request] = visits.get(request, 0) + count + span = max(visits) - min(visits) + future_bucket = ( + 2 + * self.window + * (len(visits) - 1) + / span + * sum(visits.values()) + / len(visits) + ) + # The capture estimate already includes the current length's reference. + gates = max(0, len(lengths) - 1) * eager + if ( + future_bucket * max(0, eager - rule) + > self.margin * costs.get("rule-bucket_capture", 500.0) + gates + ): + return self._result("rule-bucket", "amortized_bucket") + return self._result("eager", "cost") diff --git a/src/models/cua_s1/multimodal/graph_shared.py b/src/models/cua_s1/multimodal/graph_shared.py new file mode 100644 index 0000000..9c044d8 --- /dev/null +++ b/src/models/cua_s1/multimodal/graph_shared.py @@ -0,0 +1,163 @@ +"""Per-model shared resource ownership for exact and rule-only CUDA Graphs. + +Children are owned exclusively by this group; use its request and lifecycle +methods. Cache limits cover resident ownership, not transient candidate capture +or the model and other process allocations. +""" + +from __future__ import annotations + +import gc +import threading +from contextlib import contextmanager +from dataclasses import replace + +from .graph_admission import AdmissionPolicy +from .graph_buckets import RuleBucketRuntime +from .graph_runtime import GraphCache, GraphConfig, GraphRuntime + + +class _CacheView: + def __init__(self, group, mode): + self.group = group + self.mode = mode + + def get(self, key): + return self.group.cache.get((self.mode, key)) + + def put(self, key, entry): + entry.key = (self.mode, key) + retired = self.group.cache.put(entry.key, entry) + if retired is not None: + # Explicitly release even when diagnostics retain Python references. + # _capture subsequently applies global cooldown to these exact keys. + for previous in retired: + previous.close() + return retired + + +class _AdmissionView: + def __init__(self, group, mode): + self.group = group + self.mode = mode + + def reason(self, key): + return self.group.admission.reason((self.mode, key)) + + def start_capture(self): + return self.group.admission.start_capture() + + def finish_capture(self, ticket, elapsed_ms): + self.group.admission.finish_capture(ticket, elapsed_ms) + + def evict(self, namespaced_key): + self.group.admission.evict(namespaced_key) + + +class SharedGraphRuntime: + """One request clock, capture ledger, LRU and lock for two execution modes.""" + + def __init__(self, model, config=None): + self.model = model + self.config = config or GraphConfig() + self.cache = GraphCache(self.config.max_shapes, self.config.max_bytes) + self.admission = AdmissionPolicy(self.config) + self.lock = threading.RLock() + self._in_request = False + self._closed = False + self._requests = 0 + self._mode = "eager" + self._selected_this_request = set() + self._runtimes = {} + for mode, runtime_type in ( + ("exact", GraphRuntime), + ("rule-bucket", RuleBucketRuntime), + ): + config = replace( + self.config, + mode=mode, + bucket_width=64 if mode == "exact" else self.config.bucket_width, + ) + self._runtimes[mode] = runtime_type( + model, + config, + cache=_CacheView(self, mode), + admission=_AdmissionView(self, mode), + lock=self.lock, + ) + + @property + def stats(self): + result = {} + for runtime in self._runtimes.values(): + for key, value in runtime.stats.items(): + result[key] = result.get(key, 0) + value + result["requests"] = self._requests + return result + + def select_mode(self, mode): + if mode not in {"eager", "exact", "rule-bucket"}: + raise ValueError("mode must be eager, exact or rule-bucket") + with self.lock: + if self._closed: + raise RuntimeError("shared Graph runtime is closed") + if self._in_request: + raise RuntimeError("select a mode before beginning a request") + self._mode = mode + + @contextmanager + def request(self): + with self.lock: + if self._closed: + raise RuntimeError("shared Graph runtime is closed") + if self._in_request: + raise RuntimeError("nested Graph requests are unsupported") + self._in_request = True + self._requests += 1 + self.admission.begin_request() + self._selected_this_request.clear() + for runtime in self._runtimes.values(): + runtime._in_request = True + try: + yield + finally: + for runtime in self._runtimes.values(): + runtime._in_request = False + self._in_request = False + + def forward(self, values): + import torch + + with self.lock, torch.no_grad(): + if self._closed: + raise RuntimeError("shared Graph runtime is closed") + runtime = self._runtimes["exact" if self._mode == "eager" else self._mode] + if self._in_request and self._mode not in self._selected_this_request: + runtime.stats["requests"] += 1 + self._selected_this_request.add(self._mode) + if self._mode == "eager": + if not self._in_request: + runtime.stats["no_request"] += 1 + return runtime._eager(values) + return runtime.forward(values) + + def invalidate(self): + with self.lock: + if self._in_request: + raise RuntimeError("cannot invalidate during a Graph request") + retired = self.cache.clear() + for entry in retired: + entry.close() + self.admission.reset() + for runtime in self._runtimes.values(): + runtime.disabled.clear() + gc.collect() + + def close(self): + with self.lock: + if self._closed: + return + self.invalidate() + for runtime in self._runtimes.values(): + runtime._closed = True + self._closed = True diff --git a/src/models/cua_s1/multimodal/model.py b/src/models/cua_s1/multimodal/model.py new file mode 100644 index 0000000..f5741a6 --- /dev/null +++ b/src/models/cua_s1/multimodal/model.py @@ -0,0 +1,338 @@ +"""Direct Transformers/PEFT execution. No production dependency on cua_s1.""" + +from __future__ import annotations + +import copy +import hashlib +import json +import threading +from contextlib import nullcontext +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, + ): + self._lifecycle_lock = threading.RLock() + self._closed = False + 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 + + if graph_config.mode in {"rule-bucket", "auto"}: + from .graph_buckets import RuleBucketRuntime + from .rule_prefill import pinned_implementation + + pinned_implementation() + # Explicit adapter calls cannot honor offload/device-map hooks. + if any(p.device != self.model.device for p in self.model.parameters()): + raise ValueError( + "rule-bucket/auto requires one CUDA-resident model" + ) + if graph_config.mode == "auto": + from .graph_auto import AutoGraphRuntime + + self.graph_runtime = AutoGraphRuntime(self.model, graph_config) + else: + self.graph_runtime = RuleBucketRuntime(self.model, graph_config) + else: + 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: + with getattr(self, "_lifecycle_lock", nullcontext()): + if getattr(self, "_closed", False): + raise RuntimeError("multimodal engine is closed") + runtime = getattr(self, "graph_runtime", None) + with runtime.request() if runtime is not None else nullcontext(): + return self._predict(request) + + def close(self): + """Drain prediction before explicitly releasing captured GPU resources.""" + with getattr(self, "_lifecycle_lock", nullcontext()): + if getattr(self, "_closed", False): + return + runtime = getattr(self, "graph_runtime", None) + if runtime is not None: + runtime.close() + self._closed = True + + 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)), + } diff --git a/src/models/cua_s1/multimodal/rule_prefill.py b/src/models/cua_s1/multimodal/rule_prefill.py new file mode 100644 index 0000000..32e27eb --- /dev/null +++ b/src/models/cua_s1/multimodal/rule_prefill.py @@ -0,0 +1,64 @@ +"""Pinned Qwen3.5 cache-free prefill with an explicit DeltaNet rule callable. + +Adapted from Transformers 5.17.0 modeling_qwen3_5.Qwen3_5GatedDeltaNet.forward +(Copyright the HuggingFace Inc. team, Apache License 2.0). Operations retain the +upstream order and original token shapes; only the supplied rule may bucket. +No global functions, instance methods, weights or persistent states are changed. +""" + +from functools import lru_cache + + +@lru_cache(maxsize=1) +def pinned_implementation(): + import transformers + from transformers.models.qwen3_5 import modeling_qwen3_5 + + if transformers.__version__ != "5.17.0": + raise RuntimeError("rule-bucket requires pinned Transformers 5.17.0") + return modeling_qwen3_5 + + +def rule_prefill(module, hidden, rule): + """Dense, cache-free prefill on the supplied module's existing parameters.""" + import torch + from torch.nn import functional as F + + implementation = pinned_implementation() + batch, length, _ = hidden.shape + mixed = module.in_proj_qkv(hidden).transpose(1, 2) + z = module.in_proj_z(hidden).reshape(batch, length, -1, module.head_v_dim) + b = module.in_proj_b(hidden) + a = module.in_proj_a(hidden) + mixed = implementation.causal_conv1d_fn( + mixed, + module.conv1d.weight.squeeze(1), + module.conv1d.bias, + activation=module.activation, + ).transpose(1, 2) + query, key, value = torch.split( + mixed, [module.key_dim, module.key_dim, module.value_dim], dim=-1 + ) + query = query.reshape(batch, length, -1, module.head_k_dim) + key = key.reshape(batch, length, -1, module.head_k_dim) + value = value.reshape(batch, length, -1, module.head_v_dim) + beta = b.sigmoid() + g = -module.A_log.float().exp() * F.softplus(a.float() + module.dt_bias) + repeats = module.num_v_heads // module.num_k_heads + if repeats > 1: + query = query.repeat_interleave(repeats, dim=2) + key = key.repeat_interleave(repeats, dim=2) + output, _ = rule( + query, + key, + value, + g=g, + beta=beta, + initial_state=None, + output_final_state=False, + use_qk_l2norm_in_kernel=True, + ) + output = module.norm( + output.reshape(-1, module.head_v_dim), z.reshape(-1, module.head_v_dim) + ) + return module.out_proj(output.reshape(batch, length, -1))