diff --git a/Cargo.lock b/Cargo.lock index 4b9adda..2baa2f3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -958,6 +958,15 @@ dependencies = [ "tokio", ] +[[package]] +name = "omni-cuda" +version = "0.1.0" +dependencies = [ + "anyhow", + "libloading", + "tempfile", +] + [[package]] name = "omni-jev" version = "0.1.0" @@ -974,7 +983,9 @@ version = "0.1.0" dependencies = [ "anyhow", "half", + "libloading", "memmap2", + "omni-cuda", "safetensors 0.6.2", "serde", "serde_json", diff --git a/Cargo.toml b/Cargo.toml index 661c83e..7423291 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,3 +1,3 @@ [workspace] -members = ["src/frontend", "src/models/cua_s1/native", "src/models/laya"] +members = ["src/frontend", "src/models/cua_s1/native", "src/models/laya", "src/backends/cuda"] resolver = "3" diff --git a/recipe/laya/README.md b/recipe/laya/README.md index 254c12c..8111a0f 100644 --- a/recipe/laya/README.md +++ b/recipe/laya/README.md @@ -56,3 +56,52 @@ if the worker requires a bearer token. See the [frontend documentation](../../src/frontend/README.md) for configuration and transport behavior. + +## Native residency and workspace validation + +These opt-in Rust checks test CUDA allocations and transfers. They do not need the +Python worker or frontend and do not test inference, model outputs or latency. +The normal CPU tests skip them. + +Use a Linux host with an approved CUDA GPU, a working NVIDIA driver, the CUDA +toolkit (`nvcc`) and Rust. Build the trusted resource library from this checkout; +it uses ABI version 1 and needs neither TileLang nor cuBLAS. `LAYA_CUDA_DEVICE` is +the approved device ordinal after `CUDA_VISIBLE_DEVICES` filtering. + +```sh +export LAYA_CUDA_LIBRARY=/tmp/liblaya-resources.so +export LAYA_CUDA_DEVICE=0 +nvcc -shared -Xcompiler=-fPIC -O2 src/backends/cuda/kernels/runtime.cu \ + -o "$LAYA_CUDA_LIBRARY" +``` + +The workspace check needs no checkpoint. It writes and reads all 17 buffers twice +at `(batch, sequence) = (1, 16), (1, 512), (16, 512)`, using deterministic byte +patterns: + +```sh +cargo test --release --locked -p omni-laya --lib \ + workspace::tests::real_gpu_workspace_capacity_and_reuse \ + -- --ignored --exact --nocapture +``` + +For the residency check, use an unchanged local snapshot of +`convaiinnovations/laya` revision `55cf4c4ebb4ebe31b2550e8bdf3bd21b99753851`, +including `model.safetensors`. Generate the oracle from that same snapshot in a +Python environment with PyTorch, safetensors and NumPy: + +```sh +export LAYA_CHECKPOINT=/path/to/laya/snapshot +export LAYA_WEIGHT_ORACLE=/tmp/laya-weight-oracle.json +python recipe/laya/native/export_weights.py "$LAYA_CHECKPOINT" "$LAYA_WEIGHT_ORACLE" +cargo test --release --locked -p omni-laya --lib \ + resident::tests::real_checkpoint_residency_matches_torch \ + -- --ignored --exact --nocapture +``` + +The residency check validates all 206 checkpoint tensors, uploads the 205 used +tensors and compares readback hashes with the Torch conversion oracle. The legacy +`temperature` buffer is validated but not uploaded. Reported allocation bytes +exclude CUDA context and library overhead. See the +[model contracts](https://github.com/linear3735/system1-omni/blob/codex/laya-workspace/src/models/laya/README.md) for storage precision, +workspace layouts and ownership. diff --git a/src/backends/cuda/Cargo.toml b/src/backends/cuda/Cargo.toml new file mode 100644 index 0000000..cdc7644 --- /dev/null +++ b/src/backends/cuda/Cargo.toml @@ -0,0 +1,16 @@ +[package] +name = "omni-cuda" +version = "0.1.0" +edition = "2024" +publish = false + +[dependencies] +anyhow = "1" +libloading = "0.8" + +[dev-dependencies] +tempfile = "3" + +[[test]] +name = "runtime" +path = "../../../tests/backends/cuda/runtime.rs" diff --git a/src/backends/cuda/README.md b/src/backends/cuda/README.md index 577209d..c0e0762 100644 --- a/src/backends/cuda/README.md +++ b/src/backends/cuda/README.md @@ -1,7 +1,56 @@ # CUDA backend -Planned home for high-performance NVIDIA GPU operations and kernel integration. Implement the operations required by the first model, with hardware-specific optimizations where needed. +[`qwen3_5/`](qwen3_5/) provides the prefill-only Qwen3.5 operations used by the +Cua-S1 native worker, measured on sm_89. -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. +## Laya resources -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. +`omni-cuda` loads Laya's CUDA resource library at runtime. It owns one device and +stream per context, plus the buffers allocated through that context. Rust builds +and CPU tests need no CUDA toolkit. + +This first slice covers allocation, copies, synchronization and cleanup. Model +initialization, weights, kernel calls, Graphs and hardware-specific optimizations +remain separate. It does not yet run Laya inference. + +### Build and check + +On a machine with the CUDA toolkit, build the resource library: + +```sh +nvcc -shared -Xcompiler=-fPIC -O2 src/backends/cuda/kernels/runtime.cu -o /tmp/liblaya_cuda.so +LAYA_CUDA_LIBRARY=/tmp/liblaya_cuda.so LAYA_CUDA_DEVICE=0 \ + cargo test --locked -p omni-cuda --test runtime -- --ignored +``` + +The device is an ordinal after `CUDA_VISIBLE_DEVICES` filtering. The library +contains no generated kernels and needs neither TileLang nor cuBLAS. This command +builds only the resource slice; the complete model bundle has a separate build. + +The normal CPU tests compile a small C fixture with `cc`. They check the dynamic +loader, errors, copy bounds and resource lifetime. They do not validate CUDA or +hardware support. The ignored test exercises real allocation and copy roundtrips. + +### Ownership and ABI + +Load only a trusted library with the matching ABI. `Cuda::load(path, device)` +checks `laya_abi_version() == 1` and all required symbols before creating a stream. +The old prototype's `laya_init` library has no version symbol and is rejected. + +`Cuda` and `Buffer` stay on their creating thread. A buffer keeps its stream and +library alive even after the caller drops `Cuda`. Operations select the owning +device before using its resources. Destruction attempts synchronization and +cleanup; call `sync()` explicitly when errors need to reach the caller. + +`write` and `read` check byte limits and synchronize before returning, so borrowed +host memory cannot outlive a queued copy. They are not Graph-capture operations. +Allocation of zero bytes is rejected; empty reads and writes are no-ops. + +The native resource entry points return zero on success and CUDA error codes on +failure; code 1000 means an invalid runtime argument. `laya_error_string` explains +the code. Upload and download take the caller's stream as their last argument and +do not synchronize internally. No Hopper requirement or model initialization is +hidden in stream creation. + +These are Laya's resource entry points, not a new shared tensor interface. A common +runtime can be extracted when another model needs the same implementation. diff --git a/src/backends/cuda/kernels/runtime.cu b/src/backends/cuda/kernels/runtime.cu new file mode 100644 index 0000000..da4d69c --- /dev/null +++ b/src/backends/cuda/kernels/runtime.cu @@ -0,0 +1,65 @@ +#include +#include + +namespace { +constexpr int invalid_argument = 1000; +} + +extern "C" { +uint32_t laya_abi_version() { return 1; } + +const char* laya_error_string(int code) { + return code == invalid_argument ? "invalid runtime argument" + : cudaGetErrorString(static_cast(code)); +} + +int laya_set_device(int device) { return cudaSetDevice(device); } + +int laya_stream_create(void** stream) { + if (!stream) return invalid_argument; + *stream = nullptr; + cudaStream_t created = nullptr; + cudaError_t status = cudaStreamCreateWithFlags(&created, cudaStreamNonBlocking); + if (status == cudaSuccess) *stream = created; + return status; +} + +int laya_alloc(void** p, size_t bytes) { + if (!p) return invalid_argument; + *p = nullptr; + if (!bytes) return invalid_argument; + void* allocated = nullptr; + cudaError_t status = cudaMalloc(&allocated, bytes); + if (status == cudaSuccess) *p = allocated; + return status; +} + +int laya_free(void* p) { + if (!p) return invalid_argument; + return cudaFree(p); +} + +int laya_upload(void* dst, const void* src, size_t bytes, void* stream) { + if (!bytes) return 0; + if (!dst || !src || !stream) return invalid_argument; + return cudaMemcpyAsync(dst, src, bytes, cudaMemcpyHostToDevice, + static_cast(stream)); +} + +int laya_download(void* dst, const void* src, size_t bytes, void* stream) { + if (!bytes) return 0; + if (!dst || !src || !stream) return invalid_argument; + return cudaMemcpyAsync(dst, src, bytes, cudaMemcpyDeviceToHost, + static_cast(stream)); +} + +int laya_sync(void* stream) { + if (!stream) return invalid_argument; + return cudaStreamSynchronize(static_cast(stream)); +} + +int laya_stream_free(void* stream) { + if (!stream) return invalid_argument; + return cudaStreamDestroy(static_cast(stream)); +} +} diff --git a/src/backends/cuda/src/lib.rs b/src/backends/cuda/src/lib.rs new file mode 100644 index 0000000..4338ade --- /dev/null +++ b/src/backends/cuda/src/lib.rs @@ -0,0 +1,188 @@ +//! CUDA resources confined to one thread. CPU builds do not link CUDA. +use anyhow::{Result, anyhow, ensure}; +use libloading::Library; +use std::{ + ffi::{CStr, c_char, c_void}, + path::Path, + rc::Rc, +}; + +type Ptr = *mut c_void; + +struct Functions { + error: unsafe extern "C" fn(i32) -> *const c_char, + set_device: unsafe extern "C" fn(i32) -> i32, + stream_create: unsafe extern "C" fn(*mut Ptr) -> i32, + alloc: unsafe extern "C" fn(*mut Ptr, usize) -> i32, + free: unsafe extern "C" fn(Ptr) -> i32, + upload: unsafe extern "C" fn(Ptr, *const u8, usize, Ptr) -> i32, + download: unsafe extern "C" fn(*mut u8, Ptr, usize, Ptr) -> i32, + sync: unsafe extern "C" fn(Ptr) -> i32, + stream_free: unsafe extern "C" fn(Ptr) -> i32, +} + +impl Functions { + fn check(&self, code: i32) -> Result<()> { + if code == 0 { + return Ok(()); + } + let message = unsafe { (self.error)(code) }; + if message.is_null() { + return Err(anyhow!("CUDA error {code}")); + } + Err(anyhow!("CUDA {code}: {}", unsafe { + CStr::from_ptr(message).to_string_lossy() + })) + } +} + +struct Context { + _library: Library, + functions: Functions, + device: i32, + stream: Ptr, +} + +impl Context { + fn activate(&self) -> Result<()> { + self.functions + .check(unsafe { (self.functions.set_device)(self.device) }) + } + + fn sync(&self) -> Result<()> { + self.activate()?; + self.functions + .check(unsafe { (self.functions.sync)(self.stream) }) + } +} + +impl Drop for Context { + fn drop(&mut self) { + if self.activate().is_ok() { + unsafe { + (self.functions.sync)(self.stream); + (self.functions.stream_free)(self.stream); + } + } + } +} + +#[derive(Clone)] +pub struct Cuda { + ctx: Rc, +} + +impl Cuda { + /// # Safety + /// `path` must name a trusted library implementing the complete runtime ABI. + pub unsafe fn load(path: &Path, device: i32) -> Result { + let library = unsafe { Library::new(path) }?; + let version = + unsafe { library.get:: u32>(b"laya_abi_version\0")?() }; + ensure!(version == 1, "unsupported CUDA runtime ABI {version}"); + let functions = unsafe { + Functions { + error: *library.get(b"laya_error_string\0")?, + set_device: *library.get(b"laya_set_device\0")?, + stream_create: *library.get(b"laya_stream_create\0")?, + alloc: *library.get(b"laya_alloc\0")?, + free: *library.get(b"laya_free\0")?, + upload: *library.get(b"laya_upload\0")?, + download: *library.get(b"laya_download\0")?, + sync: *library.get(b"laya_sync\0")?, + stream_free: *library.get(b"laya_stream_free\0")?, + } + }; + functions.check(unsafe { (functions.set_device)(device) })?; + let mut stream = std::ptr::null_mut(); + functions.check(unsafe { (functions.stream_create)(&mut stream) })?; + ensure!(!stream.is_null(), "CUDA runtime returned a null stream"); + Ok(Self { + ctx: Rc::new(Context { + _library: library, + functions, + device, + stream, + }), + }) + } + + pub fn alloc(&self, bytes: usize) -> Result { + ensure!(bytes > 0, "zero CUDA allocation"); + self.ctx.activate()?; + let mut ptr = std::ptr::null_mut(); + self.ctx + .functions + .check(unsafe { (self.ctx.functions.alloc)(&mut ptr, bytes) })?; + ensure!(!ptr.is_null(), "CUDA runtime returned a null allocation"); + Ok(Buffer { + ctx: self.ctx.clone(), + ptr, + bytes, + }) + } + + pub fn upload(&self, bytes: &[u8]) -> Result { + let buffer = self.alloc(bytes.len())?; + buffer.write(bytes)?; + Ok(buffer) + } + + pub fn sync(&self) -> Result<()> { + self.ctx.sync() + } +} + +pub struct Buffer { + ctx: Rc, + ptr: Ptr, + bytes: usize, +} + +impl Buffer { + pub fn bytes(&self) -> usize { + self.bytes + } + + pub fn write(&self, bytes: &[u8]) -> Result<()> { + ensure!(bytes.len() <= self.bytes, "upload exceeds allocation"); + if bytes.is_empty() { + return Ok(()); + } + self.ctx.activate()?; + let functions = &self.ctx.functions; + let copied = + unsafe { (functions.upload)(self.ptr, bytes.as_ptr(), bytes.len(), self.ctx.stream) }; + // Even a failed copy may have queued work using the borrowed host memory. + let synced = unsafe { (functions.sync)(self.ctx.stream) }; + functions.check(copied)?; + functions.check(synced) + } + + pub fn read(&self, bytes: usize) -> Result> { + ensure!(bytes <= self.bytes, "download exceeds allocation"); + let mut data = vec![0; bytes]; + if bytes == 0 { + return Ok(data); + } + self.ctx.activate()?; + let functions = &self.ctx.functions; + let copied = + unsafe { (functions.download)(data.as_mut_ptr(), self.ptr, bytes, self.ctx.stream) }; + let synced = unsafe { (functions.sync)(self.ctx.stream) }; + functions.check(copied)?; + functions.check(synced)?; + Ok(data) + } +} + +impl Drop for Buffer { + fn drop(&mut self) { + if self.ctx.activate().is_ok() { + unsafe { + (self.ctx.functions.sync)(self.ctx.stream); + (self.ctx.functions.free)(self.ptr); + } + } + } +} diff --git a/src/models/laya/Cargo.toml b/src/models/laya/Cargo.toml index 0eba920..919da7c 100644 --- a/src/models/laya/Cargo.toml +++ b/src/models/laya/Cargo.toml @@ -8,11 +8,13 @@ publish = false anyhow = "1" half = "2" memmap2 = "0.9" +omni-cuda = { path = "../../backends/cuda" } safetensors = "0.6" serde = { version = "1", features = ["derive"] } serde_json = "1" [dev-dependencies] +libloading = "0.8" sha2 = "0.10" tempfile = "3" diff --git a/src/models/laya/README.md b/src/models/laya/README.md index 3cf9ec0..015d40f 100644 --- a/src/models/laya/README.md +++ b/src/models/laya/README.md @@ -23,6 +23,55 @@ cargo test --release --locked -p omni-laya --test weights -- --ignored These two CPU tests check all 206 tensor names and shapes, 618 conversion hashes, and the legacy temperature buffer. The normal CI job skips them because it does not download the full checkpoint. +## GPU weight residency + +`ResidentWeights::upload(&cuda, &weights)` validates the checkpoint inventory and +uploads the weights once. Embeddings use FP16; encoder norms, head norms and biases, +and the scorer input norm use FP32; other weights use BF16. Layouts stay unchanged. +The legacy `temperature` buffer is validated but not uploaded. + +`get(name)` returns the resident buffer. `bytes()` reports weight allocations only, +excluding CUDA context and allocator overhead. Buffers keep their CUDA context alive +after the caller drops the source mapping or `Cuda`. Failed loads release partial +allocations. Workspace, rotary tables and inference are separate modules. + +The ignored GPU check uploads all 205 used tensors and compares readback hashes +with the Torch conversion oracle; it does not test model outputs or latency. +Prerequisites and commands are in the +[native validation recipe](../../../recipe/laya/README.md#native-residency-and-workspace-validation). + +## Inference workspace + +`Workspace::new(&cuda, batch, sequence)` allocates fixed scratch buffers for one +shape. Batch must be 1, 2, 4, 8 or 16; sequence must be a multiple of 16 in 16..=512. +Invalid shapes fail before any allocation; a failed allocation releases the partial +workspace. Contents are uninitialized and must be written before use. + +`buffers()` borrows the named buffers without allowing allocations to be replaced. +The workspace outlives the caller's `Cuda` handle. `bytes()` reports scratch +allocations only, excluding resident weights and CUDA overhead. + +Let `B` be batch, `L` sequence, `D=1024`, and `M=MAX_MARKERS=2048`. Layouts are: + +| Buffers | Shape and dtype | +| --- | --- | +| ids / lengths / types | `[B,L]` int64 / `[B]` int32 / `[B]` int64 | +| residual / hidden / attention | `[B,L,D]` FP32 / BF16 / BF16 | +| qkv / gated / feed_forward | `[B,L,3D]` / `[B,L,2624]` / `[B,L,4096]`, BF16 | +| indices / offsets | `[M]` / `[B+1]`, int32 | +| markers / scored / logits | `[M,D]` / `[M,D]` / `[M]`, BF16 | +| features / action_hidden / actions | `[B,1028]` / `[B,256]` / `[B,2]`, BF16 | + +At `(B,L)=(1,512)`, allocations total 22,628,896 bytes; at `(16,512)`, +236,048,836 bytes. The caller must enforce at most `MAX_MARKERS` scored positions. +No Graph cache, kernel launch, cuBLAS workspace or inference is included. + +The ignored GPU check writes and reads all 17 buffers twice at three shapes, +including both capacity bounds, using deterministic byte patterns. This tests +allocation and transfer, not model numerics or latency. Prerequisites and commands +are in the +[native validation recipe](../../../recipe/laya/README.md#native-residency-and-workspace-validation). + ## Python worker The Python worker serves LAYA through laya-serve on CPU and Apple Silicon (PyTorch MPS, diff --git a/src/models/laya/src/lib.rs b/src/models/laya/src/lib.rs index fbc136f..a4be935 100644 --- a/src/models/laya/src/lib.rs +++ b/src/models/laya/src/lib.rs @@ -1,2 +1,4 @@ pub mod config; +pub mod resident; pub mod weights; +pub mod workspace; diff --git a/src/models/laya/src/resident.rs b/src/models/laya/src/resident.rs new file mode 100644 index 0000000..9c48854 --- /dev/null +++ b/src/models/laya/src/resident.rs @@ -0,0 +1,85 @@ +use crate::weights::{TensorSpec, Weights, checkpoint_tensors}; +use anyhow::{Context, Result}; +use omni_cuda::{Buffer, Cuda}; +use std::collections::HashMap; + +/// Checkpoint weights owned by one CUDA context; no inference workspace or tables. +pub struct ResidentWeights { + buffers: HashMap, + bytes: usize, +} + +impl ResidentWeights { + pub fn upload(cuda: &Cuda, source: &Weights) -> Result { + let tensors = checkpoint_tensors(); + source.validate_names(tensors.iter().map(|t| t.name.as_str()))?; + // Account for the legacy buffer, but do not keep unused calibration on GPU. + source.f32("temperature", &[3])?; + Self::upload_tensors( + cuda, + source, + tensors.iter().filter(|t| t.name != "temperature"), + ) + } + + fn upload_tensors<'a>( + cuda: &Cuda, + source: &Weights, + tensors: impl IntoIterator, + ) -> Result { + let mut resident = Self { + buffers: HashMap::new(), + bytes: 0, + }; + for spec in tensors { + let data = packed(source, spec)?; + let buffer = cuda + .upload(&data) + .with_context(|| format!("upload {}", spec.name))?; + resident.bytes += buffer.bytes(); + resident.buffers.insert(spec.name.clone(), buffer); + } + Ok(resident) + } + + pub fn get(&self, name: &str) -> Result<&Buffer> { + self.buffers + .get(name) + .with_context(|| format!("no resident weight: {name}")) + } + + /// Weight allocation bytes, excluding CUDA context and allocator overhead. + pub fn bytes(&self) -> usize { + self.bytes + } +} + +fn packed(source: &Weights, spec: &TensorSpec) -> Result> { + let name = spec.name.as_str(); + let shape = &spec.shape; + if name == "encoder.embeddings.tok_embeddings.weight" { + Ok(source + .f16(name, shape)? + .into_iter() + .flat_map(u16::to_le_bytes) + .collect()) + } else if (shape.len() == 1 && (name.starts_with("encoder.") || name.starts_with("head."))) + || name.starts_with("scorer.0.") + { + Ok(source + .f32(name, shape)? + .into_iter() + .flat_map(f32::to_le_bytes) + .collect()) + } else { + Ok(source + .bf16(name, shape)? + .into_iter() + .flat_map(u16::to_le_bytes) + .collect()) + } +} + +#[cfg(test)] +#[path = "../../../../tests/laya/unit/resident.rs"] +mod tests; diff --git a/src/models/laya/src/workspace.rs b/src/models/laya/src/workspace.rs new file mode 100644 index 0000000..da3c0f1 --- /dev/null +++ b/src/models/laya/src/workspace.rs @@ -0,0 +1,101 @@ +use anyhow::{Result, ensure}; +use omni_cuda::{Buffer, Cuda}; + +pub const MAX_MARKERS: usize = 2048; +const D: usize = 1024; + +/// Fixed-shape scratch allocations. Contents are uninitialized until written. +pub struct Workspace { + batch: usize, + sequence: usize, + bytes: usize, + buffers: WorkspaceBuffers, +} + +/// Buffer layouts consumed by Laya's encoder, decision head and scorer. +/// Access through `Workspace::buffers` keeps allocations fixed for its lifetime. +pub struct WorkspaceBuffers { + pub ids: Buffer, + pub lengths: Buffer, + pub types: Buffer, + pub residual: Buffer, + pub hidden: Buffer, + pub qkv: Buffer, + pub attention: Buffer, + pub gated: Buffer, + pub feed_forward: Buffer, + pub indices: Buffer, + pub offsets: Buffer, + pub markers: Buffer, + pub scored: Buffer, + pub logits: Buffer, + pub features: Buffer, + pub action_hidden: Buffer, + pub actions: Buffer, +} + +impl Workspace { + pub fn new(cuda: &Cuda, batch: usize, sequence: usize) -> Result { + ensure!( + batch.is_power_of_two() && batch <= 16, + "workspace batch must be 1, 2, 4, 8 or 16" + ); + ensure!( + (16..=512).contains(&sequence) && sequence.is_multiple_of(16), + "workspace sequence must be a multiple of 16 in 16..=512" + ); + let tokens = batch * sequence; + let mut bytes = 0; + let mut alloc = |size| { + let buffer = cuda.alloc(size)?; + bytes += size; + Ok::<_, anyhow::Error>(buffer) + }; + let buffers = WorkspaceBuffers { + ids: alloc(tokens * 8)?, + lengths: alloc(batch * 4)?, + types: alloc(batch * 8)?, + residual: alloc(tokens * D * 4)?, + hidden: alloc(tokens * D * 2)?, + qkv: alloc(tokens * D * 6)?, + attention: alloc(tokens * D * 2)?, + gated: alloc(tokens * 2624 * 2)?, + feed_forward: alloc(tokens * 4096 * 2)?, + indices: alloc(MAX_MARKERS * 4)?, + offsets: alloc((batch + 1) * 4)?, + markers: alloc(MAX_MARKERS * D * 2)?, + scored: alloc(MAX_MARKERS * D * 2)?, + logits: alloc(MAX_MARKERS * 2)?, + features: alloc(batch * 1028 * 2)?, + action_hidden: alloc(batch * 256 * 2)?, + actions: alloc(batch * 2 * 2)?, + }; + Ok(Self { + batch, + sequence, + bytes, + buffers, + }) + } + + pub fn batch(&self) -> usize { + self.batch + } + + pub fn sequence(&self) -> usize { + self.sequence + } + + /// Total scratch allocation bytes, excluding weights and CUDA overhead. + pub fn bytes(&self) -> usize { + self.bytes + } + + pub fn buffers(&self) -> &WorkspaceBuffers { + &self.buffers + } +} + +#[cfg(test)] +#[path = "../../../../tests/laya/unit/workspace.rs"] +mod tests; diff --git a/tests/backends/cuda/fixtures/runtime.c b/tests/backends/cuda/fixtures/runtime.c new file mode 100644 index 0000000..b37eb61 --- /dev/null +++ b/tests/backends/cuda/fixtures/runtime.c @@ -0,0 +1,116 @@ +#include +#include +#include +#include +#include + +#ifndef LAYA_TEST_ABI +#define LAYA_TEST_ABI 1 +#endif + +enum { + invalid_argument = 1000, + mode_create_error = 1, + mode_copy_error = 2, + mode_sync_error = 3, + mode_copy_and_sync_error = 4, + mode_alloc_error = 5 +}; + +typedef struct { int device, live; } Stream; +typedef struct { int device; unsigned char data[]; } Allocation; +static Stream streams[16]; +static int device = -1, mode, next_stream, live; +static int allocations_before_failure = -1; +static char trace[4096], error[64]; +static size_t trace_len, pending_bytes; +static void *pending_dst; +static const void *pending_src; + +static void record(const char *name) { + trace_len += (size_t)snprintf(trace + trace_len, sizeof(trace) - trace_len, + "%s%d ", name, device); +} +static Allocation *allocation(void *p) { + return (Allocation *)((unsigned char *)p - offsetof(Allocation, data)); +} +static int valid_stream(void *p) { + Stream *s = p; + return s && s->live && s->device == device; +} +void laya_test_mode(int value) { mode = value; } +void laya_test_fail_alloc_after(int count) { allocations_before_failure = count; } +const char *laya_test_trace(void) { return trace; } +int laya_test_live(void) { return live; } +uint32_t laya_abi_version(void) { return LAYA_TEST_ABI; } +const char *laya_error_string(int code) { + snprintf(error, sizeof(error), "fixture-error-%d", code); + return error; +} +int laya_set_device(int value) { + if (value < 0) return 11; + device = value; + record("device"); + return 0; +} +int laya_stream_create(void **out) { + record("create"); + if (mode == mode_create_error) return 23; + Stream *s = &streams[next_stream++]; + *s = (Stream){device, 1}; + *out = s; + live++; + return 0; +} +int laya_alloc(void **out, size_t bytes) { + record("alloc"); + if (mode == mode_alloc_error || allocations_before_failure == 0) return 31; + if (allocations_before_failure > 0) allocations_before_failure--; + Allocation *a = calloc(1, sizeof(*a) + bytes); + if (!a) return 32; + a->device = device; + *out = a->data; + live++; + return 0; +} +int laya_free(void *p) { + if (!p) return invalid_argument; + record("free"); + if (allocation(p)->device != device || pending_bytes) return 91; + free(allocation(p)); + live--; + return 0; +} +static int copy(void *dst, const void *src, size_t bytes, void *stream) { + if (!valid_stream(stream)) return 91; + pending_dst = dst; + pending_src = src; + pending_bytes = bytes; + return mode == mode_copy_error || mode == mode_copy_and_sync_error ? 41 : 0; +} +int laya_upload(void *dst, const unsigned char *src, size_t bytes, void *stream) { + record("upload"); + if (allocation(dst)->device != device) return 91; + return copy(dst, src, bytes, stream); +} +#ifndef LAYA_TEST_NO_DOWNLOAD +int laya_download(unsigned char *dst, void *src, size_t bytes, void *stream) { + record("download"); + if (allocation(src)->device != device) return 91; + return copy(dst, src, bytes, stream); +} +#endif +int laya_sync(void *stream) { + record("sync"); + if (!valid_stream(stream)) return 91; + if (pending_bytes) memcpy(pending_dst, pending_src, pending_bytes); + pending_bytes = 0; + return mode == mode_sync_error || mode == mode_copy_and_sync_error ? 42 : 0; +} +int laya_stream_free(void *stream) { + record("destroy"); + if (!valid_stream(stream) || pending_bytes) return 91; + ((Stream *)stream)->live = 0; + live--; + return 0; +} diff --git a/tests/backends/cuda/runtime.rs b/tests/backends/cuda/runtime.rs new file mode 100644 index 0000000..ab241d4 --- /dev/null +++ b/tests/backends/cuda/runtime.rs @@ -0,0 +1,257 @@ +#![cfg(unix)] + +use libloading::Library; +use omni_cuda::Cuda; +use std::{ffi::CStr, path::PathBuf, process::Command}; +use tempfile::{TempDir, tempdir}; + +const MODE_NORMAL: i32 = 0; +const MODE_CREATE_ERROR: i32 = 1; +const MODE_COPY_ERROR: i32 = 2; +const MODE_SYNC_ERROR: i32 = 3; +const MODE_COPY_AND_SYNC_ERROR: i32 = 4; +const MODE_ALLOC_ERROR: i32 = 5; + +struct Fixture { + _dir: TempDir, + path: PathBuf, + library: Option, +} +impl Fixture { + fn new(defines: &[&str]) -> Self { + let dir = tempdir().unwrap(); + let path = dir + .path() + .join(format!("runtime{}", std::env::consts::DLL_SUFFIX)); + let mut command = Command::new("cc"); + command.arg(if cfg!(target_os = "macos") { + "-dynamiclib" + } else { + "-shared" + }); + let result = command + .args(["-fPIC", "-std=c11", "-Wall", "-Wextra", "-Werror"]) + .args(defines) + .arg( + PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .join("../../../tests/backends/cuda/fixtures/runtime.c"), + ) + .arg("-o") + .arg(&path) + .output() + .unwrap(); + assert!( + result.status.success(), + "{}", + String::from_utf8_lossy(&result.stderr) + ); + let library = Some(unsafe { Library::new(&path) }.unwrap()); + Self { + _dir: dir, + path, + library, + } + } + fn mode(&self, mode: i32) { + unsafe { + self.library + .as_ref() + .unwrap() + .get::(b"laya_test_mode\0") + .unwrap()(mode); + } + } + fn trace(&self) -> String { + unsafe { + let f = self + .library + .as_ref() + .unwrap() + .get:: *const std::ffi::c_char>(b"laya_test_trace\0") + .unwrap(); + CStr::from_ptr(f()).to_string_lossy().into_owned() + } + } + fn live(&self) -> i32 { + unsafe { + self.library + .as_ref() + .unwrap() + .get:: i32>(b"laya_test_live\0") + .unwrap()() + } + } + fn load(&self, device: i32) -> anyhow::Result { + unsafe { Cuda::load(&self.path, device) } + } +} + +#[test] +fn loading_rejects_incompatible_libraries_before_creating_resources() { + for defines in [vec!["-DLAYA_TEST_ABI=2"], vec!["-DLAYA_TEST_NO_DOWNLOAD"]] { + let fixture = Fixture::new(&defines); + assert!(fixture.load(0).is_err()); + assert_eq!(fixture.live(), 0); + assert!(!fixture.trace().contains("create")); + } + let fixture = Fixture::new(&[]); + assert!(fixture.load(-1).is_err()); + fixture.mode(MODE_CREATE_ERROR); + let error = fixture.load(0).err().unwrap().to_string(); + assert!(error.contains("23"), "{error}"); + assert_eq!(fixture.live(), 0); + assert!(!fixture.trace().contains("destroy")); +} + +#[test] +fn copies_round_trip_and_bounds_are_checked() { + let fixture = Fixture::new(&[]); + let cuda = fixture.load(0).unwrap(); + assert!(cuda.alloc(0).is_err()); + let buffer = cuda.upload(&[1, 2, 3, 4]).unwrap(); + assert_eq!(buffer.bytes(), 4); + assert_eq!(buffer.read(4).unwrap(), [1, 2, 3, 4]); + buffer.write(&[9, 8]).unwrap(); + assert_eq!(buffer.read(4).unwrap(), [9, 8, 3, 4]); + buffer.write(&[]).unwrap(); + assert!(buffer.read(0).unwrap().is_empty()); + let before = fixture.trace(); + assert!(buffer.write(&[0; 5]).is_err()); + assert!(buffer.read(5).is_err()); + assert_eq!(fixture.trace(), before); + drop(buffer); + drop(cuda); + assert_eq!(fixture.live(), 0); +} + +#[test] +fn copy_errors_still_synchronize_and_async_errors_are_reported() { + let fixture = Fixture::new(&[]); + let cuda = fixture.load(0).unwrap(); + let buffer = cuda.alloc(3).unwrap(); + fixture.mode(MODE_COPY_ERROR); + let error = buffer.write(&[4, 5, 6]).unwrap_err().to_string(); + assert!(error.contains("fixture-error-41"), "{error}"); + assert!(fixture.trace().contains("upload0 sync0 ")); + fixture.mode(MODE_NORMAL); + assert_eq!(buffer.read(3).unwrap(), [4, 5, 6]); + fixture.mode(MODE_COPY_ERROR); + assert!(buffer.read(3).unwrap_err().to_string().contains("41")); + assert!(fixture.trace().contains("download0 sync0 ")); + fixture.mode(MODE_SYNC_ERROR); + assert!( + buffer + .write(&[7, 8, 9]) + .unwrap_err() + .to_string() + .contains("42") + ); + assert!(buffer.read(3).unwrap_err().to_string().contains("42")); + assert!(cuda.sync().unwrap_err().to_string().contains("42")); + fixture.mode(MODE_COPY_AND_SYNC_ERROR); + assert!( + buffer + .write(&[1, 2, 3]) + .unwrap_err() + .to_string() + .contains("41") + ); + fixture.mode(MODE_NORMAL); + assert_eq!(buffer.read(3).unwrap(), [1, 2, 3]); + drop(buffer); + drop(cuda); + assert_eq!(fixture.live(), 0); +} + +#[test] +fn failed_allocations_and_uploads_release_partial_resources() { + let fixture = Fixture::new(&[]); + let cuda = fixture.load(0).unwrap(); + fixture.mode(MODE_ALLOC_ERROR); + assert!(cuda.alloc(4).err().unwrap().to_string().contains("31")); + assert_eq!(fixture.live(), 1); + fixture.mode(MODE_COPY_ERROR); + assert!( + cuda.upload(&[1, 2]) + .err() + .unwrap() + .to_string() + .contains("41") + ); + assert_eq!(fixture.live(), 1); + fixture.mode(MODE_NORMAL); + drop(cuda); + assert_eq!(fixture.live(), 0); +} + +#[test] +fn buffers_keep_the_library_and_stream_alive_after_cuda_is_dropped() { + let mut fixture = Fixture::new(&[]); + let cuda = fixture.load(0).unwrap(); + let cloned = cuda.clone(); + let buffer = cuda.upload(&[7, 8, 9]).unwrap(); + // Remove the test's dlopen handle too: only the buffer may keep this library alive. + drop(fixture.library.take()); + drop(cuda); + drop(cloned); + assert_eq!(buffer.read(3).unwrap(), [7, 8, 9]); + fixture.library = Some(unsafe { Library::new(&fixture.path) }.unwrap()); + assert_eq!(fixture.live(), 2); + assert!(!fixture.trace().contains("destroy")); + drop(buffer); + assert_eq!(fixture.live(), 0); + let trace = fixture.trace(); + assert!(trace.rfind("sync0 ").unwrap() < trace.rfind("destroy0 ").unwrap()); + assert!(trace.rfind("free0 ").unwrap() < trace.rfind("destroy0 ").unwrap()); +} + +#[test] +fn operations_and_drops_restore_the_owning_device() { + let fixture = Fixture::new(&[]); + let cuda0 = fixture.load(0).unwrap(); + let buffer0 = cuda0.upload(&[1, 2]).unwrap(); + let cuda1 = fixture.load(1).unwrap(); + let buffer1 = cuda1.upload(&[3, 4]).unwrap(); + assert_eq!(buffer0.read(2).unwrap(), [1, 2]); + buffer1.write(&[5, 6]).unwrap(); + cuda0.sync().unwrap(); + drop(buffer1); + drop(cuda1); + drop(buffer0); + drop(cuda0); + assert_eq!(fixture.live(), 0); + let trace = fixture.trace(); + for event in [ + "device0 download0", + "device1 upload1", + "device0 sync0", + "device1 sync1 free1", + "device0 sync0 free0", + ] { + assert!(trace.contains(event), "missing {event}: {trace}"); + } +} + +#[test] +#[ignore = "requires an approved GPU and LAYA_CUDA_LIBRARY plus LAYA_CUDA_DEVICE"] +fn real_gpu_round_trip() { + let path = PathBuf::from(std::env::var_os("LAYA_CUDA_LIBRARY").expect("set LAYA_CUDA_LIBRARY")); + let device = std::env::var("LAYA_CUDA_DEVICE") + .expect("set LAYA_CUDA_DEVICE") + .parse() + .unwrap(); + { + let library = unsafe { Library::new(&path) }.unwrap(); + let free = unsafe { + library.get:: i32>(b"laya_free\0") + } + .unwrap(); + assert_eq!(unsafe { free(std::ptr::null_mut()) }, 1000); + } + let cuda = unsafe { Cuda::load(&path, device) }.unwrap(); + let expected: Vec = (0..4096).map(|i| (i % 251) as u8).collect(); + let buffer = cuda.upload(&expected).unwrap(); + assert_eq!(buffer.read(expected.len()).unwrap(), expected); + drop(cuda); + assert_eq!(buffer.read(expected.len()).unwrap(), expected); +} diff --git a/tests/laya/unit/resident.rs b/tests/laya/unit/resident.rs new file mode 100644 index 0000000..92bf8c0 --- /dev/null +++ b/tests/laya/unit/resident.rs @@ -0,0 +1,199 @@ +use super::*; +use libloading::Library; +use safetensors::{Dtype, tensor::TensorView}; +use sha2::{Digest, Sha256}; +use std::{fs, path::PathBuf, process::Command}; +use tempfile::{TempDir, tempdir}; + +const EMBED: &str = "encoder.embeddings.tok_embeddings.weight"; + +struct Fixture { + dir: TempDir, + library: Library, + cuda: Cuda, +} + +impl Fixture { + fn new() -> Self { + let dir = tempdir().unwrap(); + let path = dir + .path() + .join(format!("runtime{}", std::env::consts::DLL_SUFFIX)); + let output = Command::new("cc") + .args([ + if cfg!(target_os = "macos") { + "-dynamiclib" + } else { + "-shared" + }, + "-fPIC", + ]) + .arg( + PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .join("../../../tests/backends/cuda/fixtures/runtime.c"), + ) + .arg("-o") + .arg(&path) + .output() + .unwrap(); + assert!( + output.status.success(), + "{}", + String::from_utf8_lossy(&output.stderr) + ); + let library = unsafe { Library::new(&path) }.unwrap(); + let cuda = unsafe { Cuda::load(&path, 0) }.unwrap(); + Self { dir, library, cuda } + } + + fn live(&self) -> i32 { + unsafe { + self.library + .get:: i32>(b"laya_test_live\0") + .unwrap()() + } + } + + fn source(&self, specs: &[TensorSpec]) -> Weights { + let values: Vec = [1.0f32, -2.0] + .into_iter() + .flat_map(f32::to_le_bytes) + .collect(); + let tensors: Vec<_> = specs + .iter() + .map(|s| { + ( + s.name.as_str(), + TensorView::new(Dtype::F32, s.shape.clone(), &values).unwrap(), + ) + }) + .collect(); + let path = self.dir.path().join("weights.safetensors"); + safetensors::tensor::serialize_to_file(tensors, None, &path).unwrap(); + Weights::open(&path).unwrap() + } +} + +fn spec(name: &str, shape: &[usize]) -> TensorSpec { + TensorSpec { + name: name.into(), + shape: shape.into(), + } +} + +#[test] +fn precision_and_residency_preserve_expected_bytes() { + let fixture = Fixture::new(); + let specs = [ + spec(EMBED, &[1, 2]), + spec("encoder.layers.0.mlp_norm.weight", &[2]), + spec("encoder.layers.0.attn.Wqkv.weight", &[1, 2]), + spec("head.layers.0.norm1.weight", &[2]), + spec("head.layers.0.self_attn.in_proj_bias", &[2]), + spec("head.layers.0.self_attn.in_proj_weight", &[1, 2]), + spec("scorer.0.bias", &[2]), + spec("scorer.1.bias", &[2]), + spec("act_head.0.bias", &[2]), + spec("type_emb.weight", &[1, 2]), + ]; + let source = fixture.source(&specs); + let weights = ResidentWeights::upload_tensors(&fixture.cuda, &source, &specs).unwrap(); + drop(source); + let f16 = vec![0x00, 0x3c, 0x00, 0xc0]; + let bf16 = vec![0x80, 0x3f, 0x00, 0xc0]; + let f32 = vec![0x00, 0x00, 0x80, 0x3f, 0x00, 0x00, 0x00, 0xc0]; + let expected = [ + &f16, &f32, &bf16, &f32, &f32, &bf16, &f32, &bf16, &bf16, &bf16, + ]; + for (s, bytes) in specs.iter().zip(expected) { + let buffer = weights.get(&s.name).unwrap(); + assert_eq!(&buffer.read(buffer.bytes()).unwrap(), bytes, "{}", s.name); + } + assert_eq!(weights.bytes(), 56); + assert!(weights.get("missing").is_err()); + assert_eq!(fixture.live(), 11); + drop(fixture.cuda); + assert_eq!(weights.get(EMBED).unwrap().read(4).unwrap(), f16); + drop(weights); + assert_eq!( + unsafe { + fixture + .library + .get:: i32>(b"laya_test_live\0") + .unwrap()() + }, + 0 + ); +} + +#[test] +fn invalid_checkpoint_and_partial_load_release_allocations() { + let fixture = Fixture::new(); + let mut specs = [ + spec(EMBED, &[1, 2]), + spec("head.layers.0.norm1.weight", &[2]), + ]; + let source = fixture.source(&specs); + assert!(ResidentWeights::upload(&fixture.cuda, &source).is_err()); + assert_eq!(fixture.live(), 1); + specs[1].shape = vec![3]; + let error = ResidentWeights::upload_tensors(&fixture.cuda, &source, &specs) + .err() + .unwrap(); + assert!(error.to_string().contains("head.layers.0.norm1.weight")); + assert_eq!(fixture.live(), 1); + specs[1].shape = vec![2]; + unsafe { + fixture + .library + .get::(b"laya_test_mode\0") + .unwrap()(2); + } + assert!(ResidentWeights::upload_tensors(&fixture.cuda, &source, &specs).is_err()); + assert_eq!(fixture.live(), 1); +} + +#[test] +#[ignore = "requires approved GPU, LAYA_CUDA_LIBRARY, LAYA_CUDA_DEVICE, LAYA_CHECKPOINT and LAYA_WEIGHT_ORACLE"] +fn real_checkpoint_residency_matches_torch() { + let library = PathBuf::from(std::env::var_os("LAYA_CUDA_LIBRARY").unwrap()); + let device = std::env::var("LAYA_CUDA_DEVICE").unwrap().parse().unwrap(); + let checkpoint = PathBuf::from(std::env::var_os("LAYA_CHECKPOINT").unwrap()); + let rows: Vec = + serde_json::from_slice(&fs::read(std::env::var_os("LAYA_WEIGHT_ORACLE").unwrap()).unwrap()) + .unwrap(); + let cuda = unsafe { Cuda::load(&library, device) }.unwrap(); + let source = Weights::open(&checkpoint.join("model.safetensors")).unwrap(); + let weights = ResidentWeights::upload(&cuda, &source).unwrap(); + drop(source); + drop(cuda); + assert_eq!(rows.len(), 206); + assert_eq!(weights.buffers.len(), 205); + assert!(weights.get("temperature").is_err()); + let mut names = std::collections::HashSet::new(); + let mut bytes = 0; + for row in rows { + let name = row["name"].as_str().unwrap(); + assert!(names.insert(name.to_owned())); + if name == "temperature" { + continue; + } + let buffer = weights.get(name).unwrap(); + let hash = format!("{:x}", Sha256::digest(buffer.read(buffer.bytes()).unwrap())); + // The oracle contains independent Torch conversions at all three precisions. + let shape = row["shape"].as_array().unwrap(); + let dtype = if name == EMBED { + "f16" + } else if name.starts_with("scorer.0.") + || (shape.len() == 1 && (name.starts_with("encoder.") || name.starts_with("head."))) + { + "f32" + } else { + "bf16" + }; + assert_eq!(hash, row[dtype].as_str().unwrap(), "{name} {dtype}"); + bytes += buffer.bytes(); + } + assert_eq!(weights.bytes(), bytes); + println!("verified_resident_tensors=205 weight_allocation_bytes={bytes}"); +} diff --git a/tests/laya/unit/workspace.rs b/tests/laya/unit/workspace.rs new file mode 100644 index 0000000..c4f8e18 --- /dev/null +++ b/tests/laya/unit/workspace.rs @@ -0,0 +1,216 @@ +use super::*; +use libloading::Library; +use std::{path::PathBuf, process::Command}; +use tempfile::{TempDir, tempdir}; + +struct Fixture { + _dir: TempDir, + library: Library, + cuda: Cuda, +} +impl Fixture { + fn new() -> Self { + let dir = tempdir().unwrap(); + let path = dir + .path() + .join(format!("runtime{}", std::env::consts::DLL_SUFFIX)); + let output = Command::new("cc") + .args([ + if cfg!(target_os = "macos") { + "-dynamiclib" + } else { + "-shared" + }, + "-fPIC", + ]) + .arg( + PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .join("../../../tests/backends/cuda/fixtures/runtime.c"), + ) + .arg("-o") + .arg(&path) + .output() + .unwrap(); + assert!( + output.status.success(), + "{}", + String::from_utf8_lossy(&output.stderr) + ); + let library = unsafe { Library::new(&path) }.unwrap(); + let cuda = unsafe { Cuda::load(&path, 0) }.unwrap(); + Self { + _dir: dir, + library, + cuda, + } + } + fn live(&self) -> i32 { + live(&self.library) + } + fn trace(&self) -> String { + unsafe { + let trace = self + .library + .get:: *const std::ffi::c_char>(b"laya_test_trace\0") + .unwrap(); + std::ffi::CStr::from_ptr(trace()) + .to_string_lossy() + .into_owned() + } + } +} +fn live(library: &Library) -> i32 { + unsafe { + library + .get:: i32>(b"laya_test_live\0") + .unwrap()() + } +} +fn buffers(s: &Workspace) -> [&Buffer; 17] { + let s = s.buffers(); + [ + &s.ids, + &s.lengths, + &s.types, + &s.residual, + &s.hidden, + &s.qkv, + &s.attention, + &s.gated, + &s.feed_forward, + &s.indices, + &s.offsets, + &s.markers, + &s.scored, + &s.logits, + &s.features, + &s.action_hidden, + &s.actions, + ] +} + +#[test] +fn capacities_match_consumer_layouts_at_both_bounds() { + let fixture = Fixture::new(); + for (batch, sequence, expected) in [ + ( + 1, + 16, + [ + 128, 4, 8, 65536, 32768, 98304, 32768, 83968, 131072, 8192, 8, 4194304, 4194304, + 4096, 2056, 512, 4, + ], + ), + ( + 16, + 512, + [ + 65536, 64, 128, 33554432, 16777216, 50331648, 16777216, 42991616, 67108864, 8192, + 68, 4194304, 4194304, 4096, 32896, 8192, 64, + ], + ), + ] { + let workspace = Workspace::new(&fixture.cuda, batch, sequence).unwrap(); + assert_eq!(workspace.batch(), batch); + assert_eq!(workspace.sequence(), sequence); + assert_eq!(buffers(&workspace).map(Buffer::bytes), expected); + assert_eq!(workspace.bytes(), expected.iter().sum::()); + assert_eq!(fixture.live(), 18); + drop(workspace); + assert_eq!(fixture.live(), 1); + } +} + +#[test] +fn invalid_shapes_allocate_nothing() { + let fixture = Fixture::new(); + let before = fixture.trace(); + for batch in [0, 3, 17, 32, usize::MAX] { + assert!(Workspace::new(&fixture.cuda, batch, 16).is_err()); + } + for sequence in [0, 1, 15, 17, 511, 513, usize::MAX] { + assert!(Workspace::new(&fixture.cuda, 1, sequence).is_err()); + } + assert_eq!(fixture.live(), 1); + assert_eq!(fixture.trace(), before); +} + +#[test] +fn failed_allocations_release_the_partial_workspace() { + for completed in [0, 1, 8, 16] { + let fixture = Fixture::new(); + unsafe { + fixture + .library + .get::(b"laya_test_fail_alloc_after\0") + .unwrap()(completed); + } + assert!(Workspace::new(&fixture.cuda, 1, 16).is_err()); + assert_eq!(fixture.live(), 1, "failed after {completed} allocations"); + } +} + +#[test] +fn workspaces_do_not_alias_and_outlive_the_cuda_handle() { + let fixture = Fixture::new(); + let first = Workspace::new(&fixture.cuda, 1, 16).unwrap(); + let second = Workspace::new(&fixture.cuda, 2, 32).unwrap(); + assert_eq!(fixture.live(), 35); + drop(fixture.cuda); + for (i, b) in buffers(&first) + .into_iter() + .chain(buffers(&second)) + .enumerate() + { + b.write(&vec![i as u8; b.bytes().min(32)]).unwrap(); + } + for (i, b) in buffers(&first) + .into_iter() + .chain(buffers(&second)) + .enumerate() + { + assert_eq!( + b.read(b.bytes().min(32)).unwrap(), + vec![i as u8; b.bytes().min(32)] + ); + } + drop(first); + assert_eq!(live(&fixture.library), 18); + drop(second); + assert_eq!(live(&fixture.library), 0); +} + +#[test] +#[ignore = "requires approved GPU, LAYA_CUDA_LIBRARY and LAYA_CUDA_DEVICE"] +fn real_gpu_workspace_capacity_and_reuse() { + let path = PathBuf::from(std::env::var_os("LAYA_CUDA_LIBRARY").unwrap()); + let device = std::env::var("LAYA_CUDA_DEVICE").unwrap().parse().unwrap(); + for (batch, sequence, expected_bytes) in + [(1, 16, 8848032), (1, 512, 22628896), (16, 512, 236048836)] + { + let cuda = unsafe { Cuda::load(&path, device) }.unwrap(); + let workspace = Workspace::new(&cuda, batch, sequence).unwrap(); + drop(cuda); + assert_eq!(workspace.bytes(), expected_bytes); + for pass in 0..2 { + let pattern = |size, index| { + (0..size) + .map(|offset| ((offset % 251 + index * 7 + pass * 89) % 256) as u8) + .collect::>() + }; + for (i, b) in buffers(&workspace).into_iter().enumerate() { + b.write(&pattern(b.bytes(), i)).unwrap(); + } + for (i, b) in buffers(&workspace).into_iter().enumerate() { + assert_eq!( + b.read(b.bytes()).unwrap(), + pattern(b.bytes(), i), + "buffer {i}, pass {pass}" + ); + } + } + println!( + "batch={batch} sequence={sequence} buffers=17 bytes={expected_bytes} verified_passes=2" + ); + } +}