diff --git a/src/backends/cuda/kernels/runtime.cu b/src/backends/cuda/kernels/runtime.cu index 185de53d..91723c5b 100644 --- a/src/backends/cuda/kernels/runtime.cu +++ b/src/backends/cuda/kernels/runtime.cu @@ -62,14 +62,26 @@ int laya_capture_begin(void* stream) { cudaStreamCaptureModeThreadLocal); } +// End capture and release temporary/partial graph ownership on every outcome. int laya_capture_end(void* stream, void** executable) { + if (!executable) return -1; + *executable = nullptr; + if (!stream) return -1; cudaGraph_t graph = nullptr; - auto e = cudaStreamEndCapture(static_cast(stream), &graph); - if (e != cudaSuccess) - return e; - e = cudaGraphInstantiate(reinterpret_cast(executable), graph, 0); - cudaGraphDestroy(graph); - return e; + cudaGraphExec_t created = nullptr; + auto status = cudaStreamEndCapture(static_cast(stream), &graph); + if (status == cudaSuccess && graph) + status = cudaGraphInstantiate(&created, graph, 0); + if (graph) { + auto destroyed = cudaGraphDestroy(graph); + if (status == cudaSuccess) status = destroyed; + } + if (status != cudaSuccess || !created) { + if (created) cudaGraphExecDestroy(created); + return status != cudaSuccess ? status : -1; + } + *executable = created; + return 0; } int laya_graph_run(void* executable, void* stream) { diff --git a/src/backends/cuda/src/lib.rs b/src/backends/cuda/src/lib.rs index 5343723c..a72f008e 100644 --- a/src/backends/cuda/src/lib.rs +++ b/src/backends/cuda/src/lib.rs @@ -2,17 +2,26 @@ use anyhow::{Result, anyhow, ensure}; use libloading::Library; use std::{ + cell::Cell, ffi::{CStr, c_void}, path::Path, rc::Rc, }; pub type Ptr = *mut c_void; -type Kernel = unsafe extern "C" fn(*mut Ptr, i32, i32, i32, Ptr) -> i32; +type Launch = unsafe extern "C" fn(*mut Ptr, i32, i32, i32, Ptr) -> i32; struct Context { lib: Library, stream: Ptr, + capturing: Cell, } impl Context { + fn ensure_not_capturing(&self) -> Result<()> { + ensure!( + !self.capturing.get(), + "operation prohibited during CUDA capture" + ); + Ok(()) + } fn check(&self, code: i32) -> Result<()> { if code == 0 { return Ok(()); @@ -31,6 +40,7 @@ impl Context { unsafe { Ok(*self.lib.get::(name)?) } } fn sync(&self) -> Result<()> { + self.ensure_not_capturing()?; let f = self.symbol:: i32>(b"laya_sync\0")?; self.check(unsafe { f(self.stream) }) } @@ -61,11 +71,16 @@ impl Cuda { let mut stream = std::ptr::null_mut(); let init = unsafe { lib.get:: i32>(b"laya_init\0") }?; let code = unsafe { init(&mut stream) }; - let ctx = Rc::new(Context { lib, stream }); + let ctx = Rc::new(Context { + lib, + stream, + capturing: Cell::new(false), + }); ctx.check(code)?; Ok(Self { ctx }) } pub fn alloc(&self, bytes: usize) -> Result { + self.ctx.ensure_not_capturing()?; ensure!(bytes > 0, "zero CUDA allocation"); let f = self .ctx @@ -83,9 +98,59 @@ impl Cuda { b.write(bytes)?; Ok(b) } + /// Complete a prevalidated group with one synchronization before host borrows end. + /// A failed submission is also drained; this does not overlap copies and compute. + pub fn write_many(&self, writes: &[(&Buffer, &[u8])]) -> Result<()> { + self.ctx.ensure_not_capturing()?; + for (buffer, bytes) in writes { + ensure!( + Rc::ptr_eq(&self.ctx, &buffer.ctx), + "upload buffer belongs to another CUDA context" + ); + ensure!(bytes.len() <= buffer.bytes, "upload exceeds allocation"); + } + if writes.iter().all(|(_, bytes)| bytes.is_empty()) { + return Ok(()); + } + let upload = self + .ctx + .symbol:: i32>(b"laya_upload\0")?; + let sync = self + .ctx + .symbol:: i32>(b"laya_sync\0")?; + let mut copied = 0; + for (buffer, bytes) in writes { + if bytes.is_empty() { + continue; + } + copied = unsafe { upload(buffer.p, bytes.as_ptr(), bytes.len(), self.ctx.stream) }; + if copied != 0 { + break; + } + } + // Resolve both entry points before submission and keep every source borrowed + // through the drain even when a copy reports an error after queuing work. + let synced = unsafe { sync(self.ctx.stream) }; + match (self.ctx.check(copied), self.ctx.check(synced)) { + (Err(copy), Err(sync)) => Err(anyhow!( + "{copy}; stream synchronization also failed: {sync}" + )), + (Err(e), _) | (_, Err(e)) => Err(e), + (Ok(()), Ok(())) => Ok(()), + } + } pub fn sync(&self) -> Result<()> { self.ctx.sync() } + /// Resolve a kernel once while retaining the owning runtime and native code. + pub fn resolve(&self, name: &str) -> Result { + Ok(Kernel { + ctx: self.ctx.clone(), + launch: self + .ctx + .symbol::(format!("laya_{name}\0").as_bytes())?, + }) + } /// # Safety /// Tensor shape, dtype, layout, aliasing and allocation sizes must match the generated kernel. /// Buffers must belong to this context and stay alive until synchronization or graph destruction. @@ -96,7 +161,7 @@ impl Cuda { ); let k = self .ctx - .symbol::(format!("laya_{name}\0").as_bytes())?; + .symbol::(format!("laya_{name}\0").as_bytes())?; self.ctx.check(unsafe { k( args.as_ptr() as *mut Ptr, @@ -108,36 +173,31 @@ impl Cuda { }) } /// # Safety - /// Every allocation referenced by `work` must outlive the returned graph. No allocation or copy - /// that may synchronize is permitted inside work. This context is confined to one OS thread. + /// Every allocation and kernel referenced by `work` must outlive the graph. + /// Work must only launch on this stream and cannot release resources or invoke + /// capture-incompatible CUDA operations, including through raw FFI symbols. + /// This context is confined to one OS thread. Capture ends on errors or unwind. pub unsafe fn capture(&self, work: impl FnOnce() -> Result<()>) -> Result { - self.sync()?; + self.ctx.ensure_not_capturing()?; + // Resolve the complete optional API before beginning, including cleanup. + let functions = GraphFunctions { + end: self.ctx.symbol(b"laya_capture_end\0")?, + run: self.ctx.symbol(b"laya_graph_run\0")?, + free: self.ctx.symbol(b"laya_graph_free\0")?, + }; let begin = self .ctx .symbol:: i32>(b"laya_capture_begin\0")?; - let end = self - .ctx - .symbol:: i32>(b"laya_capture_end\0")?; + self.sync()?; self.ctx.check(unsafe { begin(self.ctx.stream) })?; - let result = work(); - let mut p = std::ptr::null_mut(); - let code = unsafe { end(self.ctx.stream, &mut p) }; - if result.is_err() || code != 0 { - if !p.is_null() { - let f = self - .ctx - .symbol:: i32>(b"laya_graph_free\0")?; - unsafe { - f(p); - } - } - result?; - self.ctx.check(code)?; - } - Ok(Graph { - ctx: self.ctx.clone(), - p, - }) + self.ctx.capturing.set(true); + let guard = CaptureGuard { + cuda: self, + functions, + active: true, + }; + work()?; + guard.finish() } /// # Safety /// Caller supplies the precise dimensions and allocation sizes required by this glue kernel. @@ -151,7 +211,7 @@ impl Cuda { ) -> Result<()> { let k = self .ctx - .symbol::(format!("laya_{name}\0").as_bytes())?; + .symbol::(format!("laya_{name}\0").as_bytes())?; self.ctx.check(unsafe { k( args.as_ptr() as *mut Ptr, @@ -187,15 +247,13 @@ impl Buffer { self.p } pub fn write(&self, bytes: &[u8]) -> Result<()> { - ensure!(bytes.len() <= self.bytes, "upload exceeds allocation"); - let f = self - .ctx - .symbol:: i32>(b"laya_upload\0")?; - self.ctx - .check(unsafe { f(self.p, bytes.as_ptr(), bytes.len(), self.ctx.stream) })?; - self.ctx.sync() + Cuda { + ctx: self.ctx.clone(), + } + .write_many(&[(self, bytes)]) } pub fn read(&self, bytes: usize) -> Result> { + self.ctx.ensure_not_capturing()?; ensure!(bytes <= self.bytes, "download exceeds allocation"); let mut data = vec![0; bytes]; let f = self @@ -220,28 +278,93 @@ impl Drop for Buffer { } } } +#[derive(Clone, Copy)] +struct GraphFunctions { + end: unsafe extern "C" fn(Ptr, *mut Ptr) -> i32, + run: unsafe extern "C" fn(Ptr, Ptr) -> i32, + free: unsafe extern "C" fn(Ptr) -> i32, +} pub struct Graph { ctx: Rc, + functions: GraphFunctions, p: Ptr, } impl Graph { pub fn replay(&self) -> Result<()> { - let f = self - .ctx - .symbol:: i32>(b"laya_graph_run\0")?; - self.ctx.check(unsafe { f(self.p, self.ctx.stream) }) + self.ctx.ensure_not_capturing()?; + self.ctx + .check(unsafe { (self.functions.run)(self.p, self.ctx.stream) }) } } impl Drop for Graph { fn drop(&mut self) { let _ = self.ctx.sync(); - if let Ok(f) = self - .ctx - .symbol:: i32>(b"laya_graph_free\0") - { + unsafe { + (self.functions.free)(self.p); + } + } +} +struct CaptureGuard<'a> { + cuda: &'a Cuda, + functions: GraphFunctions, + active: bool, +} +impl CaptureGuard<'_> { + fn finish(mut self) -> Result { + let mut p = std::ptr::null_mut(); + let code = unsafe { (self.functions.end)(self.cuda.ctx.stream, &mut p) }; + self.active = false; + self.cuda.ctx.capturing.set(false); + // Own partial handles before status validation, so failure also frees them. + let graph = (!p.is_null()).then(|| Graph { + ctx: self.cuda.ctx.clone(), + functions: self.functions, + p, + }); + self.cuda.ctx.check(code)?; + ensure!(graph.is_some(), "CUDA runtime returned a null graph"); + Ok(graph.unwrap()) + } +} +impl Drop for CaptureGuard<'_> { + fn drop(&mut self) { + if self.active { + let mut p = std::ptr::null_mut(); unsafe { - f(self.p); + (self.functions.end)(self.cuda.ctx.stream, &mut p); + if !p.is_null() { + (self.functions.free)(p); + } } + self.cuda.ctx.capturing.set(false); } } } + +/// A resolved entry point retaining its stream and native library. +#[derive(Clone)] +pub struct Kernel { + ctx: Rc, + launch: Launch, +} +impl Kernel { + /// # Safety + /// Shapes, dtype, layout, aliasing and pointer lifetimes must match this kernel. + /// Every pointer belongs to this context and stays alive through synchronization + /// or destruction of any graph that captures the launch. + pub unsafe fn launch(&self, args: &[Ptr], b: usize, l: usize) -> Result<()> { + ensure!( + b > 0 && b <= 16 && l > 0 && l <= 512 && l.is_multiple_of(16), + "invalid CUDA shape" + ); + self.ctx.check(unsafe { + (self.launch)( + args.as_ptr() as *mut Ptr, + b as i32, + l as i32, + (b * l) as i32, + self.ctx.stream, + ) + }) + } +} diff --git a/src/models/laya/src/model.rs b/src/models/laya/src/model.rs index ba25a769..54f6aa67 100644 --- a/src/models/laya/src/model.rs +++ b/src/models/laya/src/model.rs @@ -2,7 +2,7 @@ use crate::{config::Config, packing::Batch, weights::Weights}; use anyhow::{Context, Result, ensure}; use half::bf16; -use omni_cuda::{Buffer, Cuda, Graph, Ptr}; +use omni_cuda::{Buffer, Cuda, Graph, Kernel, Ptr}; use std::{ collections::{HashMap, VecDeque}, fs, @@ -32,9 +32,6 @@ fn bytes32(v: &[f32]) -> Vec { fn i32bytes(v: &[i32]) -> Vec { v.iter().flat_map(|x| x.to_le_bytes()).collect() } -fn i64bytes(v: &[i64]) -> Vec { - v.iter().flat_map(|x| x.to_le_bytes()).collect() -} fn decode_bf16(v: &[u8]) -> Vec { v.as_chunks::<2>() .0 @@ -120,6 +117,7 @@ struct Workspace { ids: Buffer, lens: Buffer, types: Buffer, + staging: Vec, x: Buffer, y: Buffer, qkv: Buffer, @@ -136,60 +134,431 @@ struct Workspace { actions: Buffer, } impl Workspace { - fn new(c: &Cuda, b: usize, l: usize) -> Result { + fn sizes(b: usize, l: usize) -> Result<[usize; 17]> { + ensure!( + b.is_power_of_two() && b <= 16 && (16..=512).contains(&l) && l.is_multiple_of(16), + "invalid workspace shape" + ); let m = b * l; - let mut bytes = 0; - let mut alloc = |n| { - bytes += n; - c.alloc(n) - }; + Ok([ + m * 8, + b * 4, + b * 8, + m * D * 4, + m * D * 2, + m * D * 6, + m * D * 2, + m * 2624 * 2, + m * 4096 * 2, + 2048 * 4, + 17 * 4, + 2048 * D * 2, + 2048 * D * 2, + 2048 * 2, + 16 * 1028 * 2, + 16 * 256 * 2, + 16 * 2 * 2, + ]) + } + fn required_bytes(b: usize, l: usize) -> Result { + Ok(Self::sizes(b, l)?.iter().sum()) + } + fn new(c: &Cuda, b: usize, l: usize) -> Result { + let sizes = Self::sizes(b, l)?; + let bytes = sizes.iter().sum(); + let mut sizes = sizes.into_iter(); + let mut alloc = || c.alloc(sizes.next().expect("workspace buffer layout")); Ok(Self { graph: None, b, l, - ids: alloc(m * 8)?, - lens: alloc(b * 4)?, - types: alloc(b * 8)?, - x: alloc(m * D * 4)?, - y: alloc(m * D * 2)?, - qkv: alloc(m * D * 6)?, - o: alloc(m * D * 2)?, - g: alloc(m * 2624 * 2)?, - ff: alloc(m * 4096 * 2)?, - indices: alloc(2048 * 4)?, - offsets: alloc(17 * 4)?, - markers: alloc(2048 * D * 2)?, - scored: alloc(2048 * D * 2)?, - logits: alloc(2048 * 2)?, - features: alloc(16 * 1028 * 2)?, - action_hidden: alloc(16 * 256 * 2)?, - actions: alloc(16 * 2 * 2)?, + ids: alloc()?, + lens: alloc()?, + types: alloc()?, + staging: vec![0; b * l * 8 + b * 12], + x: alloc()?, + y: alloc()?, + qkv: alloc()?, + o: alloc()?, + g: alloc()?, + ff: alloc()?, + indices: alloc()?, + offsets: alloc()?, + markers: alloc()?, + scored: alloc()?, + logits: alloc()?, + features: alloc()?, + action_hidden: alloc()?, + actions: alloc()?, bytes, }) } } +/// Limits retained workspaces. Either zero limit disables graph caching. +/// Bytes exclude weights, graph/driver overhead, host staging and one eager +/// fallback workspace used by oversized or disabled requests. +#[derive(Clone, Copy, Debug)] +pub struct CacheConfig { + pub max_shapes: usize, + pub max_bytes: usize, +} +impl Default for CacheConfig { + fn default() -> Self { + Self { + max_shapes: 4, + max_bytes: 512 * 1024 * 1024, + } + } +} + +#[derive(Clone, Copy)] +enum Scratch { + Ids, + Lens, + Types, + X, + Y, + Qkv, + O, + G, + Ff, +} +impl Scratch { + fn buffer(self, s: &Workspace) -> &Buffer { + match self { + Self::Ids => &s.ids, + Self::Lens => &s.lens, + Self::Types => &s.types, + Self::X => &s.x, + Self::Y => &s.y, + Self::Qkv => &s.qkv, + Self::O => &s.o, + Self::G => &s.g, + Self::Ff => &s.ff, + } + } +} +#[derive(Clone, Copy)] +enum Argument { + Weight(Ptr), + Scratch(Scratch), +} +impl Argument { + fn pointer(self, s: &Workspace) -> Ptr { + match self { + Self::Weight(p) => p, + Self::Scratch(slot) => slot.buffer(s).ptr(), + } + } +} +enum SelectedKernel { + Fixed(Kernel), + Attention([Option; 3]), +} +impl SelectedKernel { + fn resolve(cuda: &Cuda, name: &str) -> Result { + if name == "attn_full" || name == "attn_local" { + // Sparse trusted bundles may omit wrappers for shapes never requested. + // Preserve their eager usability and fail when a missing shape is used. + Ok(Self::Attention([ + Some(cuda.resolve(name)?), + cuda.resolve(&format!("{name}_b1_l512")).ok(), + cuda.resolve(&format!("{name}_b4_l512")).ok(), + ])) + } else { + Ok(Self::Fixed(cuda.resolve(name)?)) + } + } + fn select(&self, b: usize, l: usize) -> Result<&Kernel> { + match self { + Self::Fixed(k) => Ok(k), + Self::Attention(k) => k[if l == 512 && b == 1 { + 1 + } else if l == 512 && b == 4 { + 2 + } else { + 0 + }] + .as_ref() + .ok_or_else(|| { + anyhow::anyhow!("missing specialized attention kernel for B={b}, L={l}") + }), + } + } +} +enum Step { + Launch { + kernel: SelectedKernel, + args: [Argument; 5], + count: usize, + }, + Dump { + name: String, + slot: Scratch, + bf: bool, + }, +} + +fn prepare_encoder( + cuda: &Cuda, + weights: &HashMap, + original_rope: bool, +) -> Result> { + let steps = std::cell::RefCell::new(Vec::new()); + let weight = |name: &str| Argument::Weight(weights[name].ptr()); + let call = |name: &str, args: &[Argument]| -> Result<()> { + ensure!(args.len() <= 5, "encoder argument count"); + let mut resolved = [Argument::Weight(std::ptr::null_mut()); 5]; + resolved[..args.len()].copy_from_slice(args); + steps.borrow_mut().push(Step::Launch { + kernel: SelectedKernel::resolve(cuda, name)?, + args: resolved, + count: args.len(), + }); + Ok(()) + }; + let dump = |name: &str, slot: Scratch, bf: bool| { + steps.borrow_mut().push(Step::Dump { + name: name.to_owned(), + slot, + bf, + }); + Ok::<_, anyhow::Error>(()) + }; + let z = weight("zeros.1024"); + let attention = |label: &str| format!("attn_{label}"); + call( + "embed", + &[ + Argument::Scratch(Scratch::Ids), + weight("encoder.embeddings.tok_embeddings.weight"), + weight("encoder.embeddings.norm.weight"), + Argument::Scratch(Scratch::X), + Argument::Scratch(Scratch::Y), + ], + )?; + dump("embedding", Scratch::X, false)?; + for i in 0..28 { + let p = format!("encoder.layers.{i}"); + let w = |n: &str| weight(&format!("{p}.{n}")); + call( + "qkv", + &[ + Argument::Scratch(Scratch::Y), + w("attn.Wqkv.weight"), + weight("zeros.3072"), + Argument::Scratch(Scratch::Qkv), + ], + )?; + let kind = if i % 3 == 0 { "full" } else { "local" }; + call( + if original_rope { + "rope_original" + } else { + "rope" + }, + &[ + Argument::Scratch(Scratch::Qkv), + weight(&format!("rope_{kind}_cos")), + weight(&format!("rope_{kind}_sin")), + ], + )?; + call( + &attention(if i % 3 == 0 { "full" } else { "local" }), + &[ + Argument::Scratch(Scratch::Qkv), + Argument::Scratch(Scratch::Lens), + Argument::Scratch(Scratch::O), + ], + )?; + call( + "out", + &[ + Argument::Scratch(Scratch::O), + w("attn.Wo.weight"), + z, + Argument::Scratch(Scratch::Y), + ], + )?; + call( + "addln", + &[ + Argument::Scratch(Scratch::X), + Argument::Scratch(Scratch::Y), + w("mlp_norm.weight"), + z, + Argument::Scratch(Scratch::Y), + ], + )?; + call( + "geglu", + &[ + Argument::Scratch(Scratch::Y), + w("mlp.Wi.weight"), + Argument::Scratch(Scratch::G), + ], + )?; + call( + "down", + &[ + Argument::Scratch(Scratch::G), + w("mlp.Wo.weight"), + z, + Argument::Scratch(Scratch::Y), + ], + )?; + let next = if i < 27 { + weight(&format!("encoder.layers.{}.attn_norm.weight", i + 1)) + } else { + weight("encoder.final_norm.weight") + }; + call( + "addln", + &[ + Argument::Scratch(Scratch::X), + Argument::Scratch(Scratch::Y), + next, + z, + Argument::Scratch(Scratch::Y), + ], + )?; + if [0, 1, 2, 27].contains(&i) { + dump(&format!("encoder{i}_residual"), Scratch::X, false)?; + dump(&format!("encoder{i}_normalized"), Scratch::Y, true)?; + } + } + call( + "type", + &[ + Argument::Scratch(Scratch::Y), + weight("type_emb.weight"), + Argument::Scratch(Scratch::Types), + Argument::Scratch(Scratch::X), + ], + )?; + for i in 0..2 { + let p = format!("head.layers.{i}"); + let w = |n: &str| weight(&format!("{p}.{n}")); + call( + "ln_bias", + &[ + Argument::Scratch(Scratch::X), + Argument::Scratch(Scratch::Y), + w("norm1.weight"), + w("norm1.bias"), + Argument::Scratch(Scratch::Y), + ], + )?; + call( + "head_in", + &[ + Argument::Scratch(Scratch::Y), + w("self_attn.in_proj_weight"), + w("self_attn.in_proj_bias"), + Argument::Scratch(Scratch::Qkv), + ], + )?; + call( + &attention("full"), + &[ + Argument::Scratch(Scratch::Qkv), + Argument::Scratch(Scratch::Lens), + Argument::Scratch(Scratch::O), + ], + )?; + call( + "head_out", + &[ + Argument::Scratch(Scratch::O), + w("self_attn.out_proj.weight"), + w("self_attn.out_proj.bias"), + Argument::Scratch(Scratch::Y), + ], + )?; + call( + "addln_bias", + &[ + Argument::Scratch(Scratch::X), + Argument::Scratch(Scratch::Y), + w("norm2.weight"), + w("norm2.bias"), + Argument::Scratch(Scratch::Y), + ], + )?; + call( + "ffn1", + &[ + Argument::Scratch(Scratch::Y), + w("linear1.weight"), + w("linear1.bias"), + Argument::Scratch(Scratch::Ff), + ], + )?; + call( + "ffn2", + &[ + Argument::Scratch(Scratch::Ff), + w("linear2.weight"), + w("linear2.bias"), + Argument::Scratch(Scratch::Y), + ], + )?; + call( + "residual", + &[Argument::Scratch(Scratch::X), Argument::Scratch(Scratch::Y)], + )?; + } + Ok(steps.into_inner()) +} + pub struct Model { pub config: Config, cuda: Cuda, blas: Blas, weights: HashMap, + plan: Vec, cache: VecDeque, + cached_bytes: usize, + cache_config: CacheConfig, + eager: Option, graphs: bool, - original_rope: bool, } impl Drop for Model { fn drop(&mut self) { let _ = self.cuda.sync(); - self.cache.clear(); + self.clear_cache(); } } impl Model { + /// Release retained graph and eager workspaces while keeping this model loaded. + pub fn clear_cache(&mut self) { + self.cache.clear(); + self.cached_bytes = 0; + self.eager = None; + } + pub fn load( checkpoint: &Path, bundle: &Path, graphs: bool, original_rope: bool, + ) -> Result { + Self::load_with_cache( + checkpoint, + bundle, + graphs, + original_rope, + CacheConfig::default(), + ) + } + + pub fn load_with_cache( + checkpoint: &Path, + bundle: &Path, + graphs: bool, + original_rope: bool, + cache_config: CacheConfig, ) -> Result { let config = Config::load(checkpoint)?; crate::artifacts::validate_bundle(checkpoint, bundle)?; @@ -246,160 +615,41 @@ impl Model { weights.len() ); } + let plan = prepare_encoder(&cuda, &weights, original_rope)?; Ok(Self { config, cuda, blas, weights, + plan, cache: VecDeque::new(), + cached_bytes: 0, + cache_config, + eager: None, graphs, - original_rope, }) } fn w(&self, n: &str) -> &Buffer { &self.weights[n] } fn encode(&self, s: &Workspace) -> Result<()> { - let (b, l) = (s.b, s.l); - let z = self.w("zeros.1024").ptr(); - let attention = |label: &str| { - if l == 512 && (b == 1 || b == 4) { - format!("attn_{label}_b{b}_l512") - } else { - format!("attn_{label}") - } - }; - // All pointers refer to checked fixed-shape, resident allocations in this worker. - let call = |name: &str, args: &[Ptr]| unsafe { self.cuda.launch(name, args, b, l) }; - call( - "embed", - &[ - s.ids.ptr(), - self.w("encoder.embeddings.tok_embeddings.weight").ptr(), - self.w("encoder.embeddings.norm.weight").ptr(), - s.x.ptr(), - s.y.ptr(), - ], - )?; - self.dump("embedding", &s.x, false)?; - for i in 0..28 { - let p = format!("encoder.layers.{i}"); - let w = |n: &str| self.w(&format!("{p}.{n}")).ptr(); - call( - "qkv", - &[ - s.y.ptr(), - w("attn.Wqkv.weight"), - self.w("zeros.3072").ptr(), - s.qkv.ptr(), - ], - )?; - let kind = if i % 3 == 0 { "full" } else { "local" }; - call( - if self.original_rope { - "rope_original" - } else { - "rope" - }, - &[ - s.qkv.ptr(), - self.w(&format!("rope_{kind}_cos")).ptr(), - self.w(&format!("rope_{kind}_sin")).ptr(), - ], - )?; - call( - &attention(if i % 3 == 0 { "full" } else { "local" }), - &[s.qkv.ptr(), s.lens.ptr(), s.o.ptr()], - )?; - call("out", &[s.o.ptr(), w("attn.Wo.weight"), z, s.y.ptr()])?; - call( - "addln", - &[s.x.ptr(), s.y.ptr(), w("mlp_norm.weight"), z, s.y.ptr()], - )?; - call("geglu", &[s.y.ptr(), w("mlp.Wi.weight"), s.g.ptr()])?; - call("down", &[s.g.ptr(), w("mlp.Wo.weight"), z, s.y.ptr()])?; - let next = if i < 27 { - self.w(&format!("encoder.layers.{}.attn_norm.weight", i + 1)) - } else { - self.w("encoder.final_norm.weight") - }; - call("addln", &[s.x.ptr(), s.y.ptr(), next.ptr(), z, s.y.ptr()])?; - if [0, 1, 2, 27].contains(&i) { - self.dump(&format!("encoder{i}_residual"), &s.x, false)?; - self.dump(&format!("encoder{i}_normalized"), &s.y, true)?; + for step in &self.plan { + match step { + Step::Launch { + kernel, + args, + count, + } => { + let pointers = args.map(|arg| arg.pointer(s)); + unsafe { + kernel + .select(s.b, s.l)? + .launch(&pointers[..*count], s.b, s.l) + }?; + } + Step::Dump { name, slot, bf } => self.dump(name, slot.buffer(s), *bf)?, } } - call( - "type", - &[ - s.y.ptr(), - self.w("type_emb.weight").ptr(), - s.types.ptr(), - s.x.ptr(), - ], - )?; - for i in 0..2 { - let p = format!("head.layers.{i}"); - let w = |n: &str| self.w(&format!("{p}.{n}")).ptr(); - call( - "ln_bias", - &[ - s.x.ptr(), - s.y.ptr(), - w("norm1.weight"), - w("norm1.bias"), - s.y.ptr(), - ], - )?; - call( - "head_in", - &[ - s.y.ptr(), - w("self_attn.in_proj_weight"), - w("self_attn.in_proj_bias"), - s.qkv.ptr(), - ], - )?; - call(&attention("full"), &[s.qkv.ptr(), s.lens.ptr(), s.o.ptr()])?; - call( - "head_out", - &[ - s.o.ptr(), - w("self_attn.out_proj.weight"), - w("self_attn.out_proj.bias"), - s.y.ptr(), - ], - )?; - call( - "addln_bias", - &[ - s.x.ptr(), - s.y.ptr(), - w("norm2.weight"), - w("norm2.bias"), - s.y.ptr(), - ], - )?; - call( - "ffn1", - &[ - s.y.ptr(), - w("linear1.weight"), - w("linear1.bias"), - s.ff.ptr(), - ], - )?; - call( - "ffn2", - &[ - s.ff.ptr(), - w("linear2.weight"), - w("linear2.bias"), - s.y.ptr(), - ], - )?; - call("residual", &[s.x.ptr(), s.y.ptr()])?; - } Ok(()) } fn dump(&self, name: &str, buffer: &Buffer, bf: bool) -> Result<()> { @@ -457,19 +707,55 @@ impl Model { batch.markers.iter().map(Vec::len).sum::() <= 2048, "too many markers" ); - let found = self - .cache - .iter() - .position(|s| s.b == batch.b && s.l == batch.l); - let mut s = if let Some(i) = found { - self.cache.remove(i).unwrap() + let bytes = Workspace::required_bytes(batch.b, batch.l)?; + let cacheable = self.cache_config.max_shapes > 0 && bytes <= self.cache_config.max_bytes; + let mut s = if cacheable { + self.eager = None; + if let Some(i) = self + .cache + .iter() + .position(|s| s.b == batch.b && s.l == batch.l) + { + let workspace = self.cache.remove(i).unwrap(); + self.cached_bytes -= workspace.bytes; + workspace + } else { + // Budget the same layout used by the allocator and release LRU entries + // before allocating a miss, keeping retained admission within both limits. + while self.cache.len() >= self.cache_config.max_shapes + || self.cached_bytes > self.cache_config.max_bytes - bytes + { + let evicted = self.cache.pop_front().unwrap(); + self.cached_bytes -= evicted.bytes; + drop(evicted); + } + Workspace::new(&self.cuda, batch.b, batch.l)? + } } else { - Workspace::new(&self.cuda, batch.b, batch.l)? + let eager = self + .eager + .take() + .filter(|s| s.b == batch.b && s.l == batch.l); + // A different fallback shape is dropped before replacement allocation. + match eager { + Some(workspace) => workspace, + None => Workspace::new(&self.cuda, batch.b, batch.l)?, + } }; - s.ids.write(&i64bytes(&batch.input_ids))?; - s.lens.write(&i32bytes(&batch.lens))?; - s.types.write(&i64bytes(&batch.qtypes))?; - if self.graphs { + let (ids, rest) = s.staging.split_at_mut(batch.input_ids.len() * 8); + let (lens, types) = rest.split_at_mut(batch.lens.len() * 4); + for (bytes, value) in ids.as_chunks_mut::<8>().0.iter_mut().zip(&batch.input_ids) { + bytes.copy_from_slice(&value.to_le_bytes()); + } + for (bytes, value) in lens.as_chunks_mut::<4>().0.iter_mut().zip(&batch.lens) { + bytes.copy_from_slice(&value.to_le_bytes()); + } + for (bytes, value) in types.as_chunks_mut::<8>().0.iter_mut().zip(&batch.qtypes) { + bytes.copy_from_slice(&value.to_le_bytes()); + } + self.cuda + .write_many(&[(&s.ids, ids), (&s.lens, lens), (&s.types, types)])?; + if self.graphs && cacheable { if s.graph.is_none() { self.encode(&s)?; self.encode(&s)?; @@ -574,14 +860,12 @@ impl Model { .iter() .map(|x| [x[0], x[1]]) .collect(); - while self.cache.len() >= 4 - || self.cache.iter().map(|s| s.bytes).sum::() + s.bytes > 512 * 1024 * 1024 - { - if self.cache.pop_front().is_none() { - break; - } + if cacheable { + self.cached_bytes += s.bytes; + self.cache.push_back(s); + } else { + self.eager = Some(s); } - self.cache.push_back(s); Ok((logits, actions)) } }