diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 28aa407..510ba43 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,8 +4,15 @@ on: push: branches: [main] tags: ["v*"] + paths-ignore: + - "**/*.md" + - "LICENSE*" pull_request: branches: [main] + types: [opened, synchronize, reopened, ready_for_review] + paths-ignore: + - "**/*.md" + - "LICENSE*" permissions: contents: read @@ -32,33 +39,56 @@ jobs: needs: [trufflehog, lint] uses: ./.github/workflows/tests.yml with: - cache_scope: ${{ github.event_name }} source_sha: ${{ github.sha }} secrets: HF_TOKEN: ${{ github.event_name == 'push' && secrets.HF_TOKEN || '' }} - build: - name: Build + build_main: + name: Build main needs: tests - if: github.ref_type != 'tag' + if: github.event_name == 'push' && github.ref == 'refs/heads/main' permissions: actions: write contents: read packages: write uses: ./.github/workflows/build.yml with: - push: ${{ github.event_name == 'push' && github.ref == 'refs/heads/main' }} + push: true + source_sha: ${{ github.sha }} + + approve_pr_build: + name: Approve PR build + needs: tests + if: >- + github.event_name == 'pull_request' && + github.event.pull_request.draft == false && + github.event.pull_request.head.repo.full_name == github.repository + runs-on: ubuntu-latest + environment: pr-build + steps: + - run: echo "PR build approved" + + build_pr: + name: Build PR + needs: approve_pr_build + permissions: + actions: write + contents: read + packages: write + uses: ./.github/workflows/build.yml + with: + push: true source_sha: ${{ github.sha }} push: name: Push - needs: [tests, build] + needs: [tests, build_main] 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.result == 'success') + (github.ref_type == 'tag' || needs.build_main.result == 'success') permissions: contents: read packages: write diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index b9541b4..5c66255 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -3,9 +3,6 @@ name: Tests on: workflow_call: inputs: - cache_scope: - required: true - type: string source_sha: required: true type: string @@ -25,14 +22,19 @@ jobs: with: persist-credentials: false ref: ${{ inputs.source_sha }} - - uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9 # v6.1.0 + - name: Cache models + uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9 # v6.1.0 + with: + path: ~/.cache/huggingface/hub + key: models-v1-${{ runner.os }}-aa8c91ca-1a793eb5-e4e9ddf2 + - name: Cache Cargo + uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9 # v6.1.0 with: path: | - ~/.cache/huggingface/hub ~/.cargo/registry ~/.cargo/git - key: test-v2-${{ inputs.cache_scope }}-cpu-${{ hashFiles('Cargo.lock') }}-aa8c91ca088ec597df95a0d1c76b3063cb2ae5e8 - restore-keys: test-v2-${{ inputs.cache_scope }}-cpu- + key: test-v3-${{ runner.os }}-${{ hashFiles('Cargo.lock') }} + restore-keys: test-v3-${{ runner.os }}- - name: Test env: HF_TOKEN: ${{ secrets.HF_TOKEN }} @@ -46,14 +48,19 @@ jobs: with: persist-credentials: false ref: ${{ inputs.source_sha }} - - uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9 # v6.1.0 + - name: Cache models + uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9 # v6.1.0 + with: + path: ~/.cache/huggingface/hub + key: models-v1-${{ runner.os }}-aa8c91ca-1a793eb5-e4e9ddf2 + - name: Cache Cargo + uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9 # v6.1.0 with: path: | - ~/.cache/huggingface/hub ~/.cargo/registry ~/.cargo/git - key: test-v2-${{ inputs.cache_scope }}-metal-${{ hashFiles('Cargo.lock') }}-aa8c91ca088ec597df95a0d1c76b3063cb2ae5e8 - restore-keys: test-v2-${{ inputs.cache_scope }}-metal- + key: test-v3-${{ runner.os }}-${{ hashFiles('Cargo.lock') }} + restore-keys: test-v3-${{ runner.os }}- - name: Test env: HF_TOKEN: ${{ secrets.HF_TOKEN }} diff --git a/README.md b/README.md index cd425d2..c0a1317 100644 --- a/README.md +++ b/README.md @@ -32,7 +32,11 @@ cargo install sys1 --features cpu # cargo install sys1 --no-default-features --features cuda ``` -Then run it with [`convaiinnovations/laya`](https://huggingface.co/convaiinnovations/laya) (more models coming soon!). +Then run it with any of the supported models (more coming soon!). + +- [`convaiinnovations/laya`](https://huggingface.co/convaiinnovations/laya) for English text, guardrails, email triage +- [`convaiinnovations/laya-multilingual`](https://huggingface.co/convaiinnovations/laya-multilingual) for 100+ languages, ~2.2x faster +- [`convaiinnovations/laya-typed-decisions`](https://huggingface.co/convaiinnovations/laya-typed-decisions) for typed-decisions workflows ```bash sys1 --model-id convaiinnovations/laya --dtype auto diff --git a/src/main.rs b/src/main.rs index 49f6caa..4c543f5 100644 --- a/src/main.rs +++ b/src/main.rs @@ -48,6 +48,8 @@ struct Args { request_timeout_ms: u64, #[arg(long, default_value_t = 1_048_576)] max_request_bytes: usize, + #[arg(long)] + max_model_len: Option, #[arg(long, value_enum, default_value_t = Precision::Auto)] dtype: Precision, } @@ -72,6 +74,10 @@ impl Args { self.max_request_bytes > 0, "--max-request-bytes must be positive" ); + anyhow::ensure!( + self.max_model_len != Some(0), + "--max-model-len must be positive" + ); Ok(()) } } @@ -146,7 +152,7 @@ async fn main() -> anyhow::Result<()> { path = %model_path.display(), "loading model" ); - let model = models::load(&model_path, architecture, dtype) + let model = models::load(&model_path, architecture, dtype, args.max_model_len) .with_context(|| format!("failed to load model from {}", model_path.display()))?; info!( model = %served_model_name, @@ -236,6 +242,7 @@ mod tests { assert_eq!(args.max_queue_size, 256); assert_eq!(args.request_timeout_ms, 30_000); assert_eq!(args.max_request_bytes, 1_048_576); + assert_eq!(args.max_model_len, None); assert_eq!(args.dtype, Precision::Auto); } @@ -256,6 +263,16 @@ mod tests { assert_eq!(args.dtype, Precision::Bf16); } + #[test] + fn accepts_a_model_length_override() { + let args = Args::try_parse_from(["sys1", "--max-model-len", "8192"]).unwrap(); + assert_eq!(args.max_model_len, Some(8192)); + assert!(args.validate().is_ok()); + + let args = Args::try_parse_from(["sys1", "--max-model-len", "0"]).unwrap(); + assert!(args.validate().is_err()); + } + #[test] fn rejects_two_model_sources() { assert!( diff --git a/src/models/laya.rs b/src/models/laya.rs index 57049e9..06e32e2 100644 --- a/src/models/laya.rs +++ b/src/models/laya.rs @@ -157,12 +157,20 @@ pub struct Laya { } impl Laya { - pub fn load(path: &Path, dtype: DType) -> anyhow::Result { + pub fn load(path: &Path, dtype: DType, max_model_len: Option) -> anyhow::Result { let device = device::load()?; let (model_dtype, compute_dtype) = execution_dtypes(dtype, device.is_cuda()); - let config: LayaConfig = + 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"))?; + if let Some(max_model_len) = max_model_len { + anyhow::ensure!( + max_model_len <= encoder_config.max_position_embeddings(), + "--max-model-len {max_model_len} exceeds the encoder limit of {}", + encoder_config.max_position_embeddings() + ); + config.max_len = max_model_len; + } let weights = path.join("model.safetensors"); let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights], model_dtype, &device)? }; let encoder_vb = vb.clone().rename_f(|name| { @@ -189,15 +197,16 @@ impl Laya { }; let tokenizer_json = fs::read(tokenizer_path)?; let tokenizer = tokenizer::from_json(&tokenizer_json)?; - let token = |value: &str| { - tokenizer - .token_to_id(value) - .ok_or_else(|| anyhow::anyhow!("tokenizer is missing {value}")) + let token = |values: &[&str]| { + values + .iter() + .find_map(|value| tokenizer.token_to_id(value)) + .ok_or_else(|| anyhow::anyhow!("tokenizer is missing one of {values:?}")) }; - let pad_id = token("[PAD]")?; - let cls_id = token("[CLS]")?; - let sep_id = token("[SEP]")?; - let mask_id = token("[MASK]")?; + let pad_id = token(&["[PAD]", ""])?; + let cls_id = token(&["[CLS]", ""])?; + let sep_id = token(&["[SEP]", ""])?; + let mask_id = token(&["[MASK]", ""])?; Ok(Self { tokenizer, encoder, @@ -244,7 +253,7 @@ impl Laya { return Err(ApiError::new("questions must not be empty")); } let state = render_value(&request.state); - let sanitized_state = state.replace("[MASK]", " "); + let sanitized_state = sanitize_masks(&state); let state_ids = self.encode(&sanitized_state)?; let mut items = Vec::with_capacity(request.questions.len()); for (id, value) in request.questions { @@ -291,7 +300,7 @@ impl Laya { state_ids: &[u32], question: &Question, ) -> Result<(Vec, Vec), ApiError> { - let instructions = question.instructions.replace("[MASK]", " "); + let instructions = sanitize_masks(&question.instructions); let mut head = self.encode(&format!( "{} question: {instructions}", TYPES[question.kind] @@ -300,7 +309,7 @@ impl Laya { for option in &question.options { let mut ids = vec![self.mask_id]; ids.extend( - self.encode(&format!(" {}", option.replace("[MASK]", " ")))? + self.encode(&format!(" {}", sanitize_masks(option)))? .into_iter() .take(48), ); @@ -540,6 +549,10 @@ impl Laya { } } +fn sanitize_masks(value: &str) -> String { + value.replace("[MASK]", " ").replace("", " ") +} + fn linear_dtype( input: usize, output: usize, @@ -800,70 +813,89 @@ mod tests { #[tokio::test] async fn fp32_logits_and_probabilities() -> anyhow::Result<()> { - let path = crate::hub::download( - "convaiinnovations/laya", - "aa8c91ca088ec597df95a0d1c76b3063cb2ae5e8", - ) - .await?; - - let model = Laya::load(&path, DType::F32)?; - let request: DecisionRequest = serde_json::from_value(json!({ - "state": { - "message": "I was charged twice for invoice 4411. Please refund me today.", - "account_tier": "enterprise" - }, - "questions": { - "route": { - "type": "choice", - "instructions": "Where should this ticket go?", - "criteria": { - "billing": "payments, refunds, invoices", - "bug": "the product is broken", - "account": "login or access" - } - }, - "urgency": { - "type": "score", - "instructions": "How urgent is this message?", - "criteria": [ - "routine, no rush", - "today", - "urgent", - "critical, about to churn" - ] + let models = [ + ( + "laya", + "convaiinnovations/laya", + "aa8c91ca088ec597df95a0d1c76b3063cb2ae5e8", + ), + ( + "laya_typed_decisions", + "convaiinnovations/laya-typed-decisions", + "1a793eb568e6718f15941d08f85432581df534e3", + ), + ( + "laya_multilingual", + "convaiinnovations/laya-multilingual", + "e4e9ddf21a7b1903b7acffd8814ad4307bf63a67", + ), + ]; + + 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 request: DecisionRequest = serde_json::from_value(json!({ + "state": { + "message": "I was charged twice for invoice 4411. Please refund me today.", + "account_tier": "enterprise" }, - "escalate": { - "type": "noul", - "instructions": "Escalate to a human immediately?" + "questions": { + "route": { + "type": "choice", + "instructions": "Where should this ticket go?", + "criteria": { + "billing": "payments, refunds, invoices", + "bug": "the product is broken", + "account": "login or access" + } + }, + "urgency": { + "type": "score", + "instructions": "How urgent is this message?", + "criteria": [ + "routine, no rush", + "today", + "urgent", + "critical, about to churn" + ] + }, + "escalate": { + "type": "noul", + "instructions": "Escalate to a human immediately?" + } } - } - }))?; + }))?; - let prepared = model.prepare(request).expect("prepare the test request"); - let logits = model.forward(std::slice::from_ref(&prepared))?; - let probabilities: Vec<_> = prepared - .items - .iter() - .zip(&logits) - .map(|(item, logits)| { - let bucket = bucket(item.question.kind, logits.len()); - let temperature = model - .config - .temperature_by_options - .get(&bucket) - .copied() - .unwrap_or(model.config.temperature[item.question.kind]) - .clamp(0.5, 5.0); - probability(logits, temperature) - }) - .collect(); + let prepared = model.prepare(request).expect("prepare the test request"); + let logits = model.forward(std::slice::from_ref(&prepared))?; + let probabilities: Vec<_> = prepared + .items + .iter() + .zip(&logits) + .map(|(item, logits)| { + let bucket = bucket(item.question.kind, logits.len()); + let temperature = model + .config + .temperature_by_options + .get(&bucket) + .copied() + .unwrap_or(model.config.temperature[item.question.kind]) + .clamp(0.5, 5.0); + probability(logits, temperature) + }) + .collect(); - insta::assert_yaml_snapshot!("laya_fp32_logits", logits, { - "[][]" => insta::rounded_redaction(3), - }); - insta::assert_yaml_snapshot!("laya_fp32_probabilities", probabilities, { - "[][]" => insta::rounded_redaction(4), - }); + insta::assert_yaml_snapshot!(format!("{snapshot}_fp32_logits"), logits, { + "[][]" => insta::rounded_redaction(3), + }); + insta::assert_yaml_snapshot!( + format!("{snapshot}_fp32_probabilities"), + probabilities, + { + "[][]" => insta::rounded_redaction(4), + } + ); + } Ok(()) } diff --git a/src/models/mod.rs b/src/models/mod.rs index 9e393dd..498217d 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -11,6 +11,14 @@ use std::{fs, path::Path}; pub use laya::Laya; pub const LAYA_MODEL_ID: &str = "convaiinnovations/laya"; +pub const LAYA_TYPED_DECISIONS_MODEL_ID: &str = "convaiinnovations/laya-typed-decisions"; +pub const LAYA_MULTILINGUAL_MODEL_ID: &str = "convaiinnovations/laya-multilingual"; + +pub const LAYA_MODEL_IDS: &[&str] = &[ + LAYA_MODEL_ID, + LAYA_TYPED_DECISIONS_MODEL_ID, + LAYA_MULTILINGUAL_MODEL_ID, +]; #[derive(Clone, Copy, Debug, Eq, PartialEq)] pub enum Architecture { @@ -19,9 +27,13 @@ pub enum Architecture { impl Architecture { pub fn from_model_id(model_id: &str) -> anyhow::Result { - match model_id { - LAYA_MODEL_ID => Ok(Self::Laya), - _ => bail!("unsupported model id {model_id:?}; supported model: {LAYA_MODEL_ID}"), + if LAYA_MODEL_IDS.contains(&model_id) { + Ok(Self::Laya) + } else { + bail!( + "unsupported model id {model_id:?}; supported models: {}", + LAYA_MODEL_IDS.join(", ") + ) } } @@ -45,7 +57,10 @@ impl Architecture { .with_context(|| format!("failed to read {}", laya_config_path.display()))?, ) .with_context(|| format!("failed to parse {}", laya_config_path.display()))?; - if laya_config.model_name == "rl-agent" { + if matches!( + laya_config.model_name.as_str(), + "rl-agent" | "laya-typed-decisions" + ) { return Ok(Self::Laya); } } @@ -89,9 +104,14 @@ impl DecisionModel for Model { } } -pub fn load(path: &Path, architecture: Architecture, dtype: DType) -> anyhow::Result { +pub fn load( + path: &Path, + architecture: Architecture, + dtype: DType, + max_model_len: Option, +) -> anyhow::Result { match architecture { - Architecture::Laya => Laya::load(path, dtype).map(Model::Laya), + Architecture::Laya => Laya::load(path, dtype, max_model_len).map(Model::Laya), } } @@ -100,11 +120,13 @@ mod tests { use super::*; #[test] - fn routes_the_supported_hub_model() { - assert_eq!( - Architecture::from_model_id(LAYA_MODEL_ID).unwrap(), - Architecture::Laya - ); + fn routes_the_supported_hub_models() { + for model_id in LAYA_MODEL_IDS { + assert_eq!( + Architecture::from_model_id(model_id).unwrap(), + Architecture::Laya + ); + } assert!(Architecture::from_model_id("owner/other").is_err()); } } diff --git a/src/models/modernbert.rs b/src/models/modernbert.rs index 9b240a1..6a4e3bc 100644 --- a/src/models/modernbert.rs +++ b/src/models/modernbert.rs @@ -38,6 +38,10 @@ impl Config { self.hidden_size } + pub fn max_position_embeddings(&self) -> usize { + self.max_position_embeddings + } + fn global_rope_theta(&self) -> f64 { self.rope_parameters["full_attention"].rope_theta } diff --git a/src/models/snapshots/sys1__models__laya__tests__laya_multilingual_fp32_logits.snap b/src/models/snapshots/sys1__models__laya__tests__laya_multilingual_fp32_logits.snap new file mode 100644 index 0000000..0f0888a --- /dev/null +++ b/src/models/snapshots/sys1__models__laya__tests__laya_multilingual_fp32_logits.snap @@ -0,0 +1,13 @@ +--- +source: src/models/laya.rs +expression: logits +--- +- - 14.165 + - -4.521 + - -4.575 +- - -4.232 + - 0.933 + - 0.632 + - -1.522 +- - 0.591 + - -0.209 diff --git a/src/models/snapshots/sys1__models__laya__tests__laya_multilingual_fp32_probabilities.snap b/src/models/snapshots/sys1__models__laya__tests__laya_multilingual_fp32_probabilities.snap new file mode 100644 index 0000000..d45fc93 --- /dev/null +++ b/src/models/snapshots/sys1__models__laya__tests__laya_multilingual_fp32_probabilities.snap @@ -0,0 +1,13 @@ +--- +source: src/models/laya.rs +expression: probabilities +--- +- - 1 + - 0 + - 0 +- - 0.0031 + - 0.546 + - 0.404 + - 0.0469 +- - 0.6901 + - 0.3099 diff --git a/src/models/snapshots/sys1__models__laya__tests__laya_typed_decisions_fp32_logits.snap b/src/models/snapshots/sys1__models__laya__tests__laya_typed_decisions_fp32_logits.snap new file mode 100644 index 0000000..7f9adc2 --- /dev/null +++ b/src/models/snapshots/sys1__models__laya__tests__laya_typed_decisions_fp32_logits.snap @@ -0,0 +1,13 @@ +--- +source: src/models/laya.rs +expression: logits +--- +- - 1.85 + - -2.114 + - -2.29 +- - -2.152 + - 2.495 + - 1.213 + - -0.245 +- - -1.519 + - 0.176 diff --git a/src/models/snapshots/sys1__models__laya__tests__laya_typed_decisions_fp32_probabilities.snap b/src/models/snapshots/sys1__models__laya__tests__laya_typed_decisions_fp32_probabilities.snap new file mode 100644 index 0000000..9d05e34 --- /dev/null +++ b/src/models/snapshots/sys1__models__laya__tests__laya_typed_decisions_fp32_probabilities.snap @@ -0,0 +1,13 @@ +--- +source: src/models/laya.rs +expression: probabilities +--- +- - 0.8331 + - 0.0876 + - 0.0793 +- - 0.0163 + - 0.6688 + - 0.2401 + - 0.0749 +- - 0.2985 + - 0.7015 diff --git a/src/tokenizer.rs b/src/tokenizer.rs index 0913da5..b7701a0 100644 --- a/src/tokenizer.rs +++ b/src/tokenizer.rs @@ -95,20 +95,27 @@ fn migrate(value: &mut Value) -> anyhow::Result<()> { .filter(|token| (50254..=50279).contains(&token.id)) .map(|token| token.content.clone()) .collect(); - if removable.len() != 26 { + if !matches!(removable.len(), 0 | 26) { bail!("unexpected Laya added-token layout"); } - let vocab = value + if !removable.is_empty() { + let vocab = value + .get_mut("model") + .and_then(|model| model.get_mut("vocab")) + .and_then(Value::as_object_mut) + .context("tokenizer has no model vocab")?; + for token in removable { + vocab.remove(&token); + } + for (offset, atom) in MISSING_ATOMS.into_iter().enumerate() { + vocab.insert(atom.to_string(), Value::from(50254 + offset as u32)); + } + } + let model = value .get_mut("model") - .and_then(|model| model.get_mut("vocab")) .and_then(Value::as_object_mut) - .context("tokenizer has no model vocab")?; - for token in removable { - vocab.remove(&token); - } - for (offset, atom) in MISSING_ATOMS.into_iter().enumerate() { - vocab.insert(atom.to_string(), Value::from(50254 + offset as u32)); - } + .context("tokenizer has no model")?; + migrate_byte_fallback_tabs(model)?; canonicalize_value(value).map_err(|error| anyhow::anyhow!(error.to_string()))?; value .as_object_mut() @@ -117,6 +124,46 @@ fn migrate(value: &mut Value) -> anyhow::Result<()> { Ok(()) } +fn migrate_byte_fallback_tabs(model: &mut serde_json::Map) -> anyhow::Result<()> { + if model.get("byte_fallback").and_then(Value::as_bool) == Some(true) { + let vocab = model + .get_mut("vocab") + .and_then(Value::as_object_mut) + .context("tokenizer has no model vocab")?; + if !vocab.contains_key("<0x09>") { + let tab_tokens: Vec<_> = vocab + .keys() + .filter(|token| token.contains('\t')) + .cloned() + .collect(); + if tab_tokens.is_empty() { + bail!("byte-fallback tokenizer is missing <0x09>"); + } + for token in tab_tokens { + let id = vocab.remove(&token).unwrap(); + vocab.insert(token.replace('\t', "<0x09>"), id); + } + let merges = model + .get_mut("merges") + .and_then(Value::as_array_mut) + .context("byte-fallback tokenizer has no BPE merges")?; + for merge in merges { + let parts = merge + .as_array_mut() + .context("byte-fallback tokenizer has an invalid BPE merge")?; + for part in parts { + if let Some(token) = part.as_str() + && token.contains('\t') + { + *part = Value::String(token.replace('\t', "<0x09>")); + } + } + } + } + } + Ok(()) +} + fn read_added(value: &Value) -> anyhow::Result> { value .get("added_tokens") @@ -138,3 +185,29 @@ fn read_added(value: &Value) -> anyhow::Result> { }) .collect() } + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn migrates_tab_tokens_without_duplicate_vocabulary_ids() { + let mut model = json!({ + "byte_fallback": true, + "vocab": {"\t": 9, "\t\t": 10, "word": 11}, + "merges": [["\t", "\t"], ["\t\t", "\t"]] + }); + + migrate_byte_fallback_tabs(model.as_object_mut().unwrap()).unwrap(); + + assert_eq!( + model["vocab"], + json!({"<0x09>": 9, "<0x09><0x09>": 10, "word": 11}) + ); + assert_eq!( + model["merges"], + json!([["<0x09>", "<0x09>"], ["<0x09><0x09>", "<0x09>"]]) + ); + } +}