diff --git a/CHANGELOG.md b/CHANGELOG.md index 5164360..697e20e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,16 @@ Notable changes to vla.cpp. Format loosely follows [Keep a Changelog](https://ke ## [Unreleased] +### Added + +- **Asynchronous serving.** `vla-server` is now a ZeroMQ ROUTER with prediction + on its own thread, so it keeps receiving while the model runs and a DEALER + client can keep requests in flight; REQ clients are unchanged. `--queue latest` + (default) keeps one pending request per client and answers a superseded one + with `error="superseded"`; `--queue fifo` serves every request in order. + Replies report `latency_ms_queue`. The lerobot fork's `lerobot-vla-cpp + --mode=async` drives it with a timestep-aligned action queue. + ## [0.4.0] - 2026-09-30 ### Added diff --git a/README.md b/README.md index 9ea731d..28beab2 100644 --- a/README.md +++ b/README.md @@ -207,9 +207,13 @@ lerobot-vla-cpp --server_address=tcp://127.0.0.1:5555 --arch=smolvla "${ROBOT[@] relative actions. - `--task` must match a trained instruction exactly, and the camera keys must stay `front` and `wrist` in that order. -- `vla-server` answers one request at a time, so the loop is synchronous and - `--n_action_steps` is the feedback rate: 25 at 30 fps leaves ~0.83 s between - observations. +- `--mode=async` runs inference on a background thread and merges each new chunk + into a timestep-aligned action queue, so the arm never waits for the server. + `--chunk_size_threshold` sets how early the next observation goes out, and + `--aggregate_fn_name` how overlapping actions from two chunks are blended. The + default `--mode=sync` executes `--n_action_steps` of each chunk before the + next request, so that number is the feedback rate: 25 at 30 fps leaves ~0.83 s + between observations. The GR00T paths of the client have not been tested against real checkpoints yet. Wiring, recording, training and queue sizing are in the diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 4ce8426..f8f15e3 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -27,7 +27,8 @@ source is the detail. - `src/tokenizer.h` - SentencePiece tokenizers stored in the GGUF, for `vla-cli --text`. - `include/vla.h`, `src/vla_c_api.cpp` - the C ABI (`libvla`). -- `src/serving/` - `vla-server` (ZeroMQ + protobuf, action prediction), `vlm-server` +- `src/serving/` - `vla-server` (ZeroMQ ROUTER + protobuf, action prediction on a + worker thread behind a latest-wins request queue), `vlm-server` (chat), `vla-cli` (one-shot inference) and `vla-bench` (timing). - `src/kernels/bitvla/` - custom 1.58-bit ternary CUDA kernels for BitVLA. diff --git a/docs/USAGE.md b/docs/USAGE.md index 52b5d44..f3a981e 100644 --- a/docs/USAGE.md +++ b/docs/USAGE.md @@ -44,8 +44,8 @@ Checkpoints are cached under `$VLA_CACHE` (default `~/.cache/vla`). ## `vla-server` -`vla-server` loads the model once at startup and answers ZeroMQ REQ/REP requests -synchronously. +`vla-server` loads the model once at startup and answers ZeroMQ requests carrying +the protobuf messages in `src/serving/vla.proto`. ```bash ./build/vla-server "$VLA_GGUF" @@ -60,6 +60,30 @@ vla-server: bound to tcp://*:5555. ready. Use `--bind` to change the address and port. Stop the server with `Ctrl-C`. `vla-server` also takes `-hf user/repo[:file.gguf|:tag]` in place of a checkpoint path. +### Asynchronous clients + +The socket is a ZeroMQ ROUTER, and prediction runs on its own thread. A REQ +client sees what it always did: one request, one reply. A DEALER client can +keep several requests in flight, and the server keeps receiving while the model +runs, so a robot's control loop never has to wait for a prediction to send the +next observation. Replies carry the request's `request_id`, which is how a +client matches a chunk to the observation it came from. + +`--queue` says what happens to requests that arrive while a prediction is +running: + +- `latest` (default): one pending request per client. A newer request from the + same client replaces the pending one, which is answered at once with + `error="superseded"`. The model therefore always works on the freshest + observation each robot has sent, and several robots sharing one server are + served in turn. +- `fifo`: every request is served in arrival order, for benchmarks that want + throughput rather than freshness. + +`latency_ms_queue` in the reply is how long the request waited for the predict +thread. A malformed request is rejected on the socket thread, so an error comes +back within milliseconds even mid-prediction. + Clients: the LIBERO and SimplerEnv runners in [EVAL.md](EVAL.md), and the real-robot client in the README's [Rollout on a real robot](../README.md#rollout-on-a-real-robot). diff --git a/src/serving/request_queue.h b/src/serving/request_queue.h new file mode 100644 index 0000000..5160730 --- /dev/null +++ b/src/serving/request_queue.h @@ -0,0 +1,137 @@ +// Copyright 2026 VinRobotics +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// The hand-off between vla-server's socket thread and its predict thread. +// +// The socket thread parses and validates a request and pushes a Job; the +// predict thread pops one at a time. A policy server only ever wants the newest +// observation from each robot, so in LATEST mode a push from a client that +// already has a job waiting replaces that job and hands it back to the caller, +// who answers it with a "superseded" error. FIFO keeps every request, bounded +// by a depth cap, for benchmarks that want throughput rather than freshness. +// +// Header-only and free of ZeroMQ, protobuf and the model so tests can drive it. + +#pragma once + +#include +#include +#include +#include +#include +#include +#include + +namespace vla::serving { + +enum class QueueMode { + LATEST, ///< One pending job per client; a newer one displaces it. + FIFO, ///< Every job is served in arrival order, up to max_depth. +}; + +inline bool parse_queue_mode(const std::string & s, QueueMode & out) { + if (s == "latest") { out = QueueMode::LATEST; return true; } + if (s == "fifo") { out = QueueMode::FIFO; return true; } + return false; +} + +inline const char * queue_mode_name(QueueMode m) { + return m == QueueMode::LATEST ? "latest" : "fifo"; +} + +/// What push() did with a job. +enum class PushResult { + QUEUED, ///< Appended; nothing displaced. + REPLACED, ///< Appended, and the same client's older pending job is in `displaced`. + REJECTED, ///< Queue full (FIFO only); the job itself is handed back in `displaced`. +}; + +template +class RequestQueue { +public: + explicit RequestQueue(QueueMode mode, size_t max_depth = 64) + : mode_(mode), max_depth_(max_depth == 0 ? 1 : max_depth) {} + + QueueMode mode() const { return mode_; } + + /// Jobs are keyed by `key(job)`, the ZeroMQ routing envelope in the server. + template + PushResult push(Job job, KeyFn key, std::optional & displaced) { + std::lock_guard lk(mu_); + displaced.reset(); + if (mode_ == QueueMode::LATEST) { + const auto k = key(job); + for (auto it = q_.begin(); it != q_.end(); ++it) { + if (key(*it) == k) { + displaced = std::move(*it); + q_.erase(it); + q_.push_back(std::move(job)); + cv_.notify_one(); + return PushResult::REPLACED; + } + } + // No pending job from this client. A bounded queue still protects + // against a flood of distinct identities. + if (q_.size() >= max_depth_) { + displaced = std::move(job); + return PushResult::REJECTED; + } + q_.push_back(std::move(job)); + cv_.notify_one(); + return PushResult::QUEUED; + } + if (q_.size() >= max_depth_) { + displaced = std::move(job); + return PushResult::REJECTED; + } + q_.push_back(std::move(job)); + cv_.notify_one(); + return PushResult::QUEUED; + } + + /// Blocks until a job is available or stop() was called. Returns false on stop + /// with the queue drained. + bool pop(Job & out) { + std::unique_lock lk(mu_); + cv_.wait(lk, [&] { return stopped_ || !q_.empty(); }); + if (q_.empty()) + return false; + out = std::move(q_.front()); + q_.pop_front(); + return true; + } + + /// Wakes pop(). Jobs still queued are returned by subsequent pops until + /// empty, so the worker can answer them before it exits. + void stop() { + std::lock_guard lk(mu_); + stopped_ = true; + cv_.notify_all(); + } + + size_t size() const { + std::lock_guard lk(mu_); + return q_.size(); + } + +private: + QueueMode mode_; + size_t max_depth_; + mutable std::mutex mu_; + std::condition_variable cv_; + std::deque q_; + bool stopped_ = false; +}; + +} // namespace vla::serving diff --git a/src/serving/server.cpp b/src/serving/server.cpp index d65a6ed..c92ba91 100644 --- a/src/serving/server.cpp +++ b/src/serving/server.cpp @@ -15,6 +15,7 @@ #include "model.h" #include "options.h" #include "serving/hf_fetch.h" +#include "serving/request_queue.h" #include "serving/vla.pb.h" #define STB_IMAGE_IMPLEMENTATION @@ -30,13 +31,17 @@ #include +#include #include +#include #include #include #include #include #include +#include #include +#include #include namespace { @@ -156,18 +161,6 @@ std::string make_error_response(uint64_t request_id, const std::string & msg) { return resp.SerializeAsString(); } -// Discard frames after the first. Must run to completion: a queued frame keeps -// REP in receive state and send throws EFSM. -bool drain_extra_frames(zmq::socket_t & sock) { - bool extra = false; - while (sock.get(zmq::sockopt::rcvmore)) { - zmq::message_t junk; - (void) sock.recv(junk, zmq::recv_flags::none); - extra = true; - } - return extra; -} - int find_non_finite(const float * data, int n) { for (int i=0; i envelope; + std::string key; ///< Envelope bytes; one pending job per key in LATEST mode. + uint64_t rid = 0; + Clock::time_point t_recv; + + std::vector> u8_bufs; + std::vector> f32_bufs; + std::vector img_views; + std::vector precomputed_emb; + int precomputed_n_views = 0; + bool use_precomputed = false; + + std::vector lang_tokens; + std::vector state; + std::vector noise; + std::vector attn_mask; +}; + +// Checks the request against the model and fills the job. Returns an empty +// string on success, else the error to send back. +std::string build_job(const vla::PredictRequest & req, const vla::Config & cfg, Job & job) { + char buf[192]; + if (req.images_size() < 1 && req.precomputed_img_emb_size() == 0) { + return "PredictRequest must contain images or precomputed_img_emb"; + } + if (req.images_size() > 16) { + return "too many image views (max 16)"; + } + if (req.lang_tokens_size() < 1 || req.lang_tokens_size() > int(cfg.n_lang)) { + std::snprintf(buf, sizeof(buf), "lang_tokens length %d out of range [1, %lld]", + req.lang_tokens_size(), (long long) cfg.n_lang); + return buf; + } + for (int t=0; t 0; + if (job.use_precomputed) { + job.precomputed_n_views = static_cast(req.precomputed_img_emb_n_views()); + const int64_t per_view = cfg.n_img*cfg.hidden; + const int64_t expected = per_view * static_cast(job.precomputed_n_views); + if (job.precomputed_n_views < 1 || job.precomputed_n_views > 16) { + std::snprintf(buf, sizeof(buf), "precomputed_img_emb_n_views %d out of range [1, 16]", + job.precomputed_n_views); + return buf; + } + if (static_cast(req.precomputed_img_emb_size()) != expected) { + std::snprintf(buf, sizeof(buf), + "precomputed_img_emb size %d != %lld (n_views=%d * n_img_per_view=%lld * hidden=%lld)", + req.precomputed_img_emb_size(), (long long) expected, job.precomputed_n_views, + (long long) cfg.n_img, (long long) cfg.hidden); + return buf; + } + job.precomputed_emb.assign(req.precomputed_img_emb().begin(), req.precomputed_img_emb().end()); + const int bad = find_non_finite(job.precomputed_emb.data(), + static_cast(job.precomputed_emb.size())); + if (bad >= 0) { + std::snprintf(buf, sizeof(buf), "precomputed_img_emb[%d] = %g is not finite (NaN/Inf)", + bad, job.precomputed_emb[bad]); + return buf; + } + } else { + const int n_views = req.images_size(); + job.u8_bufs.resize(n_views); + job.f32_bufs.resize(n_views); + job.img_views.resize(n_views); + size_t total_px = 0; + for (int v=0; v kMaxTotalPixels) { + return "images exceed the per-request pixel budget"; + } + } + } + + job.lang_tokens.assign(req.lang_tokens().begin(), req.lang_tokens().end()); + job.state.assign(req.state().begin(), req.state().end()); + { + const int bad = find_non_finite(job.state.data(), static_cast(job.state.size())); + if (bad >= 0) { + std::snprintf(buf, sizeof(buf), "state[%d] = %g is not finite (NaN/Inf)", bad, job.state[bad]); + return buf; + } + } + if (req.noise_size() == expected_noise_n) { + job.noise.assign(req.noise().begin(), req.noise().end()); + const int bad = find_non_finite(job.noise.data(), static_cast(job.noise.size())); + if (bad >= 0) { + std::snprintf(buf, sizeof(buf), "noise[%d] = %g is not finite (NaN/Inf)", bad, job.noise[bad]); + return buf; + } + } + if (req.attention_mask_size() > 0) { + job.attn_mask.assign(req.attention_mask().begin(), req.attention_mask().end()); + } + return ""; +} + +// Sends envelope frames followed by body on sock. The socket thread owns the +// ROUTER socket; the worker reaches it through an inproc PUSH that the socket +// thread forwards, so both ends use this. +bool send_multipart(zmq::socket_t & sock, std::vector & envelope, + const std::string & body) { + try { + for (auto & frame : envelope) { + sock.send(frame, zmq::send_flags::sndmore); + } + sock.send(zmq::buffer(body), zmq::send_flags::none); + return true; + } catch (const zmq::error_t & e) { + std::fprintf(stderr, "vla-server: send failed: %s\n", e.what()); + return false; + } +} + +// Receives every frame of one multipart message. Returns false on no message. +bool recv_multipart(zmq::socket_t & sock, std::vector & frames) { + frames.clear(); + do { + zmq::message_t frame; + auto rr = sock.recv(frame, zmq::recv_flags::none); + if (!rr) + return !frames.empty(); + frames.push_back(std::move(frame)); + } while (sock.get(zmq::sockopt::rcvmore)); + return true; +} + +std::string envelope_key(const std::vector & envelope) { + std::string key; + for (const auto & f : envelope) { + key.append(static_cast(f.data()), f.size()); + key.push_back('\0'); + } + return key; +} + void usage(const char * prog) { std::fprintf(stderr, - "usage: %s [--bind ADDR] [--timing-detail none|phase] [--config PATH] " - "([] | -hf user/repo[:file.gguf|:tag])\n" + "usage: %s [--bind ADDR] [--queue latest|fifo] [--timing-detail none|phase] " + "[--config PATH] ([] | -hf user/repo[:file.gguf|:tag])\n" " ignored; every arch bundles its vision tower in the\n" " ckpt GGUF. Accepted so older command lines still work.\n" " -hf HuggingFace repo, user/repo[:file.gguf|:tag]; downloaded\n" @@ -188,7 +344,15 @@ void usage(const char * prog) { " SmolVLA .safetensors or .gguf, or any of the other\n" " supported architectures' .gguf; the architecture is\n" " auto-detected from the checkpoint.\n" - " --bind ADDR ZMQ bind address (default: tcp://*:5555)\n" + " --bind ADDR ZMQ bind address (default: tcp://*:5555). The socket is\n" + " a ROUTER: REQ clients get one reply per request as\n" + " before; DEALER clients may keep several in flight.\n" + " --queue MODE what to do with requests waiting for the predict\n" + " thread (default: latest)\n" + " 'latest': one pending request per client; a newer one\n" + " replaces it and the old one is answered with\n" + " error=\"superseded\"\n" + " 'fifo' : serve every request in arrival order\n" " --timing-detail LEVEL per-request timing breakdown (default: none)\n" " 'none' : single ms_inference\n" " 'phase' : ms_prefill + ms_denoise broken out\n" @@ -222,6 +386,7 @@ int main(int argc, char ** argv) { std::string hf_spec; std::string config_path; vla::TimingDetail timing_detail = vla::TimingDetail::NONE; + vla::serving::QueueMode queue_mode = vla::serving::QueueMode::LATEST; vla::Options opts; std::string opt_err; @@ -234,6 +399,13 @@ int main(int argc, char ** argv) { hf_spec = argv[++i]; } else if (a == "--config" && i+1 < argc) { config_path = argv[++i]; + } else if (a == "--queue" && i+1 < argc) { + const std::string v = argv[++i]; + if (!vla::serving::parse_queue_mode(v, queue_mode)) { + std::fprintf(stderr, "vla-server: bad --queue value '%s'\n", v.c_str()); + usage(argv[0]); + return 1; + } } else if (a == "--timing-detail" && i+1 < argc) { const std::string v = argv[++i]; if (v == "none") @@ -303,13 +475,18 @@ int main(int argc, char ** argv) { } const auto & cfg = vla::model_config(model); std::printf("vla-server: loaded. chunk_size=%lld action_dim=%lld " - "n_lang=%lld hidden=%lld expert_h=%lld timing_detail=%s\n", + "n_lang=%lld hidden=%lld expert_h=%lld timing_detail=%s queue=%s\n", (long long) cfg.n_suffix, (long long) cfg.max_action_dim, (long long) cfg.n_lang, (long long) cfg.hidden, (long long) cfg.expert_h, - timing_detail == vla::TimingDetail::PHASE ? "phase" : "none"); - - zmq::context_t zctx( 1); - zmq::socket_t sock(zctx, zmq::socket_type::rep); + timing_detail == vla::TimingDetail::PHASE ? "phase" : "none", + vla::serving::queue_mode_name(queue_mode)); + + zmq::context_t zctx(1); + // ROUTER rather than REP: replies carry the peer's identity, so the socket + // thread keeps receiving while the worker predicts, and a DEALER client may + // have more than one request in flight. A REQ client sees the same one + // request, one reply exchange as before. + zmq::socket_t sock(zctx, zmq::socket_type::router); sock.set(zmq::sockopt::linger, 0); // 64 MiB is above any real request (16 views of 512x512 F32 RGB is ~50 MiB) and // low enough to bound protobuf's expansion during ParseFromArray. @@ -320,6 +497,13 @@ int main(int argc, char ** argv) { std::fprintf(stderr, "vla-server: bind %s: %s\n", bind_addr.c_str(), e.what()); return 1; } + + // Replies travel worker -> socket thread over inproc; ZeroMQ sockets are not + // shareable between threads. + const char * reply_addr = "inproc://vla-server-replies"; + zmq::socket_t reply_pull(zctx, zmq::socket_type::pull); + reply_pull.bind(reply_addr); + std::printf("vla-server: bound to %s. ready.\n", bind_addr.c_str()); if (bind_addr.find("127.0.0.1") == std::string::npos && @@ -331,28 +515,103 @@ int main(int argc, char ** argv) { bind_addr.c_str()); } - // A failed reply desyncs the REP recv/send lockstep, so this socket cannot - // recover; log and shut down cleanly rather than crash (the old unguarded - // send) or spin forever on the wedged socket. - auto send_reply = [&sock](const std::string & body) { - try { - sock.send(zmq::buffer(body), zmq::send_flags::none); - } catch (const zmq::error_t & e) { - std::fprintf(stderr, "vla-server: reply send failed (%s); shutting down\n", e.what()); - g_shutdown.store(true, std::memory_order_relaxed); - } - }; - std::signal(SIGINT, on_signal); std::signal(SIGTERM, on_signal); - zmq::pollitem_t poll[] = {{ static_cast(sock), 0, ZMQ_POLLIN, 0 }}; + vla::serving::RequestQueue queue(queue_mode); + std::atomic served{0}; + std::atomic superseded{0}; - uint64_t served = 0; - while (!g_shutdown.load(std::memory_order_relaxed)) { + // Predict thread: the only caller of vla::predict, so the model needs no lock. + std::thread worker([&] { + zmq::socket_t reply_push(zctx, zmq::socket_type::push); + reply_push.set(zmq::sockopt::linger, 0); + reply_push.connect(reply_addr); + Job job; + while (queue.pop(job)) { + if (g_shutdown.load(std::memory_order_relaxed)) { + continue; // drain without predicting; the socket thread is gone + } + const auto t_start = Clock::now(); + const float ms_queue = std::chrono::duration(t_start - job.t_recv).count(); + + vla::Inputs in; + if (job.use_precomputed) { + in.precomputed_img_emb = job.precomputed_emb.data(); + in.n_img_views = job.precomputed_n_views; + in.images = nullptr; + in.n_images = 0; + } else { + in.images = job.img_views.data(); + in.n_images = static_cast(job.img_views.size()); + in.precomputed_img_emb = nullptr; + in.n_img_views = 0; + } + in.lang_tokens = job.lang_tokens.data(); + in.n_lang = static_cast(job.lang_tokens.size()); + in.state = job.state.data(); + in.noise = job.noise.empty() ? nullptr : job.noise.data(); + in.attention_mask = job.attn_mask.empty() ? nullptr : job.attn_mask.data(); + in.attention_mask_n = static_cast(job.attn_mask.size()); + in.timing_detail = timing_detail; + + std::vector action_chunk = vla::predict(model, in); + const auto & st = vla::last_stats(model); + + std::string body; + if (action_chunk.empty()) { + body = make_error_response(job.rid, "predict failed"); + } else { + vla::PredictResponse resp; + resp.set_request_id(job.rid); + resp.mutable_action_chunk()->Reserve(static_cast(action_chunk.size())); + for (float v : action_chunk) + resp.add_action_chunk(v); + resp.set_chunk_size(static_cast(cfg.n_suffix)); + resp.set_action_dim(static_cast(cfg.max_action_dim)); + resp.set_latency_ms_total(st.ms_total); + resp.set_latency_ms_vision(st.ms_vision); + resp.set_latency_ms_inference(st.ms_inference); + resp.set_latency_ms_prefill(st.ms_prefill); + resp.set_latency_ms_denoise(st.ms_denoise); + resp.set_latency_ms_queue(ms_queue); + body = resp.SerializeAsString(); + } + send_multipart(reply_push, job.envelope, body); + + const uint64_t n = ++served; + if (n%10 == 1) { + const float ms_other = std::max(0.f, st.ms_total-st.ms_vision-st.ms_inference); + if (timing_detail == vla::TimingDetail::PHASE) { + std::printf("vla-server: rid=%llu served=%llu superseded=%llu queue=%.1f ms " + "total=%.1f ms vision=%.1f inf=%.1f (prefill=%.1f + denoise=%.1f) other=%.1f\n", + (unsigned long long) job.rid, (unsigned long long) n, + (unsigned long long) superseded.load(), ms_queue, + st.ms_total, st.ms_vision, st.ms_inference, + st.ms_prefill, st.ms_denoise, ms_other); + } else { + std::printf("vla-server: rid=%llu served=%llu superseded=%llu queue=%.1f ms " + "total=%.1f ms vision=%.1f inf=%.1f other=%.1f\n", + (unsigned long long) job.rid, (unsigned long long) n, + (unsigned long long) superseded.load(), ms_queue, + st.ms_total, st.ms_vision, st.ms_inference, ms_other); + } + std::fflush(stdout); + } + } + }); + + zmq::pollitem_t poll[] = { + { static_cast(sock), 0, ZMQ_POLLIN, 0 }, + { static_cast(reply_pull), 0, ZMQ_POLLIN, 0 }, + }; + + // Socket thread: receive, validate, enqueue; forward the worker's replies. + std::vector frames; + while (!g_shutdown.load(std::memory_order_relaxed)) { try { - zmq::poll(poll, 1, std::chrono::milliseconds(200)); + zmq::poll(poll, 2, std::chrono::milliseconds(200)); } catch (const zmq::error_t & e) { if (e.num() == EINTR) continue; @@ -361,13 +620,32 @@ int main(int argc, char ** argv) { std::fprintf(stderr, "vla-server: zmq error: %s\n", e.what()); continue; } + + if (poll[1].revents & ZMQ_POLLIN) { + try { + // Forward every reply the worker has queued, not just one per poll. + while ((reply_pull.get(zmq::sockopt::events) & ZMQ_POLLIN) && + recv_multipart(reply_pull, frames)) { + zmq::message_t body = std::move(frames.back()); + frames.pop_back(); + // A peer that disconnected is unroutable; ROUTER drops the reply + // silently, which is what we want. + for (auto & f : frames) + sock.send(f, zmq::send_flags::sndmore); + sock.send(body, zmq::send_flags::none); + } + } catch (const zmq::error_t & e) { + if (e.num() == ETERM) + break; + std::fprintf(stderr, "vla-server: reply forward failed: %s\n", e.what()); + } + } + if (!(poll[0].revents & ZMQ_POLLIN)) continue; - zmq::message_t req_msg; try { - auto rr = sock.recv(req_msg, zmq::recv_flags::none); - if (!rr) + if (!recv_multipart(sock, frames)) continue; } catch (const zmq::error_t & e) { if (e.num() == EINTR) @@ -378,224 +656,52 @@ int main(int argc, char ** argv) { continue; } - // Without this an unauthenticated client shuts the server down with one - // two-frame request: the reply fails and send_reply sets g_shutdown. - if (drain_extra_frames(sock)) { - send_reply(make_error_response(0, "expected a single-frame request")); + // ROUTER prepends the peer identity; REQ adds an empty delimiter after it. + // Everything before the last frame is the envelope to echo. + if (frames.size() < 2) { + std::fprintf(stderr, "vla-server: dropped a request with no body\n"); continue; } + Job job; + zmq::message_t req_msg = std::move(frames.back()); + frames.pop_back(); + job.envelope = std::move(frames); + job.key = envelope_key(job.envelope); + job.t_recv = Clock::now(); + frames.clear(); vla::PredictRequest req; if (!req.ParseFromArray(req_msg.data(), static_cast(req_msg.size()))) { std::fprintf(stderr, "vla-server: PredictRequest parse failed (size=%zu)\n", req_msg.size()); - const std::string body = make_error_response(0, "request parse failed"); - send_reply(body); - continue; - } - - const uint64_t rid = req.request_id(); - - if (req.images_size() < 1 && req.precomputed_img_emb_size() == 0) { - const std::string body = make_error_response(rid, - "PredictRequest must contain images or precomputed_img_emb"); - send_reply(body); + send_multipart(sock, job.envelope, make_error_response(0, "request parse failed")); continue; } - if (req.images_size() > 16) { - send_reply(make_error_response(rid, "too many image views (max 16)")); - continue; - } - if (req.lang_tokens_size() < 1 || req.lang_tokens_size() > int(cfg.n_lang)) { - char buf[128]; std::snprintf(buf, sizeof(buf), - "lang_tokens length %d out of range [1, %lld]", - req.lang_tokens_size(), (long long) cfg.n_lang); - send_reply(make_error_response(rid, buf)); - continue; - } - { - bool tokens_ok = true; - for (int t=0; t 0; - std::vector precomputed_emb; - int precomputed_n_views = 0; - - const int n_views = req.images_size(); - std::vector> u8_bufs (n_views); - std::vector> f32_bufs(n_views); - std::vector img_views(n_views); - - if (use_precomputed) { - precomputed_n_views = static_cast(req.precomputed_img_emb_n_views()); - const int64_t per_view = cfg.n_img*cfg.hidden; - const int64_t expected = per_view * static_cast(precomputed_n_views); - if (precomputed_n_views < 1 || precomputed_n_views > 16) { - char buf[96]; std::snprintf(buf, sizeof(buf), - "precomputed_img_emb_n_views %d out of range [1, 16]", precomputed_n_views); - send_reply(make_error_response(rid, buf)); - continue; - } - if (static_cast(req.precomputed_img_emb_size()) != expected) { - char buf[160]; std::snprintf(buf, sizeof(buf), - "precomputed_img_emb size %d != %lld (n_views=%d * n_img_per_view=%lld * hidden=%lld)", - req.precomputed_img_emb_size(), (long long) expected, precomputed_n_views, - (long long) cfg.n_img, (long long) cfg.hidden); - send_reply(make_error_response(rid, buf)); - continue; - } - precomputed_emb.assign(req.precomputed_img_emb().begin(), - req.precomputed_img_emb().end()); - const int bad = find_non_finite(precomputed_emb.data(), - static_cast(precomputed_emb.size())); - if (bad >= 0) { - char buf[128]; std::snprintf(buf, sizeof(buf), - "precomputed_img_emb[%d] = %g is not finite (NaN/Inf)", - bad, precomputed_emb[bad]); - send_reply(make_error_response(rid, buf)); - continue; - } - } else { - - bool decode_ok = true; - size_t total_px = 0; - for (int v=0; v kMaxTotalPixels) { - send_reply(make_error_response(rid, "images exceed the per-request pixel budget")); - decode_ok = false; - break; - } - } - if (!decode_ok) - continue; - } - - std::vector lang_tokens(req.lang_tokens().begin(), req.lang_tokens().end()); - std::vector state_vec(req.state().begin(), req.state().end()); - { - const int bad = find_non_finite(state_vec.data(), static_cast(state_vec.size())); - if (bad >= 0) { - char buf[128]; std::snprintf(buf, sizeof(buf), - "state[%d] = %g is not finite (NaN/Inf)", bad, state_vec[bad]); - send_reply(make_error_response(rid, buf)); - continue; - } - } - std::vector noise_vec; - if (req.noise_size() == expected_noise_n) { - noise_vec.assign(req.noise().begin(), req.noise().end()); - const int bad = find_non_finite(noise_vec.data(), static_cast(noise_vec.size())); - if (bad >= 0) { - char buf[128]; std::snprintf(buf, sizeof(buf), - "noise[%d] = %g is not finite (NaN/Inf)", bad, noise_vec[bad]); - send_reply(make_error_response(rid, buf)); - continue; - } - } + job.rid = req.request_id(); - vla::Inputs in; - if (use_precomputed) { - in.precomputed_img_emb = precomputed_emb.data(); - in.n_img_views = precomputed_n_views; - - in.images = nullptr; - in.n_images = 0; - } else { - in.images = img_views.data(); - in.n_images = n_views; - in.precomputed_img_emb = nullptr; - in.n_img_views = 0; - } - std::vector attn_mask_vec; - if (req.attention_mask_size() > 0) { - attn_mask_vec.assign(req.attention_mask().begin(), req.attention_mask().end()); - } - in.lang_tokens = lang_tokens.data(); - in.n_lang = static_cast(lang_tokens.size()); - in.state = state_vec.data(); - in.noise = noise_vec.empty() ? nullptr : noise_vec.data(); - in.attention_mask = attn_mask_vec.empty() ? nullptr : attn_mask_vec.data(); - in.attention_mask_n = static_cast(attn_mask_vec.size()); - in.timing_detail = timing_detail; - - std::vector action_chunk = vla::predict(model, in); - const auto & st = vla::last_stats(model); - - if (action_chunk.empty()) { - send_reply(make_error_response(rid, "predict failed")); + const std::string err = build_job(req, cfg, job); + if (!err.empty()) { + send_multipart(sock, job.envelope, make_error_response(job.rid, err)); continue; } - vla::PredictResponse resp; - resp.set_request_id(rid); - resp.mutable_action_chunk()->Reserve(static_cast(action_chunk.size())); - for (float v : action_chunk) - resp.add_action_chunk(v); - resp.set_chunk_size(static_cast(cfg.n_suffix)); - resp.set_action_dim(static_cast(cfg.max_action_dim)); - resp.set_latency_ms_total(st.ms_total); - resp.set_latency_ms_vision(st.ms_vision); - resp.set_latency_ms_inference(st.ms_inference); - resp.set_latency_ms_prefill(st.ms_prefill); - resp.set_latency_ms_denoise(st.ms_denoise); - - const std::string body = resp.SerializeAsString(); - send_reply(body); - - ++served; - if (served%10 == 1) { - const float ms_other = std::max(0.f, - st.ms_total-st.ms_vision-st.ms_inference); - if (timing_detail == vla::TimingDetail::PHASE) { - std::printf("vla-server: rid=%llu served=%llu total=%.1f ms " - "vision=%.1f inf=%.1f (prefill=%.1f + denoise=%.1f) other=%.1f\n", - (unsigned long long) rid, (unsigned long long) served, - st.ms_total, st.ms_vision, st.ms_inference, - st.ms_prefill, st.ms_denoise, ms_other); - } else { - std::printf("vla-server: rid=%llu served=%llu total=%.1f ms " - "vision=%.1f inf=%.1f other=%.1f\n", - (unsigned long long) rid, (unsigned long long) served, - st.ms_total, st.ms_vision, st.ms_inference, ms_other); - } - std::fflush(stdout); + std::optional displaced; + const auto rc = queue.push(std::move(job), [](const Job & j) -> const std::string & { return j.key; }, + displaced); + if (rc == vla::serving::PushResult::REPLACED) { + ++superseded; + send_multipart(sock, displaced->envelope, make_error_response(displaced->rid, "superseded")); + } else if (rc == vla::serving::PushResult::REJECTED) { + send_multipart(sock, displaced->envelope, make_error_response(displaced->rid, "queue full")); } } - std::printf("vla-server: shutting down (served %llu requests)\n", - (unsigned long long) served); + std::printf("vla-server: shutting down (served %llu requests, superseded %llu)\n", + (unsigned long long) served.load(), (unsigned long long) superseded.load()); + g_shutdown.store(true, std::memory_order_relaxed); + queue.stop(); + worker.join(); + reply_pull.close(); sock.close(); zctx.close(); vla::model_free(model); diff --git a/src/serving/vla.proto b/src/serving/vla.proto index 60f64e4..6775cdd 100644 --- a/src/serving/vla.proto +++ b/src/serving/vla.proto @@ -45,6 +45,11 @@ message PredictResponse { float latency_ms_prefill = 8; float latency_ms_denoise = 9; float latency_ms_vision = 10; + // Time the request waited for the predict thread, from receipt to the start + // of prediction. Non-zero when another request was being served. + float latency_ms_queue = 11; + // Set when the request was not served. "superseded" means the same client + // sent a newer request before this one started (see vla-server --queue). string error = 7; } diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 545955f..39b810c 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -61,6 +61,15 @@ target_link_libraries(test_prompt PRIVATE vla_core) target_compile_options(test_prompt PRIVATE -Wall -Wextra) add_test(NAME prompt COMMAND test_prompt) +# vla-server's socket-thread / predict-thread hand-off. Header-only, so no +# ZeroMQ or protobuf is needed to exercise the coalescing rules. +add_executable(test_request_queue test_request_queue.cpp) +target_include_directories(test_request_queue PRIVATE ${CMAKE_SOURCE_DIR}/src) +target_compile_options(test_request_queue PRIVATE -Wall -Wextra) +find_package(Threads REQUIRED) +target_link_libraries(test_request_queue PRIVATE Threads::Threads) +add_test(NAME request_queue COMMAND test_request_queue) + # A/B harness for the two BitVLA ternary-GEMM tilings. Built so it cannot rot, # not registered with ctest: it needs a GPU and is read by hand. # VLA_BITVLA_NARROW_GEMM=1 selects the old one-tile-per-CTA kernel at runtime. diff --git a/tests/test_request_queue.cpp b/tests/test_request_queue.cpp new file mode 100644 index 0000000..6be742f --- /dev/null +++ b/tests/test_request_queue.cpp @@ -0,0 +1,147 @@ +// Copyright 2026 VinRobotics +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// The rules vla-server relies on: in LATEST mode a client's newer request +// displaces its pending one and nobody else's; FIFO keeps order and caps depth; +// stop() lets the worker drain what is queued and then returns false. + +#include "serving/request_queue.h" + +#undef NDEBUG // keep assert() live even in Release builds +#include +#include +#include +#include +#include + +using vla::serving::PushResult; +using vla::serving::QueueMode; +using vla::serving::RequestQueue; + +struct Job { + std::string client; + int rid = 0; +}; + +static const std::string & key_of(const Job & j) { return j.client; } + +static void test_latest_replaces_same_client_only() { + RequestQueue q(QueueMode::LATEST); + std::optional displaced; + + assert(q.push(Job{"a", 1}, key_of, displaced) == PushResult::QUEUED); + assert(!displaced); + assert(q.push(Job{"b", 2}, key_of, displaced) == PushResult::QUEUED); + assert(!displaced); + assert(q.size() == 2); + + // a's second request displaces a's first, and is served after b's. + assert(q.push(Job{"a", 3}, key_of, displaced) == PushResult::REPLACED); + assert(displaced && displaced->client == "a" && displaced->rid == 1); + assert(q.size() == 2); + + Job j; + assert(q.pop(j) && j.client == "b" && j.rid == 2); + assert(q.pop(j) && j.client == "a" && j.rid == 3); + assert(q.size() == 0); +} + +static void test_latest_does_not_touch_a_popped_job() { + // A job already handed to the worker is being predicted; a newer request + // from the same client queues behind it rather than cancelling it. + RequestQueue q(QueueMode::LATEST); + std::optional displaced; + Job j; + + q.push(Job{"a", 1}, key_of, displaced); + assert(q.pop(j) && j.rid == 1); + assert(q.push(Job{"a", 2}, key_of, displaced) == PushResult::QUEUED); + assert(!displaced); + assert(q.pop(j) && j.rid == 2); +} + +static void test_fifo_keeps_order_and_caps_depth() { + RequestQueue q(QueueMode::FIFO, /*max_depth=*/2); + std::optional displaced; + + assert(q.push(Job{"a", 1}, key_of, displaced) == PushResult::QUEUED); + assert(q.push(Job{"a", 2}, key_of, displaced) == PushResult::QUEUED); + assert(q.push(Job{"a", 3}, key_of, displaced) == PushResult::REJECTED); + // The rejected job is the new one, handed back so the caller can answer it. + assert(displaced && displaced->rid == 3); + assert(q.size() == 2); + + Job j; + assert(q.pop(j) && j.rid == 1); + assert(q.pop(j) && j.rid == 2); +} + +static void test_latest_caps_distinct_clients() { + RequestQueue q(QueueMode::LATEST, /*max_depth=*/2); + std::optional displaced; + + assert(q.push(Job{"a", 1}, key_of, displaced) == PushResult::QUEUED); + assert(q.push(Job{"b", 2}, key_of, displaced) == PushResult::QUEUED); + assert(q.push(Job{"c", 3}, key_of, displaced) == PushResult::REJECTED); + assert(displaced && displaced->client == "c"); + // A known client still gets to replace its own job when the queue is full. + assert(q.push(Job{"a", 4}, key_of, displaced) == PushResult::REPLACED); + assert(displaced && displaced->rid == 1); +} + +static void test_stop_drains_then_returns_false() { + RequestQueue q(QueueMode::LATEST); + std::optional displaced; + q.push(Job{"a", 1}, key_of, displaced); + q.push(Job{"b", 2}, key_of, displaced); + q.stop(); + + Job j; + assert(q.pop(j) && j.rid == 1); + assert(q.pop(j) && j.rid == 2); + assert(!q.pop(j)); + assert(!q.pop(j)); +} + +static void test_pop_blocks_until_push_or_stop() { + RequestQueue q(QueueMode::LATEST); + std::vector seen; + + std::thread worker([&] { + Job j; + while (q.pop(j)) + seen.push_back(j.rid); + }); + + std::optional displaced; + for (int i=1; i<=5; ++i) + q.push(Job{"a" + std::to_string(i), i}, key_of, displaced); + q.stop(); + worker.join(); + + assert(seen.size() == 5); + for (int i=0; i<5; ++i) + assert(seen[i] == i+1); +} + +int main() { + test_latest_replaces_same_client_only(); + test_latest_does_not_touch_a_popped_job(); + test_fifo_keeps_order_and_caps_depth(); + test_latest_caps_distinct_clients(); + test_stop_drains_then_returns_false(); + test_pop_blocks_until_push_or_stop(); + std::printf("test_request_queue: OK\n"); + return 0; +}