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
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/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..89dd01a 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,22 @@ 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|90) features=cuda,flash-attn-2 ;; \
+ *) 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 +112,17 @@ 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|90) features=cuda,flash-attn-2 ;; \
+ *) 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; \
diff --git a/README.md b/README.md
index c0a1317..a02b06d 100644
--- a/README.md
+++ b/README.md
@@ -1,14 +1,12 @@
- 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!).
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..f214926 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
@@ -103,24 +118,32 @@ impl HeadLayer {
.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 = super::modernbert::scaled_dot_product_attention(
+ &q,
+ &k,
+ &v,
+ scale,
+ AttentionOptions {
+ mask: Some(mask),
+ implementation: self.attention,
+ lengths,
+ window: None,
+ },
+ )?;
let attention = attention
.transpose(1, 2)?
.reshape((batch, length, hidden))?
@@ -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,54 @@ 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"))]
+ 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(())
+}
+
+#[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 +913,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 +1017,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..84c21db 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,161 @@ 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 {
- let scores = (q * scale)?.matmul(&k.transpose(D::Minus2, D::Minus1)?)?;
- let scores = match 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)?
+ 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,
+ implementation,
+ ),
+ }
+}
+
+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 +404,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 +435,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 +452,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 +477,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 +508,7 @@ impl Encoder {
},
local,
compute_dtype,
+ implementation,
)?);
}
Ok(Self {
@@ -343,15 +534,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 +593,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 +606,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 +628,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 +644,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 +675,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 +692,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 +701,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 +713,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));
+ }
}