From 774ab2d4db3c9f9042150abc468831eec9d80656 Mon Sep 17 00:00:00 2001 From: Alvaro Bartolome <36760800+alvarobartt@users.noreply.github.com> Date: Fri, 25 Sep 2026 07:34:52 +0000 Subject: [PATCH 1/6] Add support for `candle-flash-attn` and `candle-flash-attn-v3` --- Cargo.lock | 28 ++++ Cargo.toml | 6 +- build.rs | 45 ++++++ src/device.rs | 14 ++ src/lib.rs | 12 ++ src/main.rs | 28 +++- src/models/laya.rs | 161 ++++++++++++++++---- src/models/mod.rs | 42 +++++- src/models/modernbert.rs | 308 ++++++++++++++++++++++++++++++++++++--- 9 files changed, 589 insertions(+), 55 deletions(-) create mode 100644 build.rs diff --git a/Cargo.lock b/Cargo.lock index 2ef1fdb..6d72366 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -448,6 +448,32 @@ dependencies = [ "zip 8.6.0", ] +[[package]] +name = "candle-flash-attn" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e5f1e2f29f5123d7a627171209fdaeb4ee91d076aefeac588922aa9842e30c2c" +dependencies = [ + "anyhow", + "candle-core", + "cudaforge", + "half", +] + +[[package]] +name = "candle-flash-attn-v3" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5280d06b34f1c9370db620d6d08ea22577a5e26edb7665e6393c6bf5a4900a9f" +dependencies = [ + "anyhow", + "candle-core", + "cudaforge", + "half", + "num_cpus", + "rayon", +] + [[package]] name = "candle-kernels" version = "0.11.0" @@ -4010,6 +4036,8 @@ dependencies = [ "anyhow", "axum", "candle-core", + "candle-flash-attn", + "candle-flash-attn-v3", "candle-nn", "clap", "hf-hub", diff --git a/Cargo.toml b/Cargo.toml index e1974db..bd29c64 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -7,12 +7,14 @@ authors = ["Alvaro Bartolome <36760800+alvarobartt@users.noreply.github.com>"] default-run = "sys1" license = "Apache-2.0" repository = "https://github.com/alvarobartt/sys1" -include = ["/src/**", "/Cargo.toml", "/Cargo.lock", "/README.md", "/LICENSE"] +include = ["/src/**", "/build.rs", "/Cargo.toml", "/Cargo.lock", "/README.md", "/LICENSE"] [dependencies] anyhow = "1" axum = "0.8.9" candle-core = "0.11.0" +candle-flash-attn = { version = "0.11.0", optional = true } +candle-flash-attn-v3 = { version = "0.11.0", optional = true } candle-nn = "0.11.0" clap = { version = "4.6.7", features = ["derive"] } hf-hub = "1.0.0" @@ -41,6 +43,8 @@ default = ["cpu"] cpu = [] cuda = ["candle-core/cuda", "candle-nn/cuda"] cudnn = ["cuda", "candle-core/cudnn", "candle-nn/cudnn"] +flash-attn-2 = ["cuda", "dep:candle-flash-attn"] +flash-attn-3 = ["cuda", "dep:candle-flash-attn-v3"] metal = ["candle-core/metal", "candle-nn/metal"] [profile.release] diff --git a/build.rs b/build.rs new file mode 100644 index 0000000..deeb492 --- /dev/null +++ b/build.rs @@ -0,0 +1,45 @@ +use std::env; + +fn main() { + println!("cargo:rerun-if-env-changed=CUDA_COMPUTE_CAP"); + + let flash_attention_2 = env::var_os("CARGO_FEATURE_FLASH_ATTN_2").is_some(); + let flash_attention_3 = env::var_os("CARGO_FEATURE_FLASH_ATTN_3").is_some(); + if !flash_attention_2 && !flash_attention_3 { + return; + } + + let Ok(value) = env::var("CUDA_COMPUTE_CAP") else { + println!( + "cargo:warning=CUDA_COMPUTE_CAP is not set; the Flash Attention architecture will be checked at startup" + ); + return; + }; + let capability = parse_compute_capability(&value); + match (flash_attention_2, capability) { + (true, 80..=99) | (false, 90) => {} + (true, _) => panic!( + "`flash-attn-2` supports compute capability 8.x or 9.x; got {:?}", + value + ), + (false, _) => panic!( + "`flash-attn-3` requires Hopper compute capability 9.0; got {:?}", + value + ), + } +} + +fn parse_compute_capability(value: &str) -> u32 { + let normalized = value.trim().to_ascii_lowercase(); + let normalized = normalized + .strip_prefix("sm_") + .unwrap_or(&normalized) + .trim_end_matches(['a', 'f']) + .replace('.', ""); + normalized.parse::().unwrap_or_else(|_| { + panic!( + "invalid CUDA_COMPUTE_CAP {:?}; expected values such as 80, 89, or 90", + value + ) + }) +} diff --git a/src/device.rs b/src/device.rs index 9fa4ada..577a41b 100644 --- a/src/device.rs +++ b/src/device.rs @@ -1,5 +1,19 @@ use candle_core::Device; +pub fn validate_available() -> anyhow::Result<()> { + #[cfg(feature = "cuda")] + anyhow::ensure!( + candle_core::utils::cuda_is_available(), + "CUDA backend selected, but no CUDA device is available" + ); + #[cfg(feature = "metal")] + anyhow::ensure!( + candle_core::utils::metal_is_available(), + "Metal backend selected, but no Metal device is available" + ); + Ok(()) +} + #[cfg(all(feature = "cpu", not(any(feature = "cuda", feature = "metal"))))] pub fn load() -> anyhow::Result { Ok(Device::Cpu) diff --git a/src/lib.rs b/src/lib.rs index 1878d2e..8196a1b 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -6,6 +6,14 @@ compile_error!("enable one backend feature: cpu, cuda, or metal"); all(feature = "cuda", feature = "metal") ))] compile_error!("backend features are mutually exclusive: choose cpu, cuda, or metal"); +#[cfg(all(feature = "flash-attn-2", feature = "flash-attn-3"))] +compile_error!( + "Flash Attention features are mutually exclusive: choose flash-attn-2 or flash-attn-3" +); +#[cfg(all(feature = "metal", not(target_os = "macos")))] +compile_error!("the `metal` feature is only supported when targeting macOS"); +#[cfg(all(feature = "cuda", target_os = "macos"))] +compile_error!("the `cuda` feature is not supported when targeting macOS"); pub mod api; pub mod batching; @@ -14,3 +22,7 @@ pub mod hub; pub mod models; pub mod schema; pub mod tokenizer; + +pub fn validate_backend() -> anyhow::Result<()> { + device::validate_available() +} diff --git a/src/main.rs b/src/main.rs index 4c543f5..568a96c 100644 --- a/src/main.rs +++ b/src/main.rs @@ -52,6 +52,8 @@ struct Args { max_model_len: Option, #[arg(long, value_enum, default_value_t = Precision::Auto)] dtype: Precision, + #[arg(long, value_enum, default_value_t = models::AttentionImplementation::Eager)] + attention: models::AttentionImplementation, } impl Args { @@ -78,6 +80,7 @@ impl Args { self.max_model_len != Some(0), "--max-model-len must be positive" ); + self.attention.validate(self.dtype.resolve())?; Ok(()) } } @@ -111,6 +114,7 @@ async fn main() -> anyhow::Result<()> { .init(); let args = Args::parse(); args.validate()?; + sys1::validate_backend()?; let dtype = args.dtype.resolve(); info!( version = env!("CARGO_PKG_VERSION"), @@ -152,8 +156,14 @@ async fn main() -> anyhow::Result<()> { path = %model_path.display(), "loading model" ); - let model = models::load(&model_path, architecture, dtype, args.max_model_len) - .with_context(|| format!("failed to load model from {}", model_path.display()))?; + let model = models::load( + &model_path, + architecture, + dtype, + args.max_model_len, + args.attention, + ) + .with_context(|| format!("failed to load model from {}", model_path.display()))?; info!( model = %served_model_name, elapsed_ms = started.elapsed().as_millis(), @@ -244,6 +254,7 @@ mod tests { assert_eq!(args.max_request_bytes, 1_048_576); assert_eq!(args.max_model_len, None); assert_eq!(args.dtype, Precision::Auto); + assert_eq!(args.attention, models::AttentionImplementation::Eager); } #[test] @@ -261,6 +272,12 @@ mod tests { let args = Args::try_parse_from(["sys1", "--dtype", "bf16"]).unwrap(); assert_eq!(args.dtype, Precision::Bf16); + + let args = Args::try_parse_from(["sys1", "--attention", "flash-attn-2"]).unwrap(); + assert_eq!( + args.attention, + models::AttentionImplementation::FlashAttention2 + ); } #[test] @@ -302,4 +319,11 @@ mod tests { let args = Args::try_parse_from(["sys1", "--max-queue-size", "0"]).unwrap(); assert!(args.validate().is_err()); } + + #[test] + fn rejects_flash_attention_with_f32() { + let args = Args::try_parse_from(["sys1", "--attention", "flash-attn-3", "--dtype", "f32"]) + .unwrap(); + assert!(args.validate().is_err()); + } } diff --git a/src/models/laya.rs b/src/models/laya.rs index 06e32e2..05a97a0 100644 --- a/src/models/laya.rs +++ b/src/models/laya.rs @@ -1,5 +1,8 @@ +use super::AttentionImplementation; use super::DecisionModel; -use super::modernbert::{Config as ModernBertConfig, Encoder as ModernBertEncoder}; +use super::modernbert::{ + AttentionOptions, Config as ModernBertConfig, Encoder as ModernBertEncoder, +}; use crate::{ device, schema::{ApiError, DecisionRequest, DecisionResponse, Usage}, @@ -8,7 +11,7 @@ use crate::{ use anyhow::Context; use candle_core::{D, DType, Device, IndexOp, Tensor}; -use candle_nn::{Embedding, LayerNorm, Linear, VarBuilder, embedding, layer_norm, ops::softmax}; +use candle_nn::{Embedding, LayerNorm, Linear, VarBuilder, embedding, layer_norm}; use serde::Deserialize; use serde_json::{Map, Value, json}; use std::{collections::HashMap, fs, path::Path}; @@ -59,10 +62,16 @@ struct HeadLayer { linear2: Linear, heads: usize, compute_dtype: DType, + attention: AttentionImplementation, } impl HeadLayer { - fn load(vb: VarBuilder, hidden: usize, compute_dtype: DType) -> candle_core::Result { + fn load( + vb: VarBuilder, + hidden: usize, + compute_dtype: DType, + attention: AttentionImplementation, + ) -> candle_core::Result { Ok(Self { qkv: Linear::new( vb.get((hidden * 3, hidden), "self_attn.in_proj_weight")? @@ -79,10 +88,16 @@ impl HeadLayer { linear2: linear_dtype(hidden * 4, hidden, vb.pp("linear2"), compute_dtype)?, heads: hidden / 64, compute_dtype, + attention, }) } - fn forward(&self, xs: &Tensor, mask: &Tensor) -> candle_core::Result { + fn forward( + &self, + xs: &Tensor, + mask: &Tensor, + lengths: &[usize], + ) -> candle_core::Result { let (batch, length, hidden) = xs.dims3()?; let size = hidden / self.heads; let qkv = xs @@ -97,31 +112,39 @@ impl HeadLayer { let scale = (size as f64).powf(-0.5); #[cfg(feature = "metal")] - let attention = if xs.device().is_metal() { + let attended = if xs.device().is_metal() { let mask = mask .broadcast_as((batch, self.heads, length, length))? .contiguous()?; candle_nn::ops::sdpa(&q, &k, &v, Some(&mask), false, scale as f32, 1.0)? } else { - let scores = (&q * scale)? - .matmul(&k.transpose(D::Minus2, D::Minus1)?)? - .to_dtype(mask.dtype())? - .broadcast_add(mask)?; - softmax(&scores, D::Minus1)?.matmul(&v)? + super::modernbert::scaled_dot_product_attention( + &q, + &k, + &v, + scale, + AttentionOptions { + mask: Some(mask), + implementation: self.attention, + lengths, + window: None, + }, + )? }; #[cfg(not(feature = "metal"))] - let attention = { - let scores = (&q * scale)? - .matmul(&k.transpose(D::Minus2, D::Minus1)?)? - .broadcast_add(mask)?; - let probabilities = if scores.dtype() == DType::F16 { - softmax(&scores.to_dtype(DType::F32)?, D::Minus1)?.to_dtype(DType::F16)? - } else { - softmax(&scores, D::Minus1)? - }; - probabilities.to_dtype(v.dtype())?.matmul(&v)? - }; - let attention = attention + let attended = super::modernbert::scaled_dot_product_attention( + &q, + &k, + &v, + scale, + AttentionOptions { + mask: Some(mask), + implementation: self.attention, + lengths, + window: None, + }, + )?; + let attention = attended .transpose(1, 2)? .reshape((batch, length, hidden))? .apply(&self.projection)? @@ -157,9 +180,15 @@ pub struct Laya { } impl Laya { - pub fn load(path: &Path, dtype: DType, max_model_len: Option) -> anyhow::Result { + pub fn load( + path: &Path, + dtype: DType, + max_model_len: Option, + attention: AttentionImplementation, + ) -> anyhow::Result { let device = device::load()?; let (model_dtype, compute_dtype) = execution_dtypes(dtype, device.is_cuda()); + validate_attention(attention, &device, compute_dtype)?; let mut config: LayaConfig = serde_json::from_slice(&fs::read(path.join("rl_agent_config.json"))?)?; let encoder_config = ModernBertConfig::load(&path.join("encoder/config.json"))?; @@ -178,13 +207,15 @@ impl Laya { .map(|name| format!("encoder.{name}")) .unwrap_or_else(|| name.to_owned()) }); - let encoder = ModernBertEncoder::load(encoder_vb, &encoder_config, compute_dtype)?; + let encoder = + ModernBertEncoder::load(encoder_vb, &encoder_config, compute_dtype, attention)?; let head = (0..2) .map(|index| { HeadLayer::load( vb.pp(format!("head.layers.{index}")), encoder_config.hidden_size(), compute_dtype, + attention, ) }) .collect::>>()?; @@ -378,14 +409,15 @@ impl Laya { .to_dtype(self.compute_dtype)?; let type_ids = Tensor::from_vec(kinds, batch, &self.device)?; let has_padding = unique.iter().any(|item| item.ids.len() != length); + let lengths: Vec<_> = unique.iter().map(|item| item.ids.len()).collect(); let mut hidden = self .encoder - .forward(&ids, &attention_mask, has_padding) + .forward(&ids, &attention_mask, &lengths, has_padding) .context("encoder forward")?; hidden = hidden.broadcast_add(&type_ids.apply(&self.type_embedding)?.unsqueeze(1)?)?; for (index, layer) in self.head.iter().enumerate() { hidden = layer - .forward(&hidden, &head_mask) + .forward(&hidden, &head_mask, &lengths) .with_context(|| format!("decision head layer {index}"))?; } let output = if batch >= 8 && !self.device.is_cpu() { @@ -574,6 +606,51 @@ fn execution_dtypes(requested: DType, is_cuda: bool) -> (DType, DType) { (model, requested) } +fn validate_attention( + attention: AttentionImplementation, + _device: &Device, + compute_dtype: DType, +) -> anyhow::Result<()> { + attention.validate(compute_dtype)?; + #[cfg(any(feature = "flash-attn-2", feature = "flash-attn-3"))] + if attention != AttentionImplementation::Eager { + anyhow::ensure!( + _device.is_cuda(), + "{} requires a CUDA device", + attention.cli_name() + ); + let (major, minor) = match _device { + Device::Cuda(cuda) => cuda + .cuda_stream() + .context() + .compute_capability() + .context("failed to query CUDA compute capability")?, + _ => unreachable!(), + }; + validate_flash_capability(attention, major, minor)?; + } + Ok(()) +} + +#[cfg(any(feature = "flash-attn-2", feature = "flash-attn-3", test))] +fn validate_flash_capability( + attention: AttentionImplementation, + major: i32, + minor: i32, +) -> anyhow::Result<()> { + let supported = match attention { + AttentionImplementation::Eager => true, + AttentionImplementation::FlashAttention2 => (8..=9).contains(&major), + AttentionImplementation::FlashAttention3 => (major, minor) == (9, 0), + }; + anyhow::ensure!( + supported, + "{} does not support CUDA compute capability {major}.{minor}", + attention.cli_name() + ); + Ok(()) +} + fn deduplicate<'a>(items: &[&'a Item]) -> (Vec<&'a Item>, Vec) { let mut unique = Vec::with_capacity(items.len()); let mut indices = Vec::with_capacity(items.len()); @@ -833,7 +910,7 @@ mod tests { for (snapshot, model_id, revision) in models { let path = crate::hub::download(model_id, revision).await?; - let model = Laya::load(&path, DType::F32, None)?; + let model = Laya::load(&path, DType::F32, None, AttentionImplementation::Eager)?; let request: DecisionRequest = serde_json::from_value(json!({ "state": { "message": "I was charged twice for invoice 4411. Please refund me today.", @@ -937,4 +1014,34 @@ mod tests { ); assert_eq!(execution_dtypes(DType::F32, true), (DType::F32, DType::F32)); } + + #[test] + fn rejects_flash_attention_without_a_compatible_backend() { + assert!( + validate_attention( + AttentionImplementation::FlashAttention2, + &Device::Cpu, + DType::BF16, + ) + .is_err() + ); + validate_attention(AttentionImplementation::Eager, &Device::Cpu, DType::F32).unwrap(); + } + + #[test] + fn validates_flash_attention_compute_capabilities() { + validate_flash_capability(AttentionImplementation::FlashAttention2, 8, 9).unwrap(); + validate_flash_capability(AttentionImplementation::FlashAttention2, 9, 0).unwrap(); + validate_flash_capability(AttentionImplementation::FlashAttention3, 9, 0).unwrap(); + + assert!(validate_flash_capability(AttentionImplementation::FlashAttention2, 7, 5).is_err()); + assert!( + validate_flash_capability(AttentionImplementation::FlashAttention2, 10, 0).is_err() + ); + assert!(validate_flash_capability(AttentionImplementation::FlashAttention3, 8, 9).is_err()); + assert!(validate_flash_capability(AttentionImplementation::FlashAttention3, 9, 1).is_err()); + assert!( + validate_flash_capability(AttentionImplementation::FlashAttention3, 12, 0).is_err() + ); + } } diff --git a/src/models/mod.rs b/src/models/mod.rs index 498217d..0c0c34c 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -5,6 +5,7 @@ use crate::schema::{ApiError, DecisionRequest, DecisionResponse}; use anyhow::{Context, bail}; use candle_core::DType; +use clap::ValueEnum; use serde::Deserialize; use std::{fs, path::Path}; @@ -20,6 +21,44 @@ pub const LAYA_MODEL_IDS: &[&str] = &[ LAYA_MULTILINGUAL_MODEL_ID, ]; +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, ValueEnum)] +pub enum AttentionImplementation { + #[default] + Eager, + #[value(name = "flash-attn-2")] + FlashAttention2, + #[value(name = "flash-attn-3")] + FlashAttention3, +} + +impl AttentionImplementation { + pub fn cli_name(self) -> &'static str { + match self { + Self::Eager => "eager", + Self::FlashAttention2 => "flash-attn-2", + Self::FlashAttention3 => "flash-attn-3", + } + } + + pub fn validate(self, dtype: DType) -> anyhow::Result<()> { + let name = self.cli_name(); + let enabled = match self { + Self::Eager => return Ok(()), + Self::FlashAttention2 => cfg!(feature = "flash-attn-2"), + Self::FlashAttention3 => cfg!(feature = "flash-attn-3"), + }; + anyhow::ensure!( + enabled, + "--attention {name} requires a binary built with --features {name}" + ); + anyhow::ensure!( + matches!(dtype, DType::F16 | DType::BF16), + "--attention {name} requires --dtype f16 or --dtype bf16" + ); + Ok(()) + } +} + #[derive(Clone, Copy, Debug, Eq, PartialEq)] pub enum Architecture { Laya, @@ -109,9 +148,10 @@ pub fn load( architecture: Architecture, dtype: DType, max_model_len: Option, + attention: AttentionImplementation, ) -> anyhow::Result { match architecture { - Architecture::Laya => Laya::load(path, dtype, max_model_len).map(Model::Laya), + Architecture::Laya => Laya::load(path, dtype, max_model_len, attention).map(Model::Laya), } } diff --git a/src/models/modernbert.rs b/src/models/modernbert.rs index 6a4e3bc..b6c2d48 100644 --- a/src/models/modernbert.rs +++ b/src/models/modernbert.rs @@ -1,3 +1,4 @@ +use super::AttentionImplementation; use candle_core::{D, DType, Device, Result, Tensor}; use candle_nn::{ Embedding, LayerNorm, Linear, Module, VarBuilder, embedding, layer_norm_no_bias, ops::softmax, @@ -10,6 +11,8 @@ use std::{ sync::{Arc, Mutex}, }; +const ATTENTION_MASK_VALUE: f32 = -10_000.0; + #[derive(Deserialize)] pub struct Config { vocab_size: usize, @@ -101,6 +104,7 @@ struct Attention { head_size: usize, rotary: Arc, compute_dtype: DType, + implementation: AttentionImplementation, } impl Attention { @@ -109,6 +113,7 @@ impl Attention { config: &Config, rotary: Arc, compute_dtype: DType, + implementation: AttentionImplementation, ) -> Result { Ok(Self { qkv: linear_no_bias_dtype( @@ -127,10 +132,17 @@ impl Attention { head_size: config.hidden_size / config.num_attention_heads, rotary, compute_dtype, + implementation, }) } - fn forward(&self, xs: &Tensor, mask: Option<&Tensor>) -> Result { + fn forward( + &self, + xs: &Tensor, + mask: Option<&Tensor>, + lengths: &[usize], + window: Option, + ) -> Result { let (batch, length, hidden) = xs.dims3()?; let qkv = xs .to_dtype(self.compute_dtype)? @@ -150,10 +162,32 @@ impl Attention { .transpose()?; candle_nn::ops::sdpa(&q, &k, &v, mask.as_ref(), false, scale as f32, 1.0)? } else { - unfused_attention(&q, &k, &v, mask, scale)? + scaled_dot_product_attention( + &q, + &k, + &v, + scale, + AttentionOptions { + mask, + implementation: self.implementation, + lengths, + window, + }, + )? }; #[cfg(not(feature = "metal"))] - let attention = unfused_attention(&q, &k, &v, mask, scale)?; + let attention = scaled_dot_product_attention( + &q, + &k, + &v, + scale, + AttentionOptions { + mask, + implementation: self.implementation, + lengths, + window, + }, + )?; attention .transpose(1, 2)? @@ -162,24 +196,159 @@ impl Attention { } } -fn unfused_attention( +pub(super) struct AttentionOptions<'a> { + pub mask: Option<&'a Tensor>, + pub implementation: AttentionImplementation, + pub lengths: &'a [usize], + pub window: Option, +} + +pub(super) fn scaled_dot_product_attention( q: &Tensor, k: &Tensor, v: &Tensor, - mask: Option<&Tensor>, scale: f64, + options: AttentionOptions<'_>, ) -> Result { + if options.implementation != AttentionImplementation::Eager { + return flash_attention( + q, + k, + v, + options.lengths, + scale as f32, + options.window, + options.implementation, + ); + } let scores = (q * scale)?.matmul(&k.transpose(D::Minus2, D::Minus1)?)?; - let scores = match mask { + let scores = match options.mask { Some(mask) => scores.to_dtype(mask.dtype())?.broadcast_add(mask)?, None => scores, }; - let probabilities = if scores.dtype() == DType::F16 { - softmax(&scores.to_dtype(DType::F32)?, D::Minus1)?.to_dtype(DType::F16)? + let probabilities = attention_softmax(&scores)?; + probabilities.to_dtype(v.dtype())?.matmul(v) +} + +fn attention_softmax(scores: &Tensor) -> Result { + if matches!(scores.dtype(), DType::F16 | DType::BF16) { + softmax(&scores.to_dtype(DType::F32)?, D::Minus1)?.to_dtype(scores.dtype()) + } else { + softmax(scores, D::Minus1) + } +} + +#[cfg(any(feature = "flash-attn-2", feature = "flash-attn-3"))] +fn flash_attention( + q: &Tensor, + k: &Tensor, + v: &Tensor, + lengths: &[usize], + scale: f32, + window: Option, + implementation: AttentionImplementation, +) -> Result { + let (batch, heads, max_length, head_size) = q.dims4()?; + if lengths.len() != batch || lengths.iter().any(|&length| length > max_length) { + candle_core::bail!( + "invalid Flash Attention sequence lengths {:?} for shape {:?}", + lengths, + q.shape() + ) + } + let q = q.transpose(1, 2)?.contiguous()?; + let k = k.transpose(1, 2)?.contiguous()?; + let v = v.transpose(1, 2)?.contiguous()?; + let attention = if lengths.iter().all(|&length| length == max_length) { + match implementation { + #[cfg(feature = "flash-attn-2")] + AttentionImplementation::FlashAttention2 => { + candle_flash_attn::flash_attn_windowed(&q, &k, &v, scale, window, window)? + } + #[cfg(feature = "flash-attn-3")] + AttentionImplementation::FlashAttention3 => { + candle_flash_attn_v3::flash_attn_windowed(&q, &k, &v, scale, window, window, false)? + } + _ => candle_core::bail!("{} support is not compiled in", implementation.cli_name()), + } } else { - softmax(&scores, D::Minus1)? + let mut indices = Vec::with_capacity(lengths.iter().sum()); + let mut cumulative = Vec::with_capacity(batch + 1); + cumulative.push(0u32); + for (row, &length) in lengths.iter().enumerate() { + indices.extend((0..length).map(|column| (row * max_length + column) as u32)); + cumulative.push(cumulative.last().copied().unwrap() + length as u32); + } + let indices = Tensor::from_vec( + indices, + cumulative.last().copied().unwrap() as usize, + q.device(), + )?; + let cumulative = Tensor::from_vec(cumulative, batch + 1, q.device())?; + let pack = |tensor: &Tensor| { + tensor + .reshape((batch * max_length, heads, head_size))? + .index_select(&indices, 0) + }; + let packed_q = pack(&q)?; + let packed_k = pack(&k)?; + let packed_v = pack(&v)?; + let packed = match implementation { + #[cfg(feature = "flash-attn-2")] + AttentionImplementation::FlashAttention2 => { + candle_flash_attn::flash_attn_varlen_windowed( + &packed_q, + &packed_k, + &packed_v, + &cumulative, + &cumulative, + max_length, + max_length, + scale, + window, + window, + )? + } + #[cfg(feature = "flash-attn-3")] + AttentionImplementation::FlashAttention3 => { + candle_flash_attn_v3::flash_attn_varlen_windowed( + &packed_q, + &packed_k, + &packed_v, + &cumulative, + &cumulative, + max_length, + max_length, + scale, + window, + window, + false, + )? + } + _ => candle_core::bail!("{} support is not compiled in", implementation.cli_name()), + }; + Tensor::zeros( + (batch * max_length, heads, head_size), + packed.dtype(), + packed.device(), + )? + .index_add(&indices, &packed, 0)? + .reshape((batch, max_length, heads, head_size))? }; - probabilities.to_dtype(v.dtype())?.matmul(v) + attention.transpose(1, 2) +} + +#[cfg(not(any(feature = "flash-attn-2", feature = "flash-attn-3")))] +fn flash_attention( + _q: &Tensor, + _k: &Tensor, + _v: &Tensor, + _lengths: &[usize], + _scale: f32, + _window: Option, + implementation: AttentionImplementation, +) -> Result { + candle_core::bail!("{} support is not compiled in", implementation.cli_name()) } struct Mlp { @@ -233,9 +402,16 @@ impl Layer { rotary: Arc, uses_local_attention: bool, compute_dtype: DType, + implementation: AttentionImplementation, ) -> Result { Ok(Self { - attention: Attention::load(vb.pp("attn"), config, rotary, compute_dtype)?, + attention: Attention::load( + vb.pp("attn"), + config, + rotary, + compute_dtype, + implementation, + )?, mlp: Mlp::load(vb.pp("mlp"), config, compute_dtype)?, attention_norm: layer_norm_no_bias( config.hidden_size, @@ -257,6 +433,8 @@ impl Layer { xs: &Tensor, global_mask: Option<&Tensor>, local_mask: &Tensor, + lengths: &[usize], + local_window: usize, ) -> Result { let normalized = match &self.attention_norm { Some(norm) => xs.apply(norm)?, @@ -272,7 +450,12 @@ impl Layer { }; let attention = self .attention - .forward(&normalized, mask.as_ref())? + .forward( + &normalized, + mask.as_ref(), + lengths, + self.uses_local_attention.then_some(local_window), + )? .to_dtype(xs.dtype())?; let xs = (attention + xs)?; let mlp = xs.apply(&self.mlp_norm)?.apply(&self.mlp)?; @@ -292,7 +475,12 @@ pub struct Encoder { } impl Encoder { - pub fn load(vb: VarBuilder, config: &Config, compute_dtype: DType) -> Result { + pub fn load( + vb: VarBuilder, + config: &Config, + compute_dtype: DType, + implementation: AttentionImplementation, + ) -> Result { let global_rotary = Arc::new(RotaryEmbedding::new( vb.dtype(), config, @@ -318,6 +506,7 @@ impl Encoder { }, local, compute_dtype, + implementation, )?); } Ok(Self { @@ -343,15 +532,28 @@ impl Encoder { }) } - pub fn forward(&self, ids: &Tensor, mask: &Tensor, has_padding: bool) -> Result { + pub fn forward( + &self, + ids: &Tensor, + mask: &Tensor, + lengths: &[usize], + has_padding: bool, + ) -> Result { let length = ids.dim(1)?; let global_mask = has_padding .then(|| global_attention_mask(mask, length, self.dtype)) .transpose()?; let local_mask = self.local_mask(length, ids.device())?; + let local_window = self.local_attention_size / 2; let mut xs = ids.apply(&self.embeddings)?.apply(&self.norm)?; for layer in &self.layers { - xs = layer.forward(&xs, global_mask.as_ref(), &local_mask)?; + xs = layer.forward( + &xs, + global_mask.as_ref(), + &local_mask, + lengths, + local_window, + )?; } xs.apply(&self.final_norm) } @@ -389,7 +591,7 @@ fn global_attention_mask(mask: &Tensor, target_length: usize, dtype: DType) -> R .unsqueeze(2)? .expand((batch, 1, target_length, source_length))? .to_dtype(dtype)?; - ((1.0 - expanded)? * f32::MIN as f64)?.to_dtype(dtype) + ((1.0 - expanded)? * ATTENTION_MASK_VALUE as f64)?.to_dtype(dtype) } fn local_attention_mask( @@ -402,7 +604,7 @@ fn local_attention_mask( .flat_map(|left| { (0..length).map(move |right| { if left.abs_diff(right) > max_distance { - f32::NEG_INFINITY + ATTENTION_MASK_VALUE } else { 0.0 } @@ -424,10 +626,10 @@ mod tests { assert_eq!( mask, vec![ - vec![0.0, 0.0, f32::NEG_INFINITY, f32::NEG_INFINITY], - vec![0.0, 0.0, 0.0, f32::NEG_INFINITY], - vec![f32::NEG_INFINITY, 0.0, 0.0, 0.0], - vec![f32::NEG_INFINITY, f32::NEG_INFINITY, 0.0, 0.0], + vec![0.0, 0.0, ATTENTION_MASK_VALUE, ATTENTION_MASK_VALUE], + vec![0.0, 0.0, 0.0, ATTENTION_MASK_VALUE], + vec![ATTENTION_MASK_VALUE, 0.0, 0.0, 0.0], + vec![ATTENTION_MASK_VALUE, ATTENTION_MASK_VALUE, 0.0, 0.0], ] ); Ok(()) @@ -440,7 +642,17 @@ mod tests { .flatten_all()? .to_vec1::()?; - assert_eq!(mask, vec![0.0, 0.0, f32::MIN, 0.0, 0.0, f32::MIN]); + assert_eq!( + mask, + vec![ + 0.0, + 0.0, + ATTENTION_MASK_VALUE, + 0.0, + 0.0, + ATTENTION_MASK_VALUE + ] + ); Ok(()) } @@ -461,7 +673,7 @@ mod tests { .map(|name| format!("encoder.{name}")) .unwrap_or_else(|| name.to_owned()) }); - let encoder = Encoder::load(vb, &config, DType::F32)?; + let encoder = Encoder::load(vb, &config, DType::F32, AttentionImplementation::Eager)?; let tokenizer_path = path.join("tokenizer/tokenizer.json"); let tokenizer = crate::tokenizer::from_json(&fs::read(tokenizer_path)?)?; @@ -478,6 +690,7 @@ mod tests { let mut ids = vec![cls_id]; ids.extend(tokenizer.encode("What is Deep Learning?")?); ids.push(sep_id); + let valid_length = ids.len(); let mut mask = vec![1f32; ids.len()]; ids.resize(32, pad_id); @@ -486,7 +699,7 @@ mod tests { let length = ids.len(); let ids = Tensor::from_vec(ids, (1, length), &device)?; let mask = Tensor::from_vec(mask, (1, length), &device)?; - let hidden = encoder.forward(&ids, &mask, true)?; + let hidden = encoder.forward(&ids, &mask, &[valid_length], true)?; insta::assert_yaml_snapshot!( "modernbert_fp32_hidden_states", @@ -498,4 +711,51 @@ mod tests { Ok(()) } + + #[test] + fn eager_bf16_softmax_handles_padding_masks_without_nans() { + let device = Device::Cpu; + let scores = Tensor::new(&[[0f32, f32::NEG_INFINITY]], &device) + .unwrap() + .to_dtype(DType::BF16) + .unwrap(); + + let output = attention_softmax(&scores) + .unwrap() + .to_dtype(DType::F32) + .unwrap() + .flatten_all() + .unwrap() + .to_vec1::() + .unwrap(); + + assert!(output.iter().all(|value| value.is_finite())); + assert_eq!(output, vec![1.0, 0.0]); + } + + #[test] + fn bf16_attention_masks_remain_finite() { + let device = Device::Cpu; + let padding = Tensor::new(&[[1f32, 0.0]], &device).unwrap(); + let global = global_attention_mask(&padding, 2, DType::BF16) + .unwrap() + .to_dtype(DType::F32) + .unwrap() + .flatten_all() + .unwrap() + .to_vec1::() + .unwrap(); + let local = local_attention_mask(4, 1, DType::BF16, &device) + .unwrap() + .to_dtype(DType::F32) + .unwrap() + .flatten_all() + .unwrap() + .to_vec1::() + .unwrap(); + + assert!(global.iter().chain(&local).all(|value| value.is_finite())); + assert!(global.iter().any(|value| *value < -1_000.0)); + assert!(local.iter().any(|value| *value < -1_000.0)); + } } From 489d9d2081d3b2e9e151b4e864e81db9e8a3aef9 Mon Sep 17 00:00:00 2001 From: Alvaro Bartolome <36760800+alvarobartt@users.noreply.github.com> Date: Fri, 25 Sep 2026 07:35:00 +0000 Subject: [PATCH 2/6] Update `README.md` --- README.md | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) diff --git a/README.md b/README.md index c0a1317..a02b06d 100644 --- a/README.md +++ b/README.md @@ -1,14 +1,12 @@
System One
- System One compatible API for open decision models, e.g. - Laya, - written in Rust. + System One compatible API for open decision models, written in Rust.
@@ -18,9 +16,8 @@ - System One compatible API Spec - `candle` with [`tokenizers` release candidate](https://huggingface.co/blog/tokenizers-v1)! - Dynamic, token-based batching -- Support for ModernBert with Laya custom decision heads -- CPU, CUDA and Metal (MPS) supported -- ~14ms per query on NVIDIA RTX Pro 6000 +- SDPA on CPU, Metal, and CUDA +- Flash Attention on Ampere, Ada Lovelace, and Hopper ## Get started @@ -30,6 +27,8 @@ Install it with support for CPU, Metal or CUDA. cargo install sys1 --features cpu # cargo install sys1 --no-default-features --features metal # cargo install sys1 --no-default-features --features cuda +# cargo install sys1 --no-default-features --features cuda,flash-attn-2 # Ampere, Ada Lovelace, or Hopper +# cargo install sys1 --no-default-features --features cuda,flash-attn-3 # Hopper ``` Then run it with any of the supported models (more coming soon!). From c13bb3ad1853e3c1acc053c0be6d40221aaead0d Mon Sep 17 00:00:00 2001 From: Alvaro Bartolome <36760800+alvarobartt@users.noreply.github.com> Date: Fri, 25 Sep 2026 07:35:14 +0000 Subject: [PATCH 3/6] Update `.github/workflows/ci.yml` --- .github/workflows/ci.yml | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 510ba43..449a226 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -43,8 +43,8 @@ jobs: secrets: HF_TOKEN: ${{ github.event_name == 'push' && secrets.HF_TOKEN || '' }} - build_main: - name: Build main + build: + name: Build needs: tests if: github.event_name == 'push' && github.ref == 'refs/heads/main' permissions: @@ -82,13 +82,13 @@ jobs: push: name: Push - needs: [tests, build_main] + needs: [tests, build] if: >- always() && needs.tests.result == 'success' && github.event_name == 'push' && (github.ref == 'refs/heads/main' || github.ref_type == 'tag') && - (github.ref_type == 'tag' || needs.build_main.result == 'success') + (github.ref_type == 'tag' || needs.build.result == 'success') permissions: contents: read packages: write From e2e408e3ddff92a8ff4da41912a4a36472a23243 Mon Sep 17 00:00:00 2001 From: Alvaro Bartolome <36760800+alvarobartt@users.noreply.github.com> Date: Fri, 25 Sep 2026 07:51:37 +0000 Subject: [PATCH 4/6] Use `match` over `if` when applicable on attention --- src/models/laya.rs | 39 +++++++++++++++++++++------------------ src/models/modernbert.rs | 24 +++++++++++++----------- 2 files changed, 34 insertions(+), 29 deletions(-) diff --git a/src/models/laya.rs b/src/models/laya.rs index 05a97a0..f214926 100644 --- a/src/models/laya.rs +++ b/src/models/laya.rs @@ -112,7 +112,7 @@ impl HeadLayer { let scale = (size as f64).powf(-0.5); #[cfg(feature = "metal")] - let attended = if xs.device().is_metal() { + let attention = if xs.device().is_metal() { let mask = mask .broadcast_as((batch, self.heads, length, length))? .contiguous()?; @@ -132,7 +132,7 @@ impl HeadLayer { )? }; #[cfg(not(feature = "metal"))] - let attended = super::modernbert::scaled_dot_product_attention( + let attention = super::modernbert::scaled_dot_product_attention( &q, &k, &v, @@ -144,7 +144,7 @@ impl HeadLayer { window: None, }, )?; - let attention = attended + let attention = attention .transpose(1, 2)? .reshape((batch, length, hidden))? .apply(&self.projection)? @@ -613,21 +613,24 @@ fn validate_attention( ) -> anyhow::Result<()> { attention.validate(compute_dtype)?; #[cfg(any(feature = "flash-attn-2", feature = "flash-attn-3"))] - if attention != AttentionImplementation::Eager { - anyhow::ensure!( - _device.is_cuda(), - "{} requires a CUDA device", - attention.cli_name() - ); - let (major, minor) = match _device { - Device::Cuda(cuda) => cuda - .cuda_stream() - .context() - .compute_capability() - .context("failed to query CUDA compute capability")?, - _ => unreachable!(), - }; - validate_flash_capability(attention, major, minor)?; + match attention { + AttentionImplementation::Eager => {} + implementation => { + anyhow::ensure!( + _device.is_cuda(), + "{} requires a CUDA device", + implementation.cli_name() + ); + let (major, minor) = match _device { + Device::Cuda(cuda) => cuda + .cuda_stream() + .context() + .compute_capability() + .context("failed to query CUDA compute capability")?, + _ => unreachable!(), + }; + validate_flash_capability(implementation, major, minor)?; + } } Ok(()) } diff --git a/src/models/modernbert.rs b/src/models/modernbert.rs index b6c2d48..84c21db 100644 --- a/src/models/modernbert.rs +++ b/src/models/modernbert.rs @@ -210,24 +210,26 @@ pub(super) fn scaled_dot_product_attention( scale: f64, options: AttentionOptions<'_>, ) -> Result { - if options.implementation != AttentionImplementation::Eager { - return flash_attention( + match options.implementation { + AttentionImplementation::Eager => { + let scores = (q * scale)?.matmul(&k.transpose(D::Minus2, D::Minus1)?)?; + let scores = match options.mask { + Some(mask) => scores.to_dtype(mask.dtype())?.broadcast_add(mask)?, + None => scores, + }; + let probabilities = attention_softmax(&scores)?; + probabilities.to_dtype(v.dtype())?.matmul(v) + } + implementation => flash_attention( q, k, v, options.lengths, scale as f32, options.window, - options.implementation, - ); + implementation, + ), } - let scores = (q * scale)?.matmul(&k.transpose(D::Minus2, D::Minus1)?)?; - let scores = match options.mask { - Some(mask) => scores.to_dtype(mask.dtype())?.broadcast_add(mask)?, - None => scores, - }; - let probabilities = attention_softmax(&scores)?; - probabilities.to_dtype(v.dtype())?.matmul(v) } fn attention_softmax(scores: &Tensor) -> Result { From db79e8058071f93104148143a991f1ae4328ee52 Mon Sep 17 00:00:00 2001 From: Alvaro Bartolome <36760800+alvarobartt@users.noreply.github.com> Date: Fri, 25 Sep 2026 08:39:10 +0000 Subject: [PATCH 5/6] Add `build.rs` in Dockerfiles & pre-define features based on compute cap --- Dockerfile | 2 ++ Dockerfile.cuda | 18 +++++++++++++++--- 2 files changed, 17 insertions(+), 3 deletions(-) diff --git a/Dockerfile b/Dockerfile index e2030ae..db7a657 100644 --- a/Dockerfile +++ b/Dockerfile @@ -27,6 +27,7 @@ WORKDIR /app FROM chef AS planner COPY Cargo.toml Cargo.lock ./ +COPY build.rs ./ COPY rust-toolchain.toml ./ COPY src ./src RUN cargo chef prepare --recipe-path recipe.json @@ -55,6 +56,7 @@ RUN --mount=type=cache,id=sys1-sccache,target=/root/.cache/sccache,sharing=locke sccache --show-stats COPY Cargo.toml Cargo.lock ./ +COPY build.rs ./ COPY rust-toolchain.toml ./ COPY src ./src diff --git a/Dockerfile.cuda b/Dockerfile.cuda index f540147..dfba381 100644 --- a/Dockerfile.cuda +++ b/Dockerfile.cuda @@ -14,7 +14,7 @@ ARG TARGETARCH ENV PATH="/root/.cargo/bin:${PATH}" RUN apt-get update \ - && apt-get install -y --no-install-recommends build-essential ca-certificates curl pkg-config \ + && apt-get install -y --no-install-recommends build-essential ca-certificates curl git pkg-config \ && rm -rf /var/lib/apt/lists/* RUN curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain "${RUST_VERSION}" --profile minimal @@ -41,6 +41,7 @@ WORKDIR /app FROM chef AS planner COPY Cargo.toml Cargo.lock ./ +COPY build.rs ./ COPY rust-toolchain.toml ./ COPY src ./src @@ -74,17 +75,23 @@ RUN --mount=type=cache,id=sys1-sccache,target=/root/.cache/sccache,sharing=locke fi; \ compute_caps="$(printf '%s' "${CUDA_COMPUTE_CAPS}" | tr ';' ' ')"; \ for compute_cap in ${compute_caps}; do \ + case "${compute_cap}" in \ + 80|86|89) features=cuda,flash-attn-2 ;; \ + 90) features=cuda,flash-attn-3 ;; \ + *) features=cuda ;; \ + esac; \ CUDA_COMPUTE_CAP="${compute_cap}" \ CARGO_TARGET_DIR="target/${compute_cap}" \ cargo chef cook \ --release \ --no-default-features \ - --features cuda \ + --features "${features}" \ --recipe-path recipe.json; \ done; \ sccache --show-stats COPY Cargo.toml Cargo.lock ./ +COPY build.rs ./ COPY rust-toolchain.toml ./ COPY src ./src @@ -106,13 +113,18 @@ RUN --mount=type=cache,id=sys1-sccache,target=/root/.cache/sccache,sharing=locke set -- ${compute_caps}; \ install -d /app/bin; \ for compute_cap do \ + case "${compute_cap}" in \ + 80|86|89) features=cuda,flash-attn-2 ;; \ + 90) features=cuda,flash-attn-3 ;; \ + *) features=cuda ;; \ + esac; \ CUDA_COMPUTE_CAP="${compute_cap}" \ CARGO_TARGET_DIR="target/${compute_cap}" \ cargo build \ --release \ --locked \ --no-default-features \ - --features cuda \ + --features "${features}" \ --bin sys1; \ if [ "$#" -eq 1 ]; then \ install "target/${compute_cap}/release/sys1" /app/bin/sys1; \ From 2e821f388757d7777e67103882c391d6d958126d Mon Sep 17 00:00:00 2001 From: Alvaro Bartolome <36760800+alvarobartt@users.noreply.github.com> Date: Fri, 25 Sep 2026 09:44:55 +0000 Subject: [PATCH 6/6] Use `flash-attn-2` for Hopper in `Dockerfile.cuda` --- Dockerfile.cuda | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/Dockerfile.cuda b/Dockerfile.cuda index dfba381..89dd01a 100644 --- a/Dockerfile.cuda +++ b/Dockerfile.cuda @@ -76,8 +76,7 @@ RUN --mount=type=cache,id=sys1-sccache,target=/root/.cache/sccache,sharing=locke compute_caps="$(printf '%s' "${CUDA_COMPUTE_CAPS}" | tr ';' ' ')"; \ for compute_cap in ${compute_caps}; do \ case "${compute_cap}" in \ - 80|86|89) features=cuda,flash-attn-2 ;; \ - 90) features=cuda,flash-attn-3 ;; \ + 80|86|89|90) features=cuda,flash-attn-2 ;; \ *) features=cuda ;; \ esac; \ CUDA_COMPUTE_CAP="${compute_cap}" \ @@ -114,8 +113,7 @@ RUN --mount=type=cache,id=sys1-sccache,target=/root/.cache/sccache,sharing=locke install -d /app/bin; \ for compute_cap do \ case "${compute_cap}" in \ - 80|86|89) features=cuda,flash-attn-2 ;; \ - 90) features=cuda,flash-attn-3 ;; \ + 80|86|89|90) features=cuda,flash-attn-2 ;; \ *) features=cuda ;; \ esac; \ CUDA_COMPUTE_CAP="${compute_cap}" \