diff --git a/src/backends/cuda/src/lib.rs b/src/backends/cuda/src/lib.rs index 5343723c..9d4235e8 100644 --- a/src/backends/cuda/src/lib.rs +++ b/src/backends/cuda/src/lib.rs @@ -7,7 +7,7 @@ use std::{ 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, @@ -86,6 +86,15 @@ impl Cuda { 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 +105,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, @@ -151,7 +160,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, @@ -245,3 +254,31 @@ impl Drop for Graph { } } } + +/// 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..5b404963 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, @@ -169,14 +169,314 @@ impl Workspace { } } +#[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, graphs: bool, - original_rope: bool, } impl Drop for Model { fn drop(&mut self) { @@ -246,159 +546,37 @@ impl Model { weights.len() ); } + let plan = prepare_encoder(&cuda, &weights, original_rope)?; Ok(Self { config, cuda, blas, weights, + plan, cache: VecDeque::new(), 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}") + 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)?, } - }; - // 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)?; - } - } - 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(()) }