diff --git a/Cargo.lock b/Cargo.lock index dcba5bb4..2bb95d2f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,12 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "adler2" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" + [[package]] name = "ahash" version = "0.8.12" @@ -43,6 +49,12 @@ version = "1.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" +[[package]] +name = "autocfg" +version = "1.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" + [[package]] name = "axum" version = "0.8.9" @@ -149,6 +161,18 @@ version = "3.20.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649" +[[package]] +name = "bytemuck" +version = "1.25.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "95832e849adfb21180ccb6826a99da14e5d266ae5c2e668e1602cf234f153797" + +[[package]] +name = "byteorder-lite" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f1fe948ff07f4bd06c30984e69f5b4899c516a3ef74f34df92a2df2ab535495" + [[package]] name = "bytes" version = "1.12.1" @@ -230,6 +254,15 @@ dependencies = [ "libc", ] +[[package]] +name = "crc32fast" +version = "1.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "01a7799fd6b852db0e61728dde9a204c423b44d689dbd432522543614b490e78" +dependencies = [ + "cfg-if", +] + [[package]] name = "crossbeam-deque" version = "0.8.8" @@ -418,12 +451,32 @@ version = "2.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223" +[[package]] +name = "fdeflate" +version = "0.3.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e6853b52649d4ac5c0bd02320cddc5ba956bdb407c4b75a2c6b75bf51500f8c" +dependencies = [ + "simd-adler32", +] + [[package]] name = "find-msvc-tools" version = "0.1.13" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ef25905e51abafe4dcea6c15fec58c57b601cdbd0ee53d22ea1d3016c587d39b" +[[package]] +name = "flate2" +version = "1.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e634e2e0ebac1ee034020da1ca582e17ffe4e0f5e985823721e168928136dcb" +dependencies = [ + "crc32fast", + "miniz_oxide 0.9.1", + "zlib-rs", +] + [[package]] name = "fnv" version = "1.0.7" @@ -801,6 +854,21 @@ dependencies = [ "icu_properties", ] +[[package]] +name = "image" +version = "0.25.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85ab80394333c02fe689eaf900ab500fbd0c2213da414687ebf995a65d5a6104" +dependencies = [ + "bytemuck", + "byteorder-lite", + "moxcms", + "num-traits", + "png", + "zune-core", + "zune-jpeg", +] + [[package]] name = "indexmap" version = "2.14.2" @@ -932,6 +1000,26 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" +[[package]] +name = "miniz_oxide" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fa76a2c86f704bdb222d66965fb3d63269ce38518b83cb0575fca855ebb6316" +dependencies = [ + "adler2", + "simd-adler32", +] + +[[package]] +name = "miniz_oxide" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b63fbc4a50860e98e7b2aa7804ded1db5cbc3aff9193adaff57a6931bf7c4b4c" +dependencies = [ + "adler2", + "simd-adler32", +] + [[package]] name = "mio" version = "1.2.3" @@ -965,6 +1053,16 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "moxcms" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bb85c154ba489f01b25c0d36ae69a87e4a1c73a72631fc6c0eb6dde34a73e44b" +dependencies = [ + "num-traits", + "pxfm", +] + [[package]] name = "nom" version = "7.1.3" @@ -975,6 +1073,15 @@ dependencies = [ "minimal-lexical", ] +[[package]] +name = "num-traits" +version = "0.2.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" +dependencies = [ + "autocfg", +] + [[package]] name = "omni-clm" version = "0.1.0" @@ -994,12 +1101,16 @@ version = "0.1.0" dependencies = [ "anyhow", "axum", + "base64 0.22.1", "half", + "image", "memmap2", "omni-qwen3-5-native", "omni-runtime", "safetensors 0.8.0", + "serde", "serde_json", + "sha2", "tokenizers 0.22.2", "tokio", ] @@ -1133,6 +1244,19 @@ version = "0.3.34" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f6b464fbc74e149a392436b17d523f769e057cb6877f6a5c4618bc6f11800548" +[[package]] +name = "png" +version = "0.18.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "60769b8b31b2a9f263dae2776c37b1b28ae246943cf719eb6946a1db05128a61" +dependencies = [ + "bitflags", + "crc32fast", + "fdeflate", + "flate2", + "miniz_oxide 0.8.9", +] + [[package]] name = "potential_utf" version = "0.1.6" @@ -1160,6 +1284,12 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "pxfm" +version = "0.1.30" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d55d956fa96f5ec02be2e13af0e20391a5aa83d6a074e3ad368959d0fab299ea" + [[package]] name = "quinn" version = "0.11.12" @@ -1591,6 +1721,12 @@ dependencies = [ "libc", ] +[[package]] +name = "simd-adler32" +version = "0.3.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3a219298ac11a56ea9a6d2120044824d6f01aeb034955e7af7bc16858527deea" + [[package]] name = "slab" version = "0.4.12" @@ -2316,8 +2452,29 @@ dependencies = [ "syn 3.0.6", ] +[[package]] +name = "zlib-rs" +version = "0.6.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b268e58e7c693d7c271f93ffc4ba3b380412554231c85bf61ca7af91042a4112" + [[package]] name = "zmij" version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b" + +[[package]] +name = "zune-core" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d56377fd46368984a170bc5aac5567e52ca5da874caa60bea39fcbca78fb658b" + +[[package]] +name = "zune-jpeg" +version = "0.5.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "27bc9d5b815bc103f142aa054f561d9187d191692ec7c2d1e2b4737f8dbd7296" +dependencies = [ + "zune-core", +] diff --git a/recipe/cua_s1/export_multimodal_language.py b/recipe/cua_s1/export_multimodal_language.py new file mode 100644 index 00000000..1866003e --- /dev/null +++ b/recipe/cua_s1/export_multimodal_language.py @@ -0,0 +1,57 @@ +"""Export the pinned BF16 language model with the multimodal LoRA merged. + +Vision execution keeps its original BF16 base and FP32 adapters separately. +This one-time export requires the pinned Transformers/PEFT reference environment. +""" + +import argparse +import hashlib +import json +from pathlib import Path + +from models.cua_s1.multimodal.model import ( + ADAPTER_REVISION, + BASE_REVISION, + MultimodalEngine, +) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base", required=True) + parser.add_argument("--adapter", required=True) + parser.add_argument("--out", required=True, type=Path) + args = parser.parse_args() + if args.out.exists(): + parser.error("output already exists") + engine = MultimodalEngine(args.base, args.adapter) + model = engine.model.merge_and_unload() + model.model.language_model.save_pretrained(args.out, max_shard_size="5GB") + model.config.to_json_file(args.out / "config.json") + engine.tokenizer.save_pretrained(args.out) + (args.out / "cua_s1_language_export.json").write_text( + json.dumps( + { + "format": "cua-s1-multimodal-language-merged/1", + "base_revision": BASE_REVISION, + "adapter_revision": ADAPTER_REVISION, + "vision": "separate unmerged base and adapter", + "files": { + p.name: { + "size": p.stat().st_size, + "sha256": hashlib.file_digest( + p.open("rb"), "sha256" + ).hexdigest(), + } + for p in args.out.iterdir() + if p.is_file() + }, + }, + indent=2, + ) + + "\n" + ) + + +if __name__ == "__main__": + main() diff --git a/src/backends/cuda/qwen3_5/gemm.cu b/src/backends/cuda/qwen3_5/gemm.cu index 7de47f57..154ad698 100644 --- a/src/backends/cuda/qwen3_5/gemm.cu +++ b/src/backends/cuda/qwen3_5/gemm.cu @@ -22,7 +22,7 @@ struct Plan { cublasLtMatmulAlgo_t algo{}; }; -using Key = std::tuple; // M, N, K, ldy +using Key = std::tuple; // M, N, K, ldy, FP32, bias struct Gemm { cublasLtHandle_t handle = nullptr; @@ -41,26 +41,32 @@ void destroy(Plan& p) { p = Plan{}; } -int describe(int M, int N, int K, int ldy, Plan& p) { +int describe(int M, int N, int K, int ldy, Plan& p, bool fp32, bool bias) { cublasStatus_t s = cublasLtMatmulDescCreate(&p.op, CUBLAS_COMPUTE_32F, CUDA_R_32F); if (s != CUBLAS_STATUS_SUCCESS) return status(s); const cublasOperation_t ta = CUBLAS_OP_T, tb = CUBLAS_OP_N; cublasLtMatmulDescSetAttribute(p.op, CUBLASLT_MATMUL_DESC_TRANSA, &ta, sizeof(ta)); cublasLtMatmulDescSetAttribute(p.op, CUBLASLT_MATMUL_DESC_TRANSB, &tb, sizeof(tb)); - if ((s = cublasLtMatrixLayoutCreate(&p.a, CUDA_R_16BF, K, N, K)) != CUBLAS_STATUS_SUCCESS) return status(s); - if ((s = cublasLtMatrixLayoutCreate(&p.b, CUDA_R_16BF, K, M, K)) != CUBLAS_STATUS_SUCCESS) return status(s); - if ((s = cublasLtMatrixLayoutCreate(&p.c, CUDA_R_16BF, N, M, ldy)) != CUBLAS_STATUS_SUCCESS) return status(s); + if (bias) { + cublasLtEpilogue_t epilogue = CUBLASLT_EPILOGUE_BIAS; + if ((s = cublasLtMatmulDescSetAttribute(p.op, CUBLASLT_MATMUL_DESC_EPILOGUE, &epilogue, sizeof(epilogue))) != CUBLAS_STATUS_SUCCESS) return status(s); + } + const cudaDataType_t dtype = fp32 ? CUDA_R_32F : CUDA_R_16BF; + if ((s = cublasLtMatrixLayoutCreate(&p.a, dtype, K, N, K)) != CUBLAS_STATUS_SUCCESS) return status(s); + if ((s = cublasLtMatrixLayoutCreate(&p.b, dtype, K, M, K)) != CUBLAS_STATUS_SUCCESS) return status(s); + if ((s = cublasLtMatrixLayoutCreate(&p.c, dtype, N, M, ldy)) != CUBLAS_STATUS_SUCCESS) return status(s); return 0; } -// The heuristic's first choice, without in-place split-K reductions. -int first_choice(Gemm& g, Plan& p) { +// The heuristic's first choice. Vision disables all split-K to avoid BF16 +// intermediate reductions; existing language GEMMs exclude only in-place reductions. +int first_choice(Gemm& g, Plan& p, bool vision) { cublasLtMatmulPreference_t pref; cublasStatus_t s = cublasLtMatmulPreferenceCreate(&pref); if (s != CUBLAS_STATUS_SUCCESS) return status(s); cublasLtMatmulPreferenceSetAttribute(pref, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &g.workspace_bytes, sizeof(g.workspace_bytes)); - const uint32_t schemes = CUBLASLT_REDUCTION_SCHEME_MASK & ~CUBLASLT_REDUCTION_SCHEME_INPLACE; + const uint32_t schemes = vision ? CUBLASLT_REDUCTION_SCHEME_NONE : (CUBLASLT_REDUCTION_SCHEME_MASK & ~CUBLASLT_REDUCTION_SCHEME_INPLACE); cublasLtMatmulPreferenceSetAttribute(pref, CUBLASLT_MATMUL_PREF_REDUCTION_SCHEME_MASK, &schemes, sizeof(schemes)); cublasLtMatmulHeuristicResult_t r{}; @@ -74,16 +80,16 @@ int first_choice(Gemm& g, Plan& p) { } // The plan for a shape, created on first use. -int plan_for(Gemm& g, int M, int N, int K, int ldy, Plan*& out) { - const Key key{M, N, K, ldy}; +int plan_for(Gemm& g, int M, int N, int K, int ldy, Plan*& out, bool fp32 = false, bool bias = false) { + const Key key{M, N, K, ldy, fp32, bias}; auto it = g.plans.find(key); if (it != g.plans.end()) { out = &it->second; return 0; } Plan p; - int rc = describe(M, N, K, ldy, p); - if (rc == 0) rc = first_choice(g, p); + int rc = describe(M, N, K, ldy, p, fp32, bias); + if (rc == 0) rc = first_choice(g, p, fp32 || bias); if (rc != 0) { destroy(p); return rc; @@ -127,3 +133,31 @@ extern "C" int cs1_gemm(void* gemm, const void* x, const void* w, void* y, int M return status(cublasLtMatmul(g->handle, p->op, &alpha, w, p->a, x, p->b, &beta, y, p->c, y, p->c, &p->algo, g->workspace, g->workspace_bytes, static_cast(stream))); } + +// Vision's biased BF16 linears round only after adding bias. The patch +// convolution passes zero bias here and applies its bias after BF16 rounding. +extern "C" int cs1_vision_linear(void* gemm, const void* x, const void* w, const void* bias, + void* y, int M, int N, int K, void* stream) { + Gemm* g = static_cast(gemm); + if (!g || !bias || M <= 0 || N <= 0 || K <= 0) return cudaErrorInvalidValue; + Plan* p = nullptr; + int rc = plan_for(*g, M, N, K, N, p, false, true); + if (rc) return rc; + auto s = cublasLtMatmulDescSetAttribute(p->op, CUBLASLT_MATMUL_DESC_BIAS_POINTER, &bias, sizeof(bias)); + if (s != CUBLAS_STATUS_SUCCESS) return status(s); + const float alpha = 1.f, beta = 0.f; + return status(cublasLtMatmul(g->handle, p->op, &alpha, w, p->a, x, p->b, &beta, y, p->c, y, p->c, &p->algo, + g->workspace, g->workspace_bytes, static_cast(stream))); +} +// Separate unmerged LoRA matrices and intermediates stay FP32. No TF32 fast compute. +extern "C" int cs1_gemm_f32(void* gemm, const float* x, const float* w, float* y, + int M, int N, int K, void* stream) { + Gemm* g = static_cast(gemm); + if (!g || M <= 0 || N <= 0 || K <= 0) return cudaErrorInvalidValue; + Plan* p = nullptr; + int rc = plan_for(*g, M, N, K, N, p, true, false); + if (rc) return rc; + const float alpha = 1.f, beta = 0.f; + return status(cublasLtMatmul(g->handle, p->op, &alpha, w, p->a, x, p->b, &beta, y, p->c, y, p->c, &p->algo, + g->workspace, g->workspace_bytes, static_cast(stream))); +} diff --git a/src/backends/cuda/qwen3_5/ops.h b/src/backends/cuda/qwen3_5/ops.h index ae5ce4db..37a08a77 100644 --- a/src/backends/cuda/qwen3_5/ops.h +++ b/src/backends/cuda/qwen3_5/ops.h @@ -13,7 +13,7 @@ #include // Bumped whenever the required interface below changes. -#define CS1_ABI_VERSION 4 +#define CS1_ABI_VERSION 5 #ifdef __cplusplus extern "C" { @@ -28,6 +28,7 @@ int cs1_malloc(void** ptr, size_t bytes); int cs1_free(void* ptr); int cs1_stream_create(void** stream); int cs1_stream_sync(void* stream); +int cs1_stream_destroy(void* stream); int cs1_graph_begin(void* stream); int cs1_graph_end(void* stream, void** exec); int cs1_graph_launch(void* exec, void* stream); @@ -100,6 +101,21 @@ void cs1_gemm_destroy(void* gemm); int cs1_gemm(void* gemm, const void* x, const void* w, void* y, int M, int N, int K, int ldy, void* stream); +// ---- vision: single image, 1024 hidden, 16 heads of 64; all BF16 except explicit float pointers ---- +int cs1_vision_linear(void* gemm, const void* x, const void* w, const void* bias, void* y, int M, int N, int K, void* stream); +int cs1_gemm_f32(void* gemm, const float* x, const float* w, float* y, int M, int N, int K, void* stream); +int cs1_vision_norm(const void* x, const void* w, const void* b, void* y, int rows, int d, void* stream); +// indices/weights [N,4], 48x48 learned table; FP32 rotary cos/sin [N,32]. +int cs1_vision_position(void* x, const void* table, const int* indices, const float* weights, int n, void* stream); +int cs1_vision_rope(const void* qkv, const float* co, const float* si, void* q, void* k, int n, void* stream); +// q/k [N,1024], V is a slice in qkv [N,3072]. No causal mask, O(N) memory. +int cs1_vision_attention(const void* q, const void* k, const void* v, void* out, int n, void* stream); +int cs1_vision_bias(void* x, const void* bias, size_t n, int d, void* stream); +int cs1_vision_gelu(void* x, size_t n, int exact, void* stream); +int cs1_vision_add(void* x, const void* delta, size_t n, void* stream); +int cs1_vision_to_float(const void* x, float* out, size_t n, void* stream); +int cs1_vision_lora_add(void* x, const float* delta, size_t n, float scale, void* stream); + #ifdef __cplusplus } #endif diff --git a/src/backends/cuda/qwen3_5/runtime.cu b/src/backends/cuda/qwen3_5/runtime.cu index 02a7db6f..697c9db1 100644 --- a/src/backends/cuda/qwen3_5/runtime.cu +++ b/src/backends/cuda/qwen3_5/runtime.cu @@ -20,6 +20,8 @@ int cs1_stream_create(void** stream) { return cudaStreamCreateWithFlags(reinterpret_cast(stream), cudaStreamNonBlocking); } +int cs1_stream_destroy(void* stream) { return cudaStreamDestroy(static_cast(stream)); } + int cs1_stream_sync(void* stream) { return cudaStreamSynchronize(static_cast(stream)); } int cs1_upload(void* dst, const void* src, size_t bytes, void* stream) { diff --git a/src/backends/cuda/qwen3_5/vision.cu b/src/backends/cuda/qwen3_5/vision.cu new file mode 100644 index 00000000..d77701a5 --- /dev/null +++ b/src/backends/cuda/qwen3_5/vision.cu @@ -0,0 +1,258 @@ +// Native Qwen3.5 vision kernels. BF16 activations, FP32 normalization/rotary/LoRA. +#include "common.cuh" +#include "mma.cuh" +#include "ops.h" +namespace cs1 { namespace vision { +namespace flash { + +constexpr int D = 64, BM = 64, BN = 32, THREADS = 128; +constexpr int LDS = D + 8; // shared row stride in elements: 528 bytes keeps ldmatrix conflict-free +constexpr int SMEM_BYTES = (BM + 2 * BN) * LDS * 2; + +__global__ void __launch_bounds__(THREADS) + flash_kernel(const bf16* __restrict__ q, const bf16* __restrict__ k, const bf16* __restrict__ v, int ldv, + bf16* __restrict__ out, int T, int Hq, int Hk, float scale_log2) { + extern __shared__ __align__(16) unsigned char smem[]; + bf16* qs = reinterpret_cast(smem); + bf16* ks = qs + BM * LDS; + bf16* vs = ks + BN * LDS; + const int h = blockIdx.y, hk = h / (Hq / Hk); + const int q0 = (gridDim.x - 1 - blockIdx.x) * BM; // the longest blocks first + const int tid = threadIdx.x, warp = tid / 32, lane = tid % 32; + const int g = lane / 4, t = lane % 4; + const int row0 = q0 + warp * 16; // this warp's first query + + for (int c = tid; c < BM * (D / 8); c += THREADS) { + const int r = c / (D / 8), col = (c % (D / 8)) * 8, row = q0 + r; + cp_async16(qs + r * LDS + col, q + ((size_t)min(row, T - 1) * Hq + h) * D + col, row < T); + } + cp_async_commit(); + + float o[D / 8][4]; +#pragma unroll + for (int n = 0; n < D / 8; n++) o[n][0] = o[n][1] = o[n][2] = o[n][3] = 0.f; + float m[2] = {-INFINITY, -INFINITY}, l[2] = {0.f, 0.f}; + + const int kv_end = T; + for (int k0 = 0; k0 < kv_end; k0 += BN) { + for (int c = tid; c < BN * (D / 8); c += THREADS) { + const int r = c / (D / 8), col = (c % (D / 8)) * 8, s = k0 + r; + cp_async16(ks + r * LDS + col, k + ((size_t)min(s, T - 1) * Hk + hk) * D + col, s < T); + } + cp_async_commit(); + for (int c = tid; c < BN * (D / 8); c += THREADS) { + const int r = c / (D / 8), col = (c % (D / 8)) * 8, s = k0 + r; + cp_async16(vs + r * LDS + col, v + (size_t)min(s, T - 1) * ldv + (size_t)hk * D + col, s < T); + } + cp_async_commit(); + cp_async_wait<1>(); // Q and K + __syncthreads(); + + // keys past every query of this warp contribute nothing + const bool active = true; + float sc[BN / 8][4]; +#pragma unroll + for (int n = 0; n < BN / 8; n++) sc[n][0] = sc[n][1] = sc[n][2] = sc[n][3] = 0.f; + if (active) { +#pragma unroll + for (int kk = 0; kk < D; kk += 16) { + uint32_t a[4]; + load_a(a, qs, LDS, warp * 16, kk, lane); +#pragma unroll + for (int n = 0; n < BN / 8; n += 2) { + uint32_t b[4]; + load_b_nk(b, ks, LDS, kk, n * 8, lane); + mma16816(sc[n], a, b[0], b[1]); + mma16816(sc[n + 1], a, b[2], b[3]); + } + } + } + uint32_t p[BN / 16][4]; + if (active) { + // causal and length mask, then the online softmax in base 2 + float mx[2] = {-INFINITY, -INFINITY}; +#pragma unroll + for (int n = 0; n < BN / 8; n++) { +#pragma unroll + for (int e = 0; e < 4; e++) { + const int key = k0 + n * 8 + 2 * t + (e & 1); + sc[n][e] = (key < T) ? sc[n][e] * scale_log2 : -INFINITY; + mx[e >> 1] = fmaxf(mx[e >> 1], sc[n][e]); + } + } + float alpha[2], base[2]; +#pragma unroll + for (int r = 0; r < 2; r++) { + mx[r] = fmaxf(mx[r], __shfl_xor_sync(0xffffffffu, mx[r], 1)); + mx[r] = fmaxf(mx[r], __shfl_xor_sync(0xffffffffu, mx[r], 2)); + const float mn = fmaxf(m[r], mx[r]); + base[r] = mn == -INFINITY ? 0.f : mn; + alpha[r] = exp2f(m[r] - base[r]); + m[r] = mn; + l[r] *= alpha[r]; + } +#pragma unroll + for (int n = 0; n < BN / 8; n++) { +#pragma unroll + for (int e = 0; e < 4; e++) { + sc[n][e] = exp2f(sc[n][e] - base[e >> 1]); + l[e >> 1] += sc[n][e]; + } + } +#pragma unroll + for (int n = 0; n < D / 8; n++) { + o[n][0] *= alpha[0]; + o[n][1] *= alpha[0]; + o[n][2] *= alpha[1]; + o[n][3] *= alpha[1]; + } + // the score accumulators, two 8-key tiles at a time, are the A fragments of P*V +#pragma unroll + for (int j = 0; j < BN / 16; j++) { + p[j][0] = pack_bf16(sc[2 * j][0], sc[2 * j][1]); + p[j][1] = pack_bf16(sc[2 * j][2], sc[2 * j][3]); + p[j][2] = pack_bf16(sc[2 * j + 1][0], sc[2 * j + 1][1]); + p[j][3] = pack_bf16(sc[2 * j + 1][2], sc[2 * j + 1][3]); + } + } + cp_async_wait<0>(); // V + __syncthreads(); + if (active) { +#pragma unroll + for (int j = 0; j < BN / 16; j++) { +#pragma unroll + for (int n = 0; n < D / 8; n += 2) { + uint32_t b[4]; + load_b_kn(b, vs, LDS, j * 16, n * 8, lane); + mma16816(o[n], p[j], b[0], b[1]); + mma16816(o[n + 1], p[j], b[2], b[3]); + } + } + } + __syncthreads(); // before the next tile overwrites K and V + } + + // the four lanes of a row each summed a quarter of its keys +#pragma unroll + for (int r = 0; r < 2; r++) { + l[r] += __shfl_xor_sync(0xffffffffu, l[r], 1); + l[r] += __shfl_xor_sync(0xffffffffu, l[r], 2); + } + const float inv[2] = {1.f / l[0], 1.f / l[1]}; +#pragma unroll + for (int r = 0; r < 2; r++) { + const int row = row0 + g + r * 8; + if (row >= T) continue; + bf16* dst = out + ((size_t)row * Hq + h) * D + 2 * t; +#pragma unroll + for (int n = 0; n < D / 8; n++) + *reinterpret_cast(dst + n * 8) = pack_bf16(o[n][2 * r] * inv[r], o[n][2 * r + 1] * inv[r]); + } +} + +} // namespace flash +} } +namespace cs1 { namespace vision { +__global__ void norm_kernel(const bf16* x, const bf16* w, const bf16* b, bf16* y, int d) { + __shared__ float scratch[32]; + const size_t off = (size_t)blockIdx.x*d; + float sum = 0.f; + for (int i=threadIdx.x;i=n) return; + const int t=i/1024,d=i%1024; + float pos=0.f; + // Separate FP32 multiply and sum, then position BF16 rounding before the residual add. + for(int j=0;j<4;j++) pos=__fadd_rn(pos,__fmul_rn(f32(table[indices[t*4+j]*1024+d]),weights[t*4+j])); + x[i]=to_bf16(f32(x[i])+round_bf16(pos)); +} +__global__ void rope_kernel(const bf16* qkv,const float* co,const float* si,bf16* q,bf16* k,size_t n) { + size_t i=(size_t)blockIdx.x*blockDim.x+threadIdx.x; + if(i>=n) return; + int t=i/1024,d=i%64,channel=i%1024; + const float c=co[t*32+d%32],s=si[t*32+d%32]; + int partner=channel+(d<32?32:-32); + float sign=d<32?-1.f:1.f; + // PyTorch materializes both products in float32 (not a fused multiply-add). + q[i]=to_bf16(__fadd_rn(__fmul_rn(f32(qkv[t*3072+channel]),c),__fmul_rn(sign*f32(qkv[t*3072+partner]),s))); + k[i]=to_bf16(__fadd_rn(__fmul_rn(f32(qkv[t*3072+1024+channel]),c),__fmul_rn(sign*f32(qkv[t*3072+1024+partner]),s))); +} +__global__ void gelu_kernel(bf16* x,size_t n,int exact) { + size_t i=(size_t)blockIdx.x*blockDim.x+threadIdx.x; + if(i>=n) return; + float a=f32(x[i]); + float v=exact?0.5f*a*(1.f+erff(a*0.7071067811865475244f)): + 0.5f*a*(1.f+tanhf(0.7978845608028654f*(a+0.044715f*a*a*a))); + x[i]=to_bf16(v); +} +__global__ void add_kernel(bf16* x,const bf16* delta,size_t n) { + size_t i=(size_t)blockIdx.x*blockDim.x+threadIdx.x; + if(i>>((const bf16*)x,(const bf16*)w,(const bf16*)b,(bf16*)y,d); + return cudaGetLastError(); +} +extern "C" int cs1_vision_position(void* x,const void* table,const int* indices,const float* weights,int n,void* stream) { + if(n<=0) return cudaErrorInvalidValue; + vision::position_kernel<<<(n*1024+255)/256,256,0,(cudaStream_t)stream>>>((bf16*)x,(const bf16*)table,indices,weights,(size_t)n*1024); + return cudaGetLastError(); +} +extern "C" int cs1_vision_rope(const void* qkv,const float* co,const float* si,void* q,void* k,int n,void* stream) { + if(n<=0) return cudaErrorInvalidValue; + vision::rope_kernel<<<(n*1024+255)/256,256,0,(cudaStream_t)stream>>>((const bf16*)qkv,co,si,(bf16*)q,(bf16*)k,(size_t)n*1024); + return cudaGetLastError(); +} +extern "C" int cs1_vision_attention(const void* q,const void* k,const void* v,void* out,int n,void* stream) { + if(n<=0) return cudaErrorInvalidValue; + namespace f=vision::flash; + f::flash_kernel<<>>( + (const bf16*)q,(const bf16*)k,(const bf16*)v,3072,(bf16*)out,n,16,16,0.125f*1.4426950408889634f); + return cudaGetLastError(); +} +extern "C" int cs1_vision_gelu(void* x,size_t n,int exact,void* stream) { + if(n==0) return cudaSuccess; + vision::gelu_kernel<<<(n+255)/256,256,0,(cudaStream_t)stream>>>((bf16*)x,n,exact); return cudaGetLastError(); +} +extern "C" int cs1_vision_add(void* x,const void* delta,size_t n,void* stream) { + if(n==0) return cudaSuccess; + vision::add_kernel<<<(n+255)/256,256,0,(cudaStream_t)stream>>>((bf16*)x,(const bf16*)delta,n); return cudaGetLastError(); +} +extern "C" int cs1_vision_to_float(const void* x,float* out,size_t n,void* stream) { + if(n==0) return cudaSuccess; + vision::to_float_kernel<<<(n+255)/256,256,0,(cudaStream_t)stream>>>((const bf16*)x,out,n); return cudaGetLastError(); +} +extern "C" int cs1_vision_lora_add(void* x,const float* delta,size_t n,float scale,void* stream) { + if(n==0) return cudaSuccess; + vision::lora_add_kernel<<<(n+255)/256,256,0,(cudaStream_t)stream>>>((bf16*)x,delta,n,scale); return cudaGetLastError(); +} + +// cuDNN Conv3d rounds its convolution output before its separate bias addition. +extern "C" int cs1_vision_bias(void* x,const void* bias,size_t n,int d,void* stream) { + if(d<=0) return cudaErrorInvalidValue; + if(n==0) return cudaSuccess; + vision::bias_kernel<<<(n+255)/256,256,0,(cudaStream_t)stream>>>((bf16*)x,(const bf16*)bias,n,d); return cudaGetLastError(); +} diff --git a/src/models/cua_s1/native/Cargo.toml b/src/models/cua_s1/native/Cargo.toml index a8eb3ae1..80339951 100644 --- a/src/models/cua_s1/native/Cargo.toml +++ b/src/models/cua_s1/native/Cargo.toml @@ -11,14 +11,22 @@ path = "src/main.rs" [dependencies] anyhow = "1.0.100" +sha2 = "0.10" axum = "0.8.8" +base64 = "0.22.1" half = "2.7.1" +image = { version = "0.25", default-features = false, features = ["png", "jpeg"] } # the CUDA kernels live in libqwen3_5_cuda.so, loaded at run time omni-qwen3-5-native = { path = "../../qwen3_5/native" } omni-runtime = { path = "../../../runtime" } memmap2 = "0.9.9" safetensors = "0.8.0" +serde = { version = "1", features = ["derive"] } serde_json = { version = "1.0.149", features = ["float_roundtrip", "preserve_order"] } # the onig regex backend, as in the Python tokenizers wheel tokenizers = { version = "=0.22.2", default-features = false, features = ["onig"] } tokio = { version = "1.49.0", features = ["macros", "net", "rt-multi-thread", "sync"] } + +[[bin]] +name = "omni-cua-s1-vision" +path = "src/vision_main.rs" diff --git a/src/models/cua_s1/native/THIRD_PARTY_NOTICES.md b/src/models/cua_s1/native/THIRD_PARTY_NOTICES.md new file mode 100644 index 00000000..7d0ec38b --- /dev/null +++ b/src/models/cua_s1/native/THIRD_PARTY_NOTICES.md @@ -0,0 +1,30 @@ +# Image preprocessing attributions + +`src/image_preprocess.rs` is a Rust adaptation of the following algorithms. +It is modified for decoded interleaved RGB8 input, fixed Cua-S1 4B settings, +bounded dimensions, standard-library buffers, and a standalone CPU API. + +- PyTorch 2.14.0, `aten/src/ATen/native/cpu/UpSampleKernel.cpp` + (`_compute_indices_min_size_weights_aa`, `_compute_index_ranges_int16_weights`, + and the separable uint8 horizontal/vertical loops), plus the cubic polynomial + helpers in `aten/src/ATen/native/UpSample.h`. See [PyTorch license](licenses/PYTORCH-LICENSE) + for the retained copyright notices, redistribution conditions, and disclaimer. +- PyTorch's bicubic filter credits Pillow's `src/libImaging/Resample.c`. + The retained PIL/Pillow notice is in [Pillow license](licenses/PILLOW-LICENSE). +- Transformers 5.17.0, + `src/transformers/models/qwen2_vl/image_processing_qwen2_vl.py` (smart resize + and patch ordering) and `src/transformers/image_processing_backends.py` + (fused normalization). Copyright 2024 The Qwen team, Alibaba Group and the + HuggingFace Inc. team. All rights reserved. The backend file is + Copyright 2025 The HuggingFace Inc. team. Licensed under the + [Apache License, Version 2.0](licenses/APACHE-2.0). + +No upstream runtime or image decoder is linked by this module. The above +notices and license texts must accompany redistributed adaptations as required +by their respective licenses. + +The native vision geometry, rotary, block and merger execution in +`src/vision/` follows Transformers 5.17.0 `modeling_qwen3_5.py`, Copyright +2025 The Qwen team, Alibaba Group and the HuggingFace Inc. team, under the +[Apache License, Version 2.0](licenses/APACHE-2.0). The implementation is +adapted to Rust/CUDA and separate FP32 visual LoRA execution. diff --git a/src/models/cua_s1/native/licenses/APACHE-2.0 b/src/models/cua_s1/native/licenses/APACHE-2.0 new file mode 100644 index 00000000..68b7d66c --- /dev/null +++ b/src/models/cua_s1/native/licenses/APACHE-2.0 @@ -0,0 +1,203 @@ +Copyright 2018- The Hugging Face team. All rights reserved. + + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [yyyy] [name of copyright owner] + + 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. diff --git a/src/models/cua_s1/native/licenses/PILLOW-LICENSE b/src/models/cua_s1/native/licenses/PILLOW-LICENSE new file mode 100644 index 00000000..10dd42d9 --- /dev/null +++ b/src/models/cua_s1/native/licenses/PILLOW-LICENSE @@ -0,0 +1,30 @@ +The Python Imaging Library (PIL) is + + Copyright © 1997-2011 by Secret Labs AB + Copyright © 1995-2011 by Fredrik Lundh and contributors + +Pillow is the friendly PIL fork. It is + + Copyright © 2010 by Jeffrey A. Clark and contributors + +Like PIL, Pillow is licensed under the open source MIT-CMU License: + +By obtaining, using, and/or copying this software and/or its associated +documentation, you agree that you have read, understood, and will comply +with the following terms and conditions: + +Permission to use, copy, modify and distribute this software and its +documentation for any purpose and without fee is hereby granted, +provided that the above copyright notice appears in all copies, and that +both that copyright notice and this permission notice appear in supporting +documentation, and that the name of Secret Labs AB or the author not be +used in advertising or publicity pertaining to distribution of the software +without specific, written prior permission. + +SECRET LABS AB AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH REGARD TO THIS +SOFTWARE, INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS. +IN NO EVENT SHALL SECRET LABS AB OR THE AUTHOR BE LIABLE FOR ANY SPECIAL, +INDIRECT OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE +OR OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +PERFORMANCE OF THIS SOFTWARE. diff --git a/src/models/cua_s1/native/licenses/PYTORCH-LICENSE b/src/models/cua_s1/native/licenses/PYTORCH-LICENSE new file mode 100644 index 00000000..c23172f7 --- /dev/null +++ b/src/models/cua_s1/native/licenses/PYTORCH-LICENSE @@ -0,0 +1,84 @@ +From PyTorch: + +Copyright (c) 2016- Facebook, Inc (Adam Paszke) +Copyright (c) 2014- Facebook, Inc (Soumith Chintala) +Copyright (c) 2011-2014 Idiap Research Institute (Ronan Collobert) +Copyright (c) 2012-2014 Deepmind Technologies (Koray Kavukcuoglu) +Copyright (c) 2011-2012 NEC Laboratories America (Koray Kavukcuoglu) +Copyright (c) 2011-2013 NYU (Clement Farabet) +Copyright (c) 2006-2010 NEC Laboratories America (Ronan Collobert, Leon Bottou, Iain Melvin, Jason Weston) +Copyright (c) 2006 Idiap Research Institute (Samy Bengio) +Copyright (c) 2001-2004 Idiap Research Institute (Ronan Collobert, Samy Bengio, Johnny Mariethoz) + +From Caffe2: + +Copyright (c) 2016-present, Facebook Inc. All rights reserved. + +All contributions by Facebook: +Copyright (c) 2016 Facebook Inc. + +All contributions by Google: +Copyright (c) 2015 Google Inc. +All rights reserved. + +All contributions by Yangqing Jia: +Copyright (c) 2015 Yangqing Jia +All rights reserved. + +All contributions by Kakao Brain: +Copyright 2019-2020 Kakao Brain + +All contributions by Cruise LLC: +Copyright (c) 2022 Cruise LLC. +All rights reserved. + +All contributions by Tri Dao: +Copyright (c) 2024 Tri Dao. +All rights reserved. + +All contributions by Arm: +Copyright (c) 2021, 2023-2025 Arm Limited and/or its affiliates + +All contributions from Caffe: +Copyright(c) 2013, 2014, 2015, the respective contributors +All rights reserved. + +All other contributions: +Copyright(c) 2015, 2016 the respective contributors +All rights reserved. + +Caffe2 uses a copyright model similar to Caffe: each contributor holds +copyright over their contributions to Caffe2. The project versioning records +all such contribution and copyright details. If a contributor wants to further +mark their specific copyright on a particular contribution, they should +indicate their copyright solely in the commit message of the change when it is +committed. + +All rights reserved. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are met: + +1. Redistributions of source code must retain the above copyright + notice, this list of conditions and the following disclaimer. + +2. Redistributions in binary form must reproduce the above copyright + notice, this list of conditions and the following disclaimer in the + documentation and/or other materials provided with the distribution. + +3. Neither the names of Facebook, Deepmind Technologies, NYU, NEC Laboratories America + and IDIAP Research Institute nor the names of its contributors may be + used to endorse or promote products derived from this software without + specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT OWNER OR CONTRIBUTORS BE +LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF +SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN +CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) +ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE +POSSIBILITY OF SUCH DAMAGE. diff --git a/src/models/cua_s1/native/src/contract.rs b/src/models/cua_s1/native/src/contract.rs index 0ca10f29..23b91e6c 100644 --- a/src/models/cua_s1/native/src/contract.rs +++ b/src/models/cua_s1/native/src/contract.rs @@ -53,6 +53,7 @@ fn error(status: u16, message: impl Into) -> RequestError { } } +#[derive(Clone)] pub struct Question { pub name: String, pub goal: String, diff --git a/src/models/cua_s1/native/src/image_preprocess.rs b/src/models/cua_s1/native/src/image_preprocess.rs new file mode 100644 index 00000000..9009fb2a --- /dev/null +++ b/src/models/cua_s1/native/src/image_preprocess.rs @@ -0,0 +1,246 @@ +//! CPU preprocessing for the fixed Qwen3.5-4B / Cua-S1 image processor. +//! +//! Input is already decoded, interleaved RGB8. No image codecs, GPU, model +//! weights, or HTTP handling are involved. The output is row-major +//! `[patches, 1536]`, ready for a separate vision encoder. +//! +//! The resize algorithm is a Rust adaptation of PyTorch's CPU uint8 bicubic +//! antialias implementation (which credits Pillow). Smart resize and packing +//! follow Transformers' Qwen2VLImageProcessor. See `../THIRD_PARTY_NOTICES.md`. + +use std::borrow::Cow; + +use anyhow::{Result, ensure}; + +const PATCH_SIZE: usize = 16; +const MERGE_SIZE: usize = 2; +const FACTOR: usize = PATCH_SIZE * MERGE_SIZE; +const PATCH_VALUES: usize = 3 * 2 * PATCH_SIZE * PATCH_SIZE; +const MIN_PIXELS: usize = 65_536; +const MAX_PIXELS: usize = 16_777_216; + +/// Normalized image patches and the spatial metadata used by the vision model. +#[derive(Debug)] +pub struct ProcessedImage { + /// Contiguous row-major `[patches, 1536]` float32 values. + pub pixel_values: Vec, + /// `[1, resized_height / 16, resized_width / 16]`. + pub image_grid_thw: [usize; 3], + pub resized_width: usize, + pub resized_height: usize, +} + +impl ProcessedImage { + /// Number of image tokens after the model's 2-by-2 spatial merge. + pub fn image_tokens(&self) -> usize { + self.pixel_values.len() / PATCH_VALUES / (MERGE_SIZE * MERGE_SIZE) + } +} + +/// Preprocess a decoded RGB8 image using the fixed 4B processor settings. +/// +/// Rejects empty dimensions, sides over 2048, area over 1,048,576 pixels, +/// aspect ratios over 200, and buffers whose length is not `width * height * 3`. +/// Geometry and buffer arithmetic are checked before allocating image buffers. +pub fn preprocess_rgb8(width: usize, height: usize, rgb: &[u8]) -> Result { + ensure!(width > 0 && height > 0, "image dimensions must be nonzero"); + ensure!( + width <= 2048 && height <= 2048, + "image sides must not exceed 2048" + ); + let area = width + .checked_mul(height) + .ok_or_else(|| anyhow::anyhow!("image area overflow"))?; + ensure!( + area <= 1_048_576, + "image area must not exceed 1048576 pixels" + ); + ensure!( + width.max(height) <= width.min(height) * 200, + "image aspect ratio must not exceed 200" + ); + let input_len = area + .checked_mul(3) + .ok_or_else(|| anyhow::anyhow!("RGB buffer length overflow"))?; + ensure!( + rgb.len() == input_len, + "RGB buffer length must be {input_len}, got {}", + rgb.len() + ); + + let (resized_width, resized_height) = smart_resize(width, height); + let resized_area = resized_width + .checked_mul(resized_height) + .ok_or_else(|| anyhow::anyhow!("resized area overflow"))?; + let resized_len = resized_area + .checked_mul(3) + .ok_or_else(|| anyhow::anyhow!("resized buffer length overflow"))?; + let horizontal_len = resized_width + .checked_mul(height) + .and_then(|area| area.checked_mul(3)) + .ok_or_else(|| anyhow::anyhow!("horizontal buffer length overflow"))?; + let output_len = resized_area + .checked_mul(6) + .ok_or_else(|| anyhow::anyhow!("patch buffer length overflow"))?; + output_len + .checked_mul(std::mem::size_of::()) + .ok_or_else(|| anyhow::anyhow!("patch buffer byte length overflow"))?; + + let mut resized = Cow::Borrowed(rgb); + if resized_width != width { + let axis = AxisWeights::new(width, resized_width); + let mut horizontal = vec![0; horizontal_len]; + for y in 0..height { + for (x, kernel) in axis.kernels.iter().enumerate() { + for channel in 0..3 { + horizontal[(y * resized_width + x) * 3 + channel] = + axis.apply(kernel, |source_x| rgb[(y * width + source_x) * 3 + channel]); + } + } + } + resized = Cow::Owned(horizontal); + } + if resized_height != height { + let axis = AxisWeights::new(height, resized_height); + let mut vertical = vec![0; resized_len]; + for (y, kernel) in axis.kernels.iter().enumerate() { + for x in 0..resized_width { + for channel in 0..3 { + vertical[(y * resized_width + x) * 3 + channel] = axis + .apply(kernel, |source_y| { + resized[(source_y * resized_width + x) * 3 + channel] + }); + } + } + } + resized = Cow::Owned(vertical); + } + + let mut pixel_values = Vec::with_capacity(output_len); + for block_y in 0..resized_height / FACTOR { + for block_x in 0..resized_width / FACTOR { + for merge_y in 0..MERGE_SIZE { + for merge_x in 0..MERGE_SIZE { + for channel in 0..3 { + for _temporal in 0..2 { + for patch_y in 0..PATCH_SIZE { + for patch_x in 0..PATCH_SIZE { + let y = block_y * FACTOR + merge_y * PATCH_SIZE + patch_y; + let x = block_x * FACTOR + merge_x * PATCH_SIZE + patch_x; + let pixel = resized[(y * resized_width + x) * 3 + channel]; + // Match the fused float32 torchvision normalization, + // including its operation order (no reciprocal multiply). + pixel_values.push((f32::from(pixel) - 127.5) / 127.5); + } + } + } + } + } + } + } + } + Ok(ProcessedImage { + pixel_values, + image_grid_thw: [1, resized_height / PATCH_SIZE, resized_width / PATCH_SIZE], + resized_width, + resized_height, + }) +} + +fn smart_resize(width: usize, height: usize) -> (usize, usize) { + // Python round uses ties-to-even; Rust's ordinary round does not. + let mut w = (width as f64 / FACTOR as f64).round_ties_even() as usize * FACTOR; + let mut h = (height as f64 / FACTOR as f64).round_ties_even() as usize * FACTOR; + if w * h > MAX_PIXELS { + let beta = ((width * height) as f64 / MAX_PIXELS as f64).sqrt(); + w = ((width as f64 / beta / FACTOR as f64).floor() as usize * FACTOR).max(FACTOR); + h = ((height as f64 / beta / FACTOR as f64).floor() as usize * FACTOR).max(FACTOR); + } else if w * h < MIN_PIXELS { + let beta = (MIN_PIXELS as f64 / (width * height) as f64).sqrt(); + w = (width as f64 * beta / FACTOR as f64).ceil() as usize * FACTOR; + h = (height as f64 * beta / FACTOR as f64).ceil() as usize * FACTOR; + } + (w, h) +} + +struct Kernel { + start: usize, + weights: Vec, +} + +struct AxisWeights { + kernels: Vec, + precision: u32, +} + +impl AxisWeights { + fn new(input: usize, output: usize) -> Self { + let scale = input as f64 / output as f64; + let support = 2.0 * scale.max(1.0); + let invscale = if scale >= 1.0 { 1.0 / scale } else { 1.0 }; + let max_size = support.ceil() as usize * 2 + 1; + let mut maximum = 0.0_f64; + let mut floating = Vec::with_capacity(output); + for index in 0..output { + let center = scale * (index as f64 + 0.5); + // C++ conversion truncates toward zero before clamping the bounds. + let start = ((center - support + 0.5) as isize).max(0) as usize; + let end = ((center + support + 0.5) as usize).min(input); + let count = end.saturating_sub(start).min(max_size); + let mut weights: Vec = (0..count) + .map(|j| cubic((j as f64 + start as f64 - center + 0.5) * invscale)) + .collect(); + let total: f64 = weights.iter().sum(); + if total != 0.0 { + for weight in &mut weights { + *weight /= total; + maximum = maximum.max(*weight); + } + } + floating.push((start, weights)); + } + // One precision for the whole axis, as in PyTorch's int16 path. + let mut precision = 0; + while precision < 22 { + if (0.5 + maximum * f64::from(1 << (precision + 1))) as i32 >= (1 << 15) { + break; + } + precision += 1; + } + let multiplier = f64::from(1 << precision); + let kernels = floating + .into_iter() + .map(|(start, weights)| Kernel { + start, + weights: weights + .into_iter() + .map(|weight| { + let value = weight * multiplier; + (value + if value < 0.0 { -0.5 } else { 0.5 }) as i16 + }) + .collect(), + }) + .collect(); + Self { kernels, precision } + } + + fn apply(&self, kernel: &Kernel, pixel: impl Fn(usize) -> u8) -> u8 { + let mut accumulator = 1_i32 << (self.precision - 1); + for (offset, &weight) in kernel.weights.iter().enumerate() { + accumulator += i32::from(pixel(kernel.start + offset)) * i32::from(weight); + } + (accumulator >> self.precision).clamp(0, 255) as u8 + } +} + +fn cubic(x: f64) -> f64 { + let x = x.abs(); + const A: f64 = -0.5; + if x < 1.0 { + ((A + 2.0) * x - (A + 3.0)) * x * x + 1.0 + } else if x < 2.0 { + ((A * x - 5.0 * A) * x + 8.0 * A) * x - 4.0 * A + } else { + 0.0 + } +} diff --git a/src/models/cua_s1/native/src/image_request.rs b/src/models/cua_s1/native/src/image_request.rs new file mode 100644 index 00000000..9a626208 --- /dev/null +++ b/src/models/cua_s1/native/src/image_request.rs @@ -0,0 +1,300 @@ +//! Bounded PNG/JPEG screenshot request mapping for the native worker. +use crate::contract::Question; +use anyhow::{Context, Result, ensure}; +use base64::Engine; +use serde_json::{Map, Value}; +use std::io::Cursor; + +pub const MAX_BODY: usize = 8 * 1024 * 1024; +pub const MAX_IMAGE_BYTES: usize = 4 * 1024 * 1024; +pub const MODEL_ID: &str = + "cua-ai/cua-s1-4b-0.2@16818868b0cc7813808aae4e87b417657046ab79:multimodal"; + +pub struct ImageRequest { + pub width: usize, + pub height: usize, + pub rgb: Vec, + pub questions: Vec, +} + +pub fn parse_image_body(body: &Map) -> Result { + ensure!( + body.len() == 3 + && ["model", "state", "questions"] + .iter() + .all(|k| body.contains_key(*k)), + "request must contain model, state and questions only" + ); + let questions = body + .get("questions") + .and_then(Value::as_object) + .context("questions must be an object")?; + ensure!( + (1..=8).contains(&questions.len()), + "questions must contain 1 to 8 questions" + ); + for (name, q) in questions { + bounded_key(name)?; + let q = q.as_object().context("question must be an object")?; + ensure!( + q.keys() + .all(|k| ["type", "instructions", "criteria"].contains(&k.as_str())), + "unsupported question fields" + ); + let goal_len = match q.get("instructions") { + Some(Value::Null) => 0, + Some(v) => checked_text(v)?.chars().count(), + None => anyhow::bail!("instructions is required"), + }; + let criteria = q + .get("criteria") + .and_then(Value::as_object) + .context("criteria must be an object")?; + let mut length = goal_len; + for (key, value) in criteria { + bounded_key(key)?; + let label = if value.is_null() { + checked_text(&Value::String(key.clone()))? + } else { + checked_text(value)? + }; + // Match Python's escaped-label character budget. + length += crate::json::quote(&label).chars().count() - 2; + } + ensure!( + length <= 16384, + "combined question text exceeds 16384 characters" + ); + } + let mut mapped = body.clone(); + mapped.insert("state".into(), Value::String("image".into())); + let (_, questions) = + crate::contract::map_request(&mapped).map_err(|e| anyhow::anyhow!(e.message))?; + let state = body["state"] + .as_object() + .context("state must contain exactly one image")?; + ensure!(state.len() == 1, "state must contain exactly one image"); + let url = state + .get("image") + .and_then(Value::as_str) + .context("state.image must be a data URL")?; + let (prefix, encoded) = url.split_once(',').context("invalid image data URL")?; + let format = match prefix { + "data:image/png;base64" => image::ImageFormat::Png, + "data:image/jpeg;base64" => image::ImageFormat::Jpeg, + _ => anyhow::bail!("only inline PNG/JPEG images are supported"), + }; + ensure!( + encoded.len() <= 4 * MAX_IMAGE_BYTES.div_ceil(3), + "encoded image exceeds 4 MiB" + ); + let raw = base64::engine::general_purpose::STANDARD + .decode(encoded) + .context("invalid base64 image")?; + ensure!(raw.len() <= MAX_IMAGE_BYTES, "image exceeds 4 MiB"); + ensure!( + image::guess_format(&raw)? == format, + "image format does not match MIME type" + ); + if format == image::ImageFormat::Jpeg { + ensure_single_jpeg(&raw)?; + } + let mut limits = image::Limits::default(); + limits.max_image_width = Some(2048); + limits.max_image_height = Some(2048); + limits.max_alloc = Some(32 * 1024 * 1024); + let mut header = image::ImageReader::with_format(Cursor::new(&raw), format); + header.limits(limits.clone()); + let (width, height) = header.into_dimensions()?; + let (w, h) = (width as usize, height as usize); + ensure!( + w > 0 && h > 0 && w <= 2048 && h <= 2048 && w * h <= 1048576 && w.max(h) <= 200 * w.min(h), + "image dimensions exceed supported limits" + ); + let decoded = if format == image::ImageFormat::Png { + let decoder = image::codecs::png::PngDecoder::with_limits(Cursor::new(&raw), limits)?; + ensure!(!decoder.is_apng()?, "image must be single-frame"); + image::DynamicImage::from_decoder(decoder)? + } else { + let mut reader = image::ImageReader::with_format(Cursor::new(&raw), format); + reader.limits(limits); + reader.decode()? + }; + Ok(ImageRequest { + width: w, + height: h, + rgb: pillow_rgb( + decoded, + format == image::ImageFormat::Png && raw.get(25) == Some(&0), + ), + questions, + }) +} + +fn ensure_single_jpeg(raw: &[u8]) -> Result<()> { + // MPF APP2 identifies an MPO container. Pillow reports it as MPO rather + // than JPEG; decoding it as JPEG would silently select the first picture. + // Walk header segments only, so arbitrary metadata/entropy bytes cannot + // be mistaken for an MPF marker. + let mut offset = 2; // SOI was checked by guess_format. + while offset < raw.len() { + ensure!(raw[offset] == 0xff, "invalid JPEG marker"); + while raw.get(offset) == Some(&0xff) { + offset += 1; + } + let marker = *raw.get(offset).context("truncated JPEG marker")?; + offset += 1; + match marker { + 0xda | 0xd9 => break, // SOS / EOI; the decoder checks the rest. + 0x01 | 0xd0..=0xd8 => continue, // Standalone markers. + _ => {} + } + let length = raw + .get(offset..offset + 2) + .context("truncated JPEG segment")?; + let length = u16::from_be_bytes([length[0], length[1]]) as usize; + ensure!(length >= 2, "invalid JPEG segment length"); + let segment = raw + .get(offset + 2..offset + length) + .context("truncated JPEG segment")?; + ensure!( + marker != 0xe2 || !segment.starts_with(b"MPF\0"), + "image must be single-frame; MPF/MPO containers are unsupported" + ); + offset += length; + } + Ok(()) +} + +fn pillow_rgb(image: image::DynamicImage, png_grayscale: bool) -> Vec { + use image::DynamicImage::*; + match image { + ImageLuma16(p) => p.pixels().flat_map(|v| [v[0].min(255) as u8; 3]).collect(), + ImageLumaA16(p) if png_grayscale => { + p.pixels().flat_map(|v| [v[0].min(255) as u8; 3]).collect() + } + ImageLumaA16(p) => p.pixels().flat_map(|v| [(v[0] >> 8) as u8; 3]).collect(), + ImageRgb16(p) => p + .pixels() + .flat_map(|v| [(v[0] >> 8) as u8, (v[1] >> 8) as u8, (v[2] >> 8) as u8]) + .collect(), + ImageRgba16(p) => p + .pixels() + .flat_map(|v| [(v[0] >> 8) as u8, (v[1] >> 8) as u8, (v[2] >> 8) as u8]) + .collect(), + other => other.to_rgb8().into_raw(), + } +} + +fn bounded_key(key: &str) -> Result<()> { + ensure!( + (1..=256).contains(&key.chars().count()), + "names and option keys must contain 1 to 256 characters" + ); + Ok(()) +} + +fn checked_text(value: &Value) -> Result { + let text = match value { + Value::String(s) => s.clone(), + Value::Object(_) | Value::Array(_) => crate::json::dumps(value), + _ => anyhow::bail!("text must be a string, object or array"), + }; + ensure!( + text.chars().count() <= 16384, + "text exceeds 16384 characters" + ); + ensure!( + ![ + "<|image_pad|>", + "<|video_pad|>", + "<|vision_start|>", + "<|vision_end|>" + ] + .iter() + .any(|t| text.contains(t)), + "unsupported media control token" + ); + Ok(text) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn parse_jpeg(raw: &[u8]) -> Result { + let body = json!({ + "model": "cua-s1-4b-0.2", + "state": {"image": format!("data:image/jpeg;base64,{}", base64::engine::general_purpose::STANDARD.encode(raw))}, + "questions": {"pick": {"type": "choice", "instructions": "Choose next action", "criteria": {"a": "Continue"}}} + }); + parse_image_body(body.as_object().unwrap()) + } + + #[test] + fn request_rejects_pillow_two_frame_mpo() { + // Pillow: RGB 32x32 red.save(format="MPO", save_all=True, + // append_images=[RGB 32x32 blue]). This is an actual two-picture file. + let raw = base64::engine::general_purpose::STANDARD.decode(concat!( + "/9j/4AAQSkZJRgABAQAAAQABAAD/4gBoTVBGAElJKgAIAAAAAwAAsAcABAAAADAxMDABsAQAAQAAAAIAAAACsAcAIAAAADIAAAAA", + "AAAAAAADAO8CAAAAAAAAAAAAAAAAAACFAgAA0wIAAAAAAAAgICAgICAgICAgICAgICAg/9sAQwAIBgYHBgUIBwcHCQkICgwUDQwL", + "CwwZEhMPFB0aHx4dGhwcICQuJyAiLCMcHCg3KSwwMTQ0NB8nOT04MjwuMzQy/9sAQwEJCQkMCwwYDQ0YMiEcITIyMjIyMjIyMjIy", + "MjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIy/8AAEQgAIAAgAwEiAAIRAQMRAf/EAB8AAAEFAQEBAQEBAAAA", + "AAAAAAABAgMEBQYHCAkKC//EALUQAAIBAwMCBAMFBQQEAAABfQECAwAEEQUSITFBBhNRYQcicRQygZGhCCNCscEVUtHwJDNicoIJ", + "ChYXGBkaJSYnKCkqNDU2Nzg5OkNERUZHSElKU1RVVldYWVpjZGVmZ2hpanN0dXZ3eHl6g4SFhoeIiYqSk5SVlpeYmZqio6Slpqeo", + "qaqys7S1tre4ubrCw8TFxsfIycrS09TV1tfY2drh4uPk5ebn6Onq8fLz9PX29/j5+v/EAB8BAAMBAQEBAQEBAQEAAAAAAAABAgME", + "BQYHCAkKC//EALURAAIBAgQEAwQHBQQEAAECdwABAgMRBAUhMQYSQVEHYXETIjKBCBRCkaGxwQkjM1LwFWJy0QoWJDThJfEXGBka", + "JicoKSo1Njc4OTpDREVGR0hJSlNUVVZXWFlaY2RlZmdoaWpzdHV2d3h5eoKDhIWGh4iJipKTlJWWl5iZmqKjpKWmp6ipqrKztLW2", + "t7i5usLDxMXGx8jJytLT1NXW19jZ2uLj5OXm5+jp6vLz9PX29/j5+v/aAAwDAQACEQMRAD8A4uiiivmT9xCiiigAooooAKKKKAP/", + "2f/Y/+AAEEpGSUYAAQEAAAEAAQAA/9sAQwAIBgYHBgUIBwcHCQkICgwUDQwLCwwZEhMPFB0aHx4dGhwcICQuJyAiLCMcHCg3KSww", + "MTQ0NB8nOT04MjwuMzQy/9sAQwEJCQkMCwwYDQ0YMiEcITIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIy", + "MjIyMjIyMjIy/8AAEQgAIAAgAwEiAAIRAQMRAf/EAB8AAAEFAQEBAQEBAAAAAAAAAAABAgMEBQYHCAkKC//EALUQAAIBAwMCBAMF", + "BQQEAAABfQECAwAEEQUSITFBBhNRYQcicRQygZGhCCNCscEVUtHwJDNicoIJChYXGBkaJSYnKCkqNDU2Nzg5OkNERUZHSElKU1RV", + "VldYWVpjZGVmZ2hpanN0dXZ3eHl6g4SFhoeIiYqSk5SVlpeYmZqio6Slpqeoqaqys7S1tre4ubrCw8TFxsfIycrS09TV1tfY2drh", + "4uPk5ebn6Onq8fLz9PX29/j5+v/EAB8BAAMBAQEBAQEBAQEAAAAAAAABAgMEBQYHCAkKC//EALURAAIBAgQEAwQHBQQEAAECdwAB", + "AgMRBAUhMQYSQVEHYXETIjKBCBRCkaGxwQkjM1LwFWJy0QoWJDThJfEXGBkaJicoKSo1Njc4OTpDREVGR0hJSlNUVVZXWFlaY2Rl", + "ZmdoaWpzdHV2d3h5eoKDhIWGh4iJipKTlJWWl5iZmqKjpKWmp6ipqrKztLW2t7i5usLDxMXGx8jJytLT1NXW19jZ2uLj5OXm5+jp", + "6vLz9PX29/j5+v/aAAwDAQACEQMRAD8A8cooor9xPMCiiigAooooAKKKKAP/2Q==", + )).unwrap(); + let error = parse_jpeg(&raw) + .err() + .expect("multi-picture input must be rejected"); + assert!(error.to_string().contains("single-frame"), "{error}"); + } + + fn jpeg() -> Vec { + let mut raw = Vec::new(); + image::codecs::jpeg::JpegEncoder::new(&mut raw) + .encode(&[20, 40, 60], 1, 1, image::ExtendedColorType::Rgb8) + .unwrap(); + raw + } + + fn with_segment(marker: u8, payload: &[u8]) -> Vec { + let mut raw = jpeg(); + let mut segment = vec![0xff, marker]; + segment.extend_from_slice(&((payload.len() + 2) as u16).to_be_bytes()); + segment.extend_from_slice(payload); + raw.splice(2..2, segment); + raw + } + + #[test] + fn request_accepts_single_jpeg_and_unrelated_metadata() { + for raw in [ + jpeg(), + with_segment(0xe2, b"ICC_PROFILE\0MPF\0"), + with_segment(0xfe, b"MPF\0"), + ] { + let request = parse_jpeg(&raw).unwrap(); + assert_eq!((request.width, request.height), (1, 1)); + assert_eq!(request.rgb.len(), 3); + } + } + + #[test] + fn request_rejects_truncated_jpeg_segment() { + assert!(parse_jpeg(&[0xff, 0xd8, 0xff, 0xe2, 0, 20, b'M', b'P', b'F', 0]).is_err()); + } +} diff --git a/src/models/cua_s1/native/src/lib.rs b/src/models/cua_s1/native/src/lib.rs index a991f004..4f2e8047 100644 --- a/src/models/cua_s1/native/src/lib.rs +++ b/src/models/cua_s1/native/src/lib.rs @@ -1,9 +1,14 @@ -//! A native `/v1/systemone` worker for Cua-S1 4B 0.2 (`text` adapter): request -//! handling, tokenization and scoring in Rust, the Qwen3.5 forward pass on the CUDA -//! kernels of `src/backends/cuda/qwen3_5`, loaded at run time. +//! Native Cua-S1 text and screenshot workers using shared Qwen execution. pub mod contract; -pub use omni_qwen3_5_native::{cuda, json, model}; +pub use omni_qwen3_5_native::{cuda, inputs, json, model}; pub mod engine; pub mod executor; +pub mod image_preprocess; +pub mod image_request; +pub mod multimodal; pub mod processing; +mod provenance; +pub mod vision; +pub mod vision_engine; +pub mod vision_processing; diff --git a/src/models/cua_s1/native/src/multimodal.rs b/src/models/cua_s1/native/src/multimodal.rs new file mode 100644 index 00000000..555eda83 --- /dev/null +++ b/src/models/cua_s1/native/src/multimodal.rs @@ -0,0 +1,153 @@ +//! Single-image native prompt preparation and end-to-end execution. + +use crate::contract::{LETTERS, Question, SYSTEM_PROMPT}; +use anyhow::{Result, ensure}; + +pub const MAX_TOKENS: usize = 4096; + +pub fn chat_image(question: &Question, image_tokens: usize) -> Result { + ensure!( + (1..MAX_TOKENS).contains(&image_tokens), + "invalid image token count" + ); + ensure!( + (1..=26).contains(&question.labels.len()), + "invalid option count" + ); + let goal = if question.goal.is_empty() { + String::new() + } else { + format!("Goal: {}\n\n", question.goal) + }; + let options = LETTERS + .chars() + .zip(&question.labels) + .map(|(letter, label)| format!("{letter}. Decision \"{label}\" -> select")) + .collect::>() + .join("\n"); + Ok(format!( + "<|im_start|>system\n{SYSTEM_PROMPT}<|im_end|>\n<|im_start|>user\n<|vision_start|>{}<|vision_end|>{goal}App: Cua Driver\nTask family: closed-candidate decision\n\nThe current screenshot is attached.\n\nOptions:\n{options}\n\nAnswer with a single letter.<|im_end|>\n<|im_start|>assistant\n\n", + "<|image_pad|>".repeat(image_tokens) + )) +} + +pub fn image_positions(ids: &[u32], image_token: u32, grid: [usize; 3]) -> Result<[Vec; 3]> { + let [t, h, w] = grid; + ensure!( + t == 1 && h > 0 && w > 0 && h.is_multiple_of(2) && w.is_multiple_of(2), + "expected one image with even, nonzero spatial grid" + ); + ensure!( + !ids.is_empty() && ids.len() <= MAX_TOKENS, + "processed prompt exceeds {MAX_TOKENS} tokens or is empty" + ); + let count = (h / 2) + .checked_mul(w / 2) + .ok_or_else(|| anyhow::anyhow!("grid overflow"))?; + let start = ids + .iter() + .position(|&id| id == image_token) + .ok_or_else(|| anyhow::anyhow!("missing image placeholders"))?; + let end = start + .checked_add(count) + .ok_or_else(|| anyhow::anyhow!("grid overflow"))?; + ensure!( + end <= ids.len() + && ids[start..end].iter().all(|&id| id == image_token) + && ids[end..].iter().all(|&id| id != image_token), + "image placeholders must be one contiguous span matching the grid" + ); + let mut positions: [Vec; 3] = std::array::from_fn(|_| (0..start as i64).collect()); + for y in 0..h / 2 { + for x in 0..w / 2 { + positions[0].push(start as i64); + positions[1].push((start + y) as i64); + positions[2].push((start + x) as i64); + } + } + let next = start + h.max(w) / 2; + for axis in &mut positions { + axis.extend((next..next + ids.len() - end).map(|p| p as i64)); + } + Ok(positions) +} + +/// A prepared single-image prompt. Every prompt in a request is prepared before +/// running the shared image encoder. +pub struct ImagePrompt { + pub token_ids: Vec, + pub image_token_indices: Vec, + pub position_ids: [Vec; 3], +} + +pub fn prepare_prompt( + tokenizer: &tokenizers::Tokenizer, + question: &Question, + grid: [usize; 3], + image_token: u32, +) -> Result { + let count = grid[1] + .checked_mul(grid[2]) + .ok_or_else(|| anyhow::anyhow!("grid overflow"))? + / 4; + let encoded = tokenizer + .encode(chat_image(question, count)?, false) + .map_err(anyhow::Error::msg)?; + let token_ids = encoded.get_ids().to_vec(); + let position_ids = image_positions(&token_ids, image_token, grid)?; + let image_token_indices = token_ids + .iter() + .enumerate() + .filter_map(|(i, &id)| (id == image_token).then_some(i)) + .collect(); + Ok(ImagePrompt { + token_ids, + image_token_indices, + position_ids, + }) +} + +/// Validate public RGB callers as well as the JSON mapper before GPU work. +pub(crate) fn validate_questions(questions: &[Question]) -> Result<()> { + ensure!( + (1..=8).contains(&questions.len()), + "expected 1 to 8 questions" + ); + let mut names = std::collections::HashSet::new(); + for q in questions { + ensure!( + (1..=256).contains(&q.name.chars().count()) && names.insert(&q.name), + "invalid or duplicate question name" + ); + ensure!( + (1..=26).contains(&q.keys.len()) && q.keys.len() == q.labels.len(), + "option key/label counts must agree and be 1 to 26" + ); + let mut keys = std::collections::HashSet::new(); + ensure!( + q.keys + .iter() + .all(|k| (1..=256).contains(&k.chars().count()) && keys.insert(k)), + "invalid or duplicate option key" + ); + ensure!( + q.goal.chars().count() + q.labels.iter().map(|s| s.chars().count()).sum::() + <= 16384, + "combined question text exceeds 16384 characters" + ); + for text in std::iter::once(&q.goal).chain(&q.labels) { + ensure!( + ![ + "<|image_pad|>", + "<|video_pad|>", + "<|vision_start|>", + "<|vision_end|>" + ] + .iter() + .any(|t| text.contains(t)), + "unsupported media control token" + ); + } + } + Ok(()) +} diff --git a/src/models/cua_s1/native/src/provenance.rs b/src/models/cua_s1/native/src/provenance.rs new file mode 100644 index 00000000..4ea295bb --- /dev/null +++ b/src/models/cua_s1/native/src/provenance.rs @@ -0,0 +1,105 @@ +//! Hash checks before assigning the pinned multimodal model identity. +use anyhow::{Context, Result, ensure}; +use serde_json::Value; +use sha2::{Digest, Sha256}; +use std::{ + fs::{self, File}, + io::Read, + path::Path, +}; +const LOCK_SHA: &str = "9820bd232c5762f114e19680c0f8203d7e1faaf8a60c196cfe01964d6d8a6c09"; + +fn hash_file(path: &Path) -> Result { + let mut file = File::open(path)?; + let mut hash = Sha256::new(); + let mut buffer = vec![0; 1024 * 1024]; + loop { + let n = file.read(&mut buffer)?; + if n == 0 { + break; + } + hash.update(&buffer[..n]); + } + Ok(format!("{:x}", hash.finalize())) +} + +fn verify_files( + root: &Path, + files: &serde_json::Map, + extra: Option<&str>, +) -> Result<()> { + ensure!(!files.is_empty(), "empty artifact hash manifest"); + for (name, expected) in files { + ensure!( + Path::new(name).components().count() == 1 + && matches!( + Path::new(name).components().next(), + Some(std::path::Component::Normal(_)) + ), + "invalid artifact name" + ); + let path = root.join(name); + ensure!( + Some(fs::metadata(&path)?.len()) == expected["size"].as_u64(), + "{} size mismatch", + path.display() + ); + ensure!( + Some(hash_file(&path)?.as_str()) == expected["sha256"].as_str(), + "{} checksum mismatch", + path.display() + ); + } + for entry in fs::read_dir(root)? { + let entry = entry?; + let name = entry.file_name(); + let name = name.to_str().context("non-UTF8 artifact")?; + ensure!( + name == ".cache" || Some(name) == extra || files.contains_key(name), + "unlisted artifact: {name}" + ); + } + Ok(()) +} + +pub fn verify_sources(base: &Path, adapter: &Path) -> Result<()> { + let lock = fs::read( + base.parent() + .context("base has no parent")? + .join("weights.lock.json"), + )?; + ensure!( + format!("{:x}", Sha256::digest(&lock)) == LOCK_SHA, + "upstream weights manifest checksum mismatch" + ); + let lock: Value = serde_json::from_slice(&lock)?; + for artifact in lock["artifacts"].as_array().context("artifacts")? { + let mut selected = serde_json::Map::new(); + let is_adapter = artifact["role"] == "adapter"; + for (name, value) in artifact["files"].as_object().context("files")? { + if is_adapter { + if let Some(name) = name.strip_prefix("multimodal/") { + selected.insert(name.into(), value.clone()); + } + } else { + selected.insert(name.clone(), value.clone()); + } + } + verify_files(if is_adapter { adapter } else { base }, &selected, None)?; + } + Ok(()) +} + +pub fn verify_export(dir: &Path, marker: &Value) -> Result<()> { + let files = marker["files"] + .as_object() + .context("language export has no file hashes; rerun export_multimodal_language.py")?; + for name in ["config.json", "tokenizer.json"] { + ensure!(files.contains_key(name), "export must hash {name}"); + } + ensure!( + files.keys().any(|n| n.ends_with(".safetensors")), + "export must hash model weights" + ); + verify_files(dir, files, Some("cua_s1_language_export.json")) +} diff --git a/src/models/cua_s1/native/src/vision/geometry.rs b/src/models/cua_s1/native/src/vision/geometry.rs new file mode 100644 index 00000000..6cd9f4a4 --- /dev/null +++ b/src/models/cua_s1/native/src/vision/geometry.rs @@ -0,0 +1,60 @@ +//! Geometry in the processor's 2×2 block-major patch order. +use anyhow::{Result, ensure}; +pub(super) struct Geometry { + pub indices: Vec, + pub weights: Vec, + pub cos: Vec, + pub sin: Vec, +} +impl Geometry { + pub fn new([t, h, w]: [usize; 3]) -> Result { + ensure!( + t == 1 && h > 0 && w > 0 && h % 2 == 0 && w % 2 == 0, + "expected one image with an even, nonzero patch grid" + ); + let n = h + .checked_mul(w) + .ok_or_else(|| anyhow::anyhow!("vision grid overflow"))?; + // Processor bounds allow rounding up a 1,048,576-pixel source and narrow upscaled images. + ensure!( + n <= 4608 && h <= 512 && w <= 512, + "vision grid exceeds processor bounds" + ); + let mut g = Self { + indices: Vec::with_capacity(n * 4), + weights: Vec::with_capacity(n * 4), + cos: Vec::with_capacity(n * 32), + sin: Vec::with_capacity(n * 32), + }; + for br in 0..h / 2 { + for bc in 0..w / 2 { + for ir in 0..2 { + for ic in 0..2 { + let row = br * 2 + ir; + let col = bc * 2 + ic; + let y = (row as f32 * 47.) / (h - 1) as f32; + let x = (col as f32 * 47.) / (w - 1) as f32; + let yl = y.floor() as usize; + let xl = x.floor() as usize; + let fy = y - yl as f32; + let fx = x - xl as f32; + for (yy, wy) in [(yl, 1. - fy), ((yl + 1).min(47), fy)] { + for (xx, wx) in [(xl, 1. - fx), ((xl + 1).min(47), fx)] { + g.indices.push((yy * 48 + xx) as i32); + g.weights.push(wy * wx); + } + } + for pos in [row, col] { + for i in 0..16 { + let angle = pos as f32 / 10000f32.powf(i as f32 / 16.); + g.cos.push(angle.cos()); + g.sin.push(angle.sin()); + } + } + } + } + } + } + Ok(g) + } +} diff --git a/src/models/cua_s1/native/src/vision/mod.rs b/src/models/cua_s1/native/src/vision/mod.rs new file mode 100644 index 00000000..940e4fb5 --- /dev/null +++ b/src/models/cua_s1/native/src/vision/mod.rs @@ -0,0 +1,482 @@ +//! Structurally validated vision checkpoints and native CUDA vision execution. +//! `VisionCheckpoint` loads on CPU; `VisionModel` uploads separate base/LoRA tensors. +//! +//! Callers must keep checkpoint files immutable (including no truncation) for the +//! lifetime of the checkpoint and its borrowed views. Structural checks do not +//! verify upstream hashes or tensor values. + +use anyhow::{Context, Result, ensure}; +use memmap2::Mmap; +use safetensors::{ + Dtype, SafeTensors, + tensor::{TensorInfo, TensorView}, +}; +use serde::{ + Deserialize, + de::{self, MapAccess, Visitor}, +}; +use serde_json::Value; +use std::{ + collections::{BTreeMap, BTreeSet}, + fmt, + fs::{self, File}, + path::{Component, Path, PathBuf}, +}; + +const BASE: &str = "model.visual."; +const ADAPTER: &str = "base_model.model.model.visual."; +type Inventory = BTreeMap>; + +/// The supported Qwen3.5-4B vision architecture; all fields are validated at load. +#[derive(Debug, Deserialize)] +pub struct VisionConfig { + pub depth: usize, + pub hidden_size: usize, + pub intermediate_size: usize, + pub num_heads: usize, + pub num_position_embeddings: usize, + pub out_hidden_size: usize, + pub in_channels: usize, + pub patch_size: usize, + pub temporal_patch_size: usize, + pub spatial_merge_size: usize, + pub hidden_act: String, + pub deepstack_visual_indexes: Vec, + pub model_type: String, +} +impl VisionConfig { + fn load(dir: &Path) -> Result { + let config: Value = serde_json::from_slice(&fs::read(inside(dir, "config.json")?)?)?; + let vision: Self = + serde_json::from_value(config["vision_config"].clone()).context("vision config")?; + let sizes = [ + vision.depth, + vision.hidden_size, + vision.intermediate_size, + vision.num_heads, + vision.num_position_embeddings, + vision.out_hidden_size, + vision.in_channels, + vision.patch_size, + vision.temporal_patch_size, + vision.spatial_merge_size, + ]; + ensure!( + sizes == [24, 1024, 4096, 16, 2304, 2560, 3, 16, 2, 2] + && vision.hidden_act == "gelu_pytorch_tanh" + && vision.deepstack_visual_indexes.is_empty() + && vision.model_type == "qwen3_5" + && config["model_type"] == "qwen3_5" + && config["text_config"]["hidden_size"].as_u64() + == Some(vision.out_hidden_size as u64), + "unsupported vision/text config: expected pinned Qwen3.5-4B layout" + ); + Ok(vision) + } + fn inventory(&self) -> Inventory { + let mut tensors = Inventory::new(); + let h = self.hidden_size; + let i = self.intermediate_size; + let mut linear = |name: String, output: usize, input: Option| { + tensors.insert( + format!("{BASE}{name}.weight"), + input.map_or_else(|| vec![output], |input| vec![output, input]), + ); + tensors.insert(format!("{BASE}{name}.bias"), vec![output]); + }; + for block in 0..self.depth { + for norm in ["norm1", "norm2"] { + linear(format!("blocks.{block}.{norm}"), h, None); + } + for (name, output, input) in [ + ("attn.qkv", 3 * h, h), + ("attn.proj", h, h), + ("mlp.linear_fc1", i, h), + ("mlp.linear_fc2", h, i), + ] { + linear(format!("blocks.{block}.{name}"), output, Some(input)); + } + } + let merged = h * self.spatial_merge_size * self.spatial_merge_size; + linear("merger.norm".into(), h, None); + linear("merger.linear_fc1".into(), merged, Some(merged)); + linear( + "merger.linear_fc2".into(), + self.out_hidden_size, + Some(merged), + ); + tensors.insert( + format!("{BASE}patch_embed.proj.weight"), + vec![ + h, + self.in_channels, + self.temporal_patch_size, + self.patch_size, + self.patch_size, + ], + ); + tensors.insert(format!("{BASE}patch_embed.proj.bias"), vec![h]); + tensors.insert( + format!("{BASE}pos_embed.weight"), + vec![self.num_position_embeddings, h], + ); + tensors + } +} + +/// Inference LoRA parameters; base and adapter bytes remain separate. +#[derive(Debug)] +pub struct AdapterConfig { + pub rank: usize, + pub alpha: usize, +} +impl AdapterConfig { + pub fn scale(&self) -> f64 { + self.alpha as f64 / self.rank as f64 + } + fn load(dir: &Path) -> Result { + let config: Value = + serde_json::from_slice(&fs::read(inside(dir, "adapter_config.json")?)?)?; + let c = config + .as_object() + .context("adapter config must be an object")?; + for (key, value) in c { + let supported = match key.as_str() { + "r" => value == 16, + "lora_alpha" => value == 32, + "peft_type" => value == "LORA", + "bias" => value == "none", + "base_model_name_or_path" => value == "Qwen/Qwen3.5-4B", + "task_type" => value == "CAUSAL_LM", + "lora_bias" + | "use_dora" + | "use_rslora" + | "use_qalora" + | "fan_in_fan_out" + | "ensure_weight_tying" => value == false, + "rank_pattern" | "alpha_pattern" | "loftq_config" => { + value.as_object().is_some_and(|v| v.is_empty()) + } + "exclude_modules" + | "modules_to_save" + | "layers_to_transform" + | "layers_pattern" + | "layer_replication" + | "target_parameters" + | "trainable_token_indices" + | "alora_invocation_tokens" + | "arrow_config" + | "corda_config" + | "eva_config" + | "megatron_config" => value.is_null(), + // Training/serialization metadata does not change ordinary inference LoRA. + "auto_mapping" | "inference_mode" | "init_lora_weights" | "lora_dropout" + | "megatron_core" | "peft_version" | "qalora_group_size" | "revision" + | "target_modules" => true, + _ => false, + }; + ensure!( + supported, + "unsupported adapter config option {key}: {value}" + ); + } + for (key, value) in [ + ("r", Value::from(16)), + ("lora_alpha", Value::from(32)), + ("peft_type", Value::from("LORA")), + ("bias", Value::from("none")), + ] { + ensure!(config[key] == value, "unsupported adapter config {key}"); + } + let targets = config["target_modules"] + .as_array() + .context("adapter config target_modules must be an array")?; + let expected = BTreeSet::from([ + "up_proj", + "k_proj", + "linear_fc1", + "q_proj", + "linear_fc2", + "down_proj", + "gate_proj", + "o_proj", + "v_proj", + ]); + let actual: BTreeSet<_> = targets.iter().filter_map(Value::as_str).collect(); + ensure!( + actual == expected && targets.len() == expected.len(), + "adapter config requires the full multimodal target_modules" + ); + Ok(Self { + rank: 16, + alpha: 32, + }) + } + fn inventory(&self, base: &Inventory) -> Inventory { + let mut tensors = Inventory::new(); + for (name, shape) in base { + if name.ends_with(".weight") + && (name.contains(".linear_fc1.") || name.contains(".linear_fc2.")) + { + let module = name + .strip_prefix(BASE) + .unwrap() + .strip_suffix(".weight") + .unwrap(); + tensors.insert( + format!("{ADAPTER}{module}.lora_A.weight"), + vec![self.rank, shape[1]], + ); + tensors.insert( + format!("{ADAPTER}{module}.lora_B.weight"), + vec![shape[0], self.rank], + ); + } + } + tensors + } +} + +/// Immutable mmap storage for one base checkpoint and its multimodal adapter. +pub struct VisionCheckpoint { + config: VisionConfig, + adapter: AdapterConfig, + base: TensorStore, + lora: TensorStore, +} +impl VisionCheckpoint { + /// Loads and validates headers on CPU. Keep the files immutable while mapped. + pub fn load(base_dir: impl AsRef, adapter_dir: impl AsRef) -> Result { + let base_dir = fs::canonicalize(base_dir).context("base checkpoint directory")?; + let adapter_dir = fs::canonicalize(adapter_dir).context("adapter checkpoint directory")?; + let config = VisionConfig::load(&base_dir).context("base config")?; + let adapter = AdapterConfig::load(&adapter_dir).context("adapter config")?; + let expected = config.inventory(); + let lora_expected = adapter.inventory(&expected); + let index_path = base_dir.join("model.safetensors.index.json"); + let index = if index_path.try_exists()? { + #[derive(Deserialize)] + struct Index { + weight_map: UniqueMap, + } + let index: Index = serde_json::from_slice(&fs::read(inside( + &base_dir, + "model.safetensors.index.json", + )?)?) + .context("safetensors index")?; + let map = index.weight_map.0; + for name in map.keys().filter(|n| is_visual(n)) { + ensure!( + expected.contains_key(name), + "unexpected visual tensor in index: {name}" + ); + } + for name in expected.keys() { + ensure!( + map.contains_key(name), + "missing visual tensor in index: {name}" + ); + } + Some(map) + } else { + None + }; + let files: BTreeSet = match &index { + Some(index) => expected.keys().map(|n| index[n].clone()).collect(), + None => BTreeSet::from(["model.safetensors".into()]), + }; + let base = TensorStore::load(&base_dir, files, &expected, Dtype::BF16, index.as_ref())?; + let lora = TensorStore::load( + &adapter_dir, + BTreeSet::from(["adapter_model.safetensors".into()]), + &lora_expected, + Dtype::F32, + None, + )?; + Ok(Self { + config, + adapter, + base, + lora, + }) + } + pub fn config(&self) -> &VisionConfig { + &self.config + } + pub fn adapter(&self) -> &AdapterConfig { + &self.adapter + } + pub fn base_names(&self) -> impl Iterator { + self.base.tensors.keys().map(String::as_str) + } + pub fn adapter_names(&self) -> impl Iterator { + self.lora.tensors.keys().map(String::as_str) + } + pub fn base_tensor(&self, name: &str) -> Result> { + self.base.tensor(name) + } + pub fn adapter_tensor(&self, name: &str) -> Result> { + self.lora.tensor(name) + } +} + +struct TensorStore { + maps: Vec, + tensors: BTreeMap, +} +impl TensorStore { + fn load( + dir: &Path, + files: BTreeSet, + expected: &Inventory, + dtype: Dtype, + index: Option<&BTreeMap>, + ) -> Result { + let mut store = Self { + maps: Vec::new(), + tensors: BTreeMap::new(), + }; + for filename in files { + let path = inside(dir, &filename)?; + let file = File::open(&path).with_context(|| format!("open {}", path.display()))?; + // SAFETY: callers must not mutate or truncate checkpoint files while mapped. + let map = unsafe { Mmap::map(&file) } + .with_context(|| format!("mmap safetensors {}", path.display()))?; + validate_header(&map) + .with_context(|| format!("safetensors header {}", path.display()))?; + let (header_len, metadata) = SafeTensors::read_metadata(&map) + .with_context(|| format!("safetensors {}", path.display()))?; + for (name, info) in metadata.tensors() { + if !is_visual(&name) { + continue; + } + let shape = expected + .get(&name) + .with_context(|| format!("unexpected visual tensor {name} in {filename}"))?; + if let Some(index) = index { + ensure!( + index.get(&name) == Some(&filename), + "index mismatch for {name} in {filename}" + ); + } + ensure!( + &info.shape == shape, + "shape mismatch for {name}: {:?}, expected {shape:?}", + info.shape + ); + ensure!( + info.dtype == dtype, + "dtype mismatch for {name}: {:?}, expected {dtype:?}", + info.dtype + ); + let mut info = info.clone(); + info.data_offsets.0 += 8 + header_len; + info.data_offsets.1 += 8 + header_len; + ensure!( + store + .tensors + .insert(name.clone(), (store.maps.len(), info)) + .is_none(), + "duplicate visual tensor {name}" + ); + } + store.maps.push(map); + } + for name in expected.keys() { + ensure!( + store.tensors.contains_key(name), + "missing visual tensor {name}" + ); + } + Ok(store) + } + fn tensor(&self, name: &str) -> Result> { + let (shard, info) = self + .tensors + .get(name) + .with_context(|| format!("unknown visual tensor {name}"))?; + Ok(TensorView::new( + info.dtype, + info.shape.clone(), + &self.maps[*shard][info.data_offsets.0..info.data_offsets.1], + )?) + } +} +// Bound all offsets before safetensors 0.8 adds payload size to header size: +// its final length check uses unchecked addition, even for ignored language tensors. +fn validate_header(bytes: &[u8]) -> Result<()> { + let length_bytes = bytes.get(..8).context("missing header length")?; + let header_len = usize::try_from(u64::from_le_bytes(length_bytes.try_into()?))?; + // Match safetensors 0.8's header allocation limit. + ensure!(header_len <= 100_000_000, "header too large"); + let data_start = header_len + .checked_add(8) + .context("header length overflow")?; + let header = bytes.get(8..data_start).context("truncated header")?; + let payload_len = bytes.len() - data_start; + let entries: UniqueMap = serde_json::from_slice(header)?; + for (name, entry) in entries.0 { + if name == "__metadata__" { + continue; + } + let (start, end): (usize, usize) = serde_json::from_value(entry["data_offsets"].clone()) + .with_context(|| format!("invalid offsets for {name}"))?; + ensure!( + start <= end && end <= payload_len, + "tensor offsets exceed payload for {name}" + ); + } + Ok(()) +} +fn is_visual(name: &str) -> bool { + name.split('.').any(|part| part == "visual") +} +fn inside(dir: &Path, name: &str) -> Result { + ensure!( + !name.is_empty() + && Path::new(name) + .components() + .all(|c| matches!(c, Component::Normal(_))), + "path must remain inside checkpoint directory: {name}" + ); + let resolved = + fs::canonicalize(dir.join(name)).with_context(|| format!("checkpoint file {name}"))?; + ensure!( + resolved.starts_with(dir), + "path escapes checkpoint directory: {name}" + ); + Ok(resolved) +} + +// serde_json's ordinary maps overwrite repeated keys. Reject ambiguous headers/indexes. +struct UniqueMap(BTreeMap); +impl<'de, T: Deserialize<'de>> Deserialize<'de> for UniqueMap { + fn deserialize>( + deserializer: D, + ) -> std::result::Result { + struct UniqueVisitor(std::marker::PhantomData); + impl<'de, T: Deserialize<'de>> Visitor<'de> for UniqueVisitor { + type Value = UniqueMap; + fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("an object with unique names") + } + fn visit_map>( + self, + mut map: M, + ) -> std::result::Result { + let mut entries = BTreeMap::new(); + while let Some((key, value)) = map.next_entry::()? { + if entries.insert(key.clone(), value).is_some() { + return Err(de::Error::custom(format!("duplicate JSON key: {key}"))); + } + } + Ok(UniqueMap(entries)) + } + } + deserializer.deserialize_map(UniqueVisitor(std::marker::PhantomData)) + } +} + +mod geometry; +mod model; +pub use model::VisionModel; diff --git a/src/models/cua_s1/native/src/vision/model.rs b/src/models/cua_s1/native/src/vision/model.rs new file mode 100644 index 00000000..d89833ef --- /dev/null +++ b/src/models/cua_s1/native/src/vision/model.rs @@ -0,0 +1,407 @@ +//! GPU vision forward for the structurally validated, unmerged checkpoint. +use super::{VisionCheckpoint, geometry::Geometry}; +use crate::{ + cuda::{self, DeviceBuffer, Stream, api, check}, + image_preprocess::ProcessedImage, +}; +use anyhow::{Result, ensure}; +use half::bf16; +use std::{collections::BTreeMap, ffi::c_void, path::Path}; + +type Trace<'a> = Option<&'a mut dyn FnMut(&str, &[bf16]) -> Result<()>>; + +struct OwnedStream(Stream); +impl Drop for OwnedStream { + fn drop(&mut self) { + unsafe { (api().cs1_stream_destroy)(self.0) }; + } +} +struct Gemm(*mut c_void); +// SAFETY: the model uses its GEMM handle and stream serially. +unsafe impl Send for Gemm {} +impl Drop for Gemm { + fn drop(&mut self) { + unsafe { (api().cs1_gemm_destroy)(self.0) }; + } +} + +/// The base remains BF16; all 50 LoRA A/B pairs remain FP32, applied at scale 2. +/// A model is used by one request at a time and returns row-major [image_tokens,2560]. +pub struct VisionModel { + base: BTreeMap, + lora: BTreeMap, + stream: OwnedStream, + gemm: Gemm, +} +impl VisionModel { + pub fn load(base: impl AsRef, adapter: impl AsRef, library: &Path) -> Result { + let checkpoint = VisionCheckpoint::load(base, adapter)?; + cuda::load(library)?; + cuda::set_device(0)?; + let stream = OwnedStream(cuda::new_stream()?); + let gemm = Gemm(unsafe { (api().cs1_gemm_create)(32 << 20) }); + ensure!(!gemm.0.is_null(), "cannot create vision cuBLAS handle"); + let mut model = Self { + base: BTreeMap::new(), + lora: BTreeMap::new(), + stream, + gemm, + }; + for name in checkpoint.base_names() { + let tensor = checkpoint.base_tensor(name)?; + model.base.insert( + name.strip_prefix("model.visual.").unwrap().into(), + model.upload(tensor.data())?, + ); + } + for name in checkpoint.adapter_names() { + let tensor = checkpoint.adapter_tensor(name)?; + model.lora.insert( + name.strip_prefix("base_model.model.model.visual.") + .unwrap() + .into(), + model.upload(tensor.data())?, + ); + } + model + .base + .insert("__zero_bias".into(), model.upload(&[0; 2048])?); + Ok(model) + } + fn upload(&self, bytes: &[u8]) -> Result { + let buffer = DeviceBuffer::new(bytes.len())?; + unsafe { + cuda::upload(buffer.at(0), bytes, self.stream.0)?; + } + Ok(buffer) + } + fn weight(&self, name: &str) -> *const c_void { + self.base[name].at(0) + } + fn norm(&self, name: &str, x: &DeviceBuffer, y: &DeviceBuffer, rows: usize) -> Result<()> { + unsafe { + check( + (api().cs1_vision_norm)( + x.at(0), + self.weight(&format!("{name}.weight")), + self.weight(&format!("{name}.bias")), + y.at(0), + rows as i32, + 1024, + self.stream.0, + ), + name, + ) + } + } + #[allow(clippy::too_many_arguments)] // Mirrors the fixed-shape GEMM operation. + fn linear( + &self, + name: &str, + x: &DeviceBuffer, + y: &DeviceBuffer, + rows: usize, + output: usize, + input: usize, + work: &Work, + ) -> Result<()> { + // SAFETY: each caller provides rows*input/output sized allocations. All dimensions + // derive from validated fixed architecture and bounded image geometry. + let bias_name = if name == "patch_embed.proj" { + "__zero_bias".into() + } else { + format!("{name}.bias") + }; + unsafe { + check( + (api().cs1_vision_linear)( + self.gemm.0, + x.at(0), + self.weight(&format!("{name}.weight")), + self.weight(&bias_name), + y.at(0), + rows as i32, + output as i32, + input as i32, + self.stream.0, + ), + name, + )?; + if name == "patch_embed.proj" { + check( + (api().cs1_vision_bias)( + y.at(0), + self.weight("patch_embed.proj.bias"), + rows * output, + output as i32, + self.stream.0, + ), + "vision patch bias", + )?; + } + if let Some(a) = self.lora.get(&format!("{name}.lora_A.weight")) { + let b = &self.lora[&format!("{name}.lora_B.weight")]; + check( + (api().cs1_vision_to_float)( + x.at(0), + work.float_input.at(0).cast(), + rows * input, + self.stream.0, + ), + "vision LoRA input", + )?; + check( + (api().cs1_gemm_f32)( + self.gemm.0, + work.float_input.at(0).cast(), + a.at(0).cast(), + work.rank.at(0).cast(), + rows as i32, + 16, + input as i32, + self.stream.0, + ), + "vision LoRA A", + )?; + check( + (api().cs1_gemm_f32)( + self.gemm.0, + work.rank.at(0).cast(), + b.at(0).cast(), + work.delta.at(0).cast(), + rows as i32, + output as i32, + 16, + self.stream.0, + ), + "vision LoRA B", + )?; + check( + (api().cs1_vision_lora_add)( + y.at(0), + work.delta.at(0).cast(), + rows * output, + 2., + self.stream.0, + ), + "vision LoRA add", + )?; + } + } + Ok(()) + } + fn read(&self, x: &DeviceBuffer, n: usize) -> Result> { + let mut bytes = vec![0u8; n * 2]; + unsafe { + cuda::download(&mut bytes, x.at(0), self.stream.0)?; + } + Ok(bytes + .as_chunks::<2>() + .0 + .iter() + .map(|b| bf16::from_bits(u16::from_le_bytes([b[0], b[1]]))) + .collect()) + } + fn trace( + &self, + callback: &mut Trace<'_>, + name: &str, + x: &DeviceBuffer, + n: usize, + ) -> Result<()> { + if let Some(callback) = callback.as_mut() { + callback(name, &self.read(x, n)?)?; + } + Ok(()) + } + pub fn synchronize(&self) -> Result<()> { + cuda::set_device(0)?; + cuda::synchronize(self.stream.0) + } + + pub fn forward(&mut self, image: &ProcessedImage) -> Result> { + self.run(image, None) + } + /// Optional synchronized stage downloads for parity diagnosis; ordinary forward skips them. + pub fn forward_with_trace( + &mut self, + image: &ProcessedImage, + mut callback: impl FnMut(&str, &[bf16]) -> Result<()>, + ) -> Result> { + self.run(image, Some(&mut callback)) + } + fn run(&mut self, image: &ProcessedImage, mut callback: Trace<'_>) -> Result> { + cuda::set_device(0)?; + let geo = Geometry::new(image.image_grid_thw)?; + let n = image.image_grid_thw[1] * image.image_grid_thw[2]; + ensure!( + image.pixel_values.len() == n * 1536, + "vision pixel_values length does not match grid" + ); + ensure!( + image.resized_height == image.image_grid_thw[1] * 16 + && image.resized_width == image.image_grid_thw[2] * 16, + "vision resized geometry does not match grid" + ); + ensure!( + image.pixel_values.iter().all(|v| v.is_finite()), + "vision pixels must be finite" + ); + let pixels: Vec = image + .pixel_values + .iter() + .flat_map(|v| bf16::from_f32(*v).to_bits().to_le_bytes()) + .collect(); + let pixels = self.upload(&pixels)?; + let indices = self.upload( + &geo.indices + .iter() + .flat_map(|v| v.to_le_bytes()) + .collect::>(), + )?; + let weights = self.upload( + &geo.weights + .iter() + .flat_map(|v| v.to_le_bytes()) + .collect::>(), + )?; + let co = self.upload( + &geo.cos + .iter() + .flat_map(|v| v.to_le_bytes()) + .collect::>(), + )?; + let si = self.upload( + &geo.sin + .iter() + .flat_map(|v| v.to_le_bytes()) + .collect::>(), + )?; + let w = Work::new(n)?; + let x = DeviceBuffer::new(n * 1024 * 2)?; + let norm = DeviceBuffer::new(n * 1024 * 2)?; + let qkv = DeviceBuffer::new(n * 3072 * 2)?; + let q = DeviceBuffer::new(n * 1024 * 2)?; + let k = DeviceBuffer::new(n * 1024 * 2)?; + let attn = DeviceBuffer::new(n * 1024 * 2)?; + let delta = DeviceBuffer::new(n * 1024 * 2)?; + let mlp = DeviceBuffer::new(n * 4096 * 2)?; + self.linear("patch_embed.proj", &pixels, &x, n, 1024, 1536, &w)?; + self.trace(&mut callback, "patch_embed", &x, n * 1024)?; + unsafe { + check( + (api().cs1_vision_position)( + x.at(0), + self.weight("pos_embed.weight"), + indices.at(0).cast(), + weights.at(0).cast(), + n as i32, + self.stream.0, + ), + "vision learned positions", + )?; + } + self.trace(&mut callback, "position", &x, n * 1024)?; + for i in 0..24 { + let p = format!("blocks.{i}"); + self.norm(&format!("{p}.norm1"), &x, &norm, n)?; + self.linear(&format!("{p}.attn.qkv"), &norm, &qkv, n, 3072, 1024, &w)?; + unsafe { + check( + (api().cs1_vision_rope)( + qkv.at(0), + co.at(0).cast(), + si.at(0).cast(), + q.at(0), + k.at(0), + n as i32, + self.stream.0, + ), + "vision rotary", + )?; + check( + (api().cs1_vision_attention)( + q.at(0), + k.at(0), + qkv.at(2048 * 2), + attn.at(0), + n as i32, + self.stream.0, + ), + "vision attention", + )?; + } + self.linear(&format!("{p}.attn.proj"), &attn, &delta, n, 1024, 1024, &w)?; + unsafe { + check( + (api().cs1_vision_add)(x.at(0), delta.at(0), n * 1024, self.stream.0), + "vision attention residual", + )?; + } + self.norm(&format!("{p}.norm2"), &x, &norm, n)?; + self.linear( + &format!("{p}.mlp.linear_fc1"), + &norm, + &mlp, + n, + 4096, + 1024, + &w, + )?; + unsafe { + check( + (api().cs1_vision_gelu)(mlp.at(0), n * 4096, 0, self.stream.0), + "vision tanh GELU", + )?; + } + self.linear( + &format!("{p}.mlp.linear_fc2"), + &mlp, + &delta, + n, + 1024, + 4096, + &w, + )?; + unsafe { + check( + (api().cs1_vision_add)(x.at(0), delta.at(0), n * 1024, self.stream.0), + "vision MLP residual", + )?; + } + self.trace(&mut callback, &p, &x, n * 1024)?; + } + self.norm("merger.norm", &x, &norm, n)?; + self.trace(&mut callback, "merger.norm", &norm, n * 1024)?; + // Consecutive groups of four patches already have the required 2x2 merge order. + self.linear("merger.linear_fc1", &norm, &mlp, n / 4, 4096, 4096, &w)?; + self.trace(&mut callback, "merger.linear_fc1", &mlp, n * 1024)?; + unsafe { + check( + (api().cs1_vision_gelu)(mlp.at(0), n * 1024, 1, self.stream.0), + "vision exact GELU", + )?; + } + let out = DeviceBuffer::new(n / 4 * 2560 * 2)?; + self.linear("merger.linear_fc2", &mlp, &out, n / 4, 2560, 4096, &w)?; + let result = self.read(&out, n / 4 * 2560)?; + if let Some(callback) = callback.as_mut() { + callback("merger.output", &result)?; + } + Ok(result) + } +} +struct Work { + float_input: DeviceBuffer, + rank: DeviceBuffer, + delta: DeviceBuffer, +} +impl Work { + fn new(n: usize) -> Result { + Ok(Self { + float_input: DeviceBuffer::new(n * 4096 * 4)?, + rank: DeviceBuffer::new(n * 16 * 4)?, + delta: DeviceBuffer::new(n * 4096 * 4)?, + }) + } +} diff --git a/src/models/cua_s1/native/src/vision_engine.rs b/src/models/cua_s1/native/src/vision_engine.rs new file mode 100644 index 00000000..9d621767 --- /dev/null +++ b/src/models/cua_s1/native/src/vision_engine.rs @@ -0,0 +1,231 @@ +//! Native screenshot-to-decision orchestration, with request-local vision reuse. +pub use crate::vision_processing::PreparedRequest; +use crate::vision_processing::{VisionProcessor, probabilities}; +use crate::{ + contract::{LETTERS, Question}, + image_request::ImageRequest, + inputs::MultimodalInput, + model::Model, + multimodal::ImagePrompt, + vision::VisionModel, +}; +use anyhow::{Context, Result, ensure}; +use half::bf16; +use safetensors::{Dtype, SafeTensors}; +use serde_json::Value; +use std::sync::Arc; +use std::{fs::File, path::Path}; +use tokenizers::Tokenizer; + +pub struct VisionEngine { + pub processor: Arc, + pub vision: VisionModel, + language: Model, + letters: Vec, +} + +pub struct Readout { + pub hidden: Vec, + pub logits: Vec, + pub probabilities: Vec, +} + +impl VisionEngine { + pub fn load(base: &Path, adapter: &Path, language: &Path, library: &Path) -> Result { + let marker: Value = serde_json::from_slice( + &std::fs::read(language.join("cua_s1_language_export.json")) + .context("export the multimodal language checkpoint first")?, + )?; + ensure!( + marker["format"] == "cua-s1-multimodal-language-merged/1" + && marker["base_revision"] == "851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a" + && marker["adapter_revision"] == "16818868b0cc7813808aae4e87b417657046ab79", + "expected pinned multimodal language export" + ); + crate::provenance::verify_sources(base, adapter)?; + crate::provenance::verify_export(language, &marker)?; + let tokenizer = + Tokenizer::from_file(language.join("tokenizer.json")).map_err(anyhow::Error::msg)?; + let ids = LETTERS + .chars() + .map(|c| { + let enc = tokenizer + .encode(c.to_string(), false) + .map_err(anyhow::Error::msg)?; + ensure!(enc.len() == 1, "each candidate must be one token"); + Ok(enc.get_ids()[0]) + }) + .collect::>>()?; + let model = Model::load(language, library)?; + ensure!( + tokenizer.token_to_id("<|image_pad|>") == model.cfg.image_token_id + && model.cfg.image_token_id.is_some(), + "tokenizer/image-token mismatch" + ); + let letters = letter_rows(language, &ids, model.cfg.hidden)?; + let vision = VisionModel::load(base, adapter, library)?; + Ok(Self { + processor: Arc::new(VisionProcessor::new( + tokenizer, + model.cfg.image_token_id.unwrap(), + )), + vision, + language: model, + letters, + }) + } + + /// Validate all tokenized prompts before the first CUDA forward of a request. + pub fn prepare( + &self, + width: usize, + height: usize, + rgb: &[u8], + questions: &[Question], + ) -> Result { + self.processor.prepare(width, height, rgb, questions) + } + + fn raw_score( + &mut self, + prompt: &ImagePrompt, + features: &[bf16], + options: usize, + ) -> Result<(Vec, Vec)> { + ensure!((1..=26).contains(&options), "expected 1 to 26 options"); + let hidden = self.language.forward_multimodal(&MultimodalInput { + token_ids: &prompt.token_ids, + image_token_indices: &prompt.image_token_indices, + image_embeddings: features, + position_ids: [ + &prompt.position_ids[0], + &prompt.position_ids[1], + &prompt.position_ids[2], + ], + })?; + let logits: Vec = self + .letters + .chunks_exact(hidden.len()) + .take(options) + .map(|w| { + w.iter() + .zip(&hidden) + .map(|(&a, &b)| a as f64 * b as f64) + .sum::() as f32 + }) + .collect(); + ensure!( + logits.iter().all(|x| x.is_finite()), + "non-finite candidate logits" + ); + Ok((hidden, logits)) + } + + /// Diagnostic readout, including processor normalization of raw model logits. + pub fn score( + &mut self, + prompt: &ImagePrompt, + features: &[bf16], + options: usize, + ) -> Result { + let (hidden, logits) = self.raw_score(prompt, features, options)?; + let probabilities = probabilities(&logits)?; + Ok(Readout { + hidden, + logits, + probabilities, + }) + } + + pub fn predict(&mut self, request: &ImageRequest) -> Result { + let prepared = self.prepare( + request.width, + request.height, + &request.rgb, + &request.questions, + )?; + self.predict_prepared(&prepared) + } + + /// One admitted unit covers vision encoding and every question, keeping + /// features request-local and model resources alive until both streams finish. + pub fn execute_prepared(&mut self, prepared: &PreparedRequest) -> Result>> { + let result = (|| { + let counts: Vec = prepared.context.option_counts().collect(); + ensure!( + prepared.prompts.len() == counts.len() && !counts.is_empty(), + "question/prompt count mismatch" + ); + let features = self.vision.forward(&prepared.image)?; + prepared + .prompts + .iter() + .zip(counts) + .map(|(p, count)| { + self.raw_score(p, &features, count) + .map(|(_, logits)| logits) + }) + .collect() + })(); + // Also finish queued work on an error before releasing runtime admission. + let vision_sync = self.vision.synchronize(); + let language_sync = self.language.synchronize(); + vision_sync?; + language_sync?; + result + } + + pub fn predict_prepared(&mut self, prepared: &PreparedRequest) -> Result { + let rows = self.execute_prepared(prepared)?; + prepared.context.finish(rows) + } +} + +fn letter_rows(dir: &Path, ids: &[u32], hidden: usize) -> Result> { + let names = [ + "embed_tokens.weight", + "model.embed_tokens.weight", + "model.language_model.embed_tokens.weight", + ]; + let (path, requested) = if dir.join("model.safetensors.index.json").exists() { + let index: Value = + serde_json::from_slice(&std::fs::read(dir.join("model.safetensors.index.json"))?)?; + names + .iter() + .find_map(|name| { + index["weight_map"][name] + .as_str() + .map(|file| (dir.join(file), Some(*name))) + }) + .context("embedding missing from index")? + } else { + (dir.join("model.safetensors"), None) + }; + let file = File::open(path)?; + // SAFETY: exported checkpoint files must remain immutable during inference. + let map = unsafe { memmap2::Mmap::map(&file)? }; + let tensors = SafeTensors::deserialize(&map)?; + let name = requested + .or_else(|| names.iter().copied().find(|n| tensors.tensor(n).is_ok())) + .context("missing embedding")?; + let tensor = tensors.tensor(name)?; + ensure!( + tensor.dtype() == Dtype::BF16 && tensor.shape().len() == 2 && tensor.shape()[1] == hidden, + "embedding shape/dtype mismatch" + ); + let mut result = Vec::with_capacity(ids.len() * hidden); + for &id in ids { + ensure!( + (id as usize) < tensor.shape()[0], + "candidate outside vocabulary" + ); + let row = &tensor.data()[id as usize * hidden * 2..(id as usize + 1) * hidden * 2]; + result.extend( + row.as_chunks::<2>() + .0 + .iter() + .map(|b| bf16::from_le_bytes(*b).to_f32()), + ); + } + Ok(result) +} diff --git a/src/models/cua_s1/native/src/vision_main.rs b/src/models/cua_s1/native/src/vision_main.rs new file mode 100644 index 00000000..83664898 --- /dev/null +++ b/src/models/cua_s1/native/src/vision_main.rs @@ -0,0 +1,153 @@ +//! Screenshot `/v1/systemone` worker using native RGB, vision and language stages. +use anyhow::{Context, Result}; +use axum::{ + Json, Router, + body::Bytes, + extract::{DefaultBodyLimit, State, rejection::BytesRejection}, + http::StatusCode, + response::{IntoResponse, Response}, + routing::{get, post}, +}; +use omni_cua_s1_native::{ + contract, cuda, + image_request::{self, MODEL_ID}, + vision_engine::VisionEngine, + vision_processing::VisionProcessor, +}; +use serde_json::{Value, json}; +use std::{ + path::PathBuf, + sync::{Arc, Mutex}, +}; + +#[derive(Clone)] +struct Shared { + processor: Arc, + executor: Arc>, + scheduler: omni_runtime::SerialScheduler, +} +fn reply(status: StatusCode, value: Value) -> Response { + (status, Json(value)).into_response() +} + +async fn decide(State(engine): State, body: Result) -> Response { + let raw = match body { + Ok(b) => b, + Err(e) => return reply(e.status(), json!({"detail": e.body_text()})), + }; + let body = match contract::parse_body(&raw) { + Ok(b) => b, + Err(e) => { + return reply( + StatusCode::from_u16(e.status).unwrap(), + json!({"detail":e.message}), + ); + } + }; + // CPU decode/preparation does not hold the model mutex or admission permit. + let processor = engine.processor.clone(); + let prepared = match tokio::task::spawn_blocking(move || { + let request = image_request::parse_image_body(&body)?; + processor.prepare( + request.width, + request.height, + &request.rgb, + &request.questions, + ) + }) + .await + { + Ok(Ok(prepared)) => prepared, + Ok(Err(e)) => { + return reply( + StatusCode::UNPROCESSABLE_ENTITY, + json!({"detail": e.to_string()}), + ); + } + Err(e) => { + eprintln!("native image preparation failed: {e}"); + return reply( + StatusCode::INTERNAL_SERVER_ERROR, + json!({"detail":"inference failed"}), + ); + } + }; + let executor = engine.executor.clone(); + let result = engine + .scheduler + .run(move || { + let rows = executor + .lock() + .map_err(|_| anyhow::anyhow!("poisoned engine"))? + .execute_prepared(&prepared)?; + Ok((prepared.context, rows)) + }) + .await + .and_then(|(context, rows)| context.finish(rows)); + match result { + Ok(body) => reply(StatusCode::OK, body), + Err(e) => { + eprintln!("native image inference failed: {e:#}"); + reply( + StatusCode::INTERNAL_SERVER_ERROR, + json!({"detail":"inference failed"}), + ) + } + } +} + +#[tokio::main] +async fn main() -> Result<()> { + let path = |name| { + std::env::var_os(name) + .map(PathBuf::from) + .with_context(|| format!("set {name}")) + }; + let library = std::env::var_os("CUA_S1_CUDA_LIB") + .map(PathBuf::from) + .map_or_else(cuda::default_library, Ok)?; + let (base, adapter, language) = ( + path("CUA_S1_BASE")?, + path("CUA_S1_VISION_ADAPTER")?, + path("CUA_S1_MODEL")?, + ); + let mut executor = tokio::task::spawn_blocking(move || { + VisionEngine::load(&base, &adapter, &language, &library) + }) + .await??; + let warmup = executor.prepare( + 1, + 1, + &[0, 0, 0], + &[contract::Question { + name: "warmup".into(), + goal: String::new(), + keys: vec!["continue".into()], + labels: vec!["Continue".into()], + }], + )?; + let engine = tokio::task::spawn_blocking(move || -> Result { + executor.predict_prepared(&warmup)?; + Ok(executor) + }) + .await? + .map(|executor| Shared { + processor: executor.processor.clone(), + executor: Arc::new(Mutex::new(executor)), + scheduler: omni_runtime::SerialScheduler::default(), + })?; + let host = std::env::var("CUA_S1_HOST").unwrap_or_else(|_| "127.0.0.1".into()); + let port: u16 = std::env::var("CUA_S1_PORT").map_or(Ok(8000), |p| p.parse())?; + let app = Router::new() + .route( + "/health", + get(|| async { Json(json!({"status":"ready", "model": MODEL_ID})) }), + ) + .route("/v1/systemone", post(decide)) + .layer(DefaultBodyLimit::max(image_request::MAX_BODY)) + .with_state(engine); + let listener = tokio::net::TcpListener::bind((host.as_str(), port)).await?; + println!("native vision worker listening on {host}:{port}"); + axum::serve(listener, app).await?; + Ok(()) +} diff --git a/src/models/cua_s1/native/src/vision_processing.rs b/src/models/cua_s1/native/src/vision_processing.rs new file mode 100644 index 00000000..9d765367 --- /dev/null +++ b/src/models/cua_s1/native/src/vision_processing.rs @@ -0,0 +1,92 @@ +//! Screenshot preparation and response interpretation, independent of execution. +use crate::{ + contract::{self, Question}, + image_preprocess::{ProcessedImage, preprocess_rgb8}, + image_request::MODEL_ID, + multimodal::{ImagePrompt, prepare_prompt}, +}; +use anyhow::{Result, ensure}; +use serde_json::{Value, json}; +use tokenizers::Tokenizer; + +pub struct VisionProcessor { + tokenizer: Tokenizer, + image_token: u32, +} + +pub struct PreparedRequest { + pub image: ProcessedImage, + pub prompts: Vec, + pub context: ResponseContext, +} + +pub struct ResponseContext { + questions: Vec, + input_tokens: usize, +} + +impl VisionProcessor { + pub fn new(tokenizer: Tokenizer, image_token: u32) -> Self { + Self { + tokenizer, + image_token, + } + } + + /// Validate every question and prompt before submitting model work. + pub fn prepare( + &self, + width: usize, + height: usize, + rgb: &[u8], + questions: &[Question], + ) -> Result { + crate::multimodal::validate_questions(questions)?; + let image = preprocess_rgb8(width, height, rgb)?; + let prompts = questions + .iter() + .map(|q| prepare_prompt(&self.tokenizer, q, image.image_grid_thw, self.image_token)) + .collect::>>()?; + let input_tokens = prompts.iter().map(|p| p.token_ids.len()).sum(); + Ok(PreparedRequest { + image, + prompts, + context: ResponseContext { + questions: questions.to_vec(), + input_tokens, + }, + }) + } +} + +impl ResponseContext { + pub fn option_counts(&self) -> impl Iterator + '_ { + self.questions.iter().map(|q| q.keys.len()) + } + + pub fn finish(&self, logits: Vec>) -> Result { + ensure!( + logits.len() == self.questions.len(), + "question/output count mismatch" + ); + let mut answers = serde_json::Map::new(); + for (q, row) in self.questions.iter().zip(logits) { + ensure!(row.len() == q.keys.len(), "candidate/output count mismatch"); + answers.insert(q.name.clone(), contract::answer(q, &probabilities(&row)?)); + } + Ok( + json!({"model": MODEL_ID, "answers": answers, "usage": {"input_tokens": self.input_tokens, "output_tokens": 0}}), + ) + } +} + +pub(crate) fn probabilities(logits: &[f32]) -> Result> { + ensure!( + !logits.is_empty() && logits.iter().all(|x| x.is_finite()), + "non-finite or empty candidate logits" + ); + let max = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max) as f64; + let exps: Vec = logits.iter().map(|&x| (x as f64 - max).exp()).collect(); + let total: f64 = exps.iter().sum(); + Ok(exps.iter().map(|x| (x / total) as f32).collect()) +} diff --git a/src/models/qwen3_5/native/src/cuda.rs b/src/models/qwen3_5/native/src/cuda.rs index 7433ee67..3de05686 100644 --- a/src/models/qwen3_5/native/src/cuda.rs +++ b/src/models/qwen3_5/native/src/cuda.rs @@ -9,7 +9,7 @@ use std::sync::OnceLock; use anyhow::{Context, Result, bail, ensure}; /// `CS1_ABI_VERSION` in ops.h. -const ABI_VERSION: u32 = 4; +const ABI_VERSION: u32 = 5; pub const LIBRARY: &str = "libqwen3_5_cuda.so"; /// A `cudaStream_t`. @@ -48,6 +48,17 @@ macro_rules! api { } api! { + cs1_vision_linear(gemm: *mut c_void, x: *const c_void, w: *const c_void, bias: *const c_void, y: *mut c_void, m: c_int, n: c_int, k: c_int, stream: Stream) -> c_int; + cs1_gemm_f32(gemm: *mut c_void, x: *const f32, w: *const f32, y: *mut f32, m: c_int, n: c_int, k: c_int, stream: Stream) -> c_int; + cs1_vision_norm(x: *const c_void, w: *const c_void, b: *const c_void, y: *mut c_void, rows: c_int, d: c_int, stream: Stream) -> c_int; + cs1_vision_position(x: *mut c_void, table: *const c_void, indices: *const i32, weights: *const f32, n: c_int, stream: Stream) -> c_int; + cs1_vision_rope(qkv: *const c_void, co: *const f32, si: *const f32, q: *mut c_void, k: *mut c_void, n: c_int, stream: Stream) -> c_int; + cs1_vision_attention(q: *const c_void, k: *const c_void, v: *const c_void, out: *mut c_void, n: c_int, stream: Stream) -> c_int; + cs1_vision_bias(x: *mut c_void, bias: *const c_void, n: usize, d: c_int, stream: Stream) -> c_int; + cs1_vision_gelu(x: *mut c_void, n: usize, exact: c_int, stream: Stream) -> c_int; + cs1_vision_add(x: *mut c_void, delta: *const c_void, n: usize, stream: Stream) -> c_int; + cs1_vision_to_float(x: *const c_void, out: *mut f32, n: usize, stream: Stream) -> c_int; + cs1_vision_lora_add(x: *mut c_void, delta: *const f32, n: usize, scale: f32, stream: Stream) -> c_int; cs1_abi_version() -> u32; cs1_error_string(code: c_int) -> *const c_char; cs1_set_device(device: c_int) -> c_int; @@ -55,6 +66,7 @@ api! { cs1_free(ptr: *mut c_void) -> c_int; cs1_stream_create(stream: *mut Stream) -> c_int; cs1_stream_sync(stream: Stream) -> c_int; + cs1_stream_destroy(stream: Stream) -> c_int; cs1_graph_begin(stream: Stream) -> c_int; cs1_graph_end(stream: Stream, exec: *mut *mut c_void) -> c_int; cs1_graph_launch(exec: *mut c_void, stream: Stream) -> c_int; diff --git a/src/models/qwen3_5/native/src/inputs.rs b/src/models/qwen3_5/native/src/inputs.rs new file mode 100644 index 00000000..bf719eca --- /dev/null +++ b/src/models/qwen3_5/native/src/inputs.rs @@ -0,0 +1,95 @@ +//! Batch-one, unpadded inputs at the adapted-vision / language-model boundary. + +use anyhow::{Result, ensure}; +use half::bf16; + +/// Image rows are already adapted to the language hidden size, in placeholder +/// order. Positions are the T/H/W slices of an int64 `[3, 1, sequence]` tensor. +/// The caller owns preprocessing, vision execution and the position calculation. +pub struct MultimodalInput<'a> { + pub token_ids: &'a [u32], + pub image_token_indices: &'a [usize], + pub image_embeddings: &'a [bf16], + pub position_ids: [&'a [i64]; 3], +} + +impl MultimodalInput<'_> { + /// Check the entire boundary before allocating buffers or launching CUDA. + pub fn validate( + &self, + hidden: usize, + vocab: usize, + image_token: u32, + max_position: usize, + ) -> Result<()> { + let t = self.token_ids.len(); + ensure!(t > 0 && t <= max_position, "empty or oversized prompt"); + ensure!( + self.token_ids.iter().all(|&id| (id as usize) < vocab), + "token id outside the vocabulary" + ); + let expected: Vec = self + .token_ids + .iter() + .enumerate() + .filter_map(|(i, &id)| (id == image_token).then_some(i)) + .collect(); + ensure!( + self.image_token_indices == expected, + "image indices must exactly match the ordered placeholders" + ); + ensure!( + Some(self.image_embeddings.len()) == expected.len().checked_mul(hidden), + "image embedding shape mismatch" + ); + ensure!( + self.image_embeddings.iter().all(|x| x.is_finite()), + "non-finite image embedding" + ); + ensure!( + self.position_ids.iter().all(|axis| axis.len() == t), + "position_ids must have shape [3, 1, sequence]" + ); + ensure!( + self.position_ids + .iter() + .flat_map(|axis| axis.iter()) + .all(|&p| p >= 0 && (p as u64) < max_position as u64), + "position outside the configured range" + ); + Ok(()) + } +} + +/// Qwen3.5's interleaved recomposition: overwrite H at 1::3 and W at 2::3 up +/// to section[axis] * 3, retaining T elsewhere. The second rotary half repeats +/// these frequencies, which the existing attention-prep kernel handles. +/// Float32 inverse frequencies/products and host float64 trig preserve the +/// original native text table rounding when all three axes are equal. +pub(crate) fn rotary_tables( + positions: [&[i64]; 3], + half: usize, + theta: f64, + sections: [usize; 3], +) -> (Vec, Vec) { + let inv: Vec = (0..half) + .map(|i| 1f32 / (theta as f32).powf((2 * i) as f32 / (2 * half) as f32)) + .collect(); + let mut cos = Vec::with_capacity(positions[0].len() * half * 2); + let mut sin = Vec::with_capacity(cos.capacity()); + for (t, _) in positions[0].iter().enumerate() { + for (i, &f) in inv.iter().enumerate() { + let axis = if i % 3 == 1 && i < sections[1] * 3 { + 1 + } else if i % 3 == 2 && i < sections[2] * 3 { + 2 + } else { + 0 + }; + let angle = (f * positions[axis][t] as f32) as f64; + cos.extend(bf16::from_f32(angle.cos() as f32).to_le_bytes()); + sin.extend(bf16::from_f32(angle.sin() as f32).to_le_bytes()); + } + } + (cos, sin) +} diff --git a/src/models/qwen3_5/native/src/lib.rs b/src/models/qwen3_5/native/src/lib.rs index 929435b5..04b1f87f 100644 --- a/src/models/qwen3_5/native/src/lib.rs +++ b/src/models/qwen3_5/native/src/lib.rs @@ -3,3 +3,5 @@ pub mod cuda; pub mod json; pub mod model; + +pub mod inputs; diff --git a/src/models/qwen3_5/native/src/model.rs b/src/models/qwen3_5/native/src/model.rs index c8ab5d60..87952e6a 100644 --- a/src/models/qwen3_5/native/src/model.rs +++ b/src/models/qwen3_5/native/src/model.rs @@ -6,8 +6,8 @@ //! The order of operations follows `modeling_qwen3_5.py`, and so do the points where //! it rounds to bfloat16, except inside attention and the Gated DeltaNet prefill (see //! their kernels). Text prompts use one position per -//! token, so the multimodal rotary sections all get the same position and the -//! rotary embedding is the plain one. +//! token. The explicit multimodal boundary inserts adapted image rows and supplies +//! the interleaved temporal/height/width rotary positions to the same layer loop. use std::collections::{HashMap, VecDeque}; use std::ffi::c_void; @@ -17,6 +17,7 @@ use anyhow::{Context, Result, bail, ensure}; use serde_json::Value as Json; use crate::cuda::{self, DeviceBuffer, Stream, check}; +use crate::inputs::{MultimodalInput, rotary_tables}; const ALIGN: usize = 256; const BF16: usize = 2; @@ -35,6 +36,9 @@ pub struct Config { /// Half the number of rotary dims (rotate_half pairs dim i with dim i + half). pub rotary_half: usize, pub rope_theta: f64, + pub mrope_section: [usize; 3], + pub max_positions: usize, + pub image_token_id: Option, pub lin_k_heads: usize, pub lin_v_heads: usize, pub lin_k_dim: usize, @@ -69,6 +73,31 @@ impl Config { other => bail!("unknown layer type {other:?}"), }) .collect::>>()?; + let sections = rope + .get("mrope_section") + .cloned() + .unwrap_or(serde_json::json!([11, 11, 10])); + let sections = sections + .as_array() + .context("mrope_section must be an array")? + .iter() + .map(|v| { + v.as_u64() + .and_then(|n| usize::try_from(n).ok()) + .context("mrope_section must contain integers") + }) + .collect::>>()?; + let mrope_section: [usize; 3] = sections + .try_into() + .map_err(|_| anyhow::anyhow!("mrope_section must have three entries"))?; + let image_token_id = root + .get("image_token_id") + .map(|v| { + v.as_u64() + .and_then(|id| u32::try_from(id).ok()) + .context("image_token_id must be a u32") + }) + .transpose()?; let cfg = Config { hidden: int("hidden_size")?, intermediate: int("intermediate_size")?, @@ -81,6 +110,9 @@ impl Config { .as_f64() .or(c["rope_theta"].as_f64()) .context("rope_theta")?, + mrope_section, + max_positions: int("max_position_embeddings")?, + image_token_id, lin_k_heads: int("linear_num_key_heads")?, lin_v_heads: int("linear_num_value_heads")?, lin_k_dim: int("linear_key_head_dim")?, @@ -115,6 +147,25 @@ impl Config { cfg.lin_v_dim ); ensure!(cfg.rotary_half == 32, "{} rotary dims", 2 * cfg.rotary_half); + ensure!( + cfg.rope_theta.is_finite() && cfg.rope_theta > 0.0, + "invalid rope_theta" + ); + ensure!( + rope["mrope_interleaved"].as_bool() != Some(false), + "non-interleaved mrope is unsupported" + ); + ensure!( + cfg.mrope_section + .iter() + .try_fold(0usize, |sum, &x| sum.checked_add(x)) + == Some(cfg.rotary_half), + "mrope sections must sum to rotary_half" + ); + ensure!( + cfg.max_positions > 0 && cfg.max_positions <= (1 << 24), + "unsupported max_position_embeddings" + ); ensure!( cfg.kv_heads > 0 && cfg.heads.is_multiple_of(cfg.kv_heads), "attention heads" @@ -385,6 +436,8 @@ struct Scratch { act: usize, cos: usize, sin: usize, + custom_cos: usize, + custom_sin: usize, } impl Scratch { @@ -423,6 +476,8 @@ impl Scratch { take(cap * cfg.intermediate * BF16), take(cap * cfg.rotary_half * BF16), take(cap * cfg.rotary_half * BF16), + take(cap * cfg.rotary_half * BF16), + take(cap * cfg.rotary_half * BF16), ]; let buf = DeviceBuffer::new(next)?; let [ @@ -448,24 +503,16 @@ impl Scratch { act, cos, sin, + custom_cos, + custom_sin, ] = offsets; - // Rotary tables close to how Qwen3_5TextRotaryEmbedding builds them: inv_freq and - // freqs = inv_freq * position in float32, cos and sin rounded to bfloat16. Here - // cos and sin are taken in float64 on the host rather than in float32 on the - // GPU, so a few of the rounded values can differ by one bfloat16 step. - let half = cfg.rotary_half; - let inv: Vec = (0..half) - .map(|i| 1.0f32 / (cfg.rope_theta as f32).powf((2 * i) as f32 / (2 * half) as f32)) - .collect(); - let mut cos_t = Vec::with_capacity(cap * half * BF16); - let mut sin_t = Vec::with_capacity(cap * half * BF16); - for pos in 0..cap { - for &f in &inv { - let freq = (f * pos as f32) as f64; - cos_t.extend(half::bf16::from_f32(freq.cos() as f32).to_le_bytes()); - sin_t.extend(half::bf16::from_f32(freq.sin() as f32).to_le_bytes()); - } - } + let positions: Vec = (0..cap as i64).collect(); + let (cos_t, sin_t) = rotary_tables( + [&positions; 3], + cfg.rotary_half, + cfg.rope_theta, + cfg.mrope_section, + ); // SAFETY: both tables were laid out for cap * rotary_half bfloat16 values. unsafe { cuda::upload(buf.at(cos), &cos_t, stream)?; @@ -496,6 +543,8 @@ impl Scratch { act, cos, sin, + custom_cos, + custom_sin, }) } @@ -632,15 +681,13 @@ impl Model { ) } - /// The final-norm hidden state at the last position, as float32. - pub fn forward(&mut self, ids: &[u32]) -> Result> { - let t = ids.len(); - ensure!(t > 0, "empty prompt"); - let (vocab, h) = (self.embed.shape[0], self.cfg.hidden); - ensure!( - ids.iter().all(|&i| (i as usize) < vocab), - "token id outside the vocabulary" - ); + /// Finish queued work before releasing external execution admission. + pub fn synchronize(&self) -> Result<()> { + cuda::set_device(0)?; + cuda::synchronize(self.stream) + } + + fn prepare_scratch(&mut self, t: usize) -> Result<()> { cuda::set_device(0)?; if self.scratch.as_ref().is_none_or(|s| t > s.cap) { self.graphs.clear(); @@ -651,16 +698,33 @@ impl Model { self.stream, )?); } + Ok(()) + } + + /// The final-norm hidden state at the last position, as float32. + pub fn forward(&mut self, ids: &[u32]) -> Result> { + let t = ids.len(); + ensure!( + t > 0 && t <= self.cfg.max_positions, + "empty or oversized prompt" + ); + ensure!( + ids.iter().all(|&i| (i as usize) < self.embed.shape[0]), + "token id outside the vocabulary" + ); + self.prepare_scratch(t)?; let s = self.scratch.as_ref().unwrap(); - let ids32: Vec = ids.iter().flat_map(|&i| (i as i32).to_le_bytes()).collect(); - // SAFETY: the ids buffer holds at least t int32 values. - unsafe { cuda::upload(s.at(s.ids), &ids32, self.stream)? }; + self.embed_tokens(s, ids)?; if self.graph_enabled { - if !self.graphs.iter().any(|(length, _)| *length == t) { - // Initialize every cuBLASLt plan before stream capture. - self.run(s, t)?; + if let Some((_, graph)) = self.graphs.iter().find(|(length, _)| *length == t) { + graph.launch(self.stream)?; + } else { + // Warm GEMM plans and keep this eager result for the cache miss. + // run() advances s.res in place and no longer embeds tokens, so + // launching the new graph here would advance the residual twice. + self.run(s, t, false)?; cuda::synchronize(self.stream)?; - match cuda::Graph::capture(self.stream, || self.run(s, t)) { + match cuda::Graph::capture(self.stream, || self.run(s, t, false)) { Ok(graph) => { if self.graphs.len() == 64 { self.graphs.pop_front(); @@ -669,27 +733,111 @@ impl Model { } Err(error) => { // Capture records without executing: the eager result is valid. - // Disable graphs for this worker rather than retrying failures. eprintln!("CUDA Graph capture failed; using eager execution: {error:#}"); self.graph_enabled = false; self.graphs.clear(); } } } - if self.graph_enabled { - self.graphs - .iter() - .find(|(length, _)| *length == t) - .unwrap() - .1 - .launch(self.stream)?; - } } else { - self.run(s, t)?; + self.run(s, t, false)?; } - let mut last = vec![0u8; h * BF16]; + self.last_hidden(s, t) + } + + /// Prefill one unpadded prompt with already-adapted BF16 image embeddings and + /// explicit `[3, 1, sequence]` T/H/W positions. No vision tower runs here. + /// Load a checkpoint with the matching multimodal language adapter merged. + pub fn forward_multimodal(&mut self, input: &MultimodalInput<'_>) -> Result> { + let image_token = self + .cfg + .image_token_id + .context("checkpoint has no image_token_id")?; + input.validate( + self.cfg.hidden, + self.embed.shape[0], + image_token, + self.cfg.max_positions, + )?; + let t = input.token_ids.len(); + self.prepare_scratch(t)?; + let s = self.scratch.as_ref().unwrap(); + self.upload_positions(s, input.position_ids)?; + self.embed_tokens(s, input.token_ids)?; + let bytes: Vec = input + .image_embeddings + .iter() + .flat_map(|x| x.to_le_bytes()) + .collect(); + // Coalesce adjacent placeholders. Text rows remain those of embed_tokens. + let indices = input.image_token_indices; + let mut begin = 0; + while begin < indices.len() { + let mut end = begin + 1; + while end < indices.len() && indices[end] == indices[end - 1] + 1 { + end += 1; + } + let row_bytes = self.cfg.hidden * BF16; + // SAFETY: validated indices lie in the t-row residual buffer, and + // features contain exactly one hidden-size BF16 row per placeholder. + unsafe { + cuda::upload( + s.at(s.res + indices[begin] * row_bytes), + &bytes[begin * row_bytes..end * row_bytes], + self.stream, + )?; + } + begin = end; + } + self.run(s, t, true)?; + self.last_hidden(s, t) + } + + fn upload_positions(&self, s: &Scratch, positions: [&[i64]; 3]) -> Result<()> { + let (cos, sin) = rotary_tables( + positions, + self.cfg.rotary_half, + self.cfg.rope_theta, + self.cfg.mrope_section, + ); + // SAFETY: tables contain at most s.cap rows of rotary_half BF16 values. + unsafe { + cuda::upload(s.at(s.custom_cos), &cos, self.stream)?; + cuda::upload(s.at(s.custom_sin), &sin, self.stream)?; + } + Ok(()) + } + + fn embed_tokens(&self, s: &Scratch, ids: &[u32]) -> Result<()> { + let ids32: Vec = ids.iter().flat_map(|&i| (i as i32).to_le_bytes()).collect(); + // SAFETY: IDs were checked against the vocabulary; scratch holds t rows. + unsafe { + cuda::upload(s.at(s.ids), &ids32, self.stream)?; + check( + (cuda::api().cs1_embed)( + s.at(s.ids).cast(), + self.embed.ptr, + s.at(s.res), + ids.len() as i32, + self.cfg.hidden as i32, + self.stream, + ), + "embed", + )?; + } + Ok(()) + } + + fn last_hidden(&self, s: &Scratch, t: usize) -> Result> { + let mut last = vec![0u8; self.cfg.hidden * BF16]; // SAFETY: x holds at least t rows of the hidden size. - unsafe { cuda::download(&mut last, s.at(s.x + (t - 1) * h * BF16), self.stream)? }; + unsafe { + cuda::download( + &mut last, + s.at(s.x + (t - 1) * self.cfg.hidden * BF16), + self.stream, + )?; + } let (pairs, _) = last.as_chunks::<2>(); Ok(pairs .iter() @@ -697,9 +845,10 @@ impl Model { .collect()) } - /// Queue one forward pass over the first `t` ids in `s`. The final-norm hidden - /// states end up in `s.x`. - fn run(&self, s: &Scratch, t: usize) -> Result<()> { + /// Queue language layers over prepared embeddings in `s.res`, with rotary + /// tables in immutable text buffers or separate explicit-position buffers. + /// Final-norm hidden states end up in `s.x`. + fn run(&self, s: &Scratch, t: usize, custom_positions: bool) -> Result<()> { let cfg = &self.cfg; let st = self.stream; let (ti, hi, eps) = (t as i32, cfg.hidden as i32, cfg.eps); @@ -707,13 +856,14 @@ impl Model { let (hq, hk, hd) = (cfg.heads as i32, cfg.kv_heads as i32, cfg.head_dim as i32); let w = Widths::of(cfg); let p = |off: usize| s.at(off); - // SAFETY (every kernel call below): the pointers are weights in the arena or + let (cos, sin) = if custom_positions { + (s.custom_cos, s.custom_sin) + } else { + (s.cos, s.sin) + }; + // SAFETY (every kernel call below): pointers are weights in the arena or // scratch buffers laid out for at least t tokens with the widths used here. unsafe { - check( - (cuda::api().cs1_embed)(p(s.ids).cast(), self.embed.ptr, p(s.res), ti, hi, st), - "embed", - )?; check( (cuda::api().cs1_rms_norm)( p(s.res), @@ -814,8 +964,8 @@ impl Model { ld, fa.q_norm.ptr, fa.k_norm.ptr, - p(s.cos), - p(s.sin), + p(cos), + p(sin), p(s.aq), p(s.agate), p(s.ak),