diff --git a/.gitignore b/.gitignore index 38579376..99ba095f 100644 --- a/.gitignore +++ b/.gitignore @@ -23,4 +23,6 @@ target # Local Python worker environment .venv/ __pycache__/ +# Model weights downloaded by the recipes +weights/ .DS_Store diff --git a/Cargo.lock b/Cargo.lock index be379643..16864a16 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,41 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "ahash" +version = "0.8.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a15f179cd60c4584b8a8c596927aadc462e27f2ca70c04e0071964a73ba7a75" +dependencies = [ + "cfg-if", + "getrandom 0.3.4", + "once_cell", + "serde", + "version_check", + "zerocopy", +] + +[[package]] +name = "aho-corasick" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c982642fa9e8606056828ee9a8505737230110bb1099153c79efe865c59d12ba" +dependencies = [ + "memchr", +] + +[[package]] +name = "allocator-api2" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "683d7910e743518b0e34f1186f92494becacb047c7b6bf616c96772180fef923" + +[[package]] +name = "anyhow" +version = "1.0.104" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470" + [[package]] name = "atomic-waker" version = "1.1.2" @@ -60,6 +95,12 @@ dependencies = [ "tracing", ] +[[package]] +name = "base64" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e1b586273c5702936fe7b7d6896644d8be71e6314cfe09d3167c95f712589e8" + [[package]] name = "base64" version = "0.22.1" @@ -90,6 +131,15 @@ version = "1.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc652a48c352aef3ea3aed32080501cf3ef6ed5da78602a020c991775b0aff04" +[[package]] +name = "castaway" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dec551ab6e7578819132c713a93c022a05d60159dc86e7a7050223577484c55a" +dependencies = [ + "rustversion", +] + [[package]] name = "cc" version = "1.4.7" @@ -120,7 +170,22 @@ checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06" dependencies = [ "cfg-if", "cpufeatures", - "rand_core", + "rand_core 0.10.1", +] + +[[package]] +name = "compact_str" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9dfdd1c2274d9aa354115b09dc9a901d6c5576818cdf70d14cae2bdb47df00ab" +dependencies = [ + "castaway", + "cfg-if", + "itoa", + "rustversion", + "ryu", + "serde", + "static_assertions", ] [[package]] @@ -132,6 +197,112 @@ dependencies = [ "libc", ] +[[package]] +name = "crossbeam-deque" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "622f3fc73690be383c7214310406f28a90e6edeadc3cea882f9d71e495b9711a" +dependencies = [ + "crossbeam-epoch", + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-epoch" +version = "0.9.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc74980687109a3b14c72fd458107bf0baa1da1a1a805e178d15501ba9b86d9d" +dependencies = [ + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-utils" +version = "0.8.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a31eee39dddec8330830986fcd7625edb5a24ec90ea038215273bbc3adb08ac6" + +[[package]] +name = "crunchy" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" + +[[package]] +name = "darling" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc7f46116c46ff9ab3eb1597a45688b6715c6e628b5c133e288e709a29bcb4ee" +dependencies = [ + "darling_core", + "darling_macro", +] + +[[package]] +name = "darling_core" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d00b9596d185e565c2207a0b01f8bd1a135483d02d9b7b0a54b11da8d53412e" +dependencies = [ + "fnv", + "ident_case", + "proc-macro2", + "quote", + "strsim", + "syn 2.0.119", +] + +[[package]] +name = "darling_macro" +version = "0.20.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc34b93ccb385b40dc71c6fceac4b2ad23662c7eeb248cf10d529b7e055b6ead" +dependencies = [ + "darling_core", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "dary_heap" +version = "0.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b1e3a325bc115f096c8b77bbf027a7c2592230e70be2d985be950d3d5e60ebe" +dependencies = [ + "serde", +] + +[[package]] +name = "derive_builder" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "507dfb09ea8b7fa618fcf76e953f4f5e192547945816d5358edffe39f6f94947" +dependencies = [ + "derive_builder_macro", +] + +[[package]] +name = "derive_builder_core" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d5bcf7b024d6835cfb3d473887cd966994907effbe9227e8c8219824d06c4e8" +dependencies = [ + "darling", + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "derive_builder_macro" +version = "0.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab63b0e2bf4d5928aff72e83a7dace85d7bba5fe12dcc3c5a572d78caffd3f3c" +dependencies = [ + "derive_builder_core", + "syn 2.0.119", +] + [[package]] name = "displaydoc" version = "0.2.7" @@ -140,9 +311,21 @@ checksum = "c6232dd377dcc64799954cbd3a9bb882e9cdc1308ccd87b1c098f1fb2eaf82a8" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] +[[package]] +name = "either" +version = "1.18.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "252afb9ae5eaa683babdc6a068b3f5726eb19e05070c731f9b2a23a7c3e8ed34" + +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + [[package]] name = "errno" version = "0.3.14" @@ -153,12 +336,36 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "esaxx-rs" +version = "0.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d817e038c30374a4bcb22f94d0a8a0e216958d4c3dcde369b1439fec4bdda6e6" + +[[package]] +name = "fastrand" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223" + [[package]] name = "find-msvc-tools" version = "0.1.13" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ef25905e51abafe4dcea6c15fec58c57b601cdbd0ee53d22ea1d3016c587d39b" +[[package]] +name = "fnv" +version = "1.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" + +[[package]] +name = "foldhash" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77ce24cb58228fbb8aa041425bb1050850ac19177686ea6e0f41a70416f56fdb" + [[package]] name = "form_urlencoded" version = "1.2.2" @@ -197,7 +404,7 @@ checksum = "9fb9654ba8355388abeb8dcb4fc62f511300867002afc858860463bdd9fe0c44" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -241,6 +448,18 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "getrandom" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" +dependencies = [ + "cfg-if", + "libc", + "r-efi 5.3.0", + "wasip2", +] + [[package]] name = "getrandom" version = "0.4.3" @@ -250,11 +469,41 @@ dependencies = [ "cfg-if", "js-sys", "libc", - "r-efi", - "rand_core", + "r-efi 6.0.0", + "rand_core 0.10.1", "wasm-bindgen", ] +[[package]] +name = "half" +version = "2.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ea2d84b969582b4b1864a92dc5d27cd2b77b622a8d79306834f1be5ba20d84b" +dependencies = [ + "cfg-if", + "crunchy", + "zerocopy", +] + +[[package]] +name = "hashbrown" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" +dependencies = [ + "allocator-api2", + "equivalent", + "foldhash", + "serde", + "serde_core", +] + +[[package]] +name = "hashbrown" +version = "0.17.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" + [[package]] name = "http" version = "1.5.0" @@ -444,6 +693,12 @@ dependencies = [ "zerovec", ] +[[package]] +name = "ident_case" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9e0384b61958566e926dc50660321d12159025e767c18e043daf26b70104c39" + [[package]] name = "idna" version = "1.1.0" @@ -465,12 +720,31 @@ dependencies = [ "icu_properties", ] +[[package]] +name = "indexmap" +version = "2.14.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc4e190f5d26ca7051642629da2c52fc03bde85a03197c99408dcd291734c855" +dependencies = [ + "equivalent", + "hashbrown 0.17.1", +] + [[package]] name = "ipnet" version = "2.12.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "791930b43c0d5973160d90a8f3894509f2b273430f5c5c73b668636d0287c5c0" +[[package]] +name = "itertools" +version = "0.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2b192c782037fadd9cfa75548310488aabdbf3d2da73885b31bd0abd03351285" +dependencies = [ + "either", +] + [[package]] name = "itoa" version = "1.0.18" @@ -494,6 +768,22 @@ version = "0.2.189" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3eaf3ede3fee6db1a4c2ee091bf8a8b4dccdc6d17f656fb07896ee72867612f2" +[[package]] +name = "libloading" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7c4b02199fee7c5d21a5ae7d8cfa79a6ef5bb2fc834d6e9058e89c825efdc55" +dependencies = [ + "cfg-if", + "windows-link", +] + +[[package]] +name = "linux-raw-sys" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" + [[package]] name = "litemap" version = "0.8.3" @@ -512,6 +802,22 @@ version = "0.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4050469837a6ff301cd14c1f8f24f88549e6d548f24f64e2148eb0f72cebc51f" +[[package]] +name = "macro_rules_attribute" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b3ae8f6d608c795738406608304d30a2dfbdc8e58e44f7ba43236da5208ded3c" +dependencies = [ + "macro_rules_attribute-proc_macro", + "pastey", +] + +[[package]] +name = "macro_rules_attribute-proc_macro" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc04a4c58212d57930a24bf47d3fa87485264a3a054e9c10e042eb373573ad3c" + [[package]] name = "matchit" version = "0.8.4" @@ -524,12 +830,27 @@ version = "2.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" +[[package]] +name = "memmap2" +version = "0.9.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d1219ed1b7f229ee7104d281dd01d6802fe28bb6e95d292942c4daacdeb798c0" +dependencies = [ + "libc", +] + [[package]] name = "mime" version = "0.3.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" +[[package]] +name = "minimal-lexical" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" + [[package]] name = "mio" version = "1.2.3" @@ -541,6 +862,54 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "monostate" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3341a273f6c9d5bef1908f17b7267bbab0e95c9bf69a0d4dcf8e9e1b2c76ef67" +dependencies = [ + "monostate-impl", + "serde", + "serde_core", +] + +[[package]] +name = "monostate-impl" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e4db6d5580af57bf992f59068d4ea26fd518574ff48d7639b255a36f9de6e7e9" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "nom" +version = "7.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d273983c5a657a70a3e8f2a01329822f3b8c8172b73826411a55751e404a0a4a" +dependencies = [ + "memchr", + "minimal-lexical", +] + +[[package]] +name = "omni-cua-s1-native" +version = "0.1.0" +dependencies = [ + "anyhow", + "axum", + "half", + "libloading", + "memmap2", + "safetensors", + "serde", + "serde_json", + "tokenizers", + "tokio", +] + [[package]] name = "omni-jev" version = "0.1.0" @@ -557,6 +926,40 @@ version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" +[[package]] +name = "onig" +version = "6.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0cc3cbf698f9438986c11a880c90a6d04b9de27575afd28bbf45b154b6c709e2" +dependencies = [ + "bitflags", + "libc", + "once_cell", + "onig_sys", +] + +[[package]] +name = "onig_sys" +version = "69.9.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e68317604e77e53b85896388e1a803c1d21b74c899ec9e5e1112db90735edd7" +dependencies = [ + "cc", + "pkg-config", +] + +[[package]] +name = "paste" +version = "1.0.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" + +[[package]] +name = "pastey" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2ee67f1008b1ba2321834326597b8e186293b049a023cdef258527550b9935b4" + [[package]] name = "percent-encoding" version = "2.3.2" @@ -569,6 +972,12 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" +[[package]] +name = "pkg-config" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f6b464fbc74e149a392436b17d523f769e057cb6877f6a5c4618bc6f11800548" + [[package]] name = "potential_utf" version = "0.1.6" @@ -578,6 +987,15 @@ dependencies = [ "zerovec", ] +[[package]] +name = "ppv-lite86" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9" +dependencies = [ + "zerocopy", +] + [[package]] name = "proc-macro2" version = "1.0.107" @@ -616,7 +1034,7 @@ dependencies = [ "bytes", "getrandom 0.4.3", "lru-slab", - "rand", + "rand 0.10.3", "rand_pcg", "ring", "rustc-hash", @@ -652,12 +1070,28 @@ dependencies = [ "proc-macro2", ] +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + [[package]] name = "r-efi" version = "6.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" +[[package]] +name = "rand" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9ef1d0d795eb7d84685bca4f72f3649f064e6641543d3a8c415898726a57b41" +dependencies = [ + "rand_chacha", + "rand_core 0.9.5", +] + [[package]] name = "rand" version = "0.10.3" @@ -666,7 +1100,26 @@ checksum = "65c9fb96cbc91e3478eaae79a69fcd3f1ae4ad052e471fe6732fff548984b4af" dependencies = [ "chacha20", "getrandom 0.4.3", - "rand_core", + "rand_core 0.10.1", +] + +[[package]] +name = "rand_chacha" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" +dependencies = [ + "ppv-lite86", + "rand_core 0.9.5", +] + +[[package]] +name = "rand_core" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c" +dependencies = [ + "getrandom 0.3.4", ] [[package]] @@ -681,9 +1134,69 @@ version = "0.10.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "caa0f4137e1c0a72f4c651489402276c8e8e1cf081f3b0ba156d2cbeef09e86a" dependencies = [ - "rand_core", + "rand_core 0.10.1", +] + +[[package]] +name = "rayon" +version = "1.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fb39b166781f92d482534ef4b4b1b2568f42613b53e5b6c160e24cfbfa30926d" +dependencies = [ + "either", + "rayon-core", +] + +[[package]] +name = "rayon-cond" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2964d0cf57a3e7a06e8183d14a8b527195c706b7983549cd5462d5aa3747438f" +dependencies = [ + "either", + "itertools", + "rayon", +] + +[[package]] +name = "rayon-core" +version = "1.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22e18b0f0062d30d4230b2e85ff77fdfe4326feb054b9783a3460d8435c8ab91" +dependencies = [ + "crossbeam-deque", + "crossbeam-utils", +] + +[[package]] +name = "regex" +version = "1.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f020237b6c8eed93db2e2cb53c00c60a8e1bc73da7d073199a1180401450218d" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + +[[package]] +name = "regex-automata" +version = "0.4.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ad8553b9b26413251cbf30e620595c7a41b3887f03da04579c0e6b0d6a06b4b2" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", ] +[[package]] +name = "regex-syntax" +version = "0.8.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" + [[package]] name = "reqwest" version = "0.12.28" @@ -745,6 +1258,19 @@ version = "2.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6b1e7f9a428571be2dc5bc0505c13fb6bf936822b894ec87abf8a08a4e51742d" +[[package]] +name = "rustix" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "891efababe418670775f199f0d233d84843c227a0949a883ce15b37c78d6629d" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys 0.61.2", +] + [[package]] name = "rustls" version = "0.23.45" @@ -792,6 +1318,19 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" +[[package]] +name = "safetensors" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "79b079b829cb27a1c3c374341345ed2e8b2c0c839034522cee576c140bd7f846" +dependencies = [ + "hashbrown 0.16.1", + "libc", + "serde", + "serde_json", + "tempfile", +] + [[package]] name = "serde" version = "1.0.229" @@ -799,6 +1338,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4148590afebada386688f18773da617792bf2ef03ffc1e4cbd2b1d45b023e0ba" dependencies = [ "serde_core", + "serde_derive", ] [[package]] @@ -818,7 +1358,7 @@ checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -827,6 +1367,7 @@ version = "1.0.151" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c841b55ecdae098c80dcae9cf767f6f8a0c2cdb3416bbef72181df4d0fe73f14" dependencies = [ + "indexmap", "itoa", "memchr", "serde", @@ -895,18 +1436,53 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "spm_precompiled" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5851699c4033c63636f7ea4cf7b7c1f1bf06d0cc03cfb42e711de5a5c46cf326" +dependencies = [ + "base64 0.13.1", + "nom", + "serde", + "unicode-segmentation", +] + [[package]] name = "stable_deref_trait" version = "1.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" +[[package]] +name = "static_assertions" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2eb9349b6444b326872e140eb1cf5e7c522154d69e7a0ffb0fb81c06b37543f" + +[[package]] +name = "strsim" +version = "0.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" + [[package]] name = "subtle" version = "2.6.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" +[[package]] +name = "syn" +version = "2.0.119" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "872831b642d1a07999a962a351ed35b955ea2cfc8f3862091e2a240a84f17297" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + [[package]] name = "syn" version = "3.0.6" @@ -935,7 +1511,20 @@ checksum = "901704edd0dfe137f1987838ee4f259e4e063c31371bdb423f7ae38ec6f77f02" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", +] + +[[package]] +name = "tempfile" +version = "3.27.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" +dependencies = [ + "fastrand", + "getrandom 0.4.3", + "once_cell", + "rustix", + "windows-sys 0.61.2", ] [[package]] @@ -955,7 +1544,7 @@ checksum = "fe5197923287db20a58125f0bc85c062f7f2c892de97b18c356f9efb14b28524" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -974,6 +1563,39 @@ version = "1.13.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fd3ca314f692efd6c868f8408f53fe444634a845f96c028b97d35f6a1f79f0ee" +[[package]] +name = "tokenizers" +version = "0.22.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b238e22d44a15349529690fb07bd645cf58149a1b1e44d6cb5bd1641ff1a6223" +dependencies = [ + "ahash", + "aho-corasick", + "compact_str", + "dary_heap", + "derive_builder", + "esaxx-rs", + "getrandom 0.3.4", + "itertools", + "log", + "macro_rules_attribute", + "monostate", + "onig", + "paste", + "rand 0.9.5", + "rayon", + "rayon-cond", + "regex", + "regex-syntax", + "serde", + "serde_json", + "spm_precompiled", + "thiserror", + "unicode-normalization-alignments", + "unicode-segmentation", + "unicode_categories", +] + [[package]] name = "tokio" version = "1.53.1" @@ -998,7 +1620,7 @@ checksum = "78773a2a397f451582ce068015985c33193cf6dea8b74d2a639fe457b2f07b0e" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] @@ -1102,6 +1724,27 @@ version = "1.0.26" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d245f478577f809a851594d02313b640fb437e0bb33866753cff937863096954" +[[package]] +name = "unicode-normalization-alignments" +version = "0.1.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "43f613e4fa046e69818dd287fdc4bc78175ff20331479dab6e1b0f98d57062de" +dependencies = [ + "smallvec", +] + +[[package]] +name = "unicode-segmentation" +version = "1.13.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6f5d3c3b1bf09027a88a6bc961fc00497d651009560b5463668dc81b0fa87a8" + +[[package]] +name = "unicode_categories" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e" + [[package]] name = "untrusted" version = "0.9.0" @@ -1126,6 +1769,12 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + [[package]] name = "want" version = "0.3.1" @@ -1141,6 +1790,15 @@ version = "0.11.1+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" +[[package]] +name = "wasip2" +version = "1.0.4+wasi-0.2.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b67efb37e106e55ce722a510d6b5f9c17f083e5fc79afc2badeb12cc313d9487" +dependencies = [ + "wit-bindgen", +] + [[package]] name = "wasm-bindgen" version = "0.2.129" @@ -1184,7 +1842,7 @@ dependencies = [ "bumpalo", "proc-macro2", "quote", - "syn", + "syn 3.0.6", "wasm-bindgen-shared", ] @@ -1327,6 +1985,12 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" +[[package]] +name = "wit-bindgen" +version = "0.57.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" + [[package]] name = "writeable" version = "0.6.4" @@ -1352,10 +2016,30 @@ checksum = "33811428bee40dbceb6d545e95754741d17a6aef9a4849f0fd62e2ba4f412a78" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", "synstructure", ] +[[package]] +name = "zerocopy" +version = "0.8.59" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6df92bf3d9227be3d53173901ddbffac2babc27ae50f397776ffd6dc33f800cb" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.59" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac4f328cf2f05d084e496c3e9c3f33ed0a183656a16e1fcec4d464d8373aec82" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "zerofrom" version = "0.1.8" @@ -1373,7 +2057,7 @@ checksum = "f75b4683f6c7f45248d4d64056a24298c6281e0993356d7d1b4a1a962ef10d4a" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", "synstructure", ] @@ -1413,7 +2097,7 @@ checksum = "34df6fc39dbd26ddc9c10e6a2984476e13acce22e64e4487636ef494369225da" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.6", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 8036e4bc..8c100b4b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,3 +1,3 @@ [workspace] -members = ["src/frontend"] +members = ["src/frontend", "src/models/cua_s1/native"] resolver = "3" diff --git a/README.md b/README.md index 76944aa3..b99d91c5 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,7 @@ A community-maintained inference engine for prefill-only System1-Omni models, designed around a Rust frontend, model-owned execution, and high-performance CUDA and Metal backends. -The Rust frontend forwards requests to a separately running model worker. In-repository model engines and GPU backends are not implemented yet. +The Rust frontend forwards requests to a separately running model worker. The Cua-S1 4B 0.2 `text` adapter has a native worker with CUDA kernels in this repository; other in-repository model engines and GPU backends are not implemented yet. ## Run the frontend @@ -40,22 +40,23 @@ Implementation code lives under `src/`; recipes and documentation stay at the re | Directory | Responsibility | | --- | --- | -| [`src/frontend/`](src/frontend/) | Rust serving code and the small engine interface. | +| [`src/frontend/`](src/frontend/) | Rust serving code, Python worker adapters, and the small engine interface. | | [`src/models/`](src/models/) | Model implementations, one directory per model: preprocessing, batching, state, execution, and output processing. | | [`src/backends/cuda/`](src/backends/cuda/) | NVIDIA GPU operations and kernel integration. | | [`src/backends/metal/`](src/backends/metal/) | Apple GPU operations and kernel integration. | | [`recipe/`](recipe/) | Model setup instructions, launch commands, configuration examples, and example requests. | | [`docs/`](docs/) | Project documentation and architecture assets. | -The frontend is a Cargo workspace member. Model and backend directories currently document planned work; they do not prescribe process boundaries. +The frontend and the Cua-S1 native worker are Cargo workspace members. The other model and backend directories currently document planned work; they do not prescribe process boundaries. ## Supported models -LAYA can run as an external Python worker for text requests. Its in-repository model engine is still planned: +LAYA can run as an external Python worker for text requests; its in-repository model engine is still planned. The Cua-S1 4B 0.2 `text` adapter runs as a Python worker or as a native worker on CUDA: | Model | Status | | --- | --- | | LAYA | [External worker](recipe/laya/README.md); model engine planned | +| Cua-S1 4B 0.2 (`text` adapter) | [Python worker](recipe/cua_s1/text.md); [native worker](recipe/cua_s1/native.md), CUDA, run on sm_89 | CUDA and Metal coverage will be documented per model as implementations are added and validated. diff --git a/recipe/README.md b/recipe/README.md index 4d0bf6c4..3795fd9c 100644 --- a/recipe/README.md +++ b/recipe/README.md @@ -2,6 +2,10 @@ - [Laya text worker](laya/README.md): start the external Python worker, connect the Rust frontend and compare direct and proxied responses. +- [Cua-S1 4B 0.2 text worker](cua_s1/text.md): download the pinned weights, start + the worker and connect the Rust frontend. +- [Cua-S1 4B 0.2 native text worker](cua_s1/native.md): build the CUDA library and + the Rust worker, export the merged weights and start the worker. Recipes contain setup, launch commands and examples. Reusable implementation code belongs under `src/`. diff --git a/recipe/cua_s1/export_text_merged.py b/recipe/cua_s1/export_text_merged.py new file mode 100644 index 00000000..a2741560 --- /dev/null +++ b/recipe/cua_s1/export_text_merged.py @@ -0,0 +1,30 @@ +"""Export Qwen3.5-4B with the Cua-S1 `text` adapter merged into the bfloat16 weights, +for the native worker (recipe/cua_s1/native.md). Run it in the reference worker's +environment (recipe/cua_s1/text.md): + + PYTHONPATH=src .venv/bin/python recipe/cua_s1/export_text_merged.py \ + --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text \ + --out weights/cua-s1-4b-0.2-text-merged +""" + +import argparse +import json +from pathlib import Path + +from models.cua_s1.text.model import TextModel + +parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) +parser.add_argument("--base", required=True) +parser.add_argument("--adapter", required=True) +parser.add_argument("--out", required=True, type=Path) +args = parser.parse_args() + +loaded = TextModel(args.base, args.adapter, "cuda", "bfloat16") +loaded.model.merge_and_unload().save_pretrained(args.out, max_shard_size="5GB") +# Transformers writes the pre-tokenizer rule it uses into tokenizer.json; the native +# worker tokenizes with that file. +loaded.tokenizer.save_pretrained(args.out) +# The native worker refuses a directory without this marker, such as the base model. +(args.out / "cua_s1_export.json").write_text( + json.dumps({"format": "cua-s1-text-merged/1"}) +) diff --git a/recipe/cua_s1/native.md b/recipe/cua_s1/native.md new file mode 100644 index 00000000..cf9e9961 --- /dev/null +++ b/recipe/cua_s1/native.md @@ -0,0 +1,34 @@ +# Cua-S1 4B 0.2 native text worker + +The native worker ([`src/models/cua_s1/native/`](../../src/models/cua_s1/native/)) serves the `text` adapter like the reference worker in [`text.md`](text.md), with the Qwen3.5-4B forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../src/backends/cuda/qwen3_5/) and no Python or PyTorch. It needs an NVIDIA GPU with compute capability 8.0 or newer; only an RTX 6000 Ada (sm_89) with CUDA 13.2 has been run. + +Run the commands from the repository root. Pass your GPU's compute capability to `build.sh` (89 for Ada, 80 for A100, 90 for H100); the worker finds `libqwen3_5_cuda.so` next to its executable, or at `CUA_S1_CUDA_LIB`: + +```sh +src/backends/cuda/qwen3_5/build.sh target/release 89 # needs nvcc and cuBLASLt +cargo build --release --locked -p omni-cua-s1-native +``` + +The worker loads the weights with the `text` adapter merged in. With the reference worker's environment and weights from `text.md`, export them once (about 8.5 GB): + +```sh +PYTHONPATH=src .venv/bin/python recipe/cua_s1/export_text_merged.py \ + --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text \ + --out weights/cua-s1-4b-0.2-text-merged +``` + +Start the worker (`CUA_S1_HOST` and `CUA_S1_PORT` default to `127.0.0.1` and `8000`), then the frontend and requests as in `text.md`: + +```sh +CUA_S1_MODEL=weights/cua-s1-4b-0.2-text-merged target/release/omni-cua-s1-native +``` + +Each question is one eager forward pass over its prompt; the final hidden state at the last position times the 26 letter rows of the output projection gives the option probabilities. The probabilities are not bitwise identical to the reference worker's, since the adapter is merged and the kernels differ; they are held to the tolerance in [`src/models/cua_s1/README.md`](../../src/models/cua_s1/README.md#validation). Error messages are worded differently, and bodies nested more than 127 levels deep are refused. + +The request tests need no GPU; the kernel tests compare attention and the chunked Gated DeltaNet prefill with float64 references: + +```sh +cargo test -p omni-cua-s1-native +CUA_S1_CUDA_LIB=$PWD/target/release/libqwen3_5_cuda.so \ + cargo test --release -p omni-cua-s1-native --test kernels -- --ignored +``` diff --git a/recipe/cua_s1/requirements-text.txt b/recipe/cua_s1/requirements-text.txt new file mode 100644 index 00000000..136ef2b3 --- /dev/null +++ b/recipe/cua_s1/requirements-text.txt @@ -0,0 +1,14 @@ +# Versions match upstream's `four-b` lock (trycua/cua libs/cua-s1/python/uv.lock +# at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f), the reference environment in +# src/models/cua_s1/README.md. +torch==2.14.0 +transformers==5.17.0 +tokenizers==0.23.2 +peft==0.21.0 +accelerate==1.15.0 +safetensors==0.8.0 +huggingface-hub==1.32.0 +jinja2==3.1.6 +# HTTP serving. +fastapi==0.141.1 +uvicorn==0.54.0 diff --git a/recipe/cua_s1/text.md b/recipe/cua_s1/text.md new file mode 100644 index 00000000..1fcfcb1d --- /dev/null +++ b/recipe/cua_s1/text.md @@ -0,0 +1,46 @@ +# Cua-S1 4B 0.2 text worker + +This recipe runs the Cua-S1 4B 0.2 `text` adapter through Transformers and PEFT behind the Rust frontend. It is the correctness reference for native execution. The model is in [`src/models/cua_s1/text/`](../../src/models/cua_s1/text/), the HTTP worker in [`src/frontend/cua_s1_text.py`](../../src/frontend/cua_s1_text.py), and [`src/models/cua_s1/README.md`](../../src/models/cua_s1/README.md) documents the contract. Only `choice` questions are supported. + +Run the commands from the repository root, on Linux with an NVIDIA GPU and Python 3.12. The pinned versions match the upstream reference environment: + +```sh +python3.12 -m venv .venv +.venv/bin/python -m pip install -r recipe/cua_s1/requirements-text.txt +.venv/bin/hf download Qwen/Qwen3.5-4B \ + --revision 851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a --local-dir weights/Qwen3.5-4B +.venv/bin/hf download cua-ai/cua-s1-4b-0.2 \ + --revision 16818868b0cc7813808aae4e87b417657046ab79 --local-dir weights/cua-s1-4b-0.2 +``` + +Upstream's `libs/cua-s1/ci/fetch_pinned_weights.py --dest weights --verify-only` (in [trycua/cua](https://github.com/trycua/cua) at `0e75660ce4c2edda519e0c795fa3ad98abf4e76f`) checks every downloaded file against upstream's lock. + +Start the worker, which runs one warmup decision before it listens, then the frontend: + +```sh +PYTHONPATH=src .venv/bin/python -m frontend.cua_s1_text \ + --base weights/Qwen3.5-4B --adapter weights/cua-s1-4b-0.2/text --port 8000 +cargo build --release --locked +OMNI_JEV_BIND=127.0.0.1:8080 OMNI_JEV_BACKEND_URL=http://127.0.0.1:8000 ./target/release/omni-jev +``` + +Requests run one at a time. Bodies over 4 MiB, more than 64 questions, or a prompt over 16,384 tokens get `413`. The worker computes logits for every position, as upstream does, so memory grows with prompt length: the 15,446-token test input peaked at about 21.3 GiB in bfloat16. + +```sh +curl http://127.0.0.1:8080/v1/systemone \ + -H 'Content-Type: application/json' \ + -d '{"model":"cua-s1-4b-0.2","state":"Dialog: Delete 3 files permanently? Buttons: Delete, Cancel","questions":{"pick":{"type":"choice","instructions":"Keep the files.","criteria":{"delete":"Click Delete","cancel":"Click Cancel"}}}}' +``` + +On an RTX 6000 Ada in bfloat16, the response is: + +```json +{"model":"cua-ai/cua-s1-4b-0.2@16818868b0cc7813808aae4e87b417657046ab79:text","answers":{"pick":{"type":"choice","choice":"cancel","probabilities":{"delete":0.0024726232513785362,"cancel":0.9975274205207825},"confidence":0.9750249565060322}},"usage":{"input_tokens":153,"output_tokens":0}} +``` + +The tests need neither weights nor a GPU; with `CUA_S1_BASE=weights/Qwen3.5-4B` they also check the tokenizer: + +```sh +.venv/bin/python -m pip install pytest httpx +PYTHONPATH=src .venv/bin/python -m pytest tests/cua_s1 +``` diff --git a/src/backends/cuda/README.md b/src/backends/cuda/README.md index 1ba4c67c..577209d6 100644 --- a/src/backends/cuda/README.md +++ b/src/backends/cuda/README.md @@ -4,4 +4,4 @@ Planned home for high-performance NVIDIA GPU operations and kernel integration. Model orchestration, batching policy, state management, and kernel selection remain with the model engine. CUDA and Metal implementations do not need identical internal structures or a universal tensor abstraction. -Status: planned; no CUDA implementation or validated hardware coverage yet. +Status: [`qwen3_5/`](qwen3_5/) has the operations of a prefill-only Qwen3.5 forward pass, used by the Cua-S1 native worker and measured on sm_89. Other models are planned. diff --git a/src/backends/cuda/qwen3_5/README.md b/src/backends/cuda/qwen3_5/README.md new file mode 100644 index 00000000..b84f50b5 --- /dev/null +++ b/src/backends/cuda/qwen3_5/README.md @@ -0,0 +1,9 @@ +# Qwen3.5 prefill operations + +CUDA kernels for a prefill-only Qwen3.5 forward pass, built into `libqwen3_5_cuda.so` with a C interface ([`ops.h`](ops.h)), so that a Rust model engine loads it at run time and builds without a CUDA toolkit. The Cua-S1 native worker ([`src/models/cua_s1/native/`](../../../models/cua_s1/native/)) uses it and keeps the layer loop and buffers. + +```sh +src/backends/cuda/qwen3_5/build.sh [compute capability, default 89] +``` + +The norm, elementwise and q/k preparation kernels round to bfloat16 where Transformers (`modeling_qwen3_5.py`) does. Attention (FlashAttention-2 style, on tensor cores) and the chunked gated delta rule keep some intermediate results in bfloat16, as FlashAttention and flash-linear-attention do. GEMMs go through cuBLASLt with its first heuristic choice. Tensor-core kernels need sm_80 or newer; only sm_89 has been run. diff --git a/src/backends/cuda/qwen3_5/attention.cu b/src/backends/cuda/qwen3_5/attention.cu new file mode 100644 index 00000000..8c03666e --- /dev/null +++ b/src/backends/cuda/qwen3_5/attention.cu @@ -0,0 +1,260 @@ +// Full-attention layers of Qwen3.5: q/k preparation and causal attention. +// +// cs1_attention is a FlashAttention-2 style kernel on tensor cores (mma.sync +// m16n8k16, bfloat16 in, float32 accumulation): a block takes 64 queries of one head, +// four warps of 16 rows each, and walks the keys up to its last query in tiles of +// 32, keeping the output and the online softmax in registers. The probabilities are +// rounded to bfloat16 for the P*V product, as in flash attention; the running sums +// stay float32. +#include "common.cuh" +#include "mma.cuh" +#include "ops.h" + +namespace cs1 { +namespace { + +constexpr int DH = 256; // head dim +constexpr int PER = DH / 32; // values per lane + +// One warp per (token, head), q heads first, then k heads. Each lane holds 8 +// consecutive dims, so the rotary partner of dim i < 32 (dim i + 32) sits in lane ^ 4. +__global__ void attn_prep_kernel(const bf16* __restrict__ qg, const bf16* __restrict__ kr, int ld, + const bf16* __restrict__ qw, const bf16* __restrict__ kw, + const bf16* __restrict__ cos, const bf16* __restrict__ sin, + bf16* __restrict__ q, bf16* __restrict__ gate, bf16* __restrict__ k, int T, int Hq, + int Hk, int half, float eps) { + const int warp = blockIdx.x * (blockDim.x / 32) + threadIdx.x / 32, lane = threadIdx.x & 31; + const int heads = Hq + Hk; + if (warp >= T * heads) return; + const int t = warp / heads, hh = warp % heads; + const bool is_q = hh < Hq; + const int h = is_q ? hh : hh - Hq; + const bf16* src = is_q ? qg + (size_t)t * ld + (size_t)h * 2 * DH : kr + (size_t)t * ld + (size_t)h * DH; + const bf16* w = is_q ? qw : kw; + const int d0 = lane * PER; + + float x[PER]; + load8(src + d0, x); + float ss = 0.f; +#pragma unroll + for (int i = 0; i < PER; i++) ss += x[i] * x[i]; + const float inv = rsqrtf(warp_sum(ss) / DH + eps); + float wv[PER]; + load8(w + d0, wv); +#pragma unroll + for (int i = 0; i < PER; i++) x[i] = round_bf16(x[i] * inv * (1.f + wv[i])); + + // rotate_half on the first 2 * half dims: out = x * cos + rotate_half(x) * sin, + // each product and the sum rounded to bfloat16 as in the reference. + const int rot = 2 * half; + float y[PER]; +#pragma unroll + for (int i = 0; i < PER; i++) { + const float partner = __shfl_xor_sync(0xffffffffu, x[i], (half / PER)); + const int d = d0 + i; + y[i] = x[i]; + if (d < rot) { + const int fi = d % half; + const float c = f32(cos[(size_t)t * half + fi]), s = f32(sin[(size_t)t * half + fi]); + const float r = d < half ? -partner : partner; + y[i] = round_bf16(round_bf16(x[i] * c) + round_bf16(r * s)); + } + } + if (is_q) { + store8(q + ((size_t)t * Hq + h) * DH + d0, y); + float gv[PER]; + load8(src + DH + d0, gv); + store8(gate + ((size_t)t * Hq + h) * DH + d0, gv); + } else { + store8(k + ((size_t)t * Hk + h) * DH + d0, y); + } +} + +// ---- flash attention ---- + +namespace flash { + +constexpr int D = 256, 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 = min(T, q0 + BM); + 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 = k0 <= row0 + 15; + 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), row = row0 + g + (e >> 1) * 8; + sc[n][e] = (key <= row && 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 +} // namespace cs1 + +using namespace cs1; + +extern "C" int cs1_attn_prep(const void* qg, const void* kr, int ld, const void* qw, const void* kw, + const void* cos, const void* sin, void* q, void* gate, void* k, int T, int Hq, int Hk, + int Dh, int half, float eps, void* stream) { + if (ld % 8 != 0 || T < 0) return cudaErrorInvalidValue; + // the lane ^ (half / 8) partner exchange needs 2 * half <= 256 and half a multiple of 8 + if (Dh != DH || half % PER != 0 || 2 * half > DH || (half / PER) & ((half / PER) - 1)) return cudaErrorInvalidValue; + const int warps = T * (Hq + Hk); + if (warps == 0) return cudaSuccess; + constexpr int WARPS = 8; + attn_prep_kernel<<<(warps + WARPS - 1) / WARPS, WARPS * 32, 0, static_cast(stream)>>>( + static_cast(qg), static_cast(kr), ld, static_cast(qw), + static_cast(kw), static_cast(cos), static_cast(sin), + static_cast(q), static_cast(gate), static_cast(k), T, Hq, Hk, half, eps); + return cudaGetLastError(); +} + +extern "C" int cs1_attention(const void* q, const void* k, const void* v, int ldv, void* out, int T, int Hq, int Hk, + int Dh, float scale, void* stream) { + if (Dh != flash::D || Hk <= 0 || Hq % Hk != 0 || ldv % 8 != 0 || ldv < Hk * Dh || T < 0) + return cudaErrorInvalidValue; + if (T == 0) return cudaSuccess; + // once per process (for the device current at the first call) + static const cudaError_t configured = cudaFuncSetAttribute( + flash::flash_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, flash::SMEM_BYTES); + if (configured != cudaSuccess) return configured; + constexpr float LOG2E = 1.4426950408889634f; + flash::flash_kernel<<(stream)>>>( + static_cast(q), static_cast(k), static_cast(v), ldv, + static_cast(out), T, Hq, Hk, scale * LOG2E); + return cudaGetLastError(); +} diff --git a/src/backends/cuda/qwen3_5/build.sh b/src/backends/cuda/qwen3_5/build.sh new file mode 100755 index 00000000..ef2dad6b --- /dev/null +++ b/src/backends/cuda/qwen3_5/build.sh @@ -0,0 +1,27 @@ +#!/usr/bin/env bash +# Build libqwen3_5_cuda.so from the kernels in this directory. +# +# src/backends/cuda/qwen3_5/build.sh [compute capability, e.g. 89] +# +# Needs nvcc (from NVCC, CUDA_HOME/bin or PATH) and cuBLASLt; the compute capability +# defaults to CUDA_COMPUTE_CAP, else 89. The library holds machine code for that +# compute capability and PTX that newer GPUs can compile at load time. The CUDA +# runtime is linked statically and cuBLASLt dynamically, with an rpath to the +# toolkit's library directory when there is one next to nvcc. +set -euo pipefail +here=$(cd "$(dirname "$0")" && pwd) +out=${1:?usage: build.sh [compute capability]} +arch=${2:-${CUDA_COMPUTE_CAP:-89}} +nvcc=${NVCC:-} +if [ -z "$nvcc" ]; then + if [ -n "${CUDA_HOME:-}" ]; then nvcc=$CUDA_HOME/bin/nvcc; else nvcc=$(command -v nvcc); fi +fi +link=(-lcublasLt) +if lib=$(cd "$(dirname "$nvcc")/../lib64" 2>/dev/null && pwd); then + link=(-L"$lib" -lcublasLt -Xlinker -rpath -Xlinker "$lib") +fi +mkdir -p "$out" +"$nvcc" -O3 -std=c++17 -gencode "arch=compute_${arch},code=[sm_${arch},compute_${arch}]" \ + -shared -Xcompiler -fPIC -Xcompiler -Wall,-Wextra -I"$here" "$here"/*.cu \ + "${link[@]}" -o "$out/libqwen3_5_cuda.so" +echo "built $out/libqwen3_5_cuda.so for sm_${arch}" diff --git a/src/backends/cuda/qwen3_5/common.cuh b/src/backends/cuda/qwen3_5/common.cuh new file mode 100644 index 00000000..d70c1761 --- /dev/null +++ b/src/backends/cuda/qwen3_5/common.cuh @@ -0,0 +1,62 @@ +// Helpers shared by the Qwen3.5 kernels. +#pragma once + +#include +#include +#include + +namespace cs1 { + +using bf16 = __nv_bfloat16; + +__device__ __forceinline__ float f32(bf16 x) { return __bfloat162float(x); } +__device__ __forceinline__ bf16 to_bf16(float x) { return __float2bfloat16(x); } +// A float rounded through bfloat16, as PyTorch stores the result of each bfloat16 op. +__device__ __forceinline__ float round_bf16(float x) { return __bfloat162float(__float2bfloat16(x)); } + +__device__ __forceinline__ float warp_sum(float x) { +#pragma unroll + for (int o = 16; o > 0; o >>= 1) x += __shfl_xor_sync(0xffffffffu, x, o); + return x; +} + +// Sum over the block; `scratch` holds at least 32 floats. Every thread gets the sum. +__device__ __forceinline__ float block_sum(float x, float* scratch) { + const int lane = threadIdx.x & 31, warp = threadIdx.x >> 5, warps = (blockDim.x + 31) >> 5; + x = warp_sum(x); + if (lane == 0) scratch[warp] = x; + __syncthreads(); + if (warp == 0) { + float t = lane < warps ? scratch[lane] : 0.f; + t = warp_sum(t); + if (lane == 0) scratch[0] = t; + } + __syncthreads(); + const float total = scratch[0]; + __syncthreads(); + return total; +} + +// PyTorch's float32 SiLU and sigmoid on CUDA. +__device__ __forceinline__ float silu(float x) { return x / (1.f + expf(-x)); } +__device__ __forceinline__ float sigmoid(float x) { return 1.f / (1.f + expf(-x)); } + +// 8 bfloat16 values, one 16-byte load or store. +struct alignas(16) Pack8 { + bf16 v[8]; +}; + +__device__ __forceinline__ void load8(const bf16* p, float out[8]) { + const Pack8 pk = *reinterpret_cast(p); +#pragma unroll + for (int i = 0; i < 8; i++) out[i] = f32(pk.v[i]); +} + +__device__ __forceinline__ void store8(bf16* p, const float in[8]) { + Pack8 pk; +#pragma unroll + for (int i = 0; i < 8; i++) pk.v[i] = to_bf16(in[i]); + *reinterpret_cast(p) = pk; +} + +} // namespace cs1 diff --git a/src/backends/cuda/qwen3_5/elementwise.cu b/src/backends/cuda/qwen3_5/elementwise.cu new file mode 100644 index 00000000..3b1c01bd --- /dev/null +++ b/src/backends/cuda/qwen3_5/elementwise.cu @@ -0,0 +1,126 @@ +// Embedding lookup, the Gated DeltaNet conv and gates, and the elementwise +// activations of Qwen3.5. +#include "common.cuh" +#include "ops.h" + +namespace cs1 { +namespace { + +constexpr int THREADS = 256; + +__global__ void embed_kernel(const int32_t* __restrict__ ids, const Pack8* __restrict__ table, + Pack8* __restrict__ out, int packs) { + const size_t t = blockIdx.x; + const size_t id = ids[t]; + for (int i = threadIdx.x; i < packs; i += blockDim.x) out[t * packs + i] = table[id * packs + i]; +} + +// F.conv1d in bfloat16 (float32 accumulation, rounded), then SiLU (rounded again), +// written to three contiguous outputs. +__global__ void gdn_conv_kernel(const bf16* __restrict__ qkv, int ld, const bf16* __restrict__ w, + bf16* __restrict__ q, bf16* __restrict__ k, bf16* __restrict__ v, int T, + int key_dim, int value_dim) { + const int channels = 2 * key_dim + value_dim; + const size_t idx = (size_t)blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= (size_t)T * channels) return; + const int t = idx / channels, c = idx % channels; + float acc = 0.f; +#pragma unroll + for (int j = 0; j < 4; j++) { + const int s = t - 3 + j; + if (s >= 0) acc = fmaf(f32(w[c * 4 + j]), f32(qkv[(size_t)s * ld + c]), acc); + } + const bf16 y = to_bf16(silu(round_bf16(acc))); + if (c < key_dim) + q[(size_t)t * key_dim + c] = y; + else if (c < 2 * key_dim) + k[(size_t)t * key_dim + c - key_dim] = y; + else + v[(size_t)t * value_dim + c - 2 * key_dim] = y; +} + +// beta = sigmoid(b) in bfloat16; g = -exp(A_log) * softplus(a + dt_bias) in float32 +// (F.softplus with threshold 20). +__global__ void gdn_gates_kernel(const bf16* __restrict__ b, const bf16* __restrict__ a, int ld, + const bf16* __restrict__ A_log, const bf16* __restrict__ dt_bias, + bf16* __restrict__ beta, float* __restrict__ g, int n, int H) { + const int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i >= n) return; + const int h = i % H; + const size_t src = (size_t)(i / H) * ld + h; + beta[i] = to_bf16(sigmoid(f32(b[src]))); + const float x = f32(a[src]) + f32(dt_bias[h]); + const float sp = x > 20.f ? x : log1pf(expf(x)); + g[i] = -expf(f32(A_log[h])) * sp; +} + +// attn_output * torch.sigmoid(gate), both bfloat16. +__global__ void sigmoid_gate_kernel(bf16* __restrict__ x, const bf16* __restrict__ gate, size_t n) { + const size_t i = (size_t)blockIdx.x * blockDim.x + threadIdx.x; + if (i >= n) return; + x[i] = to_bf16(f32(x[i]) * round_bf16(sigmoid(f32(gate[i])))); +} + +// act_fn(gate_proj(x)) * up_proj(x), both bfloat16; gate and up are the two halves of +// each row of gate_up. +__global__ void silu_mul_kernel(const bf16* __restrict__ gate_up, int ld, bf16* __restrict__ out, int I, + size_t n) { + const size_t i = (size_t)blockIdx.x * blockDim.x + threadIdx.x; + if (i >= n) return; + const size_t t = i / I, j = i % I; + const bf16* row = gate_up + t * ld; + out[i] = to_bf16(round_bf16(silu(f32(row[j]))) * f32(row[I + j])); +} + +unsigned blocks(size_t n) { return (unsigned)((n + THREADS - 1) / THREADS); } + +} // namespace +} // namespace cs1 + +using namespace cs1; + +extern "C" int cs1_embed(const int32_t* ids, const void* table, void* out, int T, int D, void* stream) { + if (D % 8 != 0) return cudaErrorInvalidValue; + if (T <= 0) return cudaSuccess; + embed_kernel<<(stream)>>>( + ids, static_cast(table), static_cast(out), D / 8); + return cudaGetLastError(); +} + +extern "C" int cs1_gdn_conv(const void* qkv, int ld, const void* w, void* q, void* k, void* v, int T, int key_dim, + int value_dim, void* stream) { + if (T < 0 || key_dim < 0 || value_dim < 0 || ld < 2 * key_dim + value_dim) return cudaErrorInvalidValue; + const size_t n = (size_t)T * (2 * key_dim + value_dim); + if (n == 0) return cudaSuccess; + gdn_conv_kernel<<(stream)>>>( + static_cast(qkv), ld, static_cast(w), static_cast(q), + static_cast(k), static_cast(v), T, key_dim, value_dim); + return cudaGetLastError(); +} + +extern "C" int cs1_gdn_gates(const void* b, const void* a, int ld, const void* A_log, const void* dt_bias, + void* beta, float* g, int T, int H, void* stream) { + if (T < 0 || H < 0 || ld < H) return cudaErrorInvalidValue; + const int n = T * H; + if (n == 0) return cudaSuccess; + gdn_gates_kernel<<(stream)>>>( + static_cast(b), static_cast(a), ld, static_cast(A_log), + static_cast(dt_bias), static_cast(beta), g, n, H); + return cudaGetLastError(); +} + +extern "C" int cs1_sigmoid_gate(void* x, const void* gate, size_t n, void* stream) { + if (n == 0) return cudaSuccess; + sigmoid_gate_kernel<<(stream)>>>( + static_cast(x), static_cast(gate), n); + return cudaGetLastError(); +} + +extern "C" int cs1_silu_mul(const void* gate_up, int ld, void* out, int T, int I, void* stream) { + if (T < 0 || I < 0 || ld < 2 * I) return cudaErrorInvalidValue; + const size_t n = (size_t)T * I; + if (n == 0) return cudaSuccess; + silu_mul_kernel<<(stream)>>>( + static_cast(gate_up), ld, static_cast(out), I, n); + return cudaGetLastError(); +} diff --git a/src/backends/cuda/qwen3_5/gdn_prefill.cu b/src/backends/cuda/qwen3_5/gdn_prefill.cu new file mode 100644 index 00000000..516e51b9 --- /dev/null +++ b/src/backends/cuda/qwen3_5/gdn_prefill.cu @@ -0,0 +1,528 @@ +// Chunked Gated DeltaNet prefill (forward only, batch 1) on tensor cores, sm_80 and later. +// +// The math of torch_chunk_gated_delta_rule in Transformers (chunks of 64), split into +// three kernels the way flash-linear-attention splits its chunked forward pass: +// 1. gdn_chunk_prep, per (chunk, head): L2 norms, the cumulative decays, the pair +// products k.k and q.k (TF32 with float32 accumulation), the triangular inverse +// T = (I + A)^-1 (float32, CUDA cores), and u = T (beta v), w = T (beta exp(cum) k) +// (bfloat16 with float32 accumulation). Results are stored as bfloat16. +// 2. gdn_chunk_state, per (head, 32 value columns), over the chunks in order: keeps +// the state S in float32 registers, stores it as bfloat16 before each chunk, and +// computes v_new = u - w S and S = decay S + kd^T v_new with mma.sync. +// 3. gdn_chunk_out, per (chunk, head): o = qd S + P v_new with mma.sync. +// Transformers computes all of this in float32. Keeping the intermediate results in +// bfloat16, as flash-linear-attention does, makes this kernel less precise than that +// path; tests/kernels.rs checks it against a float64 token-by-token reference. +#include +#include +#include +#include + +#include "mma.cuh" +#include "ops.h" + +using namespace nvcuda; + +namespace { + +using cs1::bf16; +using cs1::warp_sum; + +constexpr int C = 64; // chunk length +constexpr int K = 128; // key head dim +constexpr int V = 128; // value head dim +constexpr int THREADS = 256; +// row strides: multiples of 16 bytes as WMMA needs, and not multiples of 32 floats +constexpr int KP = K + 4; // float +constexpr int CP = C + 4; // float +constexpr int HB = K + 8; // bfloat16 +constexpr int TB = C + 8; // bfloat16 + +using ATf32Row = wmma::fragment; +using BTf32Col = wmma::fragment; +using CTf32 = wmma::fragment; +using ABf16 = wmma::fragment; +using BBf16 = wmma::fragment; +using CBf16 = wmma::fragment; + +template +__device__ __forceinline__ void to_tf32(F& f) { +#pragma unroll + for (int t = 0; t < f.num_elements; t++) f.x[t] = wmma::__float_to_tf32(f.x[t]); +} + +struct Work { + bf16* u; // [H, NCC, V] + bf16* w; // [H, NCC, K] + bf16* qd; // [H, NCC, K] q * scale * exp(cum) + bf16* kd; // [H, NCC, K] k * exp(cum_last - cum) + float* p; // [H, NC, C, C] float32 scratch of q.k + bf16* pb; // [H, NC, C, C] (q.k) exp(cum_i - cum_j), j <= i + float* decay; // [H, NC] exp(cum_last) + bf16* s; // [H, NC, K, V] state before each chunk + bf16* vn; // [H, NCC, V] v_new +}; + +// Byte offsets of the workspace parts, 256-byte aligned. +struct Layout { + size_t u, w, qd, kd, p, pb, decay, s, vn, total; + Layout(int T, int H) { + const size_t NC = (T + C - 1) / C, NCC = NC * C; + size_t at = 0; + auto take = [&](size_t bytes) { + const size_t off = at; + at = (at + bytes + 255) / 256 * 256; + return off; + }; + u = take((size_t)H * NCC * V * 2); + w = take((size_t)H * NCC * K * 2); + qd = take((size_t)H * NCC * K * 2); + kd = take((size_t)H * NCC * K * 2); + p = take((size_t)H * NC * C * C * 4); + pb = take((size_t)H * NC * C * C * 2); + decay = take((size_t)H * NC * 4); + s = take((size_t)H * NC * K * V * 2); + vn = take((size_t)H * NCC * V * 2); + total = at; + } +}; + +// kernel 1 shared memory, bytes +constexpr int R1 = 0; // kn float [C][KP]; later v as bf16 [C][HB] +constexpr int R2 = R1 + C * KP * 4; // qn float [C][KP]; later T float [C][CP] + scratch, then T as bf16 +constexpr int R3 = R2 + C * KP * 4; // A float [C][CP]; later k exp(cum) as bf16 [C][HB] +constexpr int R4 = R3 + C * CP * 4; // cum, beta, exp(cum), exp(cum_last - cum) +constexpr int R5 = R4 + 4 * C * 4; // per-warp 16x16 float staging of u and w +constexpr size_t SMEM1_BYTES = R5 + (THREADS / 32) * 256 * 4; +constexpr int TBF = C * CP * 4 + 3 * 256 * 4; // offset of T as bf16 inside R2 +static_assert(C * HB * 2 <= C * CP * 4, "k exp(cum) as bf16 fits over A"); +static_assert(TBF + C * TB * 2 <= C * KP * 4, "T as bf16 fits in R2"); + +__global__ void __launch_bounds__(THREADS) gdn_chunk_prep( + const __nv_bfloat16* __restrict__ q, const __nv_bfloat16* __restrict__ k, + const __nv_bfloat16* __restrict__ v, const float* __restrict__ g, + const __nv_bfloat16* __restrict__ beta, Work ws, int T, int H, int HK, float scale) { + extern __shared__ __align__(128) unsigned char sm[]; + float* kn = reinterpret_cast(sm + R1); + float* qn = reinterpret_cast(sm + R2); + float* tm = qn; + float* sc = tm + C * CP; + float* am = reinterpret_cast(sm + R3); + float* cum = reinterpret_cast(sm + R4); + float* bet = cum + C; + float* ecum = bet + C; + float* erem = ecum + C; + __nv_bfloat16* vb = reinterpret_cast<__nv_bfloat16*>(sm + R1); + __nv_bfloat16* tb = reinterpret_cast<__nv_bfloat16*>(sm + R2 + TBF); + __nv_bfloat16* kw = reinterpret_cast<__nv_bfloat16*>(sm + R3); + + const int c = blockIdx.x, h = blockIdx.y, NC = gridDim.x, NCC = NC * C; + const int hk = h / (H / HK); + const int tid = threadIdx.x, lane = tid & 31, warp = tid >> 5; + + // 1. load q and k rows of this chunk, L2-normalize (padding rows stay zero) + for (int r = warp; r < C; r += THREADS / 32) { + const int t = c * C + r; + float kv[4], qv[4], ks = 0.f, qs = 0.f; +#pragma unroll + for (int e = 0; e < 4; e++) { + const int d = lane + 32 * e; + float kx = 0.f, qx = 0.f; + if (t < T) { + kx = __bfloat162float(k[((size_t)t * HK + hk) * K + d]); + qx = __bfloat162float(q[((size_t)t * HK + hk) * K + d]); + } + kv[e] = kx; + qv[e] = qx; + ks += kx * kx; + qs += qx * qx; + } + ks = warp_sum(ks); + qs = warp_sum(qs); + const float kinv = rsqrtf(ks + 1e-6f), qinv = rsqrtf(qs + 1e-6f) * scale; +#pragma unroll + for (int e = 0; e < 4; e++) { + const int d = lane + 32 * e; + kn[r * KP + d] = kv[e] * kinv; + qn[r * KP + d] = qv[e] * qinv; + } + } + if (tid < C) { + const int t = c * C + tid; + cum[tid] = t < T ? g[(size_t)t * H + h] : 0.f; + bet[tid] = t < T ? __bfloat162float(beta[(size_t)t * H + h]) : 0.f; + } + __syncthreads(); + if (tid == 0) { + float s = 0.f; + for (int i = 0; i < C; i++) { + s += cum[i]; + cum[i] = s; + } + ws.decay[h * NC + c] = expf(s); + } + __syncthreads(); + if (tid < C) { + ecum[tid] = expf(cum[tid]); + erem[tid] = expf(cum[C - 1] - cum[tid]); + } + __syncthreads(); + + // 2. decayed q and k for the state and output kernels + for (int x = tid; x < C * K; x += THREADS) { + const int i = x / K, d = x % K; + const size_t row = ((size_t)h * NCC + c * C + i) * K + d; + ws.qd[row] = __float2bfloat16(qn[i * KP + d] * ecum[i]); + ws.kd[row] = __float2bfloat16(kn[i * KP + d] * erem[i]); + } + + // 3. pair products on tensor cores: the 10 lower 16x16 tiles of k.k (into A) and of + // q.k (into P, in global memory), then the masks and decays elementwise + float* pout = ws.p + ((size_t)h * NC + c) * C * C; + for (int e = warp; e < 20; e += THREADS / 32) { + const int tile = e % 10; + const int it = tile < 1 ? 0 : tile < 3 ? 1 : tile < 6 ? 2 : 3; + const int jt = tile - it * (it + 1) / 2; + const float* a = (e < 10 ? kn : qn) + it * 16 * KP; + const float* b = kn + jt * 16 * KP; + CTf32 acc; + wmma::fill_fragment(acc, 0.f); +#pragma unroll 4 + for (int k0 = 0; k0 < K; k0 += 8) { + ATf32Row fa; + BTf32Col fb; + wmma::load_matrix_sync(fa, a + k0, KP); + wmma::load_matrix_sync(fb, b + k0, KP); + to_tf32(fa); + to_tf32(fb); + wmma::mma_sync(acc, fa, fb, acc); + } + if (e < 10) + wmma::store_matrix_sync(am + it * 16 * CP + jt * 16, acc, CP, wmma::mem_row_major); + else + wmma::store_matrix_sync(pout + it * 16 * C + jt * 16, acc, C, wmma::mem_row_major); + } + __syncthreads(); + __nv_bfloat16* pb = ws.pb + ((size_t)h * NC + c) * C * C; + for (int x = tid; x < C * C; x += THREADS) { + const int i = x / C, j = x % C; + const float dec = j <= i ? expf(cum[i] - cum[j]) : 0.f; + am[i * CP + j] = j < i ? bet[i] * am[i * CP + j] * dec : 0.f; + pb[i * C + j] = __float2bfloat16(j <= i ? pout[i * C + j] * dec : 0.f); + } + __syncthreads(); + + // 4. T = (I + A)^-1 in 16x16 blocks + for (int x = tid; x < C * CP; x += THREADS) tm[x] = 0.f; + __syncthreads(); + if (tid < C) { + const int base = 16 * (tid / 16), col = tid % 16; + float xs[16]; +#pragma unroll + for (int r = 0; r < 16; r++) { + float acc = r == col ? 1.f : 0.f; +#pragma unroll + for (int j = 0; j < r; j++) acc = fmaf(-am[(base + r) * CP + base + j], xs[j], acc); + xs[r] = acc; + } +#pragma unroll + for (int r = 0; r < 16; r++) tm[(base + r) * CP + base + col] = xs[r]; + } + __syncthreads(); + for (int lv = 1; lv < 4; lv++) { + const int nb = 4 - lv; + for (int x = tid; x < nb * 256; x += THREADS) { + const int bi = lv + x / 256, bj = bi - lv, r = (x % 256) / 16, cc = x % 16; + float acc = 0.f; + for (int kb = bj; kb < bi; kb++) +#pragma unroll + for (int m = 0; m < 16; m++) + acc = fmaf(am[(16 * bi + r) * CP + 16 * kb + m], tm[(16 * kb + m) * CP + 16 * bj + cc], acc); + sc[x] = acc; + } + __syncthreads(); + for (int x = tid; x < nb * 256; x += THREADS) { + const int bi = lv + x / 256, bj = bi - lv, r = (x % 256) / 16, cc = x % 16; + const float* s0 = sc + (x / 256) * 256 + cc; + float acc = 0.f; + for (int m = 0; m <= r; m++) acc = fmaf(tm[(16 * bi + r) * CP + 16 * bi + m], s0[m * 16], acc); + tm[(16 * bi + r) * CP + 16 * bj + cc] = -acc; + } + __syncthreads(); + } + + // 5. bfloat16 operands: k exp(cum) over A, T beta after T, then v over kn; + // u = T (beta v) and w = T (beta exp(cum) k) on tensor cores, stored as bfloat16 + for (int x = tid; x < C * K; x += THREADS) { + const int j = x / K, d = x % K; + kw[j * HB + d] = __float2bfloat16(kn[j * KP + d] * ecum[j]); + } + for (int x = tid; x < C * C; x += THREADS) { + const int i = x / C, j = x % C; + tb[i * TB + j] = __float2bfloat16(tm[i * CP + j] * bet[j]); + } + __syncthreads(); + for (int x = tid; x < C * V; x += THREADS) { + const int j = x / V, d = x % V; + const int t = c * C + j; + vb[j * HB + d] = t < T ? v[((size_t)t * H + h) * V + d] : __float2bfloat16(0.f); + } + __syncthreads(); + float* stage = reinterpret_cast(sm + R5) + warp * 256; + for (int e = warp; e < 64; e += THREADS / 32) { + const bool is_u = e < 32; + const int it = (e % 32) / 8, dt = e % 8; + const __nv_bfloat16* bsrc = is_u ? vb : kw; + CBf16 acc; + wmma::fill_fragment(acc, 0.f); + for (int kb = 0; kb <= it; kb++) { // T is lower triangular + ABf16 fa; + BBf16 fb; + wmma::load_matrix_sync(fa, tb + it * 16 * TB + kb * 16, TB); + wmma::load_matrix_sync(fb, bsrc + kb * 16 * HB + dt * 16, HB); + wmma::mma_sync(acc, fa, fb, acc); + } + wmma::store_matrix_sync(stage, acc, 16, wmma::mem_row_major); + __syncwarp(); + bf16* dst = (is_u ? ws.u : ws.w) + ((size_t)h * NCC + c * C + it * 16) * K + dt * 16; + for (int x = lane; x < 256; x += 32) dst[(x / 16) * K + x % 16] = __float2bfloat16(stage[x]); + __syncwarp(); + } +} + +// ---- kernel 2: the state, chunk by chunk ---- + +constexpr int BVS = 32; // value columns per block +constexpr int ST_THREADS = 128; +constexpr int WS_LD = K + 8; // bfloat16 row stride of staged w and kd +constexpr int SS_LD = BVS + 8; // bfloat16 row stride of the S copy and v_new +constexpr int STAGE = C * WS_LD; // elements of one staged w or kd +constexpr size_t SMEM2_BYTES = (4 * STAGE + K * SS_LD + C * SS_LD) * 2; + +__global__ void __launch_bounds__(ST_THREADS) gdn_chunk_state(Work ws, int NC) { + extern __shared__ __align__(128) unsigned char sm[]; + bf16* wbuf = reinterpret_cast(sm); // [2][C][WS_LD] + bf16* kbuf = wbuf + 2 * STAGE; // [2][C][WS_LD] + bf16* scopy = kbuf + 2 * STAGE; // [K][SS_LD] + bf16* vnew = scopy + K * SS_LD; // [C][SS_LD] + const int h = blockIdx.x, vb0 = blockIdx.y * BVS, NCC = NC * C; + const int tid = threadIdx.x, warp = tid / 32, lane = tid % 32, g = lane / 4, t = lane % 4; + + auto load = [&](int c, int buf) { + const size_t base = ((size_t)h * NCC + (size_t)c * C) * K; + for (int x = tid; x < C * K / 8; x += ST_THREADS) { + const int r = x / (K / 8), col = (x % (K / 8)) * 8; + cs1::cp_async16(wbuf + buf * STAGE + r * WS_LD + col, ws.w + base + (size_t)r * K + col); + cs1::cp_async16(kbuf + buf * STAGE + r * WS_LD + col, ws.kd + base + (size_t)r * K + col); + } + cs1::cp_async_commit(); + }; + + // S rows warp * 32 + mt * 16 + {g, g + 8}, columns nt * 8 + {2t, 2t + 1} + float st[2][4][4]; +#pragma unroll + for (int mt = 0; mt < 2; mt++) +#pragma unroll + for (int nt = 0; nt < 4; nt++) st[mt][nt][0] = st[mt][nt][1] = st[mt][nt][2] = st[mt][nt][3] = 0.f; + + load(0, 0); + for (int c = 0; c < NC; c++) { + const int buf = c & 1; + cs1::cp_async_wait<0>(); + __syncthreads(); // chunk c is staged, and chunk c - 1 is done with the other buffer + if (c + 1 < NC) load(c + 1, buf ^ 1); + // 1. S as bfloat16, to shared memory for w S and to global memory for the output + bf16* sg = ws.s + ((size_t)h * NC + c) * K * V + vb0; +#pragma unroll + for (int mt = 0; mt < 2; mt++) +#pragma unroll + for (int nt = 0; nt < 4; nt++) +#pragma unroll + for (int r = 0; r < 2; r++) { + const int row = warp * 32 + mt * 16 + g + r * 8, col = nt * 8 + 2 * t; + const uint32_t pk = cs1::pack_bf16(st[mt][nt][2 * r], st[mt][nt][2 * r + 1]); + *reinterpret_cast(scopy + row * SS_LD + col) = pk; + *reinterpret_cast(sg + (size_t)row * V + col) = pk; + } + __syncthreads(); + // 2. v_new = u - w S, rows warp * 16 .. + const bf16* wc = wbuf + buf * STAGE; + float acc[4][4]; +#pragma unroll + for (int nt = 0; nt < 4; nt++) acc[nt][0] = acc[nt][1] = acc[nt][2] = acc[nt][3] = 0.f; +#pragma unroll + for (int kk = 0; kk < K; kk += 16) { + uint32_t a[4]; + cs1::load_a(a, wc, WS_LD, warp * 16, kk, lane); +#pragma unroll + for (int nt = 0; nt < 4; nt += 2) { + uint32_t b[4]; + cs1::load_b_kn(b, scopy, SS_LD, kk, nt * 8, lane); + cs1::mma16816(acc[nt], a, b[0], b[1]); + cs1::mma16816(acc[nt + 1], a, b[2], b[3]); + } + } + const size_t row0 = (size_t)h * NCC + (size_t)c * C + warp * 16; +#pragma unroll + for (int nt = 0; nt < 4; nt++) +#pragma unroll + for (int r = 0; r < 2; r++) { + const int row = g + r * 8, col = nt * 8 + 2 * t; + const __nv_bfloat162 uv = + *reinterpret_cast(ws.u + (row0 + row) * V + vb0 + col); + const uint32_t pk = cs1::pack_bf16(__low2float(uv) - acc[nt][2 * r], __high2float(uv) - acc[nt][2 * r + 1]); + *reinterpret_cast(vnew + (warp * 16 + row) * SS_LD + col) = pk; + *reinterpret_cast(ws.vn + (row0 + row) * V + vb0 + col) = pk; + } + __syncthreads(); + // 3. S = decay S + kd^T v_new, rows warp * 32 .. + const float dec = ws.decay[h * NC + c]; +#pragma unroll + for (int mt = 0; mt < 2; mt++) +#pragma unroll + for (int nt = 0; nt < 4; nt++) +#pragma unroll + for (int e = 0; e < 4; e++) st[mt][nt][e] *= dec; + const bf16* kc = kbuf + buf * STAGE; +#pragma unroll + for (int kk = 0; kk < C; kk += 16) { +#pragma unroll + for (int mt = 0; mt < 2; mt++) { + uint32_t a[4]; + cs1::load_a_trans(a, kc, WS_LD, warp * 32 + mt * 16, kk, lane); +#pragma unroll + for (int nt = 0; nt < 4; nt += 2) { + uint32_t b[4]; + cs1::load_b_kn(b, vnew, SS_LD, kk, nt * 8, lane); + cs1::mma16816(st[mt][nt], a, b[0], b[1]); + cs1::mma16816(st[mt][nt + 1], a, b[2], b[3]); + } + } + } + } +} + +// ---- kernel 3: the output, per chunk ---- + +constexpr int OUT_THREADS = 128; +constexpr int O_LD = K + 8; // bfloat16 row stride of qd, S and v_new (all 128 wide) +constexpr int P_LD = C + 8; // bfloat16 row stride of P +constexpr size_t SMEM3_BYTES = (C * O_LD + K * O_LD + C * P_LD + C * O_LD) * 2; +static_assert(K == V, "S, qd and v_new share a row stride"); + +__global__ void __launch_bounds__(OUT_THREADS) gdn_chunk_out(Work ws, bf16* __restrict__ o, int T, int H) { + extern __shared__ __align__(128) unsigned char sm[]; + bf16* qs = reinterpret_cast(sm); // [C][O_LD] + bf16* ss = qs + C * O_LD; // [K][O_LD] + bf16* ps = ss + K * O_LD; // [C][P_LD] + bf16* vs = ps + C * P_LD; // [C][O_LD] + const int c = blockIdx.x, h = blockIdx.y, NC = gridDim.x, NCC = NC * C; + const int tid = threadIdx.x, warp = tid / 32, lane = tid % 32, g = lane / 4, t = lane % 4; + const size_t rows = (size_t)h * NCC + (size_t)c * C; + for (int x = tid; x < C * K / 8; x += OUT_THREADS) { + const int r = x / (K / 8), col = (x % (K / 8)) * 8; + cs1::cp_async16(qs + r * O_LD + col, ws.qd + (rows + r) * K + col); + cs1::cp_async16(vs + r * O_LD + col, ws.vn + (rows + r) * V + col); + } + for (int x = tid; x < K * V / 8; x += OUT_THREADS) { + const int r = x / (V / 8), col = (x % (V / 8)) * 8; + cs1::cp_async16(ss + r * O_LD + col, ws.s + ((size_t)h * NC + c) * K * V + (size_t)r * V + col); + } + for (int x = tid; x < C * C / 8; x += OUT_THREADS) { + const int r = x / (C / 8), col = (x % (C / 8)) * 8; + cs1::cp_async16(ps + r * P_LD + col, ws.pb + ((size_t)h * NC + c) * C * C + r * C + col); + } + cs1::cp_async_commit(); + cs1::cp_async_wait<0>(); + __syncthreads(); + + float acc[V / 8][4]; +#pragma unroll + for (int nt = 0; nt < V / 8; nt++) acc[nt][0] = acc[nt][1] = acc[nt][2] = acc[nt][3] = 0.f; + // qd S +#pragma unroll 2 + for (int kk = 0; kk < K; kk += 16) { + uint32_t a[4]; + cs1::load_a(a, qs, O_LD, warp * 16, kk, lane); +#pragma unroll + for (int nt = 0; nt < V / 8; nt += 2) { + uint32_t b[4]; + cs1::load_b_kn(b, ss, O_LD, kk, nt * 8, lane); + cs1::mma16816(acc[nt], a, b[0], b[1]); + cs1::mma16816(acc[nt + 1], a, b[2], b[3]); + } + } + // P v_new; P is lower triangular, so these rows need positions below (warp + 1) * 16 + for (int kk = 0; kk < (warp + 1) * 16; kk += 16) { + uint32_t a[4]; + cs1::load_a(a, ps, P_LD, warp * 16, kk, lane); +#pragma unroll + for (int nt = 0; nt < V / 8; nt += 2) { + uint32_t b[4]; + cs1::load_b_kn(b, vs, O_LD, kk, nt * 8, lane); + cs1::mma16816(acc[nt], a, b[0], b[1]); + cs1::mma16816(acc[nt + 1], a, b[2], b[3]); + } + } +#pragma unroll + for (int r = 0; r < 2; r++) { + const int tok = c * C + warp * 16 + g + r * 8; + if (tok >= T) continue; + bf16* dst = o + ((size_t)tok * H + h) * V + 2 * t; +#pragma unroll + for (int nt = 0; nt < V / 8; nt++) + *reinterpret_cast(dst + nt * 8) = cs1::pack_bf16(acc[nt][2 * r], acc[nt][2 * r + 1]); + } +} + +Work split(float* workspace, int T, int H) { + const Layout l(T, H); + unsigned char* b = reinterpret_cast(workspace); + Work ws; + ws.u = reinterpret_cast(b + l.u); + ws.w = reinterpret_cast(b + l.w); + ws.qd = reinterpret_cast(b + l.qd); + ws.kd = reinterpret_cast(b + l.kd); + ws.p = reinterpret_cast(b + l.p); + ws.pb = reinterpret_cast(b + l.pb); + ws.decay = reinterpret_cast(b + l.decay); + ws.s = reinterpret_cast(b + l.s); + ws.vn = reinterpret_cast(b + l.vn); + return ws; +} + +} // namespace + +extern "C" { + +size_t cs1_gdn_workspace_floats(int T, int H) { return (Layout(T, H).total + 3) / 4; } + +int cs1_gdn_prefill(const void* q, const void* k, const void* v, const float* g, const void* beta, + void* o, float* workspace, int T, int H, int HK, float scale, void* stream) { + if (T < 0 || HK <= 0 || H % HK != 0) return cudaErrorInvalidValue; + if (T == 0) return cudaSuccess; + // once per process (for the device current at the first call) + static const cudaError_t configured = [] { + cudaError_t e = cudaFuncSetAttribute(gdn_chunk_prep, cudaFuncAttributeMaxDynamicSharedMemorySize, + (int)SMEM1_BYTES); + if (e == cudaSuccess) + e = cudaFuncSetAttribute(gdn_chunk_state, cudaFuncAttributeMaxDynamicSharedMemorySize, + (int)SMEM2_BYTES); + if (e == cudaSuccess) + e = cudaFuncSetAttribute(gdn_chunk_out, cudaFuncAttributeMaxDynamicSharedMemorySize, + (int)SMEM3_BYTES); + return e; + }(); + if (configured != cudaSuccess) return configured; + const int NC = (T + C - 1) / C; + const Work ws = split(workspace, T, H); + cudaStream_t st = static_cast(stream); + gdn_chunk_prep<<>>( + static_cast(q), static_cast(k), + static_cast(v), g, static_cast(beta), ws, T, H, HK, scale); + gdn_chunk_state<<>>(ws, NC); + gdn_chunk_out<<>>(ws, static_cast(o), T, H); + return cudaGetLastError(); +} + +} // extern "C" diff --git a/src/backends/cuda/qwen3_5/gemm.cu b/src/backends/cuda/qwen3_5/gemm.cu new file mode 100644 index 00000000..7de47f57 --- /dev/null +++ b/src/backends/cuda/qwen3_5/gemm.cu @@ -0,0 +1,129 @@ +// bfloat16 GEMMs through cuBLASLt, float32 accumulation. +// +// Row-major y [M, N] = x [M, K] * w [N, K]^T is the column-major product +// y^T [N, M] = (w viewed as [K, N])^T * (x viewed as [K, M]); y's rows may be +// strided (ldy >= N), so one GEMM can fill a slice of a wider buffer. +// +// Each shape uses cuBLASLt's first heuristic choice, excluding split-K reductions that +// accumulate into the output in place, since their order, and so the rounding, is not +// fixed. +#include + +#include +#include + +#include "ops.h" + +namespace { + +struct Plan { + cublasLtMatmulDesc_t op = nullptr; + cublasLtMatrixLayout_t a = nullptr, b = nullptr, c = nullptr; + cublasLtMatmulAlgo_t algo{}; +}; + +using Key = std::tuple; // M, N, K, ldy + +struct Gemm { + cublasLtHandle_t handle = nullptr; + void* workspace = nullptr; + size_t workspace_bytes = 0; + std::map plans; +}; + +int status(cublasStatus_t s) { return s == CUBLAS_STATUS_SUCCESS ? 0 : 1000 + (int)s; } + +void destroy(Plan& p) { + if (p.a) cublasLtMatrixLayoutDestroy(p.a); + if (p.b) cublasLtMatrixLayoutDestroy(p.b); + if (p.c) cublasLtMatrixLayoutDestroy(p.c); + if (p.op) cublasLtMatmulDescDestroy(p.op); + p = Plan{}; +} + +int describe(int M, int N, int K, int ldy, Plan& p) { + 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); + return 0; +} + +// The heuristic's first choice, without in-place split-K reductions. +int first_choice(Gemm& g, Plan& p) { + 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; + cublasLtMatmulPreferenceSetAttribute(pref, CUBLASLT_MATMUL_PREF_REDUCTION_SCHEME_MASK, &schemes, + sizeof(schemes)); + cublasLtMatmulHeuristicResult_t r{}; + int found = 0; + s = cublasLtMatmulAlgoGetHeuristic(g.handle, p.op, p.a, p.b, p.c, p.c, pref, 1, &r, &found); + cublasLtMatmulPreferenceDestroy(pref); + if (s != CUBLAS_STATUS_SUCCESS) return status(s); + if (found == 0 || r.state != CUBLAS_STATUS_SUCCESS) return status(CUBLAS_STATUS_NOT_SUPPORTED); + p.algo = r.algo; + return 0; +} + +// 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}; + 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); + if (rc != 0) { + destroy(p); + return rc; + } + out = &g.plans.emplace(key, p).first->second; + return 0; +} + +} // namespace + +extern "C" void* cs1_gemm_create(size_t workspace_bytes) { + Gemm* g = new Gemm(); + if (cublasLtCreate(&g->handle) != CUBLAS_STATUS_SUCCESS || + (workspace_bytes > 0 && cudaMalloc(&g->workspace, workspace_bytes) != cudaSuccess)) { + if (g->handle) cublasLtDestroy(g->handle); + delete g; + return nullptr; + } + g->workspace_bytes = workspace_bytes; + return g; +} + +extern "C" void cs1_gemm_destroy(void* gemm) { + Gemm* g = static_cast(gemm); + if (!g) return; + for (auto& kv : g->plans) destroy(kv.second); + if (g->workspace) cudaFree(g->workspace); + cublasLtDestroy(g->handle); + delete g; +} + +extern "C" int cs1_gemm(void* gemm, const void* x, const void* w, void* y, int M, int N, int K, int ldy, + void* stream) { + Gemm* g = static_cast(gemm); + if (!g || M < 0 || N <= 0 || K <= 0 || ldy < N) return cudaErrorInvalidValue; + if (M == 0) return cudaSuccess; + Plan* p = nullptr; + const int rc = plan_for(*g, M, N, K, ldy, p); + if (rc != 0) 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/mma.cuh b/src/backends/cuda/qwen3_5/mma.cuh new file mode 100644 index 00000000..e865fb75 --- /dev/null +++ b/src/backends/cuda/qwen3_5/mma.cuh @@ -0,0 +1,74 @@ +// Tensor-core building blocks for sm_80 and later: mma.sync m16n8k16 on bfloat16 with +// float32 accumulation, ldmatrix, and cp.async. +// +// Fragment layouts (PTX ISA, mma.m16n8k16): with g = lane / 4 and t = lane % 4, an +// accumulator holds rows g (elements 0, 1) and g + 8 (elements 2, 3) at columns 2t and +// 2t + 1 of its 8-column tile; the A operand registers are (row g, cols 2t..), +// (row g + 8, cols 2t..), (row g, cols 8 + 2t..) and (row g + 8, cols 8 + 2t..). +#pragma once + +#include "common.cuh" + +namespace cs1 { + +__device__ __forceinline__ uint32_t smem_addr(const void* p) { + return static_cast(__cvta_generic_to_shared(p)); +} + +// 16-byte asynchronous copy from global to shared memory; zero-fills when !valid. +__device__ __forceinline__ void cp_async16(void* dst, const void* src, bool valid = true) { + asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" ::"r"(smem_addr(dst)), "l"(src), + "r"(valid ? 16 : 0)); +} +__device__ __forceinline__ void cp_async_commit() { asm volatile("cp.async.commit_group;\n" ::); } +template +__device__ __forceinline__ void cp_async_wait() { + asm volatile("cp.async.wait_group %0;\n" ::"n"(N)); +} + +__device__ __forceinline__ void ldmatrix_x4(uint32_t (&r)[4], const bf16* p) { + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];\n" + : "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3]) + : "r"(smem_addr(p))); +} +__device__ __forceinline__ void ldmatrix_x4_trans(uint32_t (&r)[4], const bf16* p) { + asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0,%1,%2,%3}, [%4];\n" + : "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3]) + : "r"(smem_addr(p))); +} + +// A fragment of the 16x16 tile at (row0, col0) of a row-major shared matrix. +__device__ __forceinline__ void load_a(uint32_t (&a)[4], const bf16* m, int ld, int row0, int col0, int lane) { + ldmatrix_x4(a, m + (row0 + (lane % 8) + ((lane / 8) % 2) * 8) * ld + col0 + (lane / 16) * 8); +} +// A fragment of the 16x16 tile at (row0, col0) of the transpose of a row-major shared +// matrix: rows of A are columns of m. +__device__ __forceinline__ void load_a_trans(uint32_t (&a)[4], const bf16* m, int ld, int row0, int col0, + int lane) { + ldmatrix_x4_trans(a, m + (col0 + (lane % 8) + (lane / 16) * 8) * ld + row0 + ((lane / 8) % 2) * 8); +} +// B fragments of two 16x8 tiles (k0.., n0..) and (k0.., n0 + 8..) of a row-major +// shared matrix whose rows are k: b[0], b[1] and b[2], b[3]. +__device__ __forceinline__ void load_b_kn(uint32_t (&b)[4], const bf16* m, int ld, int k0, int n0, int lane) { + ldmatrix_x4_trans(b, m + (k0 + (lane % 8) + ((lane / 8) % 2) * 8) * ld + n0 + (lane / 16) * 8); +} +// The same when the shared matrix is stored with rows n (each row contiguous in k). +__device__ __forceinline__ void load_b_nk(uint32_t (&b)[4], const bf16* m, int ld, int k0, int n0, int lane) { + ldmatrix_x4(b, m + (n0 + (lane % 8) + (lane / 16) * 8) * ld + k0 + ((lane / 8) % 2) * 8); +} + +// d += a * b for a 16x16 (row-major) by 16x8 (column-major) tile. +__device__ __forceinline__ void mma16816(float (&d)[4], const uint32_t (&a)[4], uint32_t b0, uint32_t b1) { + asm volatile( + "mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, " + "{%0,%1,%2,%3};\n" + : "+f"(d[0]), "+f"(d[1]), "+f"(d[2]), "+f"(d[3]) + : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b0), "r"(b1)); +} + +__device__ __forceinline__ uint32_t pack_bf16(float lo, float hi) { + const __nv_bfloat162 v = __floats2bfloat162_rn(lo, hi); + return *reinterpret_cast(&v); +} + +} // namespace cs1 diff --git a/src/backends/cuda/qwen3_5/norm.cu b/src/backends/cuda/qwen3_5/norm.cu new file mode 100644 index 00000000..38d6f98a --- /dev/null +++ b/src/backends/cuda/qwen3_5/norm.cu @@ -0,0 +1,106 @@ +// RMSNorm variants of Qwen3.5. +#include "common.cuh" +#include "ops.h" + +namespace cs1 { +namespace { + +constexpr int NORM_THREADS = 256; + +// Qwen3_5RMSNorm: float32 normalization, times (1 + w), rounded once to bfloat16. +__global__ void __launch_bounds__(NORM_THREADS) + rms_norm_kernel(const bf16* __restrict__ x, const bf16* __restrict__ w, bf16* __restrict__ out, int D, + float eps) { + __shared__ float scratch[32]; + const size_t row = blockIdx.x; + x += row * D; + out += row * D; + float ss = 0.f; + for (int i = threadIdx.x; i < D; i += blockDim.x) { + const float v = f32(x[i]); + ss += v * v; + } + const float inv = rsqrtf(block_sum(ss, scratch) / D + eps); + for (int i = threadIdx.x; i < D; i += blockDim.x) out[i] = to_bf16(f32(x[i]) * inv * (1.f + f32(w[i]))); +} + +// The residual add of a decoder layer (bfloat16 + bfloat16, rounded) fused with the +// following Qwen3_5RMSNorm. Each thread rereads only the elements it wrote. +__global__ void __launch_bounds__(NORM_THREADS) + add_rms_norm_kernel(bf16* __restrict__ residual, const bf16* __restrict__ delta, const bf16* __restrict__ w, + bf16* __restrict__ out, int D, float eps) { + __shared__ float scratch[32]; + const size_t row = blockIdx.x; + residual += row * D; + delta += row * D; + out += row * D; + float ss = 0.f; + for (int i = threadIdx.x; i < D; i += blockDim.x) { + const bf16 r = to_bf16(f32(residual[i]) + f32(delta[i])); + residual[i] = r; + const float v = f32(r); + ss += v * v; + } + const float inv = rsqrtf(block_sum(ss, scratch) / D + eps); + for (int i = threadIdx.x; i < D; i += blockDim.x) + out[i] = to_bf16(f32(residual[i]) * inv * (1.f + f32(w[i]))); +} + +// Qwen3_5RMSNormGated with D = 128, one warp per row: the normalized value is +// rounded to bfloat16, multiplied by w in bfloat16, then by silu(z) in float32. +__global__ void gated_rms_norm_kernel(const bf16* __restrict__ x, const bf16* __restrict__ z, int ldz, + const bf16* __restrict__ w, bf16* __restrict__ out, int rows, int H, + float eps) { + constexpr int D = 128, PER = D / 32; + const int row = blockIdx.x * (blockDim.x / 32) + threadIdx.x / 32, lane = threadIdx.x & 31; + if (row >= rows) return; + const size_t base = (size_t)row * D + lane * PER; + const size_t zbase = (size_t)(row / H) * ldz + (row % H) * D + lane * PER; + float v[PER]; + float ss = 0.f; +#pragma unroll + for (int i = 0; i < PER; i++) { + v[i] = f32(x[base + i]); + ss += v[i] * v[i]; + } + const float inv = rsqrtf(warp_sum(ss) / D + eps); +#pragma unroll + for (int i = 0; i < PER; i++) { + const float normed = round_bf16(v[i] * inv); + const float scaled = round_bf16(f32(w[lane * PER + i]) * normed); + out[base + i] = to_bf16(scaled * silu(f32(z[zbase + i]))); + } +} + +} // namespace +} // namespace cs1 + +using namespace cs1; + +extern "C" int cs1_rms_norm(const void* x, const void* w, void* out, int rows, int D, float eps, void* stream) { + if (rows <= 0) return cudaSuccess; + rms_norm_kernel<<(stream)>>>( + static_cast(x), static_cast(w), static_cast(out), D, eps); + return cudaGetLastError(); +} + +extern "C" int cs1_add_rms_norm(void* residual, const void* delta, const void* w, void* out, int rows, int D, + float eps, void* stream) { + if (rows <= 0) return cudaSuccess; + add_rms_norm_kernel<<(stream)>>>( + static_cast(residual), static_cast(delta), static_cast(w), + static_cast(out), D, eps); + return cudaGetLastError(); +} + +extern "C" int cs1_gated_rms_norm(const void* x, const void* z, int ldz, const void* w, void* out, int T, int H, + int D, float eps, void* stream) { + if (D != 128 || ldz < H * D) return cudaErrorInvalidValue; + const int rows = T * H; + if (rows <= 0) return cudaSuccess; + constexpr int WARPS = 8; + gated_rms_norm_kernel<<<(rows + WARPS - 1) / WARPS, WARPS * 32, 0, static_cast(stream)>>>( + static_cast(x), static_cast(z), ldz, static_cast(w), + static_cast(out), rows, H, eps); + return cudaGetLastError(); +} diff --git a/src/backends/cuda/qwen3_5/ops.h b/src/backends/cuda/qwen3_5/ops.h new file mode 100644 index 00000000..6f63488e --- /dev/null +++ b/src/backends/cuda/qwen3_5/ops.h @@ -0,0 +1,95 @@ +// C interface of libqwen3_5_cuda.so: the Qwen3.5 prefill operations and the few +// CUDA runtime calls a caller needs, so that it can load this one library at run time. +// +// Tensors are row-major and bfloat16 unless noted. Every operation queues work on +// `stream` (a cudaStream_t) and returns a cudaError_t, or 1000 + a cublasStatus_t for +// the GEMMs. The norm, elementwise and attention-prep operations round to bfloat16 +// where the Transformers reference (modeling_qwen3_5.py) does; attention and the +// Gated DeltaNet prefill keep some intermediate results in bfloat16, as FlashAttention +// and flash-linear-attention do (see attention.cu and gdn_prefill.cu). +#pragma once + +#include +#include + +// Bumped whenever a signature below changes. +#define CS1_ABI_VERSION 2 + +#ifdef __cplusplus +extern "C" { +#endif + +// ---- runtime ---- + +uint32_t cs1_abi_version(void); +const char* cs1_error_string(int code); +int cs1_set_device(int device); +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); +// Copy and wait for the copy. +int cs1_upload(void* dst, const void* src, size_t bytes, void* stream); +int cs1_download(void* dst, const void* src, size_t bytes, void* stream); + +// ---- operations ---- + +// out[t] = table[ids[t]], rows of D. +int cs1_embed(const int32_t* ids, const void* table, void* out, int T, int D, void* stream); + +// Zero-centred RMSNorm over rows of D: out = x / rms(x) * (1 + w), in float32. +int cs1_rms_norm(const void* x, const void* w, void* out, int rows, int D, float eps, void* stream); + +// residual = residual + delta (rounded to bfloat16), then out = cs1_rms_norm(residual). +int cs1_add_rms_norm(void* residual, const void* delta, const void* w, void* out, int rows, int D, + float eps, void* stream); + +// Gated RMSNorm of the Gated DeltaNet output x [T, H, D] (D = 128), with z [T, H*D] +// in rows of ldz: out = (w * (x / rms(x))) * silu(z). +int cs1_gated_rms_norm(const void* x, const void* z, int ldz, const void* w, void* out, int T, int H, + int D, float eps, void* stream); + +// Depthwise causal conv1d (kernel 4, no bias) and SiLU over qkv [T, ld], split into +// q [T, key_dim], k [T, key_dim] and v [T, value_dim]. w is [key_dim*2 + value_dim, 4]. +int cs1_gdn_conv(const void* qkv, int ld, const void* w, void* q, void* k, void* v, int T, int key_dim, + int value_dim, void* stream); + +// beta = sigmoid(b) (bfloat16) and g = -exp(A_log) * softplus(a + dt_bias) (float32), [T, H]; +// b and a are [T, H] in rows of ld. +int cs1_gdn_gates(const void* b, const void* a, int ld, const void* A_log, const void* dt_bias, + void* beta, float* g, int T, int H, void* stream); + +// Chunked gated delta rule, q and k L2-normalized inside, q scaled by `scale`. +// q, k [T, HK, 128], v [T, H, 128], g float [T, H], beta [T, H], o [T, H, 128]. +size_t cs1_gdn_workspace_floats(int T, int H); +int cs1_gdn_prefill(const void* q, const void* k, const void* v, const float* g, const void* beta, + void* o, float* workspace, int T, int H, int HK, float scale, void* stream); + +// Attention inputs: q and gate from qg [T, Hq, 2*Dh], k from kr [T, Hk, Dh], both in rows +// of ld; per-head zero-centred RMSNorm, then rotary embedding on the first 2*half dims +// using cos/sin [T, half] (bfloat16). Writes q [T, Hq, Dh], gate [T, Hq*Dh], k [T, Hk, Dh]. +int cs1_attn_prep(const void* qg, const void* kr, int ld, const void* qw, const void* kw, + const void* cos, const void* sin, void* q, void* gate, void* k, int T, int Hq, int Hk, + int Dh, int half, float eps, void* stream); + +// Causal attention with grouped KV heads, Dh = 256: q [T, Hq, Dh], k [T, Hk, Dh], v +// [T, Hk, Dh] in rows of ldv; out [T, Hq, Dh]. +int cs1_attention(const void* q, const void* k, const void* v, int ldv, void* out, int T, int Hq, + int Hk, int Dh, float scale, void* stream); + +// x = x * sigmoid(gate), n elements. +int cs1_sigmoid_gate(void* x, const void* gate, size_t n, void* stream); + +// out [T, I] = silu(gate) * up, from gate_up [T, 2*I] (gate first) in rows of ld. +int cs1_silu_mul(const void* gate_up, int ld, void* out, int T, int I, void* stream); + +// y [M, N] (rows of ldy) = x [M, K] * w [N, K]^T through cuBLASLt, float32 accumulation, +// with cuBLASLt's first heuristic choice for each shape (see gemm.cu). +void* cs1_gemm_create(size_t workspace_bytes); +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); + +#ifdef __cplusplus +} +#endif diff --git a/src/backends/cuda/qwen3_5/runtime.cu b/src/backends/cuda/qwen3_5/runtime.cu new file mode 100644 index 00000000..8361b83f --- /dev/null +++ b/src/backends/cuda/qwen3_5/runtime.cu @@ -0,0 +1,37 @@ +// The CUDA runtime calls the Rust side needs, so that it loads one library +// (libqwen3_5_cuda.so) and never links CUDA itself. +#include + +#include "ops.h" + +extern "C" { + +uint32_t cs1_abi_version(void) { return CS1_ABI_VERSION; } + +const char* cs1_error_string(int code) { return cudaGetErrorString(static_cast(code)); } + +int cs1_set_device(int device) { return cudaSetDevice(device); } + +int cs1_malloc(void** ptr, size_t bytes) { return cudaMalloc(ptr, bytes); } + +int cs1_free(void* ptr) { return cudaFree(ptr); } + +int cs1_stream_create(void** stream) { + return cudaStreamCreateWithFlags(reinterpret_cast(stream), cudaStreamNonBlocking); +} + +int cs1_stream_sync(void* stream) { return cudaStreamSynchronize(static_cast(stream)); } + +int cs1_upload(void* dst, const void* src, size_t bytes, void* stream) { + const cudaStream_t st = static_cast(stream); + const cudaError_t e = cudaMemcpyAsync(dst, src, bytes, cudaMemcpyHostToDevice, st); + return e != cudaSuccess ? e : cudaStreamSynchronize(st); +} + +int cs1_download(void* dst, const void* src, size_t bytes, void* stream) { + const cudaStream_t st = static_cast(stream); + const cudaError_t e = cudaMemcpyAsync(dst, src, bytes, cudaMemcpyDeviceToHost, st); + return e != cudaSuccess ? e : cudaStreamSynchronize(st); +} + +} // extern "C" diff --git a/src/frontend/cua_s1_text.py b/src/frontend/cua_s1_text.py new file mode 100644 index 00000000..ef3a330c --- /dev/null +++ b/src/frontend/cua_s1_text.py @@ -0,0 +1,110 @@ +"""HTTP worker for Cua-S1 4B 0.2 (`text` adapter): `GET /health` and `POST /v1/systemone`. + +PYTHONPATH=src python -m frontend.cua_s1_text --base --adapter /text +""" + +from __future__ import annotations + +import argparse +import asyncio +import sys +import traceback +from concurrent.futures import ThreadPoolExecutor +from typing import Any + +# FastAPI reads the handler annotations at runtime, so `Request` must be a +# module-level name while `from __future__ import annotations` is in effect. +from fastapi import FastAPI, Request +from fastapi.responses import JSONResponse + +from models.cua_s1.text.contract import ( + MODEL_ID, + RequestError, + answer, + map_request, + parse_body, +) + +MAX_BODY_BYTES = 4 << 20 +MAX_PROMPT_TOKENS = 16384 +WARMUP = ( + b'{"model": "cua-s1-4b-0.2", "state": "Dialog: Update installed.", "questions": {"q":' + b' {"type": "choice", "instructions": "Close it.", "criteria": {"ok": "OK", "wait": "Wait"}}}}' +) + + +def build_app(model: Any) -> FastAPI: + app = FastAPI() + pool = ThreadPoolExecutor(max_workers=1) # one forward pass at a time + + def decide(raw: bytes) -> dict[str, Any]: + state, questions = map_request(parse_body(raw)) + encoded = [model.encode(state, q) for q in questions] + tokens = [int(x["input_ids"].shape[1]) for x in encoded] + for q, n in zip(questions, tokens): # before any forward pass + if n > MAX_PROMPT_TOKENS: + raise RequestError( + f"question {q.name!r}: {n} prompt tokens, over {MAX_PROMPT_TOKENS}", + 413, + ) + answers = { + q.name: answer(q, model.score(x, len(q.keys))) + for q, x in zip(questions, encoded) + } + return { + "model": MODEL_ID, + "answers": answers, + "usage": {"input_tokens": sum(tokens), "output_tokens": 0}, + } + + @app.get("/health") + def health(): + return {"status": "ready", "model": MODEL_ID} + + @app.post("/v1/systemone") + async def systemone(request: Request): + raw = bytearray() + async for chunk in request.stream(): + raw += chunk + if len(raw) > MAX_BODY_BYTES: + return JSONResponse({"detail": "request body too large"}, 413) + try: + # JSONResponse, not FastAPI's encoder, which drops keys starting with "_sa". + return JSONResponse( + await asyncio.get_running_loop().run_in_executor(pool, decide, raw) + ) + except RequestError as error: + return JSONResponse({"detail": str(error)}, error.status) + except Exception: + traceback.print_exc(file=sys.stderr) + return JSONResponse({"detail": "inference failed"}, 500) + + # One decision on the worker thread before listening, so the first request does not + # pay for first-call setup there. + app.state.warmup = lambda: pool.submit(decide, WARMUP).result() + return app + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--base", required=True, help="local Qwen/Qwen3.5-4B directory") + parser.add_argument( + "--adapter", required=True, help="local text/ adapter directory" + ) + parser.add_argument("--device", default="cuda") + parser.add_argument("--dtype", default="bfloat16", choices=["bfloat16", "float32"]) + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, default=8000) + args = parser.parse_args() + + import uvicorn + + from models.cua_s1.text.model import TextModel + + app = build_app(TextModel(args.base, args.adapter, args.device, args.dtype)) + app.state.warmup() + uvicorn.run(app, host=args.host, port=args.port, log_level="warning") + + +if __name__ == "__main__": + main() diff --git a/src/models/cua_s1/README.md b/src/models/cua_s1/README.md index 6a072600..999ae0a5 100644 --- a/src/models/cua_s1/README.md +++ b/src/models/cua_s1/README.md @@ -2,7 +2,7 @@ This directory owns Cua-S1 4B 0.2 ([#10](https://github.com/ThinkFlowLab/system1-omni/issues/10)): request mapping, prompt construction, adapter selection, execution, and the answer-letter readout. This page records the pinned upstream revisions, the inference contract an implementation must match, and how its outputs will be compared with the upstream reference. -Status: planned; nothing is implemented or validated yet. The first target is the `text` adapter on CUDA, starting with a worker that loads the model directly through Hugging Face Transformers and PEFT. The `multimodal` adapter is deferred; see [Not covered yet](#not-covered-yet). +Status: a reference worker for the `text` adapter loads the model through Hugging Face Transformers and PEFT: [`text/`](text/), served by [`src/frontend/cua_s1_text.py`](../../frontend/cua_s1_text.py), with setup in [`recipe/cua_s1/text.md`](../../../recipe/cua_s1/text.md). It is the correctness reference for the native worker in [`native/`](native/): Rust, with the Qwen3.5 forward pass on the CUDA kernels in [`src/backends/cuda/qwen3_5/`](../../backends/cuda/qwen3_5/), set up as in [`recipe/cua_s1/native.md`](../../../recipe/cua_s1/native.md). The `multimodal` adapter is deferred; see [Not covered yet](#not-covered-yet). ## Pinned revisions @@ -76,19 +76,19 @@ The response `model` is `cua-ai/cua-s1-4b-0.2@:`, in An error rejects the whole request. Its body is `{"detail": ""}`, as the LAYA worker returns, and the message names the problem. -The status is `400` when the body is not a usable JSON object: invalid JSON or UTF-8, `NaN` or `Infinity`, a lone surrogate such as `\ud800`, nesting too deep to parse, or a key repeated in any object. +The status is `400` when the body is not a usable JSON object: invalid JSON or UTF-8, `NaN`, `Infinity` or a number out of range, a lone surrogate such as `\ud800`, nesting too deep to parse, or a key repeated in any object. The status is `422` when a well-formed request cannot be answered: - a `score` or `noul` question, since the adapters were trained only on closed-option choices; - a question without an `instructions` field (`null` is allowed), or a `choice` with no options or more than 26 options; - a `criteria` value that is a number or a boolean; -- an empty `state`; +- an empty `state` (`""`, `{}` or `[]`); - a `model` other than `cua-s1-4b-0.2`. ## Validation -**Inputs.** The fixed input set is upstream's two checked-in fixtures, converted to `/v1/systemone` requests with the chooser's rendered regions as `state`, plus `/v1/systemone` choice requests that will be checked in with the worker. These cover 1 to 26 options, short and long states, string and structured `state`, `instructions` and `criteria`, `null` criteria, non-ASCII text, and text that spells a special token. Each input is scored once per configuration. +**Inputs.** The fixed input set is upstream's two checked-in fixtures, converted to `/v1/systemone` requests with the chooser's rendered regions as `state`, plus `/v1/systemone` choice requests. These cover 1 to 26 options, short and long states, string and structured `state`, `instructions` and `criteria`, `null` criteria, non-ASCII text, and text that spells a special token. Each input is scored once per configuration. **Tolerances.** These are declared before any comparison is run: @@ -97,7 +97,7 @@ The status is `422` when a well-formed request cannot be answered: | Worker prompt vs upstream `build_prompt` | Same inputs | Identical token ids | | Worker vs upstream `FourBModel` | Same GPU and reference environment, bfloat16, adapter not merged, full logits, one unpadded prompt per forward pass | Identical fp32 probabilities | | Through the frontend vs direct to the worker | Same worker | Identical status, content type and body bytes | -| Native engine vs fp32 worker (later) | Same GPU; the engine runs in bfloat16; the fp32 worker runs with TF32 disabled | (1) Over the whole input set, the largest per-option probability difference is at most twice the bfloat16 worker's largest difference from the fp32 worker, plus 0.01. (2) The top option matches wherever the fp32 worker's top-two margin is at least 0.05. | +| Native engine vs fp32 worker | Same GPU; the engine runs in bfloat16; the fp32 worker runs with TF32 disabled | (1) Over the whole input set, the largest per-option probability difference is at most twice the bfloat16 worker's largest difference from the fp32 worker, plus 0.01. (2) The top option matches wherever the fp32 worker's top-two margin is at least 0.05. | The bfloat16 worker's own difference from the fp32 worker is reported next to each native-engine result. diff --git a/src/models/cua_s1/native/Cargo.toml b/src/models/cua_s1/native/Cargo.toml new file mode 100644 index 00000000..2e8caf68 --- /dev/null +++ b/src/models/cua_s1/native/Cargo.toml @@ -0,0 +1,24 @@ +[package] +name = "omni-cua-s1-native" +version = "0.1.0" +edition = "2024" +publish = false +description = "Native /v1/systemone worker for Cua-S1 4B 0.2 (text adapter)" + +[[bin]] +name = "omni-cua-s1-native" +path = "src/main.rs" + +[dependencies] +anyhow = "1.0.100" +axum = "0.8.8" +half = "2.7.1" +# the CUDA kernels live in libqwen3_5_cuda.so, loaded at run time +libloading = "0.8" +memmap2 = "0.9.9" +safetensors = "0.8.0" +serde = "1" +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"] } diff --git a/src/models/cua_s1/native/src/contract.rs b/src/models/cua_s1/native/src/contract.rs new file mode 100644 index 00000000..0ca10f29 --- /dev/null +++ b/src/models/cua_s1/native/src/contract.rs @@ -0,0 +1,252 @@ +//! Request mapping, prompts and answers for Cua-S1 4B 0.2, as in +//! `src/models/cua_s1/README.md` and the reference worker's `text/contract.py`. + +use serde_json::{Map, Value, json}; + +use crate::json::{dumps, parse, quote}; + +pub const MODEL_NAME: &str = "cua-s1-4b-0.2"; +pub const MODEL_ID: &str = "cua-ai/cua-s1-4b-0.2@16818868b0cc7813808aae4e87b417657046ab79:text"; +pub const LETTERS: &str = "ABCDEFGHIJKLMNOPQRSTUVWXYZ"; +pub const MAX_QUESTIONS: usize = 64; + +// The system message and the user message layout are copied from trycua/cua at +// 0e75660ce4c2edda519e0c795fa3ad98abf4e76f (`libs/cua-s1/python/src/cua_s1/four_b.py` +// and `libs/cua-driver/examples/jev-use/python/decision_models.py`). +// +// MIT License +// +// Copyright (c) 2025 Cua AI, Inc. +// +// Permission is hereby granted, free of charge, to any person obtaining a copy +// of this software and associated documentation files (the "Software"), to deal +// in the Software without restriction, including without limitation the rights +// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +// copies of the Software, and to permit persons to whom the Software is +// furnished to do so, subject to the following conditions: +// +// The above copyright notice and this permission notice shall be included in all +// copies or substantial portions of the Software. +// +// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +// SOFTWARE. +pub const SYSTEM_PROMPT: &str = "You are a one-pass computer-use decision model. You are shown the \ +current state of a screen and a fixed, closed list of candidate \ +(element, action) options, each given a single letter. Choose exactly \ +one option: the single best next action to take. Answer with ONLY that \ +option's letter -- no words, no punctuation, no explanation."; + +pub struct RequestError { + pub status: u16, + pub message: String, +} + +fn error(status: u16, message: impl Into) -> RequestError { + RequestError { + status, + message: message.into(), + } +} + +pub struct Question { + pub name: String, + pub goal: String, + pub keys: Vec, + pub labels: Vec, +} + +pub fn parse_body(raw: &[u8]) -> Result, RequestError> { + parse(raw).map_err(|message| error(400, message)) +} + +/// A string as is; an object or array as Python's `json.dumps` writes it. +fn text(value: &Value, place: &str) -> Result { + match value { + Value::String(s) => Ok(s.clone()), + Value::Object(_) | Value::Array(_) => Ok(dumps(value)), + _ => Err(error( + 422, + format!("{place} must be a string, an object or an array"), + )), + } +} + +pub fn map_request(body: &Map) -> Result<(String, Vec), RequestError> { + if body.get("model").and_then(Value::as_str) != Some(MODEL_NAME) { + return Err(error(422, format!("'model' must be '{MODEL_NAME}'"))); + } + let state = body.get("state").unwrap_or(&Value::Null); + if [json!(""), json!({}), json!([])].contains(state) { + return Err(error(422, "'state' must not be empty")); + } + let state = text(state, "'state'")?; + let questions = match body.get("questions") { + Some(Value::Object(q)) if !q.is_empty() => q, + _ => return Err(error(422, "'questions' must be a non-empty object")), + }; + if questions.len() > MAX_QUESTIONS { + return Err(error(413, format!("more than {MAX_QUESTIONS} questions"))); + } + let mut mapped = Vec::with_capacity(questions.len()); + for (name, q) in questions { + let place = format!("question {}", quote(name)); + let Value::Object(q) = q else { + return Err(error(422, format!("{place} must be an object"))); + }; + match q.get("type").unwrap_or(&Value::Null) { + Value::String(t) if t == "choice" => {} + Value::String(t) if t == "score" || t == "noul" => { + return Err(error(422, format!("{place}: type '{t}' is not supported"))); + } + other => return Err(error(422, format!("{place}: unknown type {other}"))), + } + let goal = match q.get("instructions") { + None => return Err(error(422, format!("{place}: 'instructions' is required"))), + Some(Value::Null) => String::new(), + Some(value) => text(value, &place)?, + }; + let criteria = match q.get("criteria") { + Some(Value::Object(c)) if (1..=LETTERS.len()).contains(&c.len()) => c, + _ => { + let message = format!("{place}: 'criteria' must be an object with 1 to 26 options"); + return Err(error(422, message)); + } + }; + let mut labels = Vec::with_capacity(criteria.len()); + for (key, value) in criteria { + let label = match value { + Value::Null => key.clone(), + value => text(value, &format!("{place}: {}", quote(key)))?, + }; + // escaped as `json.dumps(label, ensure_ascii=False)[1:-1]`, like upstream's chooser + let quoted = quote(&label); + labels.push(quoted[1..quoted.len() - 1].to_string()); + } + let keys = criteria.keys().cloned().collect(); + mapped.push(Question { + name: name.clone(), + goal, + keys, + labels, + }); + } + Ok((state, mapped)) +} + +/// The prompt the Qwen3.5 chat template renders for the system and user messages with +/// `add_generation_prompt=True`. The template trims message content, which changes +/// nothing here: the user message starts with "Goal: " or "App: " and ends with "letter.". +pub fn chat_text(state: &str, question: &Question) -> String { + let goal = match question.goal.as_str() { + "" => String::new(), + goal => format!("Goal: {goal}\n\n"), + }; + let options: Vec = LETTERS + .chars() + .zip(&question.labels) + .map(|(letter, label)| format!("{letter}. Decision \"{label}\" -> select")) + .collect(); + format!( + "<|im_start|>system\n{SYSTEM_PROMPT}<|im_end|>\n<|im_start|>user\n{goal}App: Cua Driver\n\ + Task family: closed-candidate decision\n\nAccessibility tree:\n{state}\n\nOptions:\n{}\n\n\ + Answer with a single letter.<|im_end|>\n<|im_start|>assistant\n\n", + options.join("\n") + ) +} + +/// The Jev choice answer; ties go to the earliest option. `confidence` is the +/// normalized entropy `1 - H(p) / ln(n)`, as the LAYA worker reports it. +pub fn answer(question: &Question, probabilities: &[f32]) -> Value { + let p: Vec = probabilities.iter().map(|&x| x as f64).collect(); + let best = (0..p.len()).fold(0, |b, i| if p[i] > p[b] { i } else { b }); + let entropy: f64 = -p + .iter() + .filter(|&&x| x > 0.0) + .map(|&x| x * x.ln()) + .sum::(); + let n = p.len() as f64; + let probs: Map = question + .keys + .iter() + .cloned() + .zip(p.iter().map(|&x| json!(x))) + .collect(); + json!({ + "type": "choice", + "choice": question.keys[best], + "probabilities": probs, + "confidence": if p.len() > 1 { (1.0 - entropy / n.ln()).max(0.0) } else { 1.0 }, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn map(body: &str) -> Result<(String, Vec), RequestError> { + map_request(&parse_body(body.as_bytes())?) + } + + #[test] + fn maps_a_request() { + let body = r#"{"model": "cua-s1-4b-0.2", "state": {"a": [1.5, "é"]}, "questions": {"q": {"type": "choice", "instructions": "go", "criteria": {"a": "Say \"hi\"\n", "b": null}}}}"#; + let (state, questions) = map(body).ok().unwrap(); + assert_eq!(state, r#"{"a": [1.5, "é"]}"#); + assert_eq!(questions[0].keys, ["a", "b"]); + assert_eq!(questions[0].labels, [r#"Say \"hi\"\n"#, "b"]); + let text = chat_text(&state, &questions[0]); + assert!(text.contains("<|im_start|>user\nGoal: go\n\nApp: Cua Driver\n")); + assert!(text.ends_with( + "Options:\nA. Decision \"Say \\\"hi\\\"\\n\" -> select\nB. Decision \"b\" -> select\n\n\ + Answer with a single letter.<|im_end|>\n<|im_start|>assistant\n\n" + )); + } + + #[test] + fn errors() { + let q = |question: &str| { + format!( + r#"{{"model": "cua-s1-4b-0.2", "state": "S", "questions": {{"q": {question}}}}}"# + ) + }; + let cases = [ + (r#"{"state": "S"}"#.to_string(), 422), + ( + r#"{"model": "cua-s1-4b-0.2", "state": {}}"#.to_string(), + 422, + ), + ( + r#"{"model": "cua-s1-4b-0.2", "state": 1.5}"#.to_string(), + 422, + ), + (q(r#"{"type": "noul"}"#), 422), + (q(r#"{"type": "choice"}"#), 422), + ( + q(r#"{"type": "choice", "instructions": null, "criteria": {"a": true}}"#), + 422, + ), + ( + q(r#"{"type": "choice", "instructions": null, "criteria": {}}"#), + 422, + ), + (r#"{"a": 1, "a": 2}"#.to_string(), 400), + ]; + for (body, status) in cases { + assert_eq!(map(&body).err().map(|e| e.status), Some(status), "{body}"); + } + } + + #[test] + fn answer_and_confidence() { + let (_, questions) = map(r#"{"model": "cua-s1-4b-0.2", "state": "S", "questions": {"q": {"type": "choice", "instructions": null, "criteria": {"a": "A", "b": "B"}}}}"#).ok().unwrap(); + let tie = answer(&questions[0], &[0.5, 0.5]); + assert_eq!(tie["choice"], "a"); + assert!(tie["confidence"].as_f64().unwrap().abs() < 1e-12); + assert_eq!(answer(&questions[0], &[0.25, 0.75])["choice"], "b"); + } +} diff --git a/src/models/cua_s1/native/src/cuda.rs b/src/models/cua_s1/native/src/cuda.rs new file mode 100644 index 00000000..a15fef2f --- /dev/null +++ b/src/models/cua_s1/native/src/cuda.rs @@ -0,0 +1,239 @@ +//! The CUDA side, loaded at run time from `libqwen3_5_cuda.so` (built by +//! `src/backends/cuda/qwen3_5/build.sh`): the Qwen3.5 operations and the few CUDA +//! runtime calls the model needs. Building this crate needs no CUDA toolkit. + +use std::ffi::{CStr, c_char, c_int, c_void}; +use std::path::{Path, PathBuf}; +use std::sync::OnceLock; + +use anyhow::{Context, Result, bail, ensure}; + +/// `CS1_ABI_VERSION` in ops.h. +const ABI_VERSION: u32 = 2; +pub const LIBRARY: &str = "libqwen3_5_cuda.so"; + +/// A `cudaStream_t`. +#[repr(transparent)] +#[derive(Clone, Copy)] +pub struct Stream(*mut c_void); + +// SAFETY: a stream handle may be used from any thread; the model queues work on it +// from one thread at a time. +unsafe impl Send for Stream {} + +macro_rules! api { + ($($name:ident($($arg:ident: $ty:ty),* $(,)?) $(-> $ret:ty)?;)*) => { + /// The functions of the library, as declared in ops.h. + pub struct Api { + _lib: libloading::Library, + $(pub $name: unsafe extern "C" fn($($ty),*) $(-> $ret)?,)* + } + + impl Api { + fn resolve(lib: libloading::Library) -> Result { + $( + // SAFETY: the signature is the one ops.h declares for this symbol. + let $name = unsafe { + lib.get:: $ret)?>( + concat!(stringify!($name), "\0").as_bytes(), + ) + .map(|f| *f) + } + .with_context(|| format!("{} has no {}", LIBRARY, stringify!($name)))?; + )* + Ok(Self { _lib: lib, $($name,)* }) + } + } + }; +} + +api! { + cs1_abi_version() -> u32; + cs1_error_string(code: c_int) -> *const c_char; + cs1_set_device(device: c_int) -> c_int; + cs1_malloc(ptr: *mut *mut c_void, bytes: usize) -> c_int; + cs1_free(ptr: *mut c_void) -> c_int; + cs1_stream_create(stream: *mut Stream) -> c_int; + cs1_stream_sync(stream: Stream) -> c_int; + cs1_upload(dst: *mut c_void, src: *const c_void, bytes: usize, stream: Stream) -> c_int; + cs1_download(dst: *mut c_void, src: *const c_void, bytes: usize, stream: Stream) -> c_int; + cs1_embed(ids: *const i32, table: *const c_void, out: *mut c_void, t: c_int, d: c_int, stream: Stream) -> c_int; + cs1_rms_norm( + x: *const c_void, w: *const c_void, out: *mut c_void, rows: c_int, d: c_int, eps: f32, stream: Stream, + ) -> c_int; + cs1_add_rms_norm( + residual: *mut c_void, delta: *const c_void, w: *const c_void, out: *mut c_void, rows: c_int, d: c_int, + eps: f32, stream: Stream, + ) -> c_int; + cs1_gated_rms_norm( + x: *const c_void, z: *const c_void, ldz: c_int, w: *const c_void, out: *mut c_void, t: c_int, h: c_int, + d: c_int, eps: f32, stream: Stream, + ) -> c_int; + cs1_gdn_conv( + qkv: *const c_void, ld: c_int, w: *const c_void, q: *mut c_void, k: *mut c_void, v: *mut c_void, t: c_int, + key_dim: c_int, value_dim: c_int, stream: Stream, + ) -> c_int; + cs1_gdn_gates( + b: *const c_void, a: *const c_void, ld: c_int, a_log: *const c_void, dt_bias: *const c_void, + beta: *mut c_void, g: *mut f32, t: c_int, h: c_int, stream: Stream, + ) -> c_int; + cs1_gdn_workspace_floats(t: c_int, h: c_int) -> usize; + cs1_gdn_prefill( + q: *const c_void, k: *const c_void, v: *const c_void, g: *const f32, beta: *const c_void, o: *mut c_void, + workspace: *mut f32, t: c_int, h: c_int, hk: c_int, scale: f32, stream: Stream, + ) -> c_int; + cs1_attn_prep( + qg: *const c_void, kr: *const c_void, ld: c_int, qw: *const c_void, kw: *const c_void, cos: *const c_void, + sin: *const c_void, q: *mut c_void, gate: *mut c_void, k: *mut c_void, t: c_int, hq: c_int, hk: c_int, + dh: c_int, half: c_int, eps: f32, stream: Stream, + ) -> c_int; + cs1_attention( + q: *const c_void, k: *const c_void, v: *const c_void, ldv: c_int, out: *mut c_void, t: c_int, hq: c_int, + hk: c_int, dh: c_int, scale: f32, stream: Stream, + ) -> c_int; + cs1_sigmoid_gate(x: *mut c_void, gate: *const c_void, n: usize, stream: Stream) -> c_int; + cs1_silu_mul(gate_up: *const c_void, ld: c_int, out: *mut c_void, t: c_int, i: c_int, stream: Stream) -> c_int; + cs1_gemm_create(workspace_bytes: usize) -> *mut c_void; + cs1_gemm_destroy(gemm: *mut c_void); + cs1_gemm( + gemm: *mut c_void, x: *const c_void, w: *const c_void, y: *mut c_void, m: c_int, n: c_int, k: c_int, + ldy: c_int, stream: Stream, + ) -> c_int; +} + +static API: OnceLock = OnceLock::new(); + +/// The library next to the running executable. +pub fn default_library() -> Result { + let exe = std::env::current_exe()?; + Ok(exe + .parent() + .context("the executable has no directory")? + .join(LIBRARY)) +} + +/// Load the library (once per process) and check its ABI version. +pub fn load(path: &Path) -> Result<&'static Api> { + if let Some(api) = API.get() { + return Ok(api); + } + // SAFETY: loading runs the library's initializers; it is the library these + // sources build. + let lib = unsafe { libloading::Library::new(path) }.with_context(|| { + format!( + "loading {} (build it with src/backends/cuda/qwen3_5/build.sh)", + path.display() + ) + })?; + let api = Api::resolve(lib)?; + // SAFETY: takes no arguments. + let abi = unsafe { (api.cs1_abi_version)() }; + ensure!( + abi == ABI_VERSION, + "{} has ABI version {abi}, this build needs {ABI_VERSION}; rebuild it", + path.display() + ); + Ok(API.get_or_init(|| api)) +} + +/// The loaded library; `load` must have succeeded before. +pub fn api() -> &'static Api { + API.get().expect("the CUDA library is not loaded") +} + +/// Turn a return code of the library into an error; codes from 1000 up are cuBLAS +/// statuses (see gemm.cu). +pub fn check(code: c_int, what: &str) -> Result<()> { + if code == 0 { + return Ok(()); + } + if code >= 1000 { + bail!("{what}: cuBLAS status {}", code - 1000); + } + // SAFETY: cudaGetErrorString returns a static string for any code. + let msg = unsafe { CStr::from_ptr((api().cs1_error_string)(code)) }; + bail!("{what}: {} ({code})", msg.to_string_lossy()); +} + +pub fn set_device(device: i32) -> Result<()> { + // SAFETY: plain runtime call. + check(unsafe { (api().cs1_set_device)(device) }, "cudaSetDevice") +} + +/// A device allocation, freed on drop. +pub struct DeviceBuffer { + ptr: *mut c_void, + bytes: usize, +} + +// SAFETY: the pointer is a device address; access is serialized by the owner. +unsafe impl Send for DeviceBuffer {} +unsafe impl Sync for DeviceBuffer {} + +impl DeviceBuffer { + pub fn new(bytes: usize) -> Result { + let mut ptr = std::ptr::null_mut(); + if bytes > 0 { + // SAFETY: `ptr` is a valid out-pointer. + check(unsafe { (api().cs1_malloc)(&mut ptr, bytes) }, "cudaMalloc")?; + } + Ok(Self { ptr, bytes }) + } + + /// The device address `offset` bytes into the buffer. + pub fn at(&self, offset: usize) -> *mut c_void { + debug_assert!(offset <= self.bytes); + self.ptr.wrapping_byte_add(offset) + } +} + +impl Drop for DeviceBuffer { + fn drop(&mut self) { + if !self.ptr.is_null() { + // SAFETY: allocated by cs1_malloc and not freed before. + unsafe { (api().cs1_free)(self.ptr) }; + } + } +} + +pub fn new_stream() -> Result { + let mut stream = Stream(std::ptr::null_mut()); + // SAFETY: `stream` is a valid out-pointer. + check( + unsafe { (api().cs1_stream_create)(&mut stream) }, + "cudaStreamCreateWithFlags", + )?; + Ok(stream) +} + +pub fn synchronize(stream: Stream) -> Result<()> { + // SAFETY: a Stream only comes from new_stream. + check( + unsafe { (api().cs1_stream_sync)(stream) }, + "cudaStreamSynchronize", + ) +} + +/// Copy host bytes to `dst` and wait for the copy. +/// +/// # Safety +/// `dst` must be a device allocation with room for `src.len()` bytes. +pub unsafe fn upload(dst: *mut c_void, src: &[u8], stream: Stream) -> Result<()> { + // SAFETY: see above; the library waits for the copy before returning. + check( + unsafe { (api().cs1_upload)(dst, src.as_ptr().cast(), src.len(), stream) }, + "copy to device", + ) +} + +/// Copy `dst.len()` bytes from `src` to the host, after the work queued before it. +/// +/// # Safety +/// `src` must be a device allocation holding at least `dst.len()` bytes. +pub unsafe fn download(dst: &mut [u8], src: *const c_void, stream: Stream) -> Result<()> { + // SAFETY: see above; the library waits for the copy before returning. + check( + unsafe { (api().cs1_download)(dst.as_mut_ptr().cast(), src, dst.len(), stream) }, + "copy to host", + ) +} diff --git a/src/models/cua_s1/native/src/engine.rs b/src/models/cua_s1/native/src/engine.rs new file mode 100644 index 00000000..859156da --- /dev/null +++ b/src/models/cua_s1/native/src/engine.rs @@ -0,0 +1,119 @@ +//! Tokenization and scoring: one prefill-only forward pass per question through the +//! native Qwen3.5 model, scored with the 26 letter rows of the output projection. + +use std::path::Path; +use std::sync::{Arc, Mutex}; + +use anyhow::{Context, Result, ensure}; +use serde_json::Value as Json; +use tokenizers::Tokenizer; + +use crate::contract::{LETTERS, Question, chat_text}; +use crate::model::Model; + +/// Tied with the output projection in Qwen3.5-4B. +const EMBEDDING: &str = "model.language_model.embed_tokens.weight"; + +pub struct Engine { + tokenizer: Tokenizer, + model: Arc>, + /// the letter rows of the output projection, as float32 + letters: Vec, +} + +impl Engine { + pub async fn load(dir: &Path, library: &Path) -> Result { + // Written by export_text_merged.py; without it `dir` may hold the base model alone. + ensure!( + dir.join("cua_s1_export.json").exists(), + "{} is not a merged text checkpoint; see recipe/cua_s1/native.md", + dir.display() + ); + let tokenizer = + Tokenizer::from_file(dir.join("tokenizer.json")).map_err(anyhow::Error::msg)?; + let ids = LETTERS + .chars() + .map(|c| { + tokenizer + .token_to_id(&c.to_string()) + .context("letter token") + }) + .collect::>>()?; + let (d, lib) = (dir.to_path_buf(), library.to_path_buf()); + let model = tokio::task::spawn_blocking(move || Model::load(&d, &lib)).await??; + let letters = letter_rows(dir, &ids, model.cfg.hidden)?; + Ok(Self { + tokenizer, + model: Arc::new(Mutex::new(model)), + letters, + }) + } + + pub fn encode(&self, state: &str, question: &Question) -> Result> { + let enc = self + .tokenizer + .encode(chat_text(state, question), false) + .map_err(anyhow::Error::msg)?; + Ok(enc.get_ids().to_vec()) + } + + /// Option probabilities for one prompt: the final-norm hidden state at the last + /// position times the letter rows, in float32 with float64 accumulation, then a + /// softmax over the first `n_options` letters. + pub async fn score(&self, ids: Vec, n_options: usize) -> Result> { + let model = self.model.clone(); + let last = tokio::task::spawn_blocking(move || { + model + .lock() + .map_err(|_| anyhow::anyhow!("poisoned"))? + .forward(&ids) + }) + .await??; + let logits: Vec = self + .letters + .chunks_exact(last.len()) + .take(n_options) + .map(|w| { + w.iter() + .zip(&last) + .map(|(&a, &b)| a as f64 * b as f64) + .sum::() as f32 as f64 + }) + .collect(); + let max = logits.iter().copied().fold(f64::NEG_INFINITY, f64::max); + let exps: Vec = logits.iter().map(|&l| (l - max).exp()).collect(); + let total: f64 = exps.iter().sum(); + let probs: Vec = exps.iter().map(|e| (e / total) as f32).collect(); + ensure!( + probs.iter().all(|p| p.is_finite()), + "non-finite probabilities" + ); + Ok(probs) + } +} + +/// The letter rows of the bfloat16 embedding, read from the safetensors files. +fn letter_rows(dir: &Path, ids: &[u32], hidden: usize) -> Result> { + let index: Json = serde_json::from_str(&std::fs::read_to_string( + dir.join("model.safetensors.index.json"), + )?)?; + let file = index["weight_map"][EMBEDDING].as_str().context(EMBEDDING)?; + let file = std::fs::File::open(dir.join(file))?; + // SAFETY: the checkpoint is not modified while the worker runs. + let mmap = unsafe { memmap2::Mmap::map(&file)? }; + let tensors = safetensors::SafeTensors::deserialize(&mmap)?; + let view = tensors.tensor(EMBEDDING)?; + ensure!( + view.dtype() == safetensors::Dtype::BF16 && view.shape()[1] == hidden, + "{EMBEDDING}: {:?} {:?}", + view.dtype(), + view.shape() + ); + let mut rows = Vec::with_capacity(ids.len() * hidden); + for &id in ids { + let row = &view.data()[id as usize * hidden * 2..(id as usize + 1) * hidden * 2]; + let (pairs, _) = row.as_chunks::<2>(); + rows.extend(pairs.iter().map(|&b| half::bf16::from_le_bytes(b).to_f32())); + } + Ok(rows) +} diff --git a/src/models/cua_s1/native/src/json.rs b/src/models/cua_s1/native/src/json.rs new file mode 100644 index 00000000..a0e2ed5a --- /dev/null +++ b/src/models/cua_s1/native/src/json.rs @@ -0,0 +1,245 @@ +//! Request JSON, on serde_json. +//! +//! - [`parse`] rejects what the contract rejects with a 400: invalid JSON or UTF-8, +//! `NaN`/`Infinity`, numbers out of range, lone surrogates and nesting deeper than +//! serde_json's limit (all serde_json errors), and keys repeated in an object. +//! - [`dumps`] writes `json.dumps(value, ensure_ascii=False)`: separators `, ` and +//! `: `, key order kept, floats as Python's `repr`. + +use std::fmt::Write as _; +use std::io; + +use serde::Serialize; +use serde::de::{self, Deserializer, MapAccess, SeqAccess, Visitor}; +use serde_json::{Map, Number, Value}; + +/// Decode a request body into its top-level object; the error is the 400 message. +pub fn parse(raw: &[u8]) -> Result, String> { + let mut de = serde_json::Deserializer::from_slice(raw); + let value = de + .deserialize_any(NoDuplicates) + .and_then(|v| de.end().map(|()| v)) + .map_err(|e| format!("request body is not valid JSON: {e}"))?; + match value { + Value::Object(map) => Ok(map), + _ => Err("request body must be a JSON object".into()), + } +} + +/// Builds a `Value` like serde_json does, but fails on a repeated key. +struct NoDuplicates; + +impl<'de> de::Deserialize<'de> for Wrapped { + fn deserialize>(d: D) -> Result { + d.deserialize_any(NoDuplicates).map(Wrapped) + } +} + +struct Wrapped(Value); + +impl<'de> Visitor<'de> for NoDuplicates { + type Value = Value; + + fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + f.write_str("a JSON value") + } + fn visit_unit(self) -> Result { + Ok(Value::Null) + } + fn visit_bool(self, b: bool) -> Result { + Ok(Value::Bool(b)) + } + fn visit_i64(self, n: i64) -> Result { + Ok(Value::Number(n.into())) + } + fn visit_u64(self, n: u64) -> Result { + Ok(Value::Number(n.into())) + } + fn visit_f64(self, x: f64) -> Result { + Number::from_f64(x) + .map(Value::Number) + .ok_or_else(|| E::custom("number out of range")) + } + fn visit_str(self, s: &str) -> Result { + Ok(Value::String(s.to_owned())) + } + fn visit_string(self, s: String) -> Result { + Ok(Value::String(s)) + } + fn visit_seq>(self, mut seq: A) -> Result { + let mut items = Vec::new(); + while let Some(Wrapped(v)) = seq.next_element()? { + items.push(v); + } + Ok(Value::Array(items)) + } + fn visit_map>(self, mut map: A) -> Result { + let mut obj = Map::new(); + while let Some(key) = map.next_key::()? { + let Wrapped(v) = map.next_value()?; + if obj.contains_key(&key) { + return Err(de::Error::custom(format_args!( + "duplicate key {}", + quote(&key) + ))); + } + obj.insert(key, v); + } + Ok(Value::Object(obj)) + } +} + +/// A string as a JSON literal, which is also how error messages quote names. +pub fn quote(s: &str) -> String { + serde_json::to_string(s).expect("a string serializes") +} + +/// `json.dumps(value, ensure_ascii=False)`. +pub fn dumps(value: &Value) -> String { + let mut out = Vec::new(); + let mut ser = serde_json::Serializer::with_formatter(&mut out, PyFormatter); + value.serialize(&mut ser).expect("a Value serializes"); + String::from_utf8(out).expect("serde_json writes UTF-8") +} + +/// serde_json's compact output with Python's separators and float format; its string +/// escaping (`"`, `\\` and control characters, `\u00XX` in lowercase) is Python's. +struct PyFormatter; + +impl serde_json::ser::Formatter for PyFormatter { + fn begin_array_value( + &mut self, + w: &mut W, + first: bool, + ) -> io::Result<()> { + if first { Ok(()) } else { w.write_all(b", ") } + } + fn begin_object_key( + &mut self, + w: &mut W, + first: bool, + ) -> io::Result<()> { + if first { Ok(()) } else { w.write_all(b", ") } + } + fn begin_object_value(&mut self, w: &mut W) -> io::Result<()> { + w.write_all(b": ") + } + fn write_f64(&mut self, w: &mut W, x: f64) -> io::Result<()> { + w.write_all(float_repr(x).as_bytes()) + } +} + +/// Python's `repr(float)`: the shortest digits that round-trip, in fixed notation +/// for exponents from -5 to 15 and scientific notation otherwise. (Python breaks the +/// rare exact ties between two shortest candidates to even; this does not.) +pub fn float_repr(x: f64) -> String { + if x == 0.0 { + return if x.is_sign_negative() { "-0.0" } else { "0.0" }.into(); + } + let sci = format!("{x:e}"); + let (mantissa, exp) = sci.split_once('e').expect("{:e} has an exponent"); + let exp: i32 = exp.parse().expect("integer exponent"); + let (neg, mantissa) = match mantissa.strip_prefix('-') { + Some(m) => (true, m), + None => (false, mantissa), + }; + let digits: String = mantissa.chars().filter(|c| *c != '.').collect(); + let mut out = String::new(); + if neg { + out.push('-'); + } + let decpt = exp + 1; + if decpt <= -4 || decpt > 16 { + out.push_str(&digits[..1]); + if digits.len() > 1 { + out.push('.'); + out.push_str(&digits[1..]); + } + let sign = if exp < 0 { '-' } else { '+' }; + write!(out, "e{sign}{:02}", exp.unsigned_abs()).unwrap(); + } else if decpt <= 0 { + out.push_str("0."); + out.extend(std::iter::repeat_n('0', (-decpt) as usize)); + out.push_str(&digits); + } else if decpt as usize >= digits.len() { + out.push_str(&digits); + out.extend(std::iter::repeat_n('0', decpt as usize - digits.len())); + out.push_str(".0"); + } else { + out.push_str(&digits[..decpt as usize]); + out.push('.'); + out.push_str(&digits[decpt as usize..]); + } + out +} + +#[cfg(test)] +mod tests { + use super::*; + + fn err(body: &str) -> String { + parse(body.as_bytes()).unwrap_err() + } + + #[test] + fn float_repr_matches_python_examples() { + let cases = [ + (1.0, "1.0"), + (1e16, "1e+16"), + (1e15, "1000000000000000.0"), + (1e-5, "1e-05"), + (1e-4, "0.0001"), + (-0.0, "-0.0"), + (3.14e-07, "3.14e-07"), + (5e-324, "5e-324"), + (1.7976931348623157e308, "1.7976931348623157e+308"), + (5.960464477539063e-08, "5.960464477539063e-08"), + ]; + for (x, want) in cases { + assert_eq!(float_repr(x), want, "{x:e}"); + } + } + + #[test] + fn dumps_matches_python() { + let v = Value::Object(parse(br#"{"a": [1.0, 1e16, 1e-5, 0.0001, -0.0, 123456789012345678, 3.14e-07, true, null], "b": {}}"#).unwrap()); + assert_eq!( + dumps(&v), + r#"{"a": [1.0, 1e+16, 1e-05, 0.0001, -0.0, 123456789012345678, 3.14e-07, true, null], "b": {}}"# + ); + let s = Value::String("\u{0}\u{1f}\u{7f}\u{2028}\"\\/\t\u{8}\u{c}é😀".into()); + assert_eq!( + dumps(&s), + "\"\\u0000\\u001f\u{7f}\u{2028}\\\"\\\\/\\t\\b\\fé😀\"" + ); + } + + #[test] + fn rejects_what_the_contract_rejects() { + assert_eq!(err("[]"), "request body must be a JSON object"); + for body in [ + r#"{"a": NaN}"#, + r#"{"a": 1e400}"#, + r#"{"a": "\ud800x"}"#, + r#"{"a": 1, "b": 2, "a": 3}"#, + r#"{"a": [1,]}"#, + r#"{} x"#, + "\u{feff}{}", + ] { + assert!( + err(body).starts_with("request body is not valid JSON"), + "{body}" + ); + } + assert!(err(r#"{"a": 1, "a": 2}"#).contains("duplicate key \"a\"")); + assert!( + err(&format!( + "{{\"a\": {}1{}}}", + "[".repeat(200), + "]".repeat(200) + )) + .contains("recursion limit") + ); + assert!(parse(b"{\"a\": \"\xff\"}").is_err()); + } +} diff --git a/src/models/cua_s1/native/src/lib.rs b/src/models/cua_s1/native/src/lib.rs new file mode 100644 index 00000000..0d3aa5fd --- /dev/null +++ b/src/models/cua_s1/native/src/lib.rs @@ -0,0 +1,9 @@ +//! 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. + +pub mod contract; +pub mod cuda; +pub mod engine; +pub mod json; +pub mod model; diff --git a/src/models/cua_s1/native/src/main.rs b/src/models/cua_s1/native/src/main.rs new file mode 100644 index 00000000..121a9e3d --- /dev/null +++ b/src/models/cua_s1/native/src/main.rs @@ -0,0 +1,113 @@ +//! Cua-S1 4B 0.2 (`text` adapter) `/v1/systemone` worker on the native CUDA kernels. +//! +//! CUA_S1_MODEL= omni-cua-s1-native +//! +//! `CUA_S1_CUDA_LIB` (default: next to this executable), `CUA_S1_HOST` and `CUA_S1_PORT` +//! are optional; see recipe/cua_s1/native.md. + +use std::sync::Arc; + +use anyhow::{Context, Result, ensure}; +use axum::body::Bytes; +use axum::extract::rejection::BytesRejection; +use axum::extract::{DefaultBodyLimit, State}; +use axum::http::StatusCode; +use axum::response::{IntoResponse, Response}; +use axum::routing::{get, post}; +use axum::{Json, Router}; +use serde_json::{Value, json}; + +use omni_cua_s1_native::contract::{self, MODEL_ID}; +use omni_cua_s1_native::cuda; +use omni_cua_s1_native::engine::Engine; +use omni_cua_s1_native::json::quote; + +const MAX_BODY_BYTES: usize = 4 << 20; +const MAX_PROMPT_TOKENS: usize = 16384; +const WARMUP: &[u8] = br#"{"model": "cua-s1-4b-0.2", "state": "Dialog: Update installed.", "questions": {"q": {"type": "choice", "instructions": "Close it.", "criteria": {"ok": "OK", "wait": "Wait"}}}}"#; + +fn reply(status: u16, body: Value) -> Response { + (StatusCode::from_u16(status).unwrap(), Json(body)).into_response() +} + +/// Every prompt is tokenized and checked against the limit before any forward pass. +async fn decide(engine: &Engine, raw: &[u8]) -> Response { + let (state, questions) = match contract::parse_body(raw).and_then(|b| contract::map_request(&b)) + { + Ok(request) => request, + Err(e) => return reply(e.status, json!({"detail": e.message})), + }; + let failed = |e: anyhow::Error| { + eprintln!("inference failed: {e:#}"); + reply(500, json!({"detail": "inference failed"})) + }; + let mut prompts = Vec::with_capacity(questions.len()); + for q in &questions { + let ids = match engine.encode(&state, q) { + Ok(ids) => ids, + Err(e) => return failed(e), + }; + if ids.len() > MAX_PROMPT_TOKENS { + let message = format!( + "question {}: {} prompt tokens, over {MAX_PROMPT_TOKENS}", + quote(&q.name), + ids.len() + ); + return reply(413, json!({"detail": message})); + } + prompts.push(ids); + } + let tokens: usize = prompts.iter().map(Vec::len).sum(); + let mut answers = serde_json::Map::new(); + for (q, ids) in questions.iter().zip(prompts) { + match engine.score(ids, q.keys.len()).await { + Ok(probs) => answers.insert(q.name.clone(), contract::answer(q, &probs)), + Err(e) => return failed(e), + }; + } + reply( + 200, + json!({"model": MODEL_ID, "answers": answers, "usage": {"input_tokens": tokens, "output_tokens": 0}}), + ) +} + +async fn systemone( + State(engine): State>, + body: Result, +) -> Response { + match body { + Ok(raw) => decide(&engine, &raw).await, + Err(e) => reply(e.status().as_u16(), json!({"detail": e.body_text()})), + } +} + +#[tokio::main] +async fn main() -> Result<()> { + let model = std::env::var_os("CUA_S1_MODEL").context("set CUA_S1_MODEL")?; + let library = match std::env::var_os("CUA_S1_CUDA_LIB") { + Some(path) => path.into(), + None => cuda::default_library()?, + }; + let engine = Arc::new(Engine::load(model.as_ref(), &library).await?); + // one decision before listening, so the first request does not pay for first-call setup + ensure!( + decide(&engine, WARMUP).await.status() == StatusCode::OK, + "warmup failed" + ); + 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()) + .context("CUA_S1_PORT")?; + let app = Router::new() + .route( + "/health", + get(|| async { Json(json!({"status": "ready", "model": MODEL_ID})) }), + ) + .route("/v1/systemone", post(systemone)) + .layer(DefaultBodyLimit::max(MAX_BODY_BYTES)) + .with_state(engine); + let listener = tokio::net::TcpListener::bind((host.as_str(), port)).await?; + println!("listening on {host}:{port}"); + axum::serve(listener, app).await?; + Ok(()) +} diff --git a/src/models/cua_s1/native/src/model.rs b/src/models/cua_s1/native/src/model.rs new file mode 100644 index 00000000..2b040779 --- /dev/null +++ b/src/models/cua_s1/native/src/model.rs @@ -0,0 +1,875 @@ +//! The Qwen3.5 text model (the language model of Qwen/Qwen3.5-4B), prefill only: one +//! forward pass over a prompt, returning the final-norm hidden state of the last +//! position. The layer loop and buffers live here; the operations are the CUDA +//! kernels in `src/backends/cuda/qwen3_5`. +//! +//! 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. + +use std::collections::HashMap; +use std::ffi::c_void; +use std::path::Path; + +use anyhow::{Context, Result, bail, ensure}; +use serde_json::Value as Json; + +use crate::cuda::{self, DeviceBuffer, Stream, check}; + +const ALIGN: usize = 256; +const BF16: usize = 2; +const F32: usize = 4; +const GEMM_WORKSPACE: usize = 32 << 20; + +#[derive(Debug, Clone)] +pub struct Config { + pub hidden: usize, + pub intermediate: usize, + pub eps: f32, + pub full_attention: Vec, + pub heads: usize, + pub kv_heads: usize, + pub head_dim: usize, + /// Half the number of rotary dims (rotate_half pairs dim i with dim i + half). + pub rotary_half: usize, + pub rope_theta: f64, + pub lin_k_heads: usize, + pub lin_v_heads: usize, + pub lin_k_dim: usize, + pub lin_v_dim: usize, +} + +impl Config { + pub fn load(dir: &Path) -> Result { + let path = dir.join("config.json"); + let root: Json = serde_json::from_str( + &std::fs::read_to_string(&path).with_context(|| format!("{}", path.display()))?, + )?; + let c = root.get("text_config").unwrap_or(&root); + let int = |k: &str| { + c[k].as_u64() + .map(|v| v as usize) + .with_context(|| format!("config.json: `{k}` is missing")) + }; + let rope = &c["rope_parameters"]; + let partial = rope["partial_rotary_factor"] + .as_f64() + .or(c["partial_rotary_factor"].as_f64()) + .unwrap_or(1.0); + let head_dim = int("head_dim")?; + let full_attention = c["layer_types"] + .as_array() + .context("config.json: `layer_types` is missing")? + .iter() + .map(|t| match t.as_str() { + Some("full_attention") => Ok(true), + Some("linear_attention") => Ok(false), + other => bail!("unknown layer type {other:?}"), + }) + .collect::>>()?; + let cfg = Config { + hidden: int("hidden_size")?, + intermediate: int("intermediate_size")?, + eps: c["rms_norm_eps"].as_f64().context("rms_norm_eps")? as f32, + heads: int("num_attention_heads")?, + kv_heads: int("num_key_value_heads")?, + head_dim, + rotary_half: (head_dim as f64 * partial) as usize / 2, + rope_theta: rope["rope_theta"] + .as_f64() + .or(c["rope_theta"].as_f64()) + .context("rope_theta")?, + lin_k_heads: int("linear_num_key_heads")?, + lin_v_heads: int("linear_num_value_heads")?, + lin_k_dim: int("linear_key_head_dim")?, + lin_v_dim: int("linear_value_head_dim")?, + full_attention, + }; + // What the kernels implement. + ensure!( + cfg.full_attention.len() == int("num_hidden_layers")?, + "layer_types does not match num_hidden_layers" + ); + ensure!(c["hidden_act"] == "silu", "hidden_act is not silu"); + ensure!( + c["attn_output_gate"].as_bool().unwrap_or(true), + "attention without the output gate" + ); + ensure!( + c["attention_bias"].as_bool() != Some(true), + "attention with bias" + ); + ensure!( + rope["rope_type"].as_str().unwrap_or("default") == "default", + "rope type {}", + rope["rope_type"] + ); + ensure!(int("linear_conv_kernel_dim")? == 4, "conv kernel is not 4"); + ensure!( + cfg.head_dim == 256 && cfg.lin_k_dim == 128 && cfg.lin_v_dim == 128, + "head dims {} / {} / {}", + cfg.head_dim, + cfg.lin_k_dim, + cfg.lin_v_dim + ); + ensure!(cfg.rotary_half == 32, "{} rotary dims", 2 * cfg.rotary_half); + ensure!( + cfg.kv_heads > 0 && cfg.heads.is_multiple_of(cfg.kv_heads), + "attention heads" + ); + ensure!( + cfg.lin_k_heads > 0 && cfg.lin_v_heads.is_multiple_of(cfg.lin_k_heads), + "linear attention heads" + ); + ensure!(cfg.hidden.is_multiple_of(8), "hidden size"); + Ok(cfg) + } + + fn key_dim(&self) -> usize { + self.lin_k_heads * self.lin_k_dim + } + + fn value_dim(&self) -> usize { + self.lin_v_heads * self.lin_v_dim + } +} + +/// A weight in the device arena. +#[derive(Clone)] +struct Tensor { + ptr: *const c_void, + shape: Vec, +} + +impl Tensor { + fn bytes(&self) -> usize { + self.shape.iter().product::() * BF16 + } +} + +/// Projections that run as one GEMM, in the order their rows are stacked. +const GROUPS: &[&str] = &[ + "linear_attn.in_proj_qkv.weight", + "linear_attn.in_proj_z.weight", + "linear_attn.in_proj_b.weight", + "linear_attn.in_proj_a.weight", + "self_attn.q_proj.weight", + "self_attn.k_proj.weight", + "self_attn.v_proj.weight", + "mlp.gate_proj.weight", + "mlp.up_proj.weight", +]; + +/// Upload order: by layer, and inside a layer the GROUPS members first and in order, +/// so each group's matrices sit back to back and form one [sum N, K] matrix. +fn upload_order(name: &str) -> (usize, usize, String) { + if let Some(tail) = name.strip_prefix("layers.") + && let Some((layer, rest)) = tail.split_once('.') + && let Ok(layer) = layer.parse::() + { + let rank = GROUPS + .iter() + .position(|g| *g == rest) + .unwrap_or(GROUPS.len()); + return (layer, rank, rest.to_string()); + } + (usize::MAX, 0, name.to_string()) +} + +struct Weights { + _arena: DeviceBuffer, + tensors: HashMap, + prefix: String, +} + +impl Weights { + /// Upload every bfloat16 tensor of the language model into one allocation. + fn load(dir: &Path, stream: Stream) -> Result { + let index = dir.join("model.safetensors.index.json"); + let mut files: Vec = if index.exists() { + let index: Json = serde_json::from_str(&std::fs::read_to_string(&index)?)?; + index["weight_map"] + .as_object() + .context("weight_map")? + .values() + .filter_map(|v| v.as_str().map(str::to_string)) + .collect() + } else { + vec!["model.safetensors".to_string()] + }; + files.sort(); + files.dedup(); + let maps = files + .iter() + .map(|f| { + let file = std::fs::File::open(dir.join(f)).with_context(|| f.clone())?; + // SAFETY: the checkpoint is not modified while it is loaded. + Ok(unsafe { memmap2::Mmap::map(&file)? }) + }) + .collect::>>()?; + let sts = maps + .iter() + .map(|m| safetensors::SafeTensors::deserialize(m).map_err(anyhow::Error::from)) + .collect::>>()?; + let names: Vec<(usize, String)> = sts + .iter() + .enumerate() + .flat_map(|(i, st)| st.names().into_iter().map(move |n| (i, n.to_string()))) + .collect(); + let prefix = ["model.language_model.", "model."] + .into_iter() + .find(|p| { + names + .iter() + .any(|(_, n)| *n == format!("{p}embed_tokens.weight")) + }) + .context("no embed_tokens.weight in the checkpoint")? + .to_string(); + let mut ours: Vec<(usize, String)> = names + .into_iter() + .filter(|(_, n)| n.starts_with(&prefix)) + .collect(); + ours.sort_by_key(|(_, n)| upload_order(&n[prefix.len()..])); + let mut plan = Vec::new(); + let mut total = 0usize; + for (i, name) in ours { + let view = sts[i].tensor(&name)?; + ensure!( + view.dtype() == safetensors::Dtype::BF16, + "{name} is {:?}, not bfloat16", + view.dtype() + ); + plan.push((i, name, total)); + total = (total + view.data().len()).next_multiple_of(ALIGN); + } + let arena = DeviceBuffer::new(total)?; + let mut tensors = HashMap::new(); + for (i, name, offset) in plan { + let view = sts[i].tensor(&name)?; + // SAFETY: the arena has room for every planned tensor at its offset. + unsafe { cuda::upload(arena.at(offset), view.data(), stream)? }; + tensors.insert( + name[prefix.len()..].to_string(), + Tensor { + ptr: arena.at(offset), + shape: view.shape().to_vec(), + }, + ); + } + Ok(Self { + _arena: arena, + tensors, + prefix, + }) + } + + fn get(&self, name: &str, shape: &[usize]) -> Result { + let t = self + .tensors + .get(name) + .with_context(|| format!("{}{name} is missing", self.prefix))?; + ensure!( + t.shape == shape, + "{}{name}: shape {:?}, expected {:?}", + self.prefix, + t.shape, + shape + ); + Ok(t.clone()) + } + + /// The row-stacked matrix of tensors that were uploaded back to back. + fn stacked(&self, parts: &[Tensor]) -> Result { + let k = parts[0].shape[1]; + let mut rows = 0; + for (i, p) in parts.iter().enumerate() { + ensure!( + p.shape.len() == 2 && p.shape[1] == k, + "stacked shapes differ" + ); + if i > 0 { + let prev = &parts[i - 1]; + ensure!( + p.ptr == prev.ptr.wrapping_byte_add(prev.bytes()), + "stacked weights are not contiguous" + ); + } + rows += p.shape[0]; + } + Ok(Tensor { + ptr: parts[0].ptr, + shape: vec![rows, k], + }) + } +} + +struct LinearAttention { + /// in_proj_qkv | in_proj_z | in_proj_b | in_proj_a + in_proj: Tensor, + conv: Tensor, + a_log: Tensor, + dt_bias: Tensor, + norm: Tensor, + out: Tensor, +} + +struct FullAttention { + /// q_proj (query and gate per head) | k_proj | v_proj + qkv: Tensor, + o: Tensor, + q_norm: Tensor, + k_norm: Tensor, +} + +enum Mixer { + Linear(LinearAttention), + Full(FullAttention), +} + +struct Layer { + input_norm: Tensor, + post_norm: Tensor, + mixer: Mixer, + /// gate_proj | up_proj + gate_up: Tensor, + down: Tensor, +} + +/// Row widths of the stacked projection outputs. +struct Widths { + conv: usize, + gdn_in: usize, + attn_q: usize, + attn_in: usize, +} + +impl Widths { + fn of(cfg: &Config) -> Self { + let conv = 2 * cfg.key_dim() + cfg.value_dim(); + let attn_q = cfg.heads * cfg.head_dim * 2; + Self { + conv, + gdn_in: conv + cfg.value_dim() + 2 * cfg.lin_v_heads, + attn_q, + attn_in: attn_q + 2 * cfg.kv_heads * cfg.head_dim, + } + } +} + +/// Per-request buffers for up to `cap` tokens, as byte offsets into one allocation, +/// plus the rotary tables for positions below `cap`. +struct Scratch { + cap: usize, + buf: DeviceBuffer, + ids: usize, + res: usize, + x: usize, + delta: usize, + gdn_in: usize, + beta: usize, + g: usize, + lq: usize, + lk: usize, + lv: usize, + lo: usize, + ln: usize, + workspace: usize, + attn_in: usize, + aq: usize, + agate: usize, + ak: usize, + ao: usize, + gate_up: usize, + act: usize, + cos: usize, + sin: usize, +} + +impl Scratch { + fn new(cfg: &Config, cap: usize, stream: Stream) -> Result { + let (h, kd, vd, hv) = (cfg.hidden, cfg.key_dim(), cfg.value_dim(), cfg.lin_v_heads); + let (hq, hk, hd) = (cfg.heads, cfg.kv_heads, cfg.head_dim); + let w = Widths::of(cfg); + let mut next = 0usize; + let mut take = |bytes: usize| { + let off = next; + next = (off + bytes).next_multiple_of(ALIGN); + off + }; + // SAFETY: pure function of its arguments. + let ws_floats = unsafe { (cuda::api().cs1_gdn_workspace_floats)(cap as i32, hv as i32) }; + let offsets = [ + take(cap * 4), + take(cap * h * BF16), + take(cap * h * BF16), + take(cap * h * BF16), + take(cap * w.gdn_in * BF16), + take(cap * hv * BF16), + take(cap * hv * F32), + take(cap * kd * BF16), + take(cap * kd * BF16), + take(cap * vd * BF16), + take(cap * vd * BF16), + take(cap * vd * BF16), + take(ws_floats * F32), + take(cap * w.attn_in * BF16), + take(cap * hq * hd * BF16), + take(cap * hq * hd * BF16), + take(cap * hk * hd * BF16), + take(cap * hq * hd * BF16), + take(cap * 2 * cfg.intermediate * BF16), + take(cap * cfg.intermediate * BF16), + take(cap * cfg.rotary_half * BF16), + take(cap * cfg.rotary_half * BF16), + ]; + let buf = DeviceBuffer::new(next)?; + let [ + ids, + res, + x, + delta, + gdn_in, + beta, + g, + lq, + lk, + lv, + lo, + ln, + workspace, + attn_in, + aq, + agate, + ak, + ao, + gate_up, + act, + cos, + 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()); + } + } + // SAFETY: both tables were laid out for cap * rotary_half bfloat16 values. + unsafe { + cuda::upload(buf.at(cos), &cos_t, stream)?; + cuda::upload(buf.at(sin), &sin_t, stream)?; + } + Ok(Self { + cap, + buf, + ids, + res, + x, + delta, + gdn_in, + beta, + g, + lq, + lk, + lv, + lo, + ln, + workspace, + attn_in, + aq, + agate, + ak, + ao, + gate_up, + act, + cos, + sin, + }) + } + + fn at(&self, offset: usize) -> *mut c_void { + self.buf.at(offset) + } +} + +pub struct Model { + pub cfg: Config, + _weights: Weights, + embed: Tensor, + final_norm: Tensor, + layers: Vec, + stream: Stream, + gemm: *mut c_void, + /// Buffers for the longest prompt so far; grows as needed. + scratch: Option, +} + +// SAFETY: the raw pointers are device addresses and a cuBLASLt handle owned by the +// model; the engine runs one forward pass at a time behind a mutex. +unsafe impl Send for Model {} + +impl Drop for Model { + fn drop(&mut self) { + // SAFETY: created by cs1_gemm_create and not destroyed before. + unsafe { (cuda::api().cs1_gemm_destroy)(self.gemm) }; + } +} + +impl Model { + /// Load the CUDA library and the weights. + pub fn load(dir: &Path, library: &Path) -> Result { + let cfg = Config::load(dir)?; + cuda::load(library)?; + cuda::set_device(0)?; + let stream = cuda::new_stream()?; + let weights = Weights::load(dir, stream)?; + let (h, kd, vd) = (cfg.hidden, cfg.key_dim(), cfg.value_dim()); + let embed = weights + .tensors + .get("embed_tokens.weight") + .context("embed_tokens.weight is missing")? + .clone(); + ensure!( + embed.shape.len() == 2 && embed.shape[1] == h, + "embed_tokens.weight shape {:?}", + embed.shape + ); + let final_norm = weights.get("norm.weight", &[h])?; + let mut layers = Vec::with_capacity(cfg.full_attention.len()); + for (i, &full) in cfg.full_attention.iter().enumerate() { + let w = |n: &str, s: &[usize]| weights.get(&format!("layers.{i}.{n}"), s); + let mixer = if full { + let (hq, hk, hd) = (cfg.heads, cfg.kv_heads, cfg.head_dim); + Mixer::Full(FullAttention { + qkv: weights.stacked(&[ + w("self_attn.q_proj.weight", &[hq * hd * 2, h])?, + w("self_attn.k_proj.weight", &[hk * hd, h])?, + w("self_attn.v_proj.weight", &[hk * hd, h])?, + ])?, + o: w("self_attn.o_proj.weight", &[h, hq * hd])?, + q_norm: w("self_attn.q_norm.weight", &[hd])?, + k_norm: w("self_attn.k_norm.weight", &[hd])?, + }) + } else { + let hv = cfg.lin_v_heads; + Mixer::Linear(LinearAttention { + in_proj: weights.stacked(&[ + w("linear_attn.in_proj_qkv.weight", &[2 * kd + vd, h])?, + w("linear_attn.in_proj_z.weight", &[vd, h])?, + w("linear_attn.in_proj_b.weight", &[hv, h])?, + w("linear_attn.in_proj_a.weight", &[hv, h])?, + ])?, + conv: w("linear_attn.conv1d.weight", &[2 * kd + vd, 1, 4])?, + a_log: w("linear_attn.A_log", &[hv])?, + dt_bias: w("linear_attn.dt_bias", &[hv])?, + norm: w("linear_attn.norm.weight", &[cfg.lin_v_dim])?, + out: w("linear_attn.out_proj.weight", &[h, vd])?, + }) + }; + layers.push(Layer { + input_norm: w("input_layernorm.weight", &[h])?, + post_norm: w("post_attention_layernorm.weight", &[h])?, + mixer, + gate_up: weights.stacked(&[ + w("mlp.gate_proj.weight", &[cfg.intermediate, h])?, + w("mlp.up_proj.weight", &[cfg.intermediate, h])?, + ])?, + down: w("mlp.down_proj.weight", &[h, cfg.intermediate])?, + }); + } + // SAFETY: plain allocation; checked for null below. + let gemm = unsafe { (cuda::api().cs1_gemm_create)(GEMM_WORKSPACE) }; + ensure!(!gemm.is_null(), "cuBLASLt setup failed"); + let model = Self { + cfg, + _weights: weights, + embed, + final_norm, + layers, + stream, + gemm, + scratch: None, + }; + Ok(model) + } + + fn gemm(&self, s: &Scratch, x: usize, w: &Tensor, y: usize, m: usize) -> Result<()> { + let (n, k) = (w.shape[0] as i32, w.shape[1] as i32); + // SAFETY: x and y are scratch buffers sized for m rows of w's shape. + check( + unsafe { + (cuda::api().cs1_gemm)( + self.gemm, + s.at(x), + w.ptr, + s.at(y), + m as i32, + n, + k, + n, + self.stream, + ) + }, + "gemm", + ) + } + + /// 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" + ); + cuda::set_device(0)?; + if self.scratch.as_ref().is_none_or(|s| t > s.cap) { + self.scratch = None; + self.scratch = Some(Scratch::new( + &self.cfg, + t.next_multiple_of(1024), + self.stream, + )?); + } + 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.run(s, t)?; + let mut last = vec![0u8; h * 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)? }; + let (pairs, _) = last.as_chunks::<2>(); + Ok(pairs + .iter() + .map(|&b| half::bf16::from_le_bytes(b).to_f32()) + .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<()> { + let cfg = &self.cfg; + let st = self.stream; + let (ti, hi, eps) = (t as i32, cfg.hidden as i32, cfg.eps); + let (kd, vd, hv) = (cfg.key_dim(), cfg.value_dim(), cfg.lin_v_heads); + 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 + // 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), + self.layers[0].input_norm.ptr, + p(s.x), + ti, + hi, + eps, + st, + ), + "input norm", + )?; + } + for (i, layer) in self.layers.iter().enumerate() { + match &layer.mixer { + Mixer::Linear(la) => { + self.gemm(s, s.x, &la.in_proj, s.gdn_in, t)?; + let ld = w.gdn_in as i32; + let z = s.gdn_in + w.conv * BF16; + let b = z + vd * BF16; + let a = b + hv * BF16; + unsafe { + check( + (cuda::api().cs1_gdn_conv)( + p(s.gdn_in), + ld, + la.conv.ptr, + p(s.lq), + p(s.lk), + p(s.lv), + ti, + kd as i32, + vd as i32, + st, + ), + "gdn conv", + )?; + check( + (cuda::api().cs1_gdn_gates)( + p(b), + p(a), + ld, + la.a_log.ptr, + la.dt_bias.ptr, + p(s.beta), + p(s.g).cast(), + ti, + hv as i32, + st, + ), + "gdn gates", + )?; + check( + (cuda::api().cs1_gdn_prefill)( + p(s.lq), + p(s.lk), + p(s.lv), + p(s.g).cast(), + p(s.beta), + p(s.lo), + p(s.workspace).cast(), + ti, + hv as i32, + cfg.lin_k_heads as i32, + (cfg.lin_k_dim as f32).powf(-0.5), + st, + ), + "gdn prefill", + )?; + check( + (cuda::api().cs1_gated_rms_norm)( + p(s.lo), + p(z), + ld, + la.norm.ptr, + p(s.ln), + ti, + hv as i32, + cfg.lin_v_dim as i32, + eps, + st, + ), + "gated norm", + )?; + } + self.gemm(s, s.ln, &la.out, s.delta, t)?; + } + Mixer::Full(fa) => { + self.gemm(s, s.x, &fa.qkv, s.attn_in, t)?; + let ld = w.attn_in as i32; + let k = s.attn_in + w.attn_q * BF16; + let v = k + cfg.kv_heads * cfg.head_dim * BF16; + unsafe { + check( + (cuda::api().cs1_attn_prep)( + p(s.attn_in), + p(k), + ld, + fa.q_norm.ptr, + fa.k_norm.ptr, + p(s.cos), + p(s.sin), + p(s.aq), + p(s.agate), + p(s.ak), + ti, + hq, + hk, + hd, + cfg.rotary_half as i32, + eps, + st, + ), + "attention prep", + )?; + check( + (cuda::api().cs1_attention)( + p(s.aq), + p(s.ak), + p(v), + ld, + p(s.ao), + ti, + hq, + hk, + hd, + (cfg.head_dim as f32).powf(-0.5), + st, + ), + "attention", + )?; + check( + (cuda::api().cs1_sigmoid_gate)( + p(s.ao), + p(s.agate), + t * cfg.heads * cfg.head_dim, + st, + ), + "attention gate", + )?; + } + self.gemm(s, s.ao, &fa.o, s.delta, t)?; + } + } + unsafe { + check( + (cuda::api().cs1_add_rms_norm)( + p(s.res), + p(s.delta), + layer.post_norm.ptr, + p(s.x), + ti, + hi, + eps, + st, + ), + "post-attention norm", + )?; + } + self.gemm(s, s.x, &layer.gate_up, s.gate_up, t)?; + unsafe { + check( + (cuda::api().cs1_silu_mul)( + p(s.gate_up), + (2 * cfg.intermediate) as i32, + p(s.act), + ti, + cfg.intermediate as i32, + st, + ), + "silu mul", + )?; + } + self.gemm(s, s.act, &layer.down, s.delta, t)?; + let next = self + .layers + .get(i + 1) + .map_or(&self.final_norm, |l| &l.input_norm); + unsafe { + check( + (cuda::api().cs1_add_rms_norm)( + p(s.res), + p(s.delta), + next.ptr, + p(s.x), + ti, + hi, + eps, + st, + ), + "input norm", + )?; + } + } + Ok(()) + } +} diff --git a/src/models/cua_s1/native/tests/kernels.rs b/src/models/cua_s1/native/tests/kernels.rs new file mode 100644 index 00000000..cd044c83 --- /dev/null +++ b/src/models/cua_s1/native/tests/kernels.rs @@ -0,0 +1,264 @@ +//! GPU checks of the attention and Gated DeltaNet kernels on random inputs. They need +//! a GPU and CUA_S1_CUDA_LIB pointing at libqwen3_5_cuda.so, so they only run when +//! asked for: +//! +//! CUA_S1_CUDA_LIB=$PWD/target/release/libqwen3_5_cuda.so \ +//! cargo test --release -p omni-cua-s1-native --test kernels -- --ignored + +use std::path::PathBuf; + +use half::bf16; +use omni_cua_s1_native::cuda::{self, DeviceBuffer, Stream, api, check}; + +fn setup() -> Stream { + let lib = std::env::var_os("CUA_S1_CUDA_LIB") + .map(PathBuf::from) + .expect("CUA_S1_CUDA_LIB must point at libqwen3_5_cuda.so"); + cuda::load(&lib).unwrap(); + cuda::set_device(0).unwrap(); + cuda::new_stream().unwrap() +} + +/// Uniform values in [-amp, amp), rounded to bfloat16, from a fixed seed. +fn random(n: usize, seed: u64, amp: f32) -> Vec { + let mut x = seed.wrapping_mul(0x9e37_79b9_7f4a_7c15) | 1; + (0..n) + .map(|_| { + x ^= x << 13; + x ^= x >> 7; + x ^= x << 17; + bf16::from_f32(((x >> 40) as f32 / (1u64 << 24) as f32 * 2.0 - 1.0) * amp) + }) + .collect() +} + +fn to_device(v: &[bf16], st: Stream) -> DeviceBuffer { + let bytes: Vec = v.iter().flat_map(|x| x.to_le_bytes()).collect(); + let buf = DeviceBuffer::new(bytes.len()).unwrap(); + // SAFETY: the buffer was allocated for these bytes. + unsafe { cuda::upload(buf.at(0), &bytes, st).unwrap() }; + buf +} + +fn f32_to_device(v: &[f32], st: Stream) -> DeviceBuffer { + let bytes: Vec = v.iter().flat_map(|x| x.to_le_bytes()).collect(); + let buf = DeviceBuffer::new(bytes.len()).unwrap(); + // SAFETY: the buffer was allocated for these bytes. + unsafe { cuda::upload(buf.at(0), &bytes, st).unwrap() }; + buf +} + +fn from_device(buf: &DeviceBuffer, n: usize, st: Stream) -> Vec { + let mut bytes = vec![0u8; n * 2]; + // SAFETY: the buffer holds n bfloat16 values. + unsafe { cuda::download(&mut bytes, buf.at(0), st).unwrap() }; + let (pairs, _) = bytes.as_chunks::<2>(); + pairs + .iter() + .map(|&b| bf16::from_le_bytes(b).to_f32()) + .collect() +} + +#[test] +#[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] +fn flash_attention_matches_float64_reference() { + let st = setup(); + let (hq, hk, dh) = (16usize, 4usize, 256usize); + for (t, amp) in [ + (1, 8.0), + (63, 8.0), + (65, 0.5), + (139, 2.0), + (700, 8.0), + (2048, 0.5), + ] { + let (qh, kh) = (random(t * hq * dh, 1, amp), random(t * hk * dh, 2, amp)); + // v is read in place from the q|k|v projection output, rows of 10240 as in the model + let (ldv, v_at) = (10240usize, (hq * 2 + hk) * dh); + let qkvh = random(t * ldv, 3, 1.0); + let (q, k, qkv) = (to_device(&qh, st), to_device(&kh, st), to_device(&qkvh, st)); + let out = DeviceBuffer::new(t * hq * dh * 2).unwrap(); + // SAFETY: every buffer holds t rows of the given widths. + let code = unsafe { + (api().cs1_attention)( + q.at(0), + k.at(0), + qkv.at(v_at * 2), + ldv as i32, + out.at(0), + t as i32, + hq as i32, + hk as i32, + dh as i32, + 0.0625, + st, + ) + }; + check(code, "attention").unwrap(); + let got = from_device(&out, t * hq * dh, st); + assert!( + got.iter().all(|x| x.is_finite()), + "non-finite output at t = {t}" + ); + // about 64 query rows per length, each against causal attention in float64; + // per (row, head): the largest difference over the largest magnitude + let mut worst = 0f64; + for i in (0..t).step_by(t.div_ceil(64)).chain([t - 1]) { + for h in 0..hq { + let g = h / (hq / hk); + let qi = &qh[(i * hq + h) * dh..][..dh]; + let s: Vec = (0..=i) + .map(|j| { + let kj = &kh[(j * hk + g) * dh..][..dh]; + qi.iter() + .zip(kj) + .map(|(a, b)| a.to_f64() * b.to_f64()) + .sum::() + * 0.0625 + }) + .collect(); + let m = s.iter().copied().fold(f64::NEG_INFINITY, f64::max); + let w: Vec = s.iter().map(|x| (x - m).exp()).collect(); + let z: f64 = w.iter().sum(); + let want: Vec = (0..dh) + .map(|d| { + (0..=i) + .map(|j| w[j] * qkvh[j * ldv + v_at + g * dh + d].to_f64()) + .sum::() + / z + }) + .collect(); + let row = &got[(i * hq + h) * dh..][..dh]; + let diff = row + .iter() + .zip(&want) + .map(|(&x, y)| (x as f64 - y).abs()) + .fold(0f64, f64::max); + worst = worst.max(diff / want.iter().map(|y| y.abs()).fold(1e-3, f64::max)); + } + } + eprintln!("attention t = {t}, amplitude {amp}: largest relative difference {worst:.2e}"); + assert!(worst < 1.6e-2, "t = {t}: {worst}"); + } +} + +/// Transformers' torch_recurrent_gated_delta_rule in float64, one token at a time, +/// with the L2 norms of q and k and q scaled by K^-1/2. +#[allow(clippy::too_many_arguments)] +fn gated_delta_reference( + q: &[bf16], + k: &[bf16], + v: &[bf16], + g: &[f32], + beta: &[bf16], + t: usize, + h: usize, + hk: usize, + d: usize, +) -> Vec { + let mut out = vec![0f64; t * h * d]; + for head in 0..h { + let kh = head / (h / hk); + let mut s = vec![0f64; d * d]; // [K][V] + for tok in 0..t { + let norm = |x: &[bf16]| { + let x: Vec = x.iter().map(|v| v.to_f64()).collect(); + let inv = 1.0 / (x.iter().map(|v| v * v).sum::() + 1e-6).sqrt(); + x.into_iter().map(|v| v * inv).collect::>() + }; + let qv: Vec = norm(&q[(tok * hk + kh) * d..][..d]) + .into_iter() + .map(|x| x / (d as f64).sqrt()) + .collect(); + let kv = norm(&k[(tok * hk + kh) * d..][..d]); + let vv: Vec = v[(tok * h + head) * d..][..d] + .iter() + .map(|x| x.to_f64()) + .collect(); + let decay = (g[tok * h + head] as f64).exp(); + let b = beta[tok * h + head].to_f64(); + s.iter_mut().for_each(|x| *x *= decay); + for j in 0..d { + let mem: f64 = (0..d).map(|i| kv[i] * s[i * d + j]).sum(); + let delta = (vv[j] - mem) * b; + for i in 0..d { + s[i * d + j] += kv[i] * delta; + } + } + for j in 0..d { + out[(tok * h + head) * d + j] = (0..d).map(|i| qv[i] * s[i * d + j]).sum(); + } + } + } + out +} + +#[test] +#[ignore = "needs a GPU and CUA_S1_CUDA_LIB"] +fn gated_delta_rule_matches_recurrent_reference() { + let st = setup(); + let (h, hk, d) = (4usize, 2usize, 128usize); + for t in [1usize, 64, 150] { + // q close to k, so that q.k and the outputs are of order one as in the model + let k = random(t * hk * d, 12, 1.0); + let q: Vec = k + .iter() + .zip(random(t * hk * d, 11, 1.0)) + .map(|(k, n)| bf16::from_f32(0.8 * k.to_f32() + 0.2 * n.to_f32())) + .collect(); + let v = random(t * h * d, 13, 1.0); + // log decays in (-2, 0) and learning rates in (0, 1), as sigmoid and -exp * softplus give + let g: Vec = random(t * h, 14, 1.0) + .iter() + .map(|x| x.to_f32() - 1.0) + .collect(); + let beta: Vec = random(t * h, 15, 0.5) + .iter() + .map(|x| bf16::from_f32(x.to_f32() + 0.5)) + .collect(); + let want = gated_delta_reference(&q, &k, &v, &g, &beta, t, h, hk, d); + let (qd, kd, vd, gd, bd) = ( + to_device(&q, st), + to_device(&k, st), + to_device(&v, st), + f32_to_device(&g, st), + to_device(&beta, st), + ); + let o = DeviceBuffer::new(t * h * d * 2).unwrap(); + // SAFETY: pure function of its arguments. + let floats = unsafe { (api().cs1_gdn_workspace_floats)(t as i32, h as i32) }; + let ws = DeviceBuffer::new(floats * 4).unwrap(); + // SAFETY: every buffer holds t rows of the given widths, the workspace its size. + unsafe { + check( + (api().cs1_gdn_prefill)( + qd.at(0), + kd.at(0), + vd.at(0), + gd.at(0).cast::(), + bd.at(0), + o.at(0), + ws.at(0).cast::(), + t as i32, + h as i32, + hk as i32, + (d as f32).powf(-0.5), + st, + ), + "gdn prefill", + ) + .unwrap(); + } + let got = from_device(&o, t * h * d, st); + let scale = want.iter().fold(0f64, |m, x| m.max(x.abs())); + let worst = got + .iter() + .zip(&want) + .map(|(a, b)| (*a as f64 - b).abs()) + .fold(0f64, f64::max); + eprintln!( + "gated delta t = {t}: largest difference {worst:.2e}, largest |reference| {scale:.2}" + ); + assert!(worst <= 2e-2 * scale, "t = {t}: {worst} vs scale {scale}"); + } +} diff --git a/src/models/cua_s1/text/contract.py b/src/models/cua_s1/text/contract.py new file mode 100644 index 00000000..326e758d --- /dev/null +++ b/src/models/cua_s1/text/contract.py @@ -0,0 +1,160 @@ +"""Request mapping, prompts and answers for Cua-S1 4B 0.2, following +`src/models/cua_s1/README.md`. No torch imports, so it can be tested without weights. +""" + +from __future__ import annotations + +import json +import math +from dataclasses import dataclass +from typing import Any + +MODEL_NAME = "cua-s1-4b-0.2" +MODEL_ID = "cua-ai/cua-s1-4b-0.2@16818868b0cc7813808aae4e87b417657046ab79:text" +LETTERS = "ABCDEFGHIJKLMNOPQRSTUVWXYZ" +MAX_QUESTIONS = 64 + +# The system message and the user message layout are copied from trycua/cua at +# 0e75660ce4c2edda519e0c795fa3ad98abf4e76f (`libs/cua-s1/python/src/cua_s1/four_b.py` +# and `libs/cua-driver/examples/jev-use/python/decision_models.py`). +# +# MIT License +# +# Copyright (c) 2025 Cua AI, Inc. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. +SYSTEM_PROMPT = ( + "You are a one-pass computer-use decision model. You are shown the " + "current state of a screen and a fixed, closed list of candidate " + "(element, action) options, each given a single letter. Choose exactly " + "one option: the single best next action to take. Answer with ONLY that " + "option's letter -- no words, no punctuation, no explanation." +) + + +class RequestError(ValueError): + def __init__(self, message: str, status: int = 422) -> None: + super().__init__(message) + self.status = status + + +@dataclass(frozen=True) +class Question: + name: str + goal: str + keys: tuple[str, ...] + labels: tuple[str, ...] + + +def _unique_keys(pairs: list[tuple[str, Any]]) -> dict[str, Any]: + obj = {} + for key, value in pairs: + if key in obj: + raise RequestError(f"duplicate key {key!r}", 400) + obj[key] = value + return obj + + +def parse_body(raw: bytes) -> dict[str, Any]: + try: + body = json.loads(raw.decode(), object_pairs_hook=_unique_keys) + # NaN, Infinity, numbers out of range and lone surrogates fail here. + json.dumps(body, ensure_ascii=False, allow_nan=False).encode() + except (ValueError, RecursionError) as error: + if isinstance(error, RequestError): + raise + raise RequestError(f"request body is not valid JSON: {error}", 400) from error + if not isinstance(body, dict): + raise RequestError("request body must be a JSON object", 400) + return body + + +def _text(value: Any, where: str) -> str: + """A string as is; an object or array as Python's json.dumps writes it.""" + if not isinstance(value, (str, dict, list)): + raise RequestError(f"{where} must be a string, an object or an array") + return value if isinstance(value, str) else json.dumps(value, ensure_ascii=False) + + +def map_request(body: dict[str, Any]) -> tuple[str, list[Question]]: + if body.get("model") != MODEL_NAME: + raise RequestError(f"'model' must be {MODEL_NAME!r}") + if body.get("state") in ("", {}, []): + raise RequestError("'state' must not be empty") + state = _text(body.get("state"), "'state'") + questions = body.get("questions") + if not isinstance(questions, dict) or not questions: + raise RequestError("'questions' must be a non-empty object") + if len(questions) > MAX_QUESTIONS: + raise RequestError(f"more than {MAX_QUESTIONS} questions", 413) + mapped = [] + for name, q in questions.items(): + where = f"question {name!r}" + if not isinstance(q, dict): + raise RequestError(f"{where} must be an object") + if q.get("type") in ("score", "noul"): + raise RequestError(f"{where}: type {q['type']!r} is not supported") + if q.get("type") != "choice": + raise RequestError(f"{where}: unknown type {q.get('type')!r}") + if "instructions" not in q: + raise RequestError(f"{where}: 'instructions' is required") + goal = "" if q["instructions"] is None else _text(q["instructions"], where) + criteria = q.get("criteria") + if not isinstance(criteria, dict) or not 1 <= len(criteria) <= len(LETTERS): + raise RequestError( + f"{where}: 'criteria' must be an object with 1 to 26 options" + ) + labels = tuple( + json.dumps( + key if value is None else _text(value, f"{where}: {key!r}"), + ensure_ascii=False, + )[1:-1] + for key, value in criteria.items() + ) + mapped.append(Question(name, goal, tuple(criteria), labels)) + return state, mapped + + +def build_messages(state: str, question: Question) -> list[dict[str, str]]: + options = "\n".join( + f'{letter}. Decision "{label}" -> select' + for letter, label in zip(LETTERS, question.labels) + ) + user = ( + (f"Goal: {question.goal}\n\n" if question.goal else "") + + "App: Cua Driver\nTask family: closed-candidate decision\n\n" + + f"Accessibility tree:\n{state}\n\nOptions:\n{options}\n\nAnswer with a single letter." + ) + return [ + {"role": "system", "content": SYSTEM_PROMPT}, + {"role": "user", "content": user}, + ] + + +def answer(question: Question, probabilities: list[float]) -> dict[str, Any]: + """The Jev choice answer; ties go to the earliest option. `confidence` is the + normalized entropy `1 - H(p) / ln(n)`, as the LAYA worker reports it.""" + n = len(probabilities) + entropy = -sum(p * math.log(p) for p in probabilities if p > 0) + return { + "type": "choice", + "choice": question.keys[max(range(n), key=probabilities.__getitem__)], + "probabilities": dict(zip(question.keys, probabilities)), + "confidence": max(0.0, 1 - entropy / math.log(n)) if n > 1 else 1.0, + } diff --git a/src/models/cua_s1/text/model.py b/src/models/cua_s1/text/model.py new file mode 100644 index 00000000..ea0415af --- /dev/null +++ b/src/models/cua_s1/text/model.py @@ -0,0 +1,45 @@ +"""Qwen3.5-4B with the Cua-S1 `text` adapter, loaded and scored as upstream +`cua_s1.four_b.FourBModel` does: an unmerged PEFT adapter, the chat template with its +generation prompt, full logits, and a float32 softmax over the option letters at the +last position. That keeps the probabilities bitwise identical to the reference. +""" + +from __future__ import annotations + +import json +from pathlib import Path + +import torch +from peft import PeftModel +from transformers import AutoModelForCausalLM, AutoTokenizer + +from .contract import LETTERS, Question, build_messages + + +class TextModel: + def __init__(self, base: str, adapter: str, device: str, dtype: str) -> None: + # PEFT only warns about keys it cannot place: refuse the multimodal adapter. + config = json.loads((Path(adapter) / "adapter_config.json").read_text()) + if "linear_fc1" in config["target_modules"]: + raise ValueError( + f"{adapter} is the multimodal adapter; pass its text/ directory" + ) + self.tokenizer = AutoTokenizer.from_pretrained(base) + model = AutoModelForCausalLM.from_pretrained( + base, dtype=getattr(torch, dtype), device_map=device + ) + self.model = PeftModel.from_pretrained(model, adapter).eval() + self.letter_ids = self.tokenizer.convert_tokens_to_ids(list(LETTERS)) + + def encode(self, state: str, question: Question): + text = self.tokenizer.apply_chat_template( + build_messages(state, question), tokenize=False, add_generation_prompt=True + ) + return self.tokenizer(text, return_tensors="pt") + + @torch.no_grad() + def score(self, inputs, n_options: int) -> list[float]: + logits = self.model(**inputs.to(self.model.device)).logits[0, -1] + return torch.softmax( + logits[self.letter_ids[:n_options]].float(), dim=-1 + ).tolist() diff --git a/tests/cua_s1/test_text_contract.py b/tests/cua_s1/test_text_contract.py new file mode 100644 index 00000000..870c50f2 --- /dev/null +++ b/tests/cua_s1/test_text_contract.py @@ -0,0 +1,183 @@ +"""Contract tests without weights or torch. The tokenizer test also runs when +CUA_S1_BASE points to a local Qwen/Qwen3.5-4B directory (tokenizer files only). + +PYTHONPATH=src python -m pytest tests/cua_s1 +""" + +import json +import math +import os + +import pytest + +from models.cua_s1.text.contract import ( + RequestError, + answer, + build_messages, + map_request, + parse_body, +) + + +def body(state="Screen", **question): + q = { + "type": "choice", + "instructions": "Pick one.", + "criteria": {"a": "A", "b": "B"}, + } + q.update(question) + return {"model": "cua-s1-4b-0.2", "state": state, "questions": {"q": q}} + + +def mapped(request): + return map_request(parse_body(json.dumps(request).encode())) + + +def reject(request, status=422): + raw = request if isinstance(request, bytes) else json.dumps(request).encode() + with pytest.raises(RequestError) as info: + map_request(parse_body(raw)) + assert info.value.status == status + return str(info.value) + + +# upstream's libs/cua-driver/examples/jev-use/fixtures/jev-choice-request-v1.json, +# with the chooser's rendered region as `state`, and the user message upstream builds for it +FIXTURE = body( + 'Visual-region-derived observation for capture "capture-fixture-1":\n' + "\"submit-text\": text 'Submit' at (300,240,100,40) confidence=0.96 interactive=true", + instructions="Submit the verified form.", + criteria={ + "submit-form": "Submit using the unique validated visual region.", + "reobserve": "Discard this decision set and obtain a fresh observation.", + "abstain": "Stop without acting if no supplied action is safe.", + }, +) +FIXTURE_USER = ( + "Goal: Submit the verified form.\n\n" + "App: Cua Driver\nTask family: closed-candidate decision\n\n" + "Accessibility tree:\n" + 'Visual-region-derived observation for capture "capture-fixture-1":\n' + "\"submit-text\": text 'Submit' at (300,240,100,40) confidence=0.96 interactive=true\n\n" + "Options:\n" + 'A. Decision "Submit using the unique validated visual region." -> select\n' + 'B. Decision "Discard this decision set and obtain a fresh observation." -> select\n' + 'C. Decision "Stop without acting if no supplied action is safe." -> select\n\n' + "Answer with a single letter." +) + + +def test_fixture_prompt_matches_upstream(): + state, (question,) = mapped(FIXTURE) + system, user = build_messages(state, question) + assert user == {"role": "user", "content": FIXTURE_USER} + assert system["content"].startswith( + "You are a one-pass computer-use decision model." + ) + + +@pytest.mark.skipif(not os.environ.get("CUA_S1_BASE"), reason="set CUA_S1_BASE to run") +def test_tokenizer(): + from transformers import AutoTokenizer + + tokenizer = AutoTokenizer.from_pretrained(os.environ["CUA_S1_BASE"]) + assert tokenizer.convert_tokens_to_ids(list("ABCDEFGHIJKLMNOPQRSTUVWXYZ")) == list( + range(32, 58) + ) + state, (question,) = mapped(FIXTURE) + text = tokenizer.apply_chat_template( + build_messages(state, question), tokenize=False, add_generation_prompt=True + ) + assert text.endswith("<|im_start|>assistant\n\n") + ids = tokenizer(text)["input_ids"] + assert len(ids) == 218 + assert ids == tokenizer(text, add_special_tokens=False)["input_ids"] + + +def test_goal_left_out_when_empty_or_null(): + for goal in ["", None]: + state, (question,) = mapped(body(instructions=goal)) + assert build_messages(state, question)[1]["content"].startswith( + "App: Cua Driver\n" + ) + + +def test_structured_values_escaping_and_null_label(): + tree = { + "app": "Settings", + "elements": [{"id": "e1", "label": "Location", "on": True}], + } + state, (question,) = mapped( + body( + tree, + instructions={"question": "Which one?"}, + criteria={ + "e1": {"action": "click", "element": "e1"}, + "e2": ["click", "e2"], + "quote": 'Click "Submit"\n(tab\there) C:\\Users', + "save": "点击「保存」", + "abstain": None, + }, + ) + ) + assert state == json.dumps(tree, ensure_ascii=False) + assert question.goal == '{"question": "Which one?"}' + assert question.labels == ( + '{\\"action\\": \\"click\\", \\"element\\": \\"e1\\"}', + '[\\"click\\", \\"e2\\"]', + 'Click \\"Submit\\"\\n(tab\\there) C:\\\\Users', + "点击「保存」", + "abstain", + ) + + +def test_request_errors(): + assert "not supported" in reject(body(type="score", criteria=["low", "high"])) + assert "not supported" in reject(body(type="noul")) + assert "unknown type" in reject(body(type="rank")) + assert "1 to 26 options" in reject(body(criteria={})) + assert "1 to 26 options" in reject(body(criteria={f"o{i}": "x" for i in range(27)})) + assert ( + len(mapped(body(criteria={f"o{i}": "x" for i in range(26)}))[1][0].keys) == 26 + ) + no_instructions = body() + del no_instructions["questions"]["q"]["instructions"] + assert "'instructions' is required" in reject(no_instructions) + for value in [1, 2.5, True]: + reject(body(criteria={"a": value})) + for state in ["", {}, [], None, 3, True]: + reject(body(state)) + assert "'model'" in reject({**body(), "model": "english"}) + many = body() + many["questions"] = {f"q{i}": many["questions"]["q"] for i in range(65)} + reject(many, status=413) + + +@pytest.mark.parametrize( + "raw", + [ + b'{"model": "cua-s1-4b-0.2", "state": {"x": 1, "x": 2}}', + b'{"model": "cua-s1-4b-0.2", "state": NaN}', + b'{"model": "cua-s1-4b-0.2", "state": {"x": 1e400}}', + b'{"model": "cua-s1-4b-0.2", "state": {"x": ' + b"9" * 5000 + b"}}", + b'{"model": "cua-s1-4b-0.2", "state": "\\ud800"}', + b"[" * 100000 + b"]" * 100000, + b"\xff\xfe", + b"\xef\xbb\xbf{}", + b"not json", + b"[1, 2]", + ], +) +def test_malformed_bodies_are_400(raw): + reject(raw, status=400) + + +def test_answer(): + _, (question,) = mapped(body()) + tie = answer(question, [0.5, 0.5]) + assert tie["choice"] == "a" and tie["confidence"] == pytest.approx(0.0, abs=1e-12) + result = answer(question, [0.12, 0.88]) + assert result["choice"] == "b" + assert list(result["probabilities"]) == ["a", "b"] + h = -(0.12 * math.log(0.12) + 0.88 * math.log(0.88)) + assert result["confidence"] == pytest.approx(1 - h / math.log(2)) diff --git a/tests/cua_s1/test_text_server.py b/tests/cua_s1/test_text_server.py new file mode 100644 index 00000000..bc4be9c9 --- /dev/null +++ b/tests/cua_s1/test_text_server.py @@ -0,0 +1,85 @@ +"""HTTP tests for the worker with a fake model: no weights, no torch.""" + +import json + +import pytest + +pytest.importorskip("fastapi") +pytest.importorskip("httpx") +from fastapi.testclient import TestClient # noqa: E402 + +from frontend.cua_s1_text import build_app # noqa: E402 + + +class Ids: + def __init__(self, n): + self.shape = (1, n) + + +class FakeModel: + def __init__(self, tokens=100, error=None, nan=False): + self.tokens, self.error, self.nan, self.calls = tokens, error, nan, 0 + + def encode(self, state, question): + return {"input_ids": Ids(self.tokens)} + + def score(self, inputs, n_options): + self.calls += 1 + if self.error: + raise self.error + p = [0.1] * (n_options - 1) + [1 - 0.1 * (n_options - 1)] + return [float("nan")] + p[1:] if self.nan else p + + +BODY = { + "model": "cua-s1-4b-0.2", + "state": "Screen", + "questions": { + "_sa": { + "type": "choice", + "instructions": "Pick.", + "criteria": {"_x": "A", "b": "B"}, + } + }, +} + + +def post(model=None, **kwargs): + return TestClient(build_app(model or FakeModel())).post("/v1/systemone", **kwargs) + + +def test_health_and_answer(): + app = build_app(FakeModel()) + assert TestClient(app).get("/health").json()["status"] == "ready" + response = post(json=BODY) + assert response.status_code == 200, response.text + reply = response.json() + assert reply["answers"]["_sa"]["choice"] == "b" + assert list(reply["answers"]["_sa"]["probabilities"]) == ["_x", "b"] + assert reply["usage"] == {"input_tokens": 100, "output_tokens": 0} + + +def test_errors(): + bad = json.loads(json.dumps(BODY)) + bad["questions"]["_sa"]["type"] = "noul" + assert post(json=bad).status_code == 422 + assert post(content=b"{").status_code == 400 + raw = b" " * (4 << 20) + json.dumps(BODY).encode() + assert post(content=iter([raw[:10], raw[10:]])).status_code == 413 + model = FakeModel(tokens=20000) + assert post(model, json=BODY).status_code == 413 and model.calls == 0 + + +@pytest.mark.parametrize( + "model", [FakeModel(error=RuntimeError("CUDA out of memory")), FakeModel(nan=True)] +) +def test_model_failure_is_500(model): + response = post(model, json=BODY) + assert response.status_code == 500 + assert response.json() == {"detail": "inference failed"} + + +def test_warmup(): + model = FakeModel() + build_app(model).state.warmup() + assert model.calls == 1