diff --git a/.github/actions/install-build-dependencies/action.yml b/.github/actions/install-build-dependencies/action.yml index 64c84668..706de7b5 100644 --- a/.github/actions/install-build-dependencies/action.yml +++ b/.github/actions/install-build-dependencies/action.yml @@ -1,5 +1,5 @@ name: install-build-dependencies -description: Install OS-specific system packages needed to build llama-cpp-bindings (CMake, libclang, GNU make). +description: Install the Linux system packages needed to build llama-cpp-bindings (CMake, libclang). runs: using: composite @@ -8,15 +8,3 @@ runs: if: runner.os == 'Linux' shell: bash run: sudo apt-get update && sudo apt-get install -y cmake libclang-dev - - - name: install windows build dependencies - if: runner.os == 'Windows' - shell: bash - run: choco install -y make - - - name: set windows libclang path - if: runner.os == 'Windows' - shell: bash - env: - WINDOWS_LIBCLANG_PATH: 'C:\Program Files\LLVM\bin' - run: echo "LIBCLANG_PATH=$WINDOWS_LIBCLANG_PATH" >> "$GITHUB_ENV" diff --git a/.github/workflows/unit-tests.yml b/.github/workflows/unit-tests.yml index f4d1bf12..93bcddd0 100644 --- a/.github/workflows/unit-tests.yml +++ b/.github/workflows/unit-tests.yml @@ -29,7 +29,7 @@ jobs: strategy: fail-fast: false matrix: - os: [ubuntu-latest, windows-latest, macos-latest] + os: [ubuntu-latest, macos-latest] defaults: run: shell: bash diff --git a/llama-cpp-bindings-build/src/android_ndk.rs b/llama-cpp-bindings-build/src/android_ndk.rs index 247d45db..d6734fbc 100644 --- a/llama-cpp-bindings-build/src/android_ndk.rs +++ b/llama-cpp-bindings-build/src/android_ndk.rs @@ -152,8 +152,6 @@ const fn detect_host_tag() -> Result<&'static str, AndroidNdkDetectionError> { Ok("darwin-x86_64") } else if cfg!(target_os = "linux") { Ok("linux-x86_64") - } else if cfg!(target_os = "windows") { - Ok("windows-x86_64") } else { Err(AndroidNdkDetectionError::UnsupportedHostPlatform) } diff --git a/llama-cpp-bindings-build/src/bindgen_config.rs b/llama-cpp-bindings-build/src/bindgen_config.rs index 90694d18..521cb055 100644 --- a/llama-cpp-bindings-build/src/bindgen_config.rs +++ b/llama-cpp-bindings-build/src/bindgen_config.rs @@ -154,10 +154,6 @@ pub fn generate_bindings( builder = configure_android_bindgen(builder, ndk, target_triple); } - if target_os.is_msvc() { - builder = configure_msvc_bindgen(builder, target_triple)?; - } - let bindings = builder.generate().map_err(BuildError::Bindgen)?; callbacks.verify_every_privatized_field_was_found()?; @@ -180,6 +176,7 @@ fn create_base_builder(llama_src: &Path, callbacks: BindingCallbacks) -> bindgen .derive_partialeq(true) .allowlist_function("ggml_.*") .allowlist_type("ggml_.*") + .allowlist_var("GGML_MAX_DIMS") .allowlist_function("gguf_.*") .allowlist_type("gguf_.*") .allowlist_function("llama_.*") @@ -192,7 +189,6 @@ fn create_base_builder(llama_src: &Path, callbacks: BindingCallbacks) -> bindgen .blocklist_function("llama_model_load_from_file_ptr") .blocklist_type("FILE") .blocklist_type("_IO_.*") - .blocklist_type("_iobuf") .prepend_enum_name(false); for function in DEPRECATED_FUNCTIONS { @@ -228,41 +224,3 @@ fn configure_android_bindgen( builder.clang_arg(format!("--target={target_triple}")) } - -fn configure_msvc_bindgen( - mut builder: bindgen::Builder, - target_triple: &str, -) -> Result { - let compiler = cc::Build::new() - .try_get_compiler() - .map_err(BuildError::NativeCompiler)?; - - let msvc_include_paths = compiler - .env() - .iter() - .find(|(key, _)| key.eq_ignore_ascii_case("INCLUDE")) - .map(|(_, value)| value.clone()); - - if let Some(include_paths) = msvc_include_paths { - for include_path in include_paths - .to_string_lossy() - .split(';') - .filter(|path| !path.is_empty()) - { - builder = builder.clang_arg("-isystem").clang_arg(include_path); - debug_log!("Added MSVC include path: {}", include_path); - } - } - - builder = builder - .clang_arg(format!("--target={target_triple}")) - .clang_arg("-fms-compatibility") - .clang_arg("-fms-extensions"); - - debug_log!( - "Configured bindgen with MSVC toolchain for target: {}", - target_triple - ); - - Ok(builder) -} diff --git a/llama-cpp-bindings-build/src/cmake_config.rs b/llama-cpp-bindings-build/src/cmake_config.rs index 8058b3a9..742b0523 100644 --- a/llama-cpp-bindings-build/src/cmake_config.rs +++ b/llama-cpp-bindings-build/src/cmake_config.rs @@ -31,9 +31,6 @@ pub fn configure_and_build(context: &BuildContext) -> Result Option<&'static str> { } } -const fn msvc_exception_handling_flag(target_os: TargetOs) -> Option<&'static str> { - if target_os.is_msvc() { - Some("/EHsc") - } else { - None - } -} - -fn configure_msvc_exception_handling(config: &mut Config, target_os: TargetOs) { - let Some(flag) = msvc_exception_handling_flag(target_os) else { - return; - }; - - config.cxxflag(flag); -} - -fn msvc_config_flags(target_os: TargetOs, profile: &str) -> Option<&'static str> { - if !target_os.is_msvc() { - return None; - } - - match profile { - "Debug" => Some("/Ob0 /Od /RTC1"), - "MinSizeRel" => Some("/O1 /Ob1 /DNDEBUG"), - "Release" => Some("/O2 /Ob2 /DNDEBUG"), - "RelWithDebInfo" => Some("/O2 /Ob1 /DNDEBUG"), - _ => None, - } -} - -fn configure_msvc_config_flags(config: &mut Config, target_os: TargetOs, profile: &str) { - let Some(flags) = msvc_config_flags(target_os, profile) else { - return; - }; - let config_suffix = profile.to_uppercase(); - - config.define(format!("CMAKE_C_FLAGS_{config_suffix}"), flags); - config.define(format!("CMAKE_CXX_FLAGS_{config_suffix}"), flags); -} - fn configure_shared_libs(config: &mut Config, build_shared_libs: bool) { config.define( "BUILD_SHARED_LIBS", @@ -207,7 +164,7 @@ fn configure_platform_specific( configure_android_cmake(config, ndk, target_triple); } } - _ => {} + TargetOs::Linux => {} } } @@ -261,16 +218,6 @@ fn configure_gpu_backends(config: &mut Config, target_os: TargetOs) -> Result<() fn configure_vulkan_linking(target_os: TargetOs) -> Result<(), BuildError> { match target_os { - TargetOs::Windows(_) => { - let vulkan_path = env::var("VULKAN_SDK").map_err(|source| BuildError::Environment { - name: "VULKAN_SDK", - source, - })?; - let vulkan_lib_path = Path::new(&vulkan_path).join("Lib"); - - println!("cargo:rustc-link-search={}", vulkan_lib_path.display()); - println!("cargo:rustc-link-lib=vulkan-1"); - } TargetOs::Linux => { match env::var("VULKAN_SDK") { Ok(vulkan_path) => { @@ -289,7 +236,7 @@ fn configure_vulkan_linking(target_os: TargetOs) -> Result<(), BuildError> { println!("cargo:rustc-link-lib=vulkan"); } - _ => (), + TargetOs::Apple(_) | TargetOs::Android => {} } Ok(()) @@ -315,67 +262,6 @@ fn configure_system_ggml(config: &mut Config) -> Result<(), BuildError> { Ok(()) } -#[cfg(test)] -mod msvc_exception_handling_tests { - use crate::target_os::TargetOs; - use crate::windows_variant::WindowsVariant; - - use super::msvc_exception_handling_flag; - - #[test] - fn msvc_targets_compile_llama_cpp_with_unwind_semantics() { - assert_eq!( - msvc_exception_handling_flag(TargetOs::Windows(WindowsVariant::Msvc)), - Some("/EHsc") - ); - } - - #[test] - fn targets_without_msvc_keep_their_toolchain_default_exception_handling() { - assert_eq!(msvc_exception_handling_flag(TargetOs::Linux), None); - assert_eq!( - msvc_exception_handling_flag(TargetOs::Windows(WindowsVariant::Other)), - None - ); - } -} - -#[cfg(test)] -mod msvc_config_flag_tests { - use crate::target_os::TargetOs; - use crate::windows_variant::WindowsVariant; - - use super::msvc_config_flags; - - #[test] - fn every_msvc_configuration_keeps_llama_cpp_assertions_compiled_out() { - let msvc = TargetOs::Windows(WindowsVariant::Msvc); - - assert_eq!(msvc_config_flags(msvc, "Debug"), Some("/Ob0 /Od /RTC1")); - assert_eq!( - msvc_config_flags(msvc, "MinSizeRel"), - Some("/O1 /Ob1 /DNDEBUG") - ); - assert_eq!( - msvc_config_flags(msvc, "Release"), - Some("/O2 /Ob2 /DNDEBUG") - ); - assert_eq!( - msvc_config_flags(msvc, "RelWithDebInfo"), - Some("/O2 /Ob1 /DNDEBUG") - ); - } - - #[test] - fn targets_and_profiles_without_msvc_defaults_keep_the_flags_cmake_chose() { - assert_eq!(msvc_config_flags(TargetOs::Linux, "Release"), None); - assert_eq!( - msvc_config_flags(TargetOs::Windows(WindowsVariant::Msvc), "Fastest"), - None - ); - } -} - #[cfg(test)] mod cpu_feature_mapping_tests { use super::map_cpu_feature_to_ggml; diff --git a/llama-cpp-bindings-build/src/cpp_wrapper.rs b/llama-cpp-bindings-build/src/cpp_wrapper.rs index bcdb7772..2c1bc17c 100644 --- a/llama-cpp-bindings-build/src/cpp_wrapper.rs +++ b/llama-cpp-bindings-build/src/cpp_wrapper.rs @@ -20,11 +20,6 @@ pub fn compile_cpp_wrappers(target_os: TargetOs) -> Result<(), BuildError> { build.file(source); } - if target_os.is_msvc() { - build.flag(format!("/std:{CPP_STANDARD}")); - build.flag("/EHsc"); - } - if target_os.is_android() && cfg!(feature = "static-stdcxx") { build.cpp_link_stdlib(None); } diff --git a/llama-cpp-bindings-build/src/lib.rs b/llama-cpp-bindings-build/src/lib.rs index 84b17d65..a690ab14 100644 --- a/llama-cpp-bindings-build/src/lib.rs +++ b/llama-cpp-bindings-build/src/lib.rs @@ -11,7 +11,6 @@ mod library_linking; mod native_library; mod rebuild_tracking; mod target_os; -mod windows_variant; use std::env; use std::path::{Path, PathBuf}; @@ -42,8 +41,6 @@ pub enum BuildError { }, #[error("generated bindings could not be written: {0}")] BindingsWrite(#[source] std::io::Error), - #[error("native compiler setup failed: {0}")] - NativeCompiler(#[source] cc::Error), #[error("native wrapper compilation failed: {0}")] NativeWrapper(#[source] cc::Error), #[error("filesystem operation failed for {path}: {source}")] @@ -88,7 +85,6 @@ pub struct BuildContext { pub cargo_cfg_target_env: String, pub build_shared_libs: bool, pub profile: String, - pub static_crt: bool, pub android_ndk: Option, } @@ -97,20 +93,17 @@ impl BuildContext { let target_triple = required_env("TARGET")?; let cargo_cfg_target_os = required_env("CARGO_CFG_TARGET_OS")?; let cargo_cfg_target_env = optional_env("CARGO_CFG_TARGET_ENV")?.unwrap_or_default(); - let target_os = TargetOs::from_cargo_cfg(&cargo_cfg_target_os, &cargo_cfg_target_env) - .ok_or_else(|| BuildError::UnsupportedTargetOs { + let target_os = TargetOs::from_cargo_cfg(&cargo_cfg_target_os).ok_or_else(|| { + BuildError::UnsupportedTargetOs { cargo_cfg_target_os: cargo_cfg_target_os.clone(), - })?; + } + })?; let out_dir = PathBuf::from(required_env("OUT_DIR")?); let manifest_dir = required_env("CARGO_MANIFEST_DIR")?; let llama_src = Path::new(&manifest_dir).join("llama.cpp"); let build_shared_libs = cfg!(feature = "dynamic-link"); let profile = native_profile(&required_env("PROFILE")?); - let static_crt = optional_env("CARGO_CFG_TARGET_FEATURE")? - .unwrap_or_default() - .split(',') - .any(|feature| feature == "crt-static"); let cargo_cfg_target_arch = required_env("CARGO_CFG_TARGET_ARCH")?; let android_ndk = if target_os.is_android() { @@ -137,7 +130,6 @@ impl BuildContext { cargo_cfg_target_env, build_shared_libs, profile, - static_crt, android_ndk, }) } diff --git a/llama-cpp-bindings-build/src/library_linking.rs b/llama-cpp-bindings-build/src/library_linking.rs index 302e351e..c779bfc3 100644 --- a/llama-cpp-bindings-build/src/library_linking.rs +++ b/llama-cpp-bindings-build/src/library_linking.rs @@ -6,7 +6,6 @@ use crate::apple_variant::AppleVariant; use crate::debug_log; use crate::native_library::NativeLibrary; use crate::target_os::TargetOs; -use crate::windows_variant::WindowsVariant; pub fn link_libraries( cmake_dir: &Path, @@ -19,7 +18,7 @@ pub fn link_libraries( emit_search_paths(cmake_dir, build_dir); link_system_ggml_paths()?; link_cmake_built_libraries(cmake_dir, build_shared_libs, profile); - link_cuda_libraries(target_os, build_shared_libs); + link_cuda_libraries(build_shared_libs); link_rocm_libraries(build_shared_libs)?; link_openmp(cargo_cfg_target_env); link_platform_system_libraries(target_os); @@ -190,7 +189,7 @@ fn emit_search_path_with_profile(lib_dir: &Path, profile: &str) { println!("cargo:rustc-link-search=native={}", profile_dir.display()); } -fn link_cuda_libraries(target_os: TargetOs, build_shared_libs: bool) { +fn link_cuda_libraries(build_shared_libs: bool) { if !cfg!(feature = "cuda") || build_shared_libs { return; } @@ -201,23 +200,6 @@ fn link_cuda_libraries(target_os: TargetOs, build_shared_libs: bool) { println!("cargo:rustc-link-search=native={}", lib_dir.display()); } - match target_os { - TargetOs::Windows(_) => link_cuda_windows(), - _ => link_cuda_unix(), - } -} - -fn link_cuda_windows() { - println!("cargo:rustc-link-lib=cudart"); - println!("cargo:rustc-link-lib=cublas"); - println!("cargo:rustc-link-lib=cublasLt"); - - if !cfg!(feature = "cuda-no-vmm") { - println!("cargo:rustc-link-lib=cuda"); - } -} - -fn link_cuda_unix() { println!("cargo:rustc-link-lib=static=cudart_static"); println!("cargo:rustc-link-lib=static=cublas_static"); println!("cargo:rustc-link-lib=static=cublasLt_static"); @@ -265,9 +247,6 @@ fn link_openmp(cargo_cfg_target_env: &str) { fn link_platform_system_libraries(target_os: TargetOs) { match target_os { - TargetOs::Windows(WindowsVariant::Msvc) => { - println!("cargo:rustc-link-lib=advapi32"); - } TargetOs::Linux => { println!("cargo:rustc-link-lib=dylib=stdc++"); } @@ -277,9 +256,6 @@ fn link_platform_system_libraries(target_os: TargetOs) { TargetOs::Android => { link_android_cpp_stdlib(); } - TargetOs::Windows(WindowsVariant::Other) => { - println!("cargo:rustc-link-lib=stdc++"); - } } } diff --git a/llama-cpp-bindings-build/src/target_os.rs b/llama-cpp-bindings-build/src/target_os.rs index ece70f63..240a3de8 100644 --- a/llama-cpp-bindings-build/src/target_os.rs +++ b/llama-cpp-bindings-build/src/target_os.rs @@ -1,9 +1,7 @@ use crate::apple_variant::AppleVariant; -use crate::windows_variant::WindowsVariant; #[derive(Debug, Clone, Copy, Eq, PartialEq)] pub enum TargetOs { - Windows(WindowsVariant), Apple(AppleVariant), Linux, Android, @@ -11,13 +9,8 @@ pub enum TargetOs { impl TargetOs { #[must_use] - pub fn from_cargo_cfg(cargo_cfg_target_os: &str, cargo_cfg_target_env: &str) -> Option { + pub fn from_cargo_cfg(cargo_cfg_target_os: &str) -> Option { match cargo_cfg_target_os { - "windows" => Some(Self::Windows(if cargo_cfg_target_env == "msvc" { - WindowsVariant::Msvc - } else { - WindowsVariant::Other - })), "macos" => Some(Self::Apple(AppleVariant::MacOS)), "ios" | "tvos" | "watchos" | "visionos" => Some(Self::Apple(AppleVariant::Other)), "android" => Some(Self::Android), @@ -30,41 +23,23 @@ impl TargetOs { pub const fn is_android(self) -> bool { matches!(self, Self::Android) } - - #[must_use] - pub const fn is_msvc(self) -> bool { - matches!(self, Self::Windows(WindowsVariant::Msvc)) - } } #[cfg(test)] mod tests { use super::TargetOs; use crate::apple_variant::AppleVariant; - use crate::windows_variant::WindowsVariant; - - #[test] - fn windows_is_split_by_its_target_environment() { - assert_eq!( - TargetOs::from_cargo_cfg("windows", "msvc"), - Some(TargetOs::Windows(WindowsVariant::Msvc)) - ); - assert_eq!( - TargetOs::from_cargo_cfg("windows", "gnu"), - Some(TargetOs::Windows(WindowsVariant::Other)) - ); - } #[test] fn macos_is_distinguished_from_the_other_apple_platforms() { assert_eq!( - TargetOs::from_cargo_cfg("macos", ""), + TargetOs::from_cargo_cfg("macos"), Some(TargetOs::Apple(AppleVariant::MacOS)) ); for apple_os in ["ios", "tvos", "watchos", "visionos"] { assert_eq!( - TargetOs::from_cargo_cfg(apple_os, ""), + TargetOs::from_cargo_cfg(apple_os), Some(TargetOs::Apple(AppleVariant::Other)), "{apple_os} must classify as a non-macOS Apple target" ); @@ -73,21 +48,8 @@ mod tests { #[test] fn android_is_not_mistaken_for_linux() { - assert_eq!( - TargetOs::from_cargo_cfg("android", ""), - Some(TargetOs::Android) - ); - assert_eq!( - TargetOs::from_cargo_cfg("linux", "gnu"), - Some(TargetOs::Linux) - ); - } - - #[test] - fn only_the_msvc_windows_variant_needs_msvc_compiler_flags() { - assert!(TargetOs::Windows(WindowsVariant::Msvc).is_msvc()); - assert!(!TargetOs::Windows(WindowsVariant::Other).is_msvc()); - assert!(!TargetOs::Linux.is_msvc()); + assert_eq!(TargetOs::from_cargo_cfg("android"), Some(TargetOs::Android)); + assert_eq!(TargetOs::from_cargo_cfg("linux"), Some(TargetOs::Linux)); } #[test] @@ -98,6 +60,8 @@ mod tests { #[test] fn an_unsupported_target_os_is_rejected() { - assert_eq!(TargetOs::from_cargo_cfg("freebsd", ""), None); + for unsupported_target_os in ["freebsd", "windows"] { + assert_eq!(TargetOs::from_cargo_cfg(unsupported_target_os), None); + } } } diff --git a/llama-cpp-bindings-build/src/windows_variant.rs b/llama-cpp-bindings-build/src/windows_variant.rs deleted file mode 100644 index 03a20b6f..00000000 --- a/llama-cpp-bindings-build/src/windows_variant.rs +++ /dev/null @@ -1,5 +0,0 @@ -#[derive(Debug, Clone, Copy, Eq, PartialEq)] -pub enum WindowsVariant { - Msvc, - Other, -} diff --git a/llama-cpp-bindings-sys/Cargo.toml b/llama-cpp-bindings-sys/Cargo.toml index 4a2a0345..b4ce0791 100644 --- a/llama-cpp-bindings-sys/Cargo.toml +++ b/llama-cpp-bindings-sys/Cargo.toml @@ -36,6 +36,9 @@ include = [ "/llama.cpp/tools/mtmd/CMakeLists.txt", "/llama.cpp/convert_hf_to_gguf.py", + "/llama.cpp/conversion/**/*.py", + "/llama.cpp/gguf-py/gguf/**/*.py", + "/llama.cpp/gguf-py/LICENSE", "/llama.cpp/common/build-info.cpp.in", "/llama.cpp/ggml/src/ggml-version.h.in", "/llama.cpp/src/llama-version.h.in", diff --git a/llama-cpp-bindings-sys/wrapper_context.cpp b/llama-cpp-bindings-sys/wrapper_context.cpp index 5129711e..e452c28c 100644 --- a/llama-cpp-bindings-sys/wrapper_context.cpp +++ b/llama-cpp-bindings-sys/wrapper_context.cpp @@ -1,8 +1,29 @@ #include "wrapper_context.h" +#include +#include + #include "llama.cpp/include/llama.h" #include "llama.cpp/src/llama-context.h" +#include "llama.cpp/src/llama-ext.h" +#include "llama.cpp/src/llama-model.h" extern "C" auto llama_rs_context_decodes_batches_in_one_micro_batch(const struct llama_context * ctx) -> bool { return llama_get_memory(ctx) == nullptr || !ctx->get_cparams().causal_attn; } + +extern "C" auto llama_rs_context_embedding_row_length(const struct llama_context * ctx) -> size_t { + return ctx->get_model().hparams.n_embd_out(); +} + +extern "C" void llama_rs_set_embeddings_nextn(struct llama_context * ctx, bool value, bool masked) { + llama_set_embeddings_nextn(ctx, value, masked); +} + +extern "C" auto llama_rs_get_embeddings_nextn_ith(struct llama_context * ctx, int32_t token_index) -> float * { + return llama_get_embeddings_nextn_ith(ctx, token_index); +} + +extern "C" auto llama_rs_model_causal_attn(const struct llama_model * model) -> bool { + return model->hparams.causal_attn; +} diff --git a/llama-cpp-bindings-sys/wrapper_context.h b/llama-cpp-bindings-sys/wrapper_context.h index 943cbd91..d5dfbce6 100644 --- a/llama-cpp-bindings-sys/wrapper_context.h +++ b/llama-cpp-bindings-sys/wrapper_context.h @@ -1,15 +1,26 @@ #pragma once #include +#include +#include #ifdef __cplusplus extern "C" { #endif struct llama_context; +struct llama_model; bool llama_rs_context_decodes_batches_in_one_micro_batch(const struct llama_context * ctx); +size_t llama_rs_context_embedding_row_length(const struct llama_context * ctx); + +void llama_rs_set_embeddings_nextn(struct llama_context * ctx, bool value, bool masked); + +float * llama_rs_get_embeddings_nextn_ith(struct llama_context * ctx, int32_t token_index); + +bool llama_rs_model_causal_attn(const struct llama_model * model); + #ifdef __cplusplus } #endif diff --git a/llama-cpp-bindings-tests/src/prime_kv_cache_with.rs b/llama-cpp-bindings-tests/src/prime_kv_cache_with.rs index cbca9334..cf3b992d 100644 --- a/llama-cpp-bindings-tests/src/prime_kv_cache_with.rs +++ b/llama-cpp-bindings-tests/src/prime_kv_cache_with.rs @@ -2,6 +2,7 @@ use anyhow::Result; use llama_cpp_bindings::context::LlamaContext; use llama_cpp_bindings::llama_batch::LlamaBatch; use llama_cpp_bindings::model::AddBos; +use llama_cpp_bindings::model::ParseSpecialTokens; use llama_cpp_test_harness::LlamaFixture; /// # Errors @@ -12,7 +13,9 @@ pub fn prime_kv_cache_with( text: &str, batch_capacity: usize, ) -> Result<()> { - let tokens = fixture.model.str_to_token(text, AddBos::Always)?; + let tokens = fixture + .model + .str_to_token(text, AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(batch_capacity, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; diff --git a/llama-cpp-bindings-tests/tests/context_state.rs b/llama-cpp-bindings-tests/tests/context_state.rs index 04963898..65cc6662 100644 --- a/llama-cpp-bindings-tests/tests/context_state.rs +++ b/llama-cpp-bindings-tests/tests/context_state.rs @@ -1,9 +1,11 @@ use std::num::NonZeroU8; +use std::num::NonZeroU32; use std::sync::Arc; use std::sync::atomic::AtomicBool; use anyhow::Result; use llama_cpp_bindings::DecodeError; +use llama_cpp_bindings::LlamaContextLoadError; use llama_cpp_bindings::LogitsError; use llama_cpp_bindings::context::LlamaContext; use llama_cpp_bindings::context::params::LlamaAttentionType; @@ -13,6 +15,7 @@ use llama_cpp_bindings::error::KvCacheSeqAddError; use llama_cpp_bindings::error::KvCacheSeqDivError; use llama_cpp_bindings::llama_batch::LlamaBatch; use llama_cpp_bindings::model::AddBos; +use llama_cpp_bindings::model::ParseSpecialTokens; use llama_cpp_bindings::model::lora_adapter_scale::LoraAdapterScale; use llama_cpp_bindings::token::LlamaToken; use llama_cpp_bindings::token::data::LlamaTokenData; @@ -105,6 +108,65 @@ fn new_context_with_huge_ctx_returns_null_error(fixture: &LlamaFixture<'_>) -> R Ok(()) } +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 256, + n_batch = 1, + n_ubatch = 1, +)] +fn new_context_rejects_more_sequences_than_one_batch_can_output( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let context_load = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params) + .into_llama_context_params() + .with_n_seq_max(2), + ); + + assert_eq!( + context_load.err(), + Some(LlamaContextLoadError::SequencesExceedOutputCapacity { + n_seq_max: 2, + output_capacity: 1, + }) + ); + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 256, + n_batch = 2, + n_ubatch = 2, +)] +fn new_context_rejects_more_sequences_than_a_causal_context_can_output( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let context_load = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params) + .into_llama_context_params() + .with_n_ctx(NonZeroU32::new(1)) + .with_n_seq_max(2), + ); + + assert_eq!( + context_load.err(), + Some(LlamaContextLoadError::SequencesExceedOutputCapacity { + n_seq_max: 2, + output_capacity: 1, + }) + ); + Ok(()) +} + #[llama_test( model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), n_gpu_layers = 999, @@ -159,7 +221,9 @@ fn decode_and_get_logits(fixture: &LlamaFixture<'_>) -> Result<()> { fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -208,7 +272,9 @@ fn token_data_array_has_entries_after_decode(fixture: &LlamaFixture<'_>) -> Resu fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -234,7 +300,9 @@ fn get_logits_ith_returns_valid_slice(fixture: &LlamaFixture<'_>) -> Result<()> fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let last_index = i32::try_from(tokens.len() - 1)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -261,7 +329,9 @@ fn token_data_array_ith_returns_valid_data(fixture: &LlamaFixture<'_>) -> Result fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let last_index = i32::try_from(tokens.len() - 1)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -337,7 +407,9 @@ fn candidates_returns_n_vocab_entries(fixture: &LlamaFixture<'_>) -> Result<()> fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -384,7 +456,9 @@ fn candidates_ith_returns_n_vocab_entries(fixture: &LlamaFixture<'_>) -> Result< fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let last_index = i32::try_from(tokens.len() - 1)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -468,7 +542,9 @@ fn encode_on_non_encoder_model_returns_error(fixture: &LlamaFixture<'_>) -> Resu fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -496,7 +572,9 @@ fn embeddings_seq_ith_returns_null_embedding_error_for_invalid_seq( fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -548,7 +626,9 @@ fn set_abort_flag_aborts_decode(fixture: &LlamaFixture<'_>) -> Result<()> { let abort_flag = Arc::new(AtomicBool::new(true)); context.set_abort_flag(abort_flag); - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -576,7 +656,9 @@ fn set_abort_flag_false_allows_decode(fixture: &LlamaFixture<'_>) -> Result<()> let abort_flag = Arc::new(AtomicBool::new(false)); context.set_abort_flag(abort_flag); - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -605,7 +687,9 @@ fn clear_abort_callback_allows_decode_with_flag_true(fixture: &LlamaFixture<'_>) context.set_abort_flag(abort_flag); context.clear_abort_callback(); - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -1579,7 +1663,10 @@ fn kv_cache_seq_div_rejects_p1_exceeding_i32_max(fixture: &LlamaFixture<'_>) -> fn save_and_load_session_file(fixture: &LlamaFixture<'_>) -> Result<()> { let mut context = fixture.build_context()?; - let tokens = fixture.model.str_to_token("Hello world", AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token("Hello world", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -1670,7 +1757,10 @@ fn get_state_size_is_positive(fixture: &LlamaFixture<'_>) -> Result<()> { fn state_seq_save_and_load_file_roundtrip(fixture: &LlamaFixture<'_>) -> Result<()> { let mut context = fixture.build_context()?; - let tokens = fixture.model.str_to_token("Hello world", AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token("Hello world", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -1727,7 +1817,10 @@ fn state_seq_save_and_load_file_roundtrip(fixture: &LlamaFixture<'_>) -> Result< fn set_state_data_rejects_a_truncated_snapshot(fixture: &LlamaFixture<'_>) -> Result<()> { let mut context = fixture.build_context()?; - let tokens = fixture.model.str_to_token("Hello world", AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token("Hello world", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -1760,7 +1853,10 @@ fn set_state_data_rejects_a_truncated_snapshot(fixture: &LlamaFixture<'_>) -> Re fn copy_state_data_and_set_state_data_roundtrip(fixture: &LlamaFixture<'_>) -> Result<()> { let mut context = fixture.build_context()?; - let tokens = fixture.model.str_to_token("Hello world", AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token("Hello world", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -1988,7 +2084,10 @@ fn state_seq_save_file_to_invalid_directory_returns_failed_to_save( fn state_load_file_with_zero_max_tokens_returns_error(fixture: &LlamaFixture<'_>) -> Result<()> { let mut context = fixture.build_context()?; - let tokens = fixture.model.str_to_token("Hello world", AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token("Hello world", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -2041,7 +2140,10 @@ fn state_seq_load_file_with_zero_max_tokens_returns_error( ) -> Result<()> { let mut context = fixture.build_context()?; - let tokens = fixture.model.str_to_token("Hello world", AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token("Hello world", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -2097,6 +2199,7 @@ fn state_load_file_with_insufficient_max_tokens_returns_length_error( let tokens = fixture.model.str_to_token( "Hello world this is a longer string for more tokens", AddBos::Always, + ParseSpecialTokens::Always, )?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -2153,6 +2256,7 @@ fn state_seq_load_file_with_insufficient_max_tokens_returns_length_error( let tokens = fixture.model.str_to_token( "Hello world this is a longer string for more tokens", AddBos::Always, + ParseSpecialTokens::Always, )?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -2169,7 +2273,6 @@ fn state_seq_load_file_with_insufficient_max_tokens_returns_length_error( Ok(()) } -#[cfg(unix)] #[llama_test( model_source = HuggingFace("unsloth/DeepSeek-R1-Distill-Llama-8B-GGUF", "DeepSeek-R1-Distill-Llama-8B-Q4_K_M.gguf"), n_gpu_layers = 999, @@ -2216,7 +2319,6 @@ fn state_save_file_with_non_utf8_path_returns_error(fixture: &LlamaFixture<'_>) Ok(()) } -#[cfg(unix)] #[llama_test( model_source = HuggingFace("unsloth/DeepSeek-R1-Distill-Llama-8B-GGUF", "DeepSeek-R1-Distill-Llama-8B-Q4_K_M.gguf"), n_gpu_layers = 999, @@ -2263,7 +2365,6 @@ fn state_load_file_with_non_utf8_path_returns_error(fixture: &LlamaFixture<'_>) Ok(()) } -#[cfg(unix)] #[llama_test( model_source = HuggingFace("unsloth/DeepSeek-R1-Distill-Llama-8B-GGUF", "DeepSeek-R1-Distill-Llama-8B-Q4_K_M.gguf"), n_gpu_layers = 999, @@ -2310,7 +2411,6 @@ fn state_seq_save_file_with_non_utf8_path_returns_error(fixture: &LlamaFixture<' Ok(()) } -#[cfg(unix)] #[llama_test( model_source = HuggingFace("unsloth/DeepSeek-R1-Distill-Llama-8B-GGUF", "DeepSeek-R1-Distill-Llama-8B-Q4_K_M.gguf"), n_gpu_layers = 999, @@ -2572,7 +2672,10 @@ fn state_seq_get_size_ext_returns_size_for_decoded_sequence( let mut context = fixture.build_context()?; - let tokens = fixture.model.str_to_token("Hello world", AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token("Hello world", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -2622,7 +2725,10 @@ fn state_seq_get_data_ext_and_set_data_ext_round_trip(fixture: &LlamaFixture<'_> let mut context = fixture.build_context()?; - let tokens = fixture.model.str_to_token("Hello world", AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token("Hello world", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -2661,7 +2767,7 @@ fn qwen35_refilled_candidate_array_matches_a_freshly_built_one( fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let prompt_tokens = model.str_to_token("Hello", AddBos::Never)?; + let prompt_tokens = model.str_to_token("Hello", AddBos::Never, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(64, 1)?; batch.add_sequence(&prompt_tokens, 0, false)?; @@ -2699,9 +2805,11 @@ fn decoding_more_tokens_than_the_micro_batch_with_non_causal_attention_returns_a .into_llama_context_params() .with_attention_type(LlamaAttentionType::NonCausal), )?; - let tokens = fixture - .model - .str_to_token(&"hello ".repeat(100), AddBos::Always)?; + let tokens = fixture.model.str_to_token( + &"hello ".repeat(100), + AddBos::Always, + ParseSpecialTokens::Always, + )?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; let n_tokens = batch.n_tokens(); diff --git a/llama-cpp-bindings-tests/tests/embedding_models.rs b/llama-cpp-bindings-tests/tests/embedding_models.rs index 9586a4ac..8fe47ec2 100644 --- a/llama-cpp-bindings-tests/tests/embedding_models.rs +++ b/llama-cpp-bindings-tests/tests/embedding_models.rs @@ -16,6 +16,7 @@ use llama_cpp_bindings::context::LlamaContext; use llama_cpp_bindings::ggml_time_us; use llama_cpp_bindings::llama_batch::LlamaBatch; use llama_cpp_bindings::model::AddBos; +use llama_cpp_bindings::model::ParseSpecialTokens; use llama_cpp_bindings_tests::prime_kv_cache::prime_kv_cache; use llama_cpp_test_harness::LlamaFixture; use llama_cpp_test_harness::llama_test; @@ -59,7 +60,7 @@ fn embedding_generation_produces_vectors(fixture: &LlamaFixture<'_>) -> Result<( let prompt = "Hello my name is"; let tokens = model - .str_to_token(prompt, AddBos::Always) + .str_to_token(prompt, AddBos::Always, ParseSpecialTokens::Always) .with_context(|| format!("failed to tokenize {prompt}"))?; let prompt_token_count = u64::try_from(tokens.len())?; @@ -159,7 +160,7 @@ fn reranking_produces_scores(fixture: &LlamaFixture<'_>) -> Result<()> { let tokens_lines_list = prompt_lines .iter() - .map(|line| model.str_to_token(line, AddBos::Always)) + .map(|line| model.str_to_token(line, AddBos::Always, ParseSpecialTokens::Always)) .collect::, _>>() .with_context(|| "failed to tokenize prompts")?; @@ -257,7 +258,9 @@ fn decode_with_embeddings_enabled(fixture: &LlamaFixture<'_>) -> Result<()> { fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -283,7 +286,9 @@ fn embeddings_seq_ith_returns_valid_embeddings(fixture: &LlamaFixture<'_>) -> Re fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -323,7 +328,10 @@ fn multi_sequence_embeddings_returns_one_embedding_per_sequence( let mut batch = LlamaBatch::new(64, 4)?; for (sequence_index, text) in inputs.iter().enumerate() { - let tokens = fixture.model.str_to_token(text, AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token(text, AddBos::Always, ParseSpecialTokens::Always)?; let sequence_id = i32::try_from(sequence_index)?; batch.add_sequence(&tokens, sequence_id, true)?; @@ -395,7 +403,10 @@ fn embeddings_returns_distinct_values_when_reused_batch_has_extra_capacity( for iteration_inputs in iterations { for (sequence_index, text) in iteration_inputs.iter().enumerate() { - let tokens = fixture.model.str_to_token(text, AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token(text, AddBos::Always, ParseSpecialTokens::Always)?; let sequence_id = i32::try_from(sequence_index)?; batch.add_sequence(&tokens, sequence_id, true)?; @@ -453,7 +464,9 @@ fn embeddings_ith_returns_valid_embeddings(fixture: &LlamaFixture<'_>) -> Result fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; let last_index = i32::try_from(tokens.len() - 1)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -506,7 +519,9 @@ fn encode_succeeds_with_encoder_model(fixture: &LlamaFixture<'_>) -> Result<()> fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("hello", AddBos::Never)?; + let tokens = fixture + .model + .str_to_token("hello", AddBos::Never, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -674,9 +689,11 @@ fn decoding_more_tokens_than_the_micro_batch_on_an_encoder_returns_an_error( fixture: &LlamaFixture<'_>, ) -> Result<()> { let mut context = fixture.build_context()?; - let tokens = fixture - .model - .str_to_token(&"hello ".repeat(100), AddBos::Always)?; + let tokens = fixture.model.str_to_token( + &"hello ".repeat(100), + AddBos::Always, + ParseSpecialTokens::Always, + )?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; let n_tokens = batch.n_tokens(); @@ -705,9 +722,11 @@ fn encoding_more_tokens_than_the_micro_batch_returns_an_error( fixture: &LlamaFixture<'_>, ) -> Result<()> { let mut context = fixture.build_context()?; - let tokens = fixture - .model - .str_to_token(&"hello ".repeat(100), AddBos::Never)?; + let tokens = fixture.model.str_to_token( + &"hello ".repeat(100), + AddBos::Never, + ParseSpecialTokens::Always, + )?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; let n_tokens = batch.n_tokens(); diff --git a/llama-cpp-bindings-tests/tests/generation_control.rs b/llama-cpp-bindings-tests/tests/generation_control.rs index 5cf1e6f8..200d0c9d 100644 --- a/llama-cpp-bindings-tests/tests/generation_control.rs +++ b/llama-cpp-bindings-tests/tests/generation_control.rs @@ -19,6 +19,7 @@ use llama_cpp_bindings::llama_batch::LlamaBatch; use llama_cpp_bindings::llguidance_sampler::create_llg_sampler; use llama_cpp_bindings::model::AddBos; use llama_cpp_bindings::model::LlamaChatMessage; +use llama_cpp_bindings::model::ParseSpecialTokens; use llama_cpp_bindings::sampled_token_section::SampledTokenSection; use llama_cpp_bindings::sampling::LlamaSampler; use llama_cpp_bindings::token::LlamaToken; @@ -66,7 +67,7 @@ fn sample_returns_result_and_succeeds_with_valid_index(fixture: &LlamaFixture<'_ (*fixture.context_params).into_llama_context_params(), )?; - let tokens = model.str_to_token("Hello", AddBos::Always)?; + let tokens = model.str_to_token("Hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -123,7 +124,7 @@ fn grammar_sampler_constrains_output_to_yes_or_no(fixture: &LlamaFixture<'_>) -> )?; let prompt = "<|im_start|>user\nIs the sky blue? Answer yes or no.<|im_end|>\n<|im_start|>assistant\n\n\n\n\n"; - let tokens = model.str_to_token(prompt, AddBos::Always)?; + let tokens = model.str_to_token(prompt, AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -217,7 +218,7 @@ fn json_schema_grammar_sampler_constrains_output_to_json(fixture: &LlamaFixture< )?; let prompt = "<|im_start|>user\nWhat is 2+2? Respond with a JSON object.<|im_end|>\n<|im_start|>assistant\n\n\n\n\n"; - let tokens = model.str_to_token(prompt, AddBos::Always)?; + let tokens = model.str_to_token(prompt, AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -310,7 +311,7 @@ fn sample_with_grammar_produces_constrained_output_in_loop( )?; let prompt = "<|im_start|>user\nIs the sky blue? yes or no<|im_end|>\n<|im_start|>assistant\n\n\n\n\n"; - let tokens = model.str_to_token(prompt, AddBos::Always)?; + let tokens = model.str_to_token(prompt, AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; @@ -401,7 +402,7 @@ fn sample_without_grammar_produces_multiple_tokens(fixture: &LlamaFixture<'_>) - let prompt = "<|im_start|>user\nSay hello<|im_end|>\n<|im_start|>assistant\n\n\n\n\n"; - let tokens = model.str_to_token(prompt, AddBos::Always)?; + let tokens = model.str_to_token(prompt, AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; @@ -570,7 +571,10 @@ fn grammar_lazy_with_null_byte_in_pattern_returns_error(fixture: &LlamaFixture<' fn grammar_lazy_returns_sampler_for_valid_grammar_with_trigger_tokens( fixture: &LlamaFixture<'_>, ) -> Result<()> { - let trigger_tokens = fixture.model.str_to_token("{", AddBos::Never)?; + let trigger_tokens = + fixture + .model + .str_to_token("{", AddBos::Never, ParseSpecialTokens::Always)?; assert!( !trigger_tokens.is_empty(), @@ -768,7 +772,9 @@ fn apply_runs_sampler_over_token_data_array(fixture: &LlamaFixture<'_>) -> Resul fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("Hi", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("Hi", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -794,7 +800,9 @@ fn sample_returns_token_after_decode(fixture: &LlamaFixture<'_>) -> Result<()> { fixture.backend, (*fixture.context_params).into_llama_context_params(), )?; - let tokens = fixture.model.str_to_token("Hello", AddBos::Always)?; + let tokens = fixture + .model + .str_to_token("Hello", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -854,7 +862,7 @@ fn raw_prompt_completion_with_timing(fixture: &LlamaFixture<'_>) -> Result<()> { let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let tokens_list = model - .str_to_token(prompt, AddBos::Always) + .str_to_token(prompt, AddBos::Always, ParseSpecialTokens::Always) .with_context(|| format!("failed to tokenize {prompt}"))?; let prompt_token_count = u64::try_from(tokens_list.len())?; @@ -999,7 +1007,7 @@ fn chat_inference_produces_coherent_output(fixture: &LlamaFixture<'_>) -> Result let prompt = model.apply_chat_template(&chat_template, &messages, true, true)?; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let tokens = model.str_to_token(&prompt, AddBos::Always)?; + let tokens = model.str_to_token(&prompt, AddBos::Always, ParseSpecialTokens::Always)?; let prompt_token_count = u64::try_from(tokens.len())?; let mut batch = LlamaBatch::new(512, 1)?; @@ -1117,7 +1125,7 @@ fn json_schema_constrains_output(fixture: &LlamaFixture<'_>) -> Result<()> { (*fixture.context_params).into_llama_context_params(), )?; - let tokens_list = model.str_to_token(prompt, AddBos::Always)?; + let tokens_list = model.str_to_token(prompt, AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; let last_index = i32::try_from(tokens_list.len())? - 1; @@ -1552,7 +1560,7 @@ fn samples_token_constrained_by_grammar(fixture: &LlamaFixture<'_>) -> Result<() )?; let prompt = "Answer yes or no:"; - let tokens = model.str_to_token(prompt, AddBos::Always)?; + let tokens = model.str_to_token(prompt, AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -1604,7 +1612,10 @@ fn samples_token_constrained_by_grammar(fixture: &LlamaFixture<'_>) -> Result<() fn reset_rolls_back_accepted_tokens(fixture: &LlamaFixture<'_>) -> Result<()> { let mut sampler = create_llg_sampler(fixture.model, "regex", REGEX_GRAMMAR)?; - let yes_tokens = fixture.model.str_to_token("yes", AddBos::Never)?; + let yes_tokens = + fixture + .model + .str_to_token("yes", AddBos::Never, ParseSpecialTokens::Always)?; let first_allowed_token = *yes_tokens .first() .ok_or_else(|| anyhow::anyhow!("the tokenizer must produce a token for \"yes\""))?; @@ -1776,7 +1787,7 @@ fn llguidance_chain_samples_a_valid_token(fixture: &LlamaFixture<'_>) -> Result< (*fixture.context_params).into_llama_context_params(), )?; - let tokens = model.str_to_token("Answer:", AddBos::Always)?; + let tokens = model.str_to_token("Answer:", AddBos::Always, ParseSpecialTokens::Always)?; let mut batch = LlamaBatch::new(512, 1)?; batch.add_sequence(&tokens, 0, false)?; context.decode(&mut batch)?; @@ -1915,7 +1926,9 @@ fn finishing_releases_an_unmatched_token_with_its_visible_and_raw_piece( ) -> Result<()> { let model = fixture.model; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let [ordinary_token] = model.str_to_token("hello", AddBos::Never)?[..] else { + let [ordinary_token] = + model.str_to_token("hello", AddBos::Never, ParseSpecialTokens::Always)?[..] + else { bail!("\"hello\" must be a single token"); }; @@ -1971,7 +1984,9 @@ fn ingest_counts_every_token_through_the_end_of_generation( ) -> Result<()> { let model = fixture.model; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let [ordinary_token] = model.str_to_token("hello", AddBos::Never)?[..] else { + let [ordinary_token] = + model.str_to_token("hello", AddBos::Never, ParseSpecialTokens::Always)?[..] + else { bail!("\"hello\" must be a single token"); }; let mut outcomes = Vec::new(); @@ -2293,7 +2308,9 @@ fn qwen35_grammar_with_a_token_reference_constrains_the_first_token( fixture: &LlamaFixture<'_>, ) -> Result<()> { let model = fixture.model; - let [think_token] = model.str_to_token("", AddBos::Never)?[..] else { + let [think_token] = + model.str_to_token("", AddBos::Never, ParseSpecialTokens::Always)?[..] + else { bail!(" must be a single token"); }; let mut context = LlamaContext::from_model( @@ -2304,6 +2321,7 @@ fn qwen35_grammar_with_a_token_reference_constrains_the_first_token( let prompt_tokens = model.str_to_token( "<|im_start|>user\nSay hi<|im_end|>\n<|im_start|>assistant\n", AddBos::Never, + ParseSpecialTokens::Always, )?; let mut batch = LlamaBatch::new(128, 1)?; diff --git a/llama-cpp-bindings-tests/tests/hidden_states.rs b/llama-cpp-bindings-tests/tests/hidden_states.rs new file mode 100644 index 00000000..bcc6f946 --- /dev/null +++ b/llama-cpp-bindings-tests/tests/hidden_states.rs @@ -0,0 +1,285 @@ +use anyhow::Result; +use llama_cpp_bindings::EmbeddingsError; +use llama_cpp_bindings::SampledToken; +use llama_cpp_bindings::context::LlamaContext; +use llama_cpp_bindings::context::params::LlamaPoolingType; +use llama_cpp_bindings::llama_batch::LlamaBatch; +use llama_cpp_bindings::model::AddBos; +use llama_cpp_bindings::model::ParseSpecialTokens; +use llama_cpp_bindings::token::LlamaToken; +use llama_cpp_test_harness::LlamaFixture; +use llama_cpp_test_harness::llama_test; + +const STATE: &str = "Shoes arrived two weeks late and in the wrong size."; +const FIRST_BRANCH: &str = " Which team should handle this ticket?"; +const SECOND_BRANCH: &str = " Is the customer angry about the delay?"; + +fn tokenize(fixture: &LlamaFixture<'_>, text: &str) -> Result> { + Ok(fixture + .model + .str_to_token(text, AddBos::Never, ParseSpecialTokens::Never)?) +} + +fn hidden_state_context<'fixture>( + fixture: &'fixture LlamaFixture<'_>, +) -> Result> { + let mut context = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params) + .into_llama_context_params() + .with_pooling_type(LlamaPoolingType::None) + .with_kv_unified(true), + )?; + + context.enable_masked_nextn_embeddings()?; + + Ok(context) +} + +fn add_tokens( + batch: &mut LlamaBatch<'_>, + tokens: &[LlamaToken], + first_position: i32, + sequence_id: i32, + output_on_last_token: bool, +) -> Result<()> { + let last_index = tokens.len() - 1; + + for (index, token) in tokens.iter().enumerate() { + batch.add( + &SampledToken::Content(*token), + first_position + i32::try_from(index)?, + &[sequence_id], + output_on_last_token && index == last_index, + )?; + } + + Ok(()) +} + +fn hidden_state_after( + fixture: &LlamaFixture<'_>, + state: &[LlamaToken], + branch: &[LlamaToken], +) -> Result> { + let mut context = hidden_state_context(fixture)?; + let mut state_batch = LlamaBatch::new(512, 1)?; + + add_tokens(&mut state_batch, state, 0, 0, false)?; + context.decode(&mut state_batch)?; + + let mut branch_batch = LlamaBatch::new(512, 1)?; + + add_tokens( + &mut branch_batch, + branch, + i32::try_from(state.len())?, + 0, + true, + )?; + context.decode(&mut branch_batch)?; + + Ok(context + .nextn_embeddings_ith(branch_batch.n_tokens() - 1)? + .to_vec()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 512, + n_batch = 512, + n_ubatch = 512, +)] +fn masked_nextn_embeddings_equal_the_final_hidden_states(fixture: &LlamaFixture<'_>) -> Result<()> { + let tokens = tokenize(fixture, STATE)?; + let last_index = i32::try_from(tokens.len() - 1)?; + + let mut embedding_context = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params) + .into_llama_context_params() + .with_embeddings(true) + .with_pooling_type(LlamaPoolingType::None), + )?; + let mut embedding_batch = LlamaBatch::new(512, 1)?; + + embedding_batch.add_sequence(&tokens, 0, true)?; + embedding_context.decode(&mut embedding_batch)?; + + let mut nextn_context = hidden_state_context(fixture)?; + let mut nextn_batch = LlamaBatch::new(512, 1)?; + + nextn_batch.add_sequence(&tokens, 0, true)?; + nextn_context.decode(&mut nextn_batch)?; + + assert_eq!( + nextn_context.nextn_embeddings_ith(last_index)?, + embedding_context.embeddings_ith(last_index)?, + ); + + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 512, + n_batch = 512, + n_ubatch = 512, +)] +fn masked_nextn_embeddings_exist_only_for_output_tokens(fixture: &LlamaFixture<'_>) -> Result<()> { + let tokens = tokenize(fixture, STATE)?; + let mut context = hidden_state_context(fixture)?; + let mut batch = LlamaBatch::new(512, 1)?; + + add_tokens(&mut batch, &tokens, 0, 0, true)?; + context.decode(&mut batch)?; + + assert_eq!( + context.nextn_embeddings_ith(0), + Err(EmbeddingsError::NextnEmbeddingUnavailable { token_index: 0 }) + ); + + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 512, + n_batch = 512, + n_ubatch = 512, +)] +fn nextn_embeddings_refuse_a_pooling_context(fixture: &LlamaFixture<'_>) -> Result<()> { + let mut context = LlamaContext::from_model( + fixture.model, + fixture.backend, + (*fixture.context_params) + .into_llama_context_params() + .with_pooling_type(LlamaPoolingType::Mean), + )?; + + assert_eq!( + context.enable_masked_nextn_embeddings(), + Err(EmbeddingsError::NextnEmbeddingsRequireNonePooling { + pooling_type: LlamaPoolingType::Mean, + }) + ); + + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 512, + n_batch = 512, + n_ubatch = 512, +)] +fn nextn_embeddings_are_unavailable_until_enabled(fixture: &LlamaFixture<'_>) -> Result<()> { + let context = fixture.build_context()?; + + assert_eq!( + context.nextn_embeddings_ith(0), + Err(EmbeddingsError::NextnEmbeddingsNotEnabled) + ); + + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 512, + n_batch = 512, + n_ubatch = 512, + n_seq_max = 2, +)] +fn a_released_fork_leaves_its_shared_state_as_an_independent_decode_would( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let state = tokenize(fixture, STATE)?; + let first_branch = tokenize(fixture, FIRST_BRANCH)?; + let second_branch = tokenize(fixture, SECOND_BRANCH)?; + let branch_position = i32::try_from(state.len())?; + + let mut context = hidden_state_context(fixture)?; + let mut state_batch = LlamaBatch::new(512, 1)?; + + add_tokens(&mut state_batch, &state, 0, 0, false)?; + context.decode(&mut state_batch)?; + context.copy_kv_cache_seq(0, 1, None, None)?; + + let mut first_branch_batch = LlamaBatch::new(512, 1)?; + + add_tokens( + &mut first_branch_batch, + &first_branch, + branch_position, + 1, + true, + )?; + context.decode(&mut first_branch_batch)?; + + let first_branch_hidden_state = context + .nextn_embeddings_ith(first_branch_batch.n_tokens() - 1)? + .to_vec(); + + context.clear_kv_cache_seq(Some(1), None, None)?; + + let mut second_branch_batch = LlamaBatch::new(512, 1)?; + + add_tokens( + &mut second_branch_batch, + &second_branch, + branch_position, + 0, + true, + )?; + context.decode(&mut second_branch_batch)?; + + assert_eq!( + first_branch_hidden_state, + hidden_state_after(fixture, &state, &first_branch)?, + ); + assert_eq!( + context.nextn_embeddings_ith(second_branch_batch.n_tokens() - 1)?, + hidden_state_after(fixture, &state, &second_branch)?.as_slice(), + ); + + Ok(()) +} + +#[llama_test( + model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), + n_gpu_layers = 999, + load_mode = Mmap, + n_ctx = 512, + n_batch = 512, + n_ubatch = 512, +)] +fn tokenizing_without_special_token_parsing_keeps_control_markup_as_text( + fixture: &LlamaFixture<'_>, +) -> Result<()> { + let markup = "<|fim_prefix|>"; + let parsed = fixture + .model + .str_to_token(markup, AddBos::Never, ParseSpecialTokens::Always)?; + let plain = fixture + .model + .str_to_token(markup, AddBos::Never, ParseSpecialTokens::Never)?; + + assert_eq!(parsed.len(), 1); + assert!(plain.len() > 1); + assert!(!plain.contains(&parsed[0])); + + Ok(()) +} diff --git a/llama-cpp-bindings-tests/tests/main.rs b/llama-cpp-bindings-tests/tests/main.rs index f147a0f9..cae2ba09 100644 --- a/llama-cpp-bindings-tests/tests/main.rs +++ b/llama-cpp-bindings-tests/tests/main.rs @@ -7,6 +7,7 @@ mod chat_protocol; mod context_state; mod embedding_models; mod generation_control; +mod hidden_states; mod model_derived_data_caching; mod model_introspection; mod model_loading_errors; diff --git a/llama-cpp-bindings-tests/tests/model_introspection.rs b/llama-cpp-bindings-tests/tests/model_introspection.rs index 0ec046ac..c8e8b12c 100644 --- a/llama-cpp-bindings-tests/tests/model_introspection.rs +++ b/llama-cpp-bindings-tests/tests/model_introspection.rs @@ -8,6 +8,7 @@ use llama_cpp_bindings::SampledToken; use llama_cpp_bindings::context::params::LlamaContextParams; use llama_cpp_bindings::max_devices; use llama_cpp_bindings::model::AddBos; +use llama_cpp_bindings::model::ParseSpecialTokens; use llama_cpp_bindings::model::params::LlamaModelParams; use llama_cpp_bindings::model::params::fit_extra_model::FitExtraModel; use llama_cpp_test_harness::LlamaFixture; @@ -1184,7 +1185,7 @@ fn token_attr_returns_attrs_for_bos(fixture: &LlamaFixture<'_>) -> Result<()> { )] fn str_to_token_roundtrip(fixture: &LlamaFixture<'_>) -> Result<()> { let model = fixture.model; - let tokens = model.str_to_token("hello world", AddBos::Never)?; + let tokens = model.str_to_token("hello world", AddBos::Never, ParseSpecialTokens::Always)?; assert!(!tokens.is_empty()); let mut decoder = encoding_rs::UTF_8.new_decoder(); let piece = @@ -1231,9 +1232,10 @@ fn str_to_token_grows_buffer_when_initial_estimation_too_small( fixture: &LlamaFixture<'_>, ) -> Result<()> { let many_short_chars = "a b c d e f g h i j k l"; - let tokens = fixture - .model - .str_to_token(many_short_chars, AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token(many_short_chars, AddBos::Always, ParseSpecialTokens::Always)?; assert!( tokens.len() > 8, @@ -1278,8 +1280,10 @@ fn str_to_token_grows_buffer_when_initial_estimation_too_small( )] fn str_to_token_with_add_bos_never(fixture: &LlamaFixture<'_>) -> Result<()> { let model = fixture.model; - let tokens_with_bos = model.str_to_token("hello", AddBos::Always)?; - let tokens_without_bos = model.str_to_token("hello", AddBos::Never)?; + let tokens_with_bos = + model.str_to_token("hello", AddBos::Always, ParseSpecialTokens::Always)?; + let tokens_without_bos = + model.str_to_token("hello", AddBos::Never, ParseSpecialTokens::Always)?; assert!(tokens_with_bos.len() >= tokens_without_bos.len()); @@ -1326,7 +1330,10 @@ fn str_to_token_with_many_tokens_triggers_buffer_resize(fixture: &LlamaFixture<' accumulator }); - let tokens = fixture.model.str_to_token(&many_numbers, AddBos::Always)?; + let tokens = + fixture + .model + .str_to_token(&many_numbers, AddBos::Always, ParseSpecialTokens::Always)?; assert!(tokens.len() > many_numbers.len() / 2); @@ -1367,7 +1374,7 @@ fn str_to_token_with_many_tokens_triggers_buffer_resize(fixture: &LlamaFixture<' )] fn token_to_piece_bytes_returns_bytes_for_known_token(fixture: &LlamaFixture<'_>) -> Result<()> { let model = fixture.model; - let tokens = model.str_to_token("hello", AddBos::Never)?; + let tokens = model.str_to_token("hello", AddBos::Never, ParseSpecialTokens::Always)?; let bytes = model.token_to_piece_bytes(tokens[0], 32, false, None)?; assert!(!bytes.is_empty()); @@ -1455,7 +1462,7 @@ fn token_to_piece_bytes_insufficient_buffer_returns_error( fixture: &LlamaFixture<'_>, ) -> Result<()> { let model = fixture.model; - let tokens = model.str_to_token("hello", AddBos::Never)?; + let tokens = model.str_to_token("hello", AddBos::Never, ParseSpecialTokens::Always)?; let result = model.token_to_piece_bytes(tokens[0], 1, false, None); assert!( @@ -1502,7 +1509,7 @@ fn token_to_piece_bytes_insufficient_buffer_returns_error( fn token_to_piece_with_lstrip(fixture: &LlamaFixture<'_>) -> Result<()> { let model = fixture.model; let mut decoder = encoding_rs::UTF_8.new_decoder(); - let tokens = model.str_to_token("hello", AddBos::Never)?; + let tokens = model.str_to_token("hello", AddBos::Never, ParseSpecialTokens::Always)?; let result = model.token_to_piece( &SampledToken::Content(tokens[0]), &mut decoder, @@ -1549,7 +1556,7 @@ fn token_to_piece_with_lstrip(fixture: &LlamaFixture<'_>) -> Result<()> { fn token_to_piece_decodes_reasoning_variant(fixture: &LlamaFixture<'_>) -> Result<()> { let model = fixture.model; let mut decoder = encoding_rs::UTF_8.new_decoder(); - let tokens = model.str_to_token("hi", AddBos::Never)?; + let tokens = model.str_to_token("hi", AddBos::Never, ParseSpecialTokens::Always)?; let piece = model.token_to_piece( &SampledToken::Reasoning(tokens[0]), @@ -1597,7 +1604,7 @@ fn token_to_piece_decodes_reasoning_variant(fixture: &LlamaFixture<'_>) -> Resul fn token_to_piece_decodes_tool_call_variant(fixture: &LlamaFixture<'_>) -> Result<()> { let model = fixture.model; let mut decoder = encoding_rs::UTF_8.new_decoder(); - let tokens = model.str_to_token("hi", AddBos::Never)?; + let tokens = model.str_to_token("hi", AddBos::Never, ParseSpecialTokens::Always)?; let piece = model.token_to_piece(&SampledToken::ToolCall(tokens[0]), &mut decoder, true, None)?; @@ -1641,7 +1648,7 @@ fn token_to_piece_decodes_tool_call_variant(fixture: &LlamaFixture<'_>) -> Resul fn token_to_piece_decodes_undeterminable_variant(fixture: &LlamaFixture<'_>) -> Result<()> { let model = fixture.model; let mut decoder = encoding_rs::UTF_8.new_decoder(); - let tokens = model.str_to_token("hi", AddBos::Never)?; + let tokens = model.str_to_token("hi", AddBos::Never, ParseSpecialTokens::Always)?; let piece = model.token_to_piece( &SampledToken::Undeterminable(tokens[0]), diff --git a/llama-cpp-bindings-tests/tests/model_loading_errors.rs b/llama-cpp-bindings-tests/tests/model_loading_errors.rs index 8fdc86bd..e8b151a7 100644 --- a/llama-cpp-bindings-tests/tests/model_loading_errors.rs +++ b/llama-cpp-bindings-tests/tests/model_loading_errors.rs @@ -54,7 +54,6 @@ fn load_model_with_invalid_file_content_returns_unloadable_or_reported( Ok(()) } -#[cfg(unix)] #[llama_test( model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), n_gpu_layers = 999, @@ -121,7 +120,6 @@ fn lora_adapter_init_with_invalid_gguf_returns_unloadable( Ok(()) } -#[cfg(unix)] #[llama_test( model_source = HuggingFace("unsloth/Qwen3.5-0.8B-GGUF", "Qwen3.5-0.8B-Q4_K_M.gguf"), n_gpu_layers = 999, diff --git a/llama-cpp-bindings-tests/tests/structured_chat_output.rs b/llama-cpp-bindings-tests/tests/structured_chat_output.rs index 55756a8f..a8c9f94f 100644 --- a/llama-cpp-bindings-tests/tests/structured_chat_output.rs +++ b/llama-cpp-bindings-tests/tests/structured_chat_output.rs @@ -15,6 +15,7 @@ use llama_cpp_bindings::llama_batch::LlamaBatch; use llama_cpp_bindings::model::AddBos; use llama_cpp_bindings::model::LlamaChatMessage; use llama_cpp_bindings::model::LlamaModel; +use llama_cpp_bindings::model::ParseSpecialTokens; use llama_cpp_bindings::sampling::LlamaSampler; use llama_cpp_bindings_tests::classify_sample_loop::ClassifySampleLoop; use llama_cpp_bindings_tests::classify_sample_loop::ClassifySampleLoopOutcome; @@ -68,8 +69,11 @@ fn deepseek_r1_8b_classifier_does_not_emit_reasoning_for_thinking_disabled_promp let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = - model.str_to_token(DEEPSEEK_R1_8B_THINKING_DISABLED_PROMPT, AddBos::Never)?; + let prompt_tokens = model.str_to_token( + DEEPSEEK_R1_8B_THINKING_DISABLED_PROMPT, + AddBos::Never, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -178,7 +182,11 @@ fn deepseek_r1_8b_classifier_emits_reasoning_for_thinking_enabled_prompt( let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(DEEPSEEK_R1_8B_THINKING_PROMPT, AddBos::Never)?; + let prompt_tokens = model.str_to_token( + DEEPSEEK_R1_8B_THINKING_PROMPT, + AddBos::Never, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -631,7 +639,11 @@ fn gemma4_classifier_does_not_emit_reasoning_for_thinking_disabled_prompt( let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(GEMMA4_THINKING_DISABLED_PROMPT, AddBos::Never)?; + let prompt_tokens = model.str_to_token( + GEMMA4_THINKING_DISABLED_PROMPT, + AddBos::Never, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -731,7 +743,11 @@ fn gemma4_classifier_emits_reasoning_for_thinking_prompt(fixture: &LlamaFixture< let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(GEMMA4_THINKING_PROMPT, AddBos::Never)?; + let prompt_tokens = model.str_to_token( + GEMMA4_THINKING_PROMPT, + AddBos::Never, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -936,7 +952,11 @@ What is 2 + 2? let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(GLM47_THINKING_DISABLED_PROMPT, AddBos::Never)?; + let prompt_tokens = model.str_to_token( + GLM47_THINKING_DISABLED_PROMPT, + AddBos::Never, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -1015,7 +1035,11 @@ What is 2 + 2? let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(GLM47_THINKING_PROMPT, AddBos::Never)?; + let prompt_tokens = model.str_to_token( + GLM47_THINKING_PROMPT, + AddBos::Never, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -1192,7 +1216,11 @@ fn mistral3_classifier_does_not_emit_reasoning_for_thinking_disabled_prompt( let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(MISTRAL3_THINKING_DISABLED_PROMPT, AddBos::Always)?; + let prompt_tokens = model.str_to_token( + MISTRAL3_THINKING_DISABLED_PROMPT, + AddBos::Always, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -1268,7 +1296,11 @@ to the user.[/THINK]Here, provide a self-contained response.[/SYSTEM_PROMPT]\ let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(MISTRAL3_THINKING_PROMPT, AddBos::Always)?; + let prompt_tokens = model.str_to_token( + MISTRAL3_THINKING_PROMPT, + AddBos::Always, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -1402,7 +1434,7 @@ fn qwen35_chat_inference_emits_reasoning_when_template_auto_opens( let prompt = model.apply_chat_template(&chat_template, &messages, true, true)?; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let tokens = model.str_to_token(&prompt, AddBos::Always)?; + let tokens = model.str_to_token(&prompt, AddBos::Always, ParseSpecialTokens::Always)?; let prompt_token_count = u64::try_from(tokens.len())?; let mut batch = LlamaBatch::new(512, 1)?; @@ -1523,7 +1555,11 @@ What is 2 + 2?<|im_end|> let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(QWEN35_THINKING_DISABLED_PROMPT, AddBos::Never)?; + let prompt_tokens = model.str_to_token( + QWEN35_THINKING_DISABLED_PROMPT, + AddBos::Never, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -1602,7 +1638,11 @@ What is 2 + 2?<|im_end|> let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(QWEN35_THINKING_PROMPT, AddBos::Never)?; + let prompt_tokens = model.str_to_token( + QWEN35_THINKING_PROMPT, + AddBos::Never, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -1983,7 +2023,7 @@ fn qwen36_chat_inference_emits_reasoning_when_template_auto_opens( let prompt = model.apply_chat_template(&chat_template, &messages, true, true)?; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let tokens = model.str_to_token(&prompt, AddBos::Always)?; + let tokens = model.str_to_token(&prompt, AddBos::Always, ParseSpecialTokens::Always)?; let prompt_token_count = u64::try_from(tokens.len())?; let mut batch = LlamaBatch::new(512, 1)?; @@ -2052,7 +2092,11 @@ What is 2 + 2?<|im_end|> let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(QWEN36_THINKING_DISABLED_PROMPT, AddBos::Never)?; + let prompt_tokens = model.str_to_token( + QWEN36_THINKING_DISABLED_PROMPT, + AddBos::Never, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -2131,7 +2175,11 @@ What is 2 + 2?<|im_end|> let backend = fixture.backend; let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; - let prompt_tokens = model.str_to_token(QWEN36_THINKING_PROMPT, AddBos::Never)?; + let prompt_tokens = model.str_to_token( + QWEN36_THINKING_PROMPT, + AddBos::Never, + ParseSpecialTokens::Always, + )?; let prompt_token_count = u64::try_from(prompt_tokens.len())?; let mut batch = LlamaBatch::new(2048, 1)?; @@ -2199,7 +2247,7 @@ fn visible_text_of_generation_ending_after( let mut classifier = model.sampled_token_classifier(BareJsonToolCalls::Detect)?; let mut outcomes = Vec::new(); - for token in model.str_to_token(generated_text, AddBos::Never)? { + for token in model.str_to_token(generated_text, AddBos::Never, ParseSpecialTokens::Always)? { assert_eq!( classifier.ingest(token, &mut outcomes)?, GenerationProgress::Continues diff --git a/llama-cpp-bindings/src/context.rs b/llama-cpp-bindings/src/context.rs index 7dfc01b2..bc89a1ad 100644 --- a/llama-cpp-bindings/src/context.rs +++ b/llama-cpp-bindings/src/context.rs @@ -10,6 +10,8 @@ use std::sync::atomic::Ordering; use llama_cpp_ffi_status::read_and_free_cpp_string; use crate::context::params::LlamaContextParams; +use crate::context::params::LlamaPoolingType; +use crate::context::sequence_output_capacity::SequenceOutputCapacity; use crate::llama_backend::LlamaBackend; use crate::llama_batch::LlamaBatch; use crate::model::LlamaModel; @@ -245,6 +247,7 @@ pub mod params; pub mod rope_scaling_type; pub mod save_seq_state_error; pub mod save_session_error; +mod sequence_output_capacity; pub mod session; pub mod state_data_error; @@ -260,6 +263,7 @@ pub struct LlamaContext<'model> { abort_flag: Option>, initialized_logits: Vec, embeddings_enabled: bool, + nextn_embeddings_enabled: bool, } impl Debug for LlamaContext<'_> { @@ -283,18 +287,40 @@ impl<'model> LlamaContext<'model> { abort_flag: None, initialized_logits: Vec::new(), embeddings_enabled, + nextn_embeddings_enabled: false, } } /// # Errors /// - /// Returns [`LlamaContextLoadError`] when llama.cpp fails to allocate the context. + /// Returns [`LlamaContextLoadError`] when llama.cpp fails to allocate the context, or when + /// the context has more sequences than one batch can output. pub fn from_model( model: &'model LlamaModel, _backend: &LlamaBackend, params: LlamaContextParams, ) -> Result { let context_params = params.context_params; + + SequenceOutputCapacity { + attention_type: params.attention_type(), + model_causal_attention: unsafe { + llama_cpp_bindings_sys::llama_rs_model_causal_attn(model.model.as_ptr()) + }, + model_has_encoder: unsafe { + llama_cpp_bindings_sys::llama_model_has_encoder(model.model.as_ptr()) + }, + model_n_ctx_train: unsafe { + llama_cpp_bindings_sys::llama_model_n_ctx_train(model.model.as_ptr()) + } + .cast_unsigned(), + n_batch: context_params.n_batch, + n_ctx: context_params.n_ctx, + n_outputs_max: context_params.n_outputs_max, + n_seq_max: context_params.n_seq_max, + } + .validate()?; + let mut out_ctx: *mut llama_cpp_bindings_sys::llama_context = std::ptr::null_mut(); let mut out_error: *mut std::os::raw::c_char = std::ptr::null_mut(); let status = unsafe { @@ -495,6 +521,58 @@ impl<'model> LlamaContext<'model> { } } + /// # Errors + /// + /// - When the context pools its outputs, because llama.cpp only extracts `NextN` embeddings + /// with `LLAMA_POOLING_TYPE_NONE`. + pub fn enable_masked_nextn_embeddings(&mut self) -> Result<(), EmbeddingsError> { + let pooling_type = LlamaPoolingType::from(unsafe { + llama_cpp_bindings_sys::llama_pooling_type(self.context.as_ptr()) + }); + + if pooling_type != LlamaPoolingType::None { + return Err(EmbeddingsError::NextnEmbeddingsRequireNonePooling { pooling_type }); + } + + unsafe { + llama_cpp_bindings_sys::llama_rs_set_embeddings_nextn( + self.context.as_ptr(), + true, + true, + ); + } + + self.nextn_embeddings_enabled = true; + + Ok(()) + } + + /// # Errors + /// + /// - When masked `NextN` embeddings were not enabled on this context. + /// - When the given token was not marked as an output of the last decoded batch. + pub fn nextn_embeddings_ith(&self, token_index: i32) -> Result<&[f32], EmbeddingsError> { + if !self.nextn_embeddings_enabled { + return Err(EmbeddingsError::NextnEmbeddingsNotEnabled); + } + + let row_length = unsafe { + llama_cpp_bindings_sys::llama_rs_context_embedding_row_length(self.context.as_ptr()) + }; + let embedding = unsafe { + llama_cpp_bindings_sys::llama_rs_get_embeddings_nextn_ith( + self.context.as_ptr(), + token_index, + ) + }; + + if embedding.is_null() { + Err(EmbeddingsError::NextnEmbeddingUnavailable { token_index }) + } else { + Ok(unsafe { slice::from_raw_parts(embedding, row_length) }) + } + } + /// # Errors /// Returns `LogitsError` if logits are null or `n_vocab` overflows. pub fn candidates(&self) -> Result + '_, LogitsError> { diff --git a/llama-cpp-bindings/src/context/sequence_output_capacity.rs b/llama-cpp-bindings/src/context/sequence_output_capacity.rs new file mode 100644 index 00000000..5192bf13 --- /dev/null +++ b/llama-cpp-bindings/src/context/sequence_output_capacity.rs @@ -0,0 +1,170 @@ +use crate::LlamaContextLoadError; +use crate::context::params::LlamaAttentionType; + +pub struct SequenceOutputCapacity { + pub attention_type: LlamaAttentionType, + pub model_causal_attention: bool, + pub model_has_encoder: bool, + pub model_n_ctx_train: u32, + pub n_batch: u32, + pub n_ctx: u32, + pub n_outputs_max: u32, + pub n_seq_max: u32, +} + +impl SequenceOutputCapacity { + /// # Errors + /// Returns [`LlamaContextLoadError::SequencesExceedOutputCapacity`] when there are more + /// sequences than one batch can output. + pub fn validate(&self) -> Result<(), LlamaContextLoadError> { + let n_ctx = if self.n_ctx == 0 { + self.model_n_ctx_train + } else { + self.n_ctx + }; + let causal_attention = match self.attention_type { + LlamaAttentionType::Unspecified => self.model_causal_attention, + LlamaAttentionType::Causal => true, + LlamaAttentionType::NonCausal => false, + }; + let n_batch = if causal_attention { + n_ctx.min(self.n_batch) + } else { + self.n_batch + }; + let output_capacity = if self.n_outputs_max == 0 || self.model_has_encoder { + n_batch + } else { + self.n_outputs_max + }; + let n_seq_max = self.n_seq_max.max(1); + + if n_seq_max > output_capacity { + return Err(LlamaContextLoadError::SequencesExceedOutputCapacity { + n_seq_max, + output_capacity, + }); + } + + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::SequenceOutputCapacity; + use crate::LlamaContextLoadError; + use crate::context::params::LlamaAttentionType; + + const TWO_SEQUENCES_PAST_ONE_OUTPUT: Result<(), LlamaContextLoadError> = + Err(LlamaContextLoadError::SequencesExceedOutputCapacity { + n_seq_max: 2, + output_capacity: 1, + }); + + fn two_sequences() -> SequenceOutputCapacity { + SequenceOutputCapacity { + attention_type: LlamaAttentionType::NonCausal, + model_causal_attention: false, + model_has_encoder: false, + model_n_ctx_train: 4096, + n_batch: 2, + n_ctx: 4096, + n_outputs_max: 0, + n_seq_max: 2, + } + } + + #[test] + fn the_default_output_capacity_is_one_batch() { + assert_eq!( + SequenceOutputCapacity { + n_batch: 1, + ..two_sequences() + } + .validate(), + TWO_SEQUENCES_PAST_ONE_OUTPUT + ); + } + + #[test] + fn an_explicit_output_capacity_admits_more_sequences_than_one_batch() { + assert_eq!( + SequenceOutputCapacity { + n_batch: 1, + n_outputs_max: 2, + ..two_sequences() + } + .validate(), + Ok(()) + ); + } + + #[test] + fn an_encoder_model_outputs_at_most_one_batch() { + assert_eq!( + SequenceOutputCapacity { + model_has_encoder: true, + n_batch: 1, + n_outputs_max: 2, + ..two_sequences() + } + .validate(), + TWO_SEQUENCES_PAST_ONE_OUTPUT + ); + } + + #[test] + fn a_causal_context_outputs_at_most_its_context_size() { + assert_eq!( + SequenceOutputCapacity { + attention_type: LlamaAttentionType::Causal, + n_ctx: 1, + ..two_sequences() + } + .validate(), + TWO_SEQUENCES_PAST_ONE_OUTPUT + ); + } + + #[test] + fn a_non_causal_context_outputs_a_whole_batch_whatever_its_context_size() { + assert_eq!( + SequenceOutputCapacity { + model_causal_attention: true, + n_ctx: 1, + ..two_sequences() + } + .validate(), + Ok(()) + ); + } + + #[test] + fn an_unspecified_attention_type_follows_the_model() { + assert_eq!( + SequenceOutputCapacity { + attention_type: LlamaAttentionType::Unspecified, + model_causal_attention: true, + n_ctx: 1, + ..two_sequences() + } + .validate(), + TWO_SEQUENCES_PAST_ONE_OUTPUT + ); + } + + #[test] + fn a_context_without_a_size_takes_the_training_context_size() { + assert_eq!( + SequenceOutputCapacity { + attention_type: LlamaAttentionType::Causal, + model_n_ctx_train: 1, + n_ctx: 0, + ..two_sequences() + } + .validate(), + TWO_SEQUENCES_PAST_ONE_OUTPUT + ); + } +} diff --git a/llama-cpp-bindings/src/error/embeddings_error.rs b/llama-cpp-bindings/src/error/embeddings_error.rs index 9555f196..3bf58bc6 100644 --- a/llama-cpp-bindings/src/error/embeddings_error.rs +++ b/llama-cpp-bindings/src/error/embeddings_error.rs @@ -8,4 +8,16 @@ pub enum EmbeddingsError { NonePoolType, #[error("Invalid embedding dimension: {0}")] InvalidEmbeddingDimension(#[source] std::num::TryFromIntError), + #[error( + "NextN embeddings need a context with LLAMA_POOLING_TYPE_NONE, but it pools with {pooling_type:?}" + )] + NextnEmbeddingsRequireNonePooling { + pooling_type: crate::context::params::LlamaPoolingType, + }, + #[error("NextN embeddings weren't enabled on this context")] + NextnEmbeddingsNotEnabled, + #[error( + "No NextN embedding exists for token {token_index}; it was not marked as an output of the last decoded batch" + )] + NextnEmbeddingUnavailable { token_index: i32 }, } diff --git a/llama-cpp-bindings/src/error/llama_context_load_error.rs b/llama-cpp-bindings/src/error/llama_context_load_error.rs index 145ad622..b0145ce3 100644 --- a/llama-cpp-bindings/src/error/llama_context_load_error.rs +++ b/llama-cpp-bindings/src/error/llama_context_load_error.rs @@ -12,4 +12,11 @@ pub enum LlamaContextLoadError { LlamaCppOutOfMemory, #[error("{message}")] Reported { message: String }, + #[error( + "{n_seq_max} sequences need an output each, but a batch holds at most {output_capacity} outputs" + )] + SequencesExceedOutputCapacity { + n_seq_max: u32, + output_capacity: u32, + }, } diff --git a/llama-cpp-bindings/src/gguf_context.rs b/llama-cpp-bindings/src/gguf_context.rs index 743f28f3..ca86dbc2 100644 --- a/llama-cpp-bindings/src/gguf_context.rs +++ b/llama-cpp-bindings/src/gguf_context.rs @@ -1,13 +1,21 @@ use std::ffi::{CStr, CString}; -use std::path::Path; +use std::fs::File; +use std::io; +use std::os::unix::fs::FileExt as _; +use std::path::{Path, PathBuf}; use std::ptr::NonNull; +use std::slice; use crate::gguf_context_error::GgufContextError; +use crate::gguf_tensor_f32::GgufTensorF32; use crate::gguf_type::GgufType; +const F32_SIZE_IN_BYTES: usize = size_of::(); + #[derive(Debug)] pub struct GgufContext { context: NonNull, + path: PathBuf, } impl GgufContext { @@ -33,7 +41,10 @@ impl GgufContext { let context = NonNull::new(raw) .ok_or_else(|| GgufContextError::InitFailed(path_ref.to_path_buf()))?; - Ok(Self { context }) + Ok(Self { + context, + path: path_ref.to_path_buf(), + }) } #[must_use] @@ -59,14 +70,13 @@ impl GgufContext { Ok(index) } - /// # Safety considerations - /// - /// The caller must ensure `key_id` is in range `[0, n_kv())`. - /// /// # Errors /// + /// Returns [`GgufContextError::KeyIdOutOfRange`] if the file holds no key with this id. /// Returns [`GgufContextError::Utf8Error`] if the key name is not valid UTF-8. pub fn key_at(&self, key_id: i64) -> Result<&str, GgufContextError> { + self.require_key_id_in_range(key_id)?; + let c_str = unsafe { CStr::from_ptr(llama_cpp_bindings_sys::gguf_get_key( self.context.as_ptr(), @@ -77,49 +87,67 @@ impl GgufContext { Ok(c_str.to_str()?) } - /// # Safety considerations + /// # Errors /// - /// The caller must ensure `key_id` is in range `[0, n_kv())`. - #[must_use] - pub fn kv_type(&self, key_id: i64) -> Option { - let raw = + /// Returns [`GgufContextError::KeyIdOutOfRange`] if the file holds no key with this id. + /// Returns [`GgufContextError::UnknownValueType`] if the value type is not a known GGUF type. + pub fn kv_type(&self, key_id: i64) -> Result { + self.require_key_id_in_range(key_id)?; + + let raw_type = unsafe { llama_cpp_bindings_sys::gguf_get_kv_type(self.context.as_ptr(), key_id) }; - GgufType::from_raw(raw) + GgufType::from_raw(raw_type).ok_or(GgufContextError::UnknownValueType { key_id, raw_type }) } - /// # Safety considerations + /// # Errors /// - /// The caller must ensure the key at `key_id` has type [`GgufType::Uint32`]. - #[must_use] - pub fn val_u32(&self, key_id: i64) -> u32 { - unsafe { llama_cpp_bindings_sys::gguf_get_val_u32(self.context.as_ptr(), key_id) } + /// Returns [`GgufContextError::KeyIdOutOfRange`] or [`GgufContextError::ValueTypeMismatch`] + /// if the key does not hold a [`GgufType::Uint32`] value. + pub fn val_u32(&self, key_id: i64) -> Result { + self.require_value_type(key_id, GgufType::Uint32)?; + + Ok(unsafe { llama_cpp_bindings_sys::gguf_get_val_u32(self.context.as_ptr(), key_id) }) } - /// # Safety considerations + /// # Errors /// - /// The caller must ensure the key at `key_id` has type [`GgufType::Int32`]. - #[must_use] - pub fn val_i32(&self, key_id: i64) -> i32 { - unsafe { llama_cpp_bindings_sys::gguf_get_val_i32(self.context.as_ptr(), key_id) } + /// Returns [`GgufContextError::KeyIdOutOfRange`] or [`GgufContextError::ValueTypeMismatch`] + /// if the key does not hold a [`GgufType::Int32`] value. + pub fn val_i32(&self, key_id: i64) -> Result { + self.require_value_type(key_id, GgufType::Int32)?; + + Ok(unsafe { llama_cpp_bindings_sys::gguf_get_val_i32(self.context.as_ptr(), key_id) }) } - /// # Safety considerations + /// # Errors /// - /// The caller must ensure the key at `key_id` has type [`GgufType::Uint64`]. - #[must_use] - pub fn val_u64(&self, key_id: i64) -> u64 { - unsafe { llama_cpp_bindings_sys::gguf_get_val_u64(self.context.as_ptr(), key_id) } + /// Returns [`GgufContextError::KeyIdOutOfRange`] or [`GgufContextError::ValueTypeMismatch`] + /// if the key does not hold a [`GgufType::Uint64`] value. + pub fn val_u64(&self, key_id: i64) -> Result { + self.require_value_type(key_id, GgufType::Uint64)?; + + Ok(unsafe { llama_cpp_bindings_sys::gguf_get_val_u64(self.context.as_ptr(), key_id) }) } - /// # Safety considerations - /// - /// The caller must ensure the key at `key_id` has type [`GgufType::String`]. + /// # Errors /// + /// Returns [`GgufContextError::KeyIdOutOfRange`] or [`GgufContextError::ValueTypeMismatch`] + /// if the key does not hold a [`GgufType::Float32`] value. + pub fn val_f32(&self, key_id: i64) -> Result { + self.require_value_type(key_id, GgufType::Float32)?; + + Ok(unsafe { llama_cpp_bindings_sys::gguf_get_val_f32(self.context.as_ptr(), key_id) }) + } + /// # Errors /// + /// Returns [`GgufContextError::KeyIdOutOfRange`] or [`GgufContextError::ValueTypeMismatch`] + /// if the key does not hold a [`GgufType::String`] value. /// Returns [`GgufContextError::Utf8Error`] if the string value is not valid UTF-8. pub fn val_str(&self, key_id: i64) -> Result<&str, GgufContextError> { + self.require_value_type(key_id, GgufType::String)?; + let c_str = unsafe { CStr::from_ptr(llama_cpp_bindings_sys::gguf_get_val_str( self.context.as_ptr(), @@ -134,6 +162,116 @@ impl GgufContext { pub fn n_tensors(&self) -> i64 { unsafe { llama_cpp_bindings_sys::gguf_get_n_tensors(self.context.as_ptr()) } } + + /// # Errors + /// + /// Returns [`GgufContextError::TensorNotFound`] if the file holds no tensor with this name. + /// Returns [`GgufContextError::TensorIsNotF32`] if the tensor is stored in another type. + /// Returns [`GgufContextError::TensorReadFailed`] if the tensor data cannot be read. + /// Returns [`GgufContextError::NulError`] if the name contains a null byte. + pub fn read_tensor_f32(&self, name: &str) -> Result { + let c_name = CString::new(name)?; + let tensor_id = unsafe { + llama_cpp_bindings_sys::gguf_find_tensor(self.context.as_ptr(), c_name.as_ptr()) + }; + + if tensor_id < 0 { + return Err(GgufContextError::TensorNotFound { + name: name.to_owned(), + }); + } + + let ggml_type = unsafe { + llama_cpp_bindings_sys::gguf_get_tensor_type(self.context.as_ptr(), tensor_id) + }; + + if ggml_type != llama_cpp_bindings_sys::GGML_TYPE_F32 { + return Err(GgufContextError::TensorIsNotF32 { + name: name.to_owned(), + ggml_type, + }); + } + + Ok(GgufTensorF32 { + shape: self.tensor_shape(tensor_id), + values: self.tensor_f32_values(name, tensor_id)?, + }) + } + + fn require_key_id_in_range(&self, key_id: i64) -> Result<(), GgufContextError> { + let n_kv = self.n_kv(); + + if (0..n_kv).contains(&key_id) { + Ok(()) + } else { + Err(GgufContextError::KeyIdOutOfRange { key_id, n_kv }) + } + } + + fn require_value_type(&self, key_id: i64, expected: GgufType) -> Result<(), GgufContextError> { + let actual = self.kv_type(key_id)?; + + if actual == expected { + Ok(()) + } else { + Err(GgufContextError::ValueTypeMismatch { + key_id, + expected, + actual, + }) + } + } + + fn tensor_shape( + &self, + tensor_id: i64, + ) -> [i64; llama_cpp_bindings_sys::GGML_MAX_DIMS as usize] { + let element_counts = unsafe { + slice::from_raw_parts( + llama_cpp_bindings_sys::gguf_get_tensor_ne(self.context.as_ptr(), tensor_id), + llama_cpp_bindings_sys::GGML_MAX_DIMS as usize, + ) + }; + let mut shape = [0; llama_cpp_bindings_sys::GGML_MAX_DIMS as usize]; + + shape.copy_from_slice(element_counts); + + shape + } + + fn tensor_f32_values(&self, name: &str, tensor_id: i64) -> Result, GgufContextError> { + let data_offset = + unsafe { llama_cpp_bindings_sys::gguf_get_data_offset(self.context.as_ptr()) }; + let tensor_offset = unsafe { + llama_cpp_bindings_sys::gguf_get_tensor_offset(self.context.as_ptr(), tensor_id) + }; + let tensor_size = unsafe { + llama_cpp_bindings_sys::gguf_get_tensor_size(self.context.as_ptr(), tensor_id) + }; + let read_failed = |source: io::Error| GgufContextError::TensorReadFailed { + name: name.to_owned(), + path: self.path.clone(), + source, + }; + + let file = File::open(&self.path).map_err(read_failed)?; + let mut bytes = vec![0; tensor_size]; + + file.read_exact_at(&mut bytes, (data_offset + tensor_offset) as u64) + .map_err(read_failed)?; + + Ok(bytes + .chunks_exact(F32_SIZE_IN_BYTES) + .map(|value_bytes| { + f32::from_ne_bytes([ + value_bytes[0], + value_bytes[1], + value_bytes[2], + value_bytes[3], + ]) + }) + .collect()) + } } impl Drop for GgufContext { @@ -152,6 +290,9 @@ mod tests { use crate::gguf_context_error::GgufContextError; use crate::gguf_type::GgufType; + const GGUF_ALIGNMENT: usize = 32; + const GGML_TYPE_F16: u32 = 1; + fn fixture_path() -> PathBuf { PathBuf::from(env!("CARGO_MANIFEST_DIR")) .join("fixtures") @@ -171,17 +312,158 @@ mod tests { std::mem::discriminant(&GgufContextError::NulError(nul_err)) } - #[cfg(unix)] fn path_to_str_error_disc() -> Discriminant { std::mem::discriminant(&GgufContextError::PathToStrError(PathBuf::new())) } + fn key_id_out_of_range_disc() -> Discriminant { + std::mem::discriminant(&GgufContextError::KeyIdOutOfRange { key_id: 0, n_kv: 0 }) + } + + fn value_type_mismatch_disc() -> Discriminant { + std::mem::discriminant(&GgufContextError::ValueTypeMismatch { + key_id: 0, + expected: GgufType::Uint32, + actual: GgufType::Uint32, + }) + } + + fn tensor_not_found_disc() -> Discriminant { + std::mem::discriminant(&GgufContextError::TensorNotFound { + name: String::new(), + }) + } + + fn tensor_is_not_f32_disc() -> Discriminant { + std::mem::discriminant(&GgufContextError::TensorIsNotF32 { + name: String::new(), + ggml_type: 0, + }) + } + + fn tensor_read_failed_disc() -> Discriminant { + std::mem::discriminant(&GgufContextError::TensorReadFailed { + name: String::new(), + path: PathBuf::new(), + source: std::io::Error::other("discriminant"), + }) + } + fn utf8_error_disc() -> Discriminant { let invalid_utf8_bytes: Vec = vec![0xFF]; let utf8_err = std::str::from_utf8(&invalid_utf8_bytes).unwrap_err(); std::mem::discriminant(&GgufContextError::Utf8Error(utf8_err)) } + struct SyntheticTensor { + name: Vec, + shape: Vec, + ggml_type: u32, + data: Vec, + } + + #[derive(Default)] + struct SyntheticGgufBuilder { + key_values: Vec, + key_value_count: u64, + tensors: Vec, + } + + impl SyntheticGgufBuilder { + fn value(mut self, key: &[u8], gguf_type: GgufType, value: &[u8]) -> Self { + self.key_values + .extend_from_slice(&(key.len() as u64).to_le_bytes()); + self.key_values.extend_from_slice(key); + self.key_values + .extend_from_slice(&gguf_type.to_raw().to_le_bytes()); + self.key_values.extend_from_slice(value); + self.key_value_count += 1; + self + } + + fn string_value(self, key: &[u8], value: &[u8]) -> Self { + let mut encoded = (value.len() as u64).to_le_bytes().to_vec(); + encoded.extend_from_slice(value); + self.value(key, GgufType::String, &encoded) + } + + fn tensor(mut self, name: &[u8], shape: &[u64], ggml_type: u32, data: Vec) -> Self { + self.tensors.push(SyntheticTensor { + name: name.to_vec(), + shape: shape.to_vec(), + ggml_type, + data, + }); + self + } + + fn write(self, test_name: &str) -> SyntheticGgufFile { + let mut bytes: Vec = Vec::new(); + bytes.extend_from_slice(b"GGUF"); + bytes.extend_from_slice(&3u32.to_le_bytes()); + bytes.extend_from_slice(&(self.tensors.len() as u64).to_le_bytes()); + bytes.extend_from_slice(&self.key_value_count.to_le_bytes()); + bytes.extend_from_slice(&self.key_values); + + let mut data_section: Vec = Vec::new(); + + for tensor in &self.tensors { + bytes.extend_from_slice(&(tensor.name.len() as u64).to_le_bytes()); + bytes.extend_from_slice(&tensor.name); + bytes.extend_from_slice(&u32::try_from(tensor.shape.len()).unwrap().to_le_bytes()); + for dimension in &tensor.shape { + bytes.extend_from_slice(&dimension.to_le_bytes()); + } + bytes.extend_from_slice(&tensor.ggml_type.to_le_bytes()); + bytes.extend_from_slice(&(data_section.len() as u64).to_le_bytes()); + data_section.extend_from_slice(&tensor.data); + data_section.resize(data_section.len().next_multiple_of(GGUF_ALIGNMENT), 0); + } + + if !self.tensors.is_empty() { + bytes.resize(bytes.len().next_multiple_of(GGUF_ALIGNMENT), 0); + bytes.extend_from_slice(&data_section); + } + + SyntheticGgufFile::from_bytes(test_name, &bytes) + } + } + + struct SyntheticGgufFile { + path: PathBuf, + } + + impl SyntheticGgufFile { + fn from_bytes(test_name: &str, bytes: &[u8]) -> Self { + use std::io::Write as _; + + let path = std::env::temp_dir().join(format!( + "llama_cpp_bindings_synthetic_{}_{}.gguf", + std::process::id(), + test_name, + )); + + let mut file = std::fs::File::create(&path).unwrap(); + file.write_all(bytes).unwrap(); + + Self { path } + } + } + + impl Drop for SyntheticGgufFile { + fn drop(&mut self) { + std::fs::remove_file(&self.path) + .unwrap_or_else(|error| panic!("failed to remove synthetic GGUF: {error}")); + } + } + + fn f32_bytes(values: &[f32]) -> Vec { + values + .iter() + .flat_map(|value| value.to_ne_bytes()) + .collect() + } + #[test] fn from_file_opens_valid_gguf() { let context = GgufContext::from_file(fixture_path()); @@ -203,13 +485,6 @@ mod tests { assert!(context.n_kv() > 0); } - #[test] - fn n_tensors_returns_count() { - let context = GgufContext::from_file(fixture_path()).unwrap(); - - assert!(context.n_tensors() >= 0); - } - #[test] fn find_key_returns_valid_index_for_known_key() { let context = GgufContext::from_file(fixture_path()).unwrap(); @@ -236,13 +511,28 @@ mod tests { assert_eq!(key_name, "general.architecture"); } + #[test] + fn key_at_rejects_a_key_id_past_the_last_key() { + let context = GgufContext::from_file(fixture_path()).unwrap(); + let err = context.key_at(context.n_kv()).unwrap_err(); + + assert_eq!(std::mem::discriminant(&err), key_id_out_of_range_disc()); + } + + #[test] + fn typed_value_rejects_a_negative_key_id() { + let context = GgufContext::from_file(fixture_path()).unwrap(); + let err = context.val_u32(-1).unwrap_err(); + + assert_eq!(std::mem::discriminant(&err), key_id_out_of_range_disc()); + } + #[test] fn kv_type_returns_expected_type_for_string_key() { let context = GgufContext::from_file(fixture_path()).unwrap(); let index = context.find_key("general.architecture").unwrap(); - let value_type = context.kv_type(index); - assert_eq!(value_type, Some(GgufType::String)); + assert_eq!(context.kv_type(index).unwrap(), GgufType::String); } #[test] @@ -254,7 +544,6 @@ mod tests { assert!(!value.is_empty()); } - #[cfg(unix)] #[test] fn from_file_non_utf8_path_returns_error() { use std::ffi::OsStr; @@ -282,135 +571,177 @@ mod tests { } #[test] - fn val_u32_returns_value_for_uint32_key() { - let context = GgufContext::from_file(fixture_path()).unwrap(); + fn typed_values_round_trip_through_synthetic_fixture() { + let fixture = SyntheticGgufBuilder::default() + .value( + b"synthetic.i32_value", + GgufType::Int32, + &(-12345i32).to_le_bytes(), + ) + .value( + b"synthetic.u32_value", + GgufType::Uint32, + &4096u32.to_le_bytes(), + ) + .value( + b"synthetic.u64_value", + GgufType::Uint64, + &987_654_321u64.to_le_bytes(), + ) + .value( + b"synthetic.f32_value", + GgufType::Float32, + &2.41f32.to_le_bytes(), + ) + .write("typed_values_round_trip"); + let context = GgufContext::from_file(&fixture.path).unwrap(); - let key_id = (0..context.n_kv()) - .find(|&id| context.kv_type(id) == Some(GgufType::Uint32)) - .expect("fixture must contain at least one uint32 key"); + let i32_index = context.find_key("synthetic.i32_value").unwrap(); + let u32_index = context.find_key("synthetic.u32_value").unwrap(); + let u64_index = context.find_key("synthetic.u64_value").unwrap(); + let f32_index = context.find_key("synthetic.f32_value").unwrap(); - let _ = context.val_u32(key_id); + assert_eq!(context.val_i32(i32_index).unwrap(), -12345); + assert_eq!(context.val_u32(u32_index).unwrap(), 4096); + assert_eq!(context.val_u64(u64_index).unwrap(), 987_654_321); + assert!((context.val_f32(f32_index).unwrap() - 2.41).abs() < f32::EPSILON); } - struct SyntheticGgufFile { - path: PathBuf, + #[test] + fn typed_values_reject_a_key_holding_another_type() { + let fixture = SyntheticGgufBuilder::default() + .value(b"synthetic.bool_value", GgufType::Bool, &[1]) + .write("typed_values_reject_another_type"); + let context = GgufContext::from_file(&fixture.path).unwrap(); + let bool_index = context.find_key("synthetic.bool_value").unwrap(); + let mismatches = [ + context.val_u32(bool_index).unwrap_err(), + context.val_i32(bool_index).unwrap_err(), + context.val_u64(bool_index).unwrap_err(), + context.val_f32(bool_index).unwrap_err(), + context.val_str(bool_index).unwrap_err(), + ]; + + for mismatch in mismatches { + assert_eq!( + std::mem::discriminant(&mismatch), + value_type_mismatch_disc() + ); + } } - impl SyntheticGgufFile { - fn from_bytes(test_name: &str, bytes: &[u8]) -> Self { - use std::io::Write as _; + #[test] + fn val_str_returns_utf8_error_for_non_utf8_value() { + let fixture = SyntheticGgufBuilder::default() + .string_value(b"synthetic.str_value", &[0xFF, 0xFE]) + .write("val_str_returns_utf8_error_for_non_utf8_value"); + let context = GgufContext::from_file(&fixture.path).unwrap(); - let path = std::env::temp_dir().join(format!( - "llama_cpp_bindings_synthetic_{}_{}.gguf", - std::process::id(), - test_name, - )); + let value_index = context.find_key("synthetic.str_value").unwrap(); + let err = context.val_str(value_index).unwrap_err(); - let mut file = std::fs::File::create(&path).unwrap(); - file.write_all(bytes).unwrap(); + assert_eq!(std::mem::discriminant(&err), utf8_error_disc()); + } - Self { path } - } + #[test] + fn key_at_returns_utf8_error_for_non_utf8_key() { + let fixture = SyntheticGgufBuilder::default() + .value(&[0xFF, 0xFE], GgufType::Int32, &42i32.to_le_bytes()) + .write("key_at_returns_utf8_error_for_non_utf8_key"); + let context = GgufContext::from_file(&fixture.path).unwrap(); - fn new(test_name: &str) -> Self { - let mut bytes: Vec = Vec::new(); - bytes.extend_from_slice(b"GGUF"); - bytes.extend_from_slice(&3u32.to_le_bytes()); - bytes.extend_from_slice(&0u64.to_le_bytes()); - bytes.extend_from_slice(&3u64.to_le_bytes()); - - let arch_key = b"general.architecture"; - bytes.extend_from_slice(&(arch_key.len() as u64).to_le_bytes()); - bytes.extend_from_slice(arch_key); - bytes.extend_from_slice(&8u32.to_le_bytes()); - let arch_val = b"synthetic"; - bytes.extend_from_slice(&(arch_val.len() as u64).to_le_bytes()); - bytes.extend_from_slice(arch_val); - - let i32_key = b"synthetic.i32_value"; - bytes.extend_from_slice(&(i32_key.len() as u64).to_le_bytes()); - bytes.extend_from_slice(i32_key); - bytes.extend_from_slice(&5u32.to_le_bytes()); - bytes.extend_from_slice(&(-12345i32).to_le_bytes()); - - let u64_key = b"synthetic.u64_value"; - bytes.extend_from_slice(&(u64_key.len() as u64).to_le_bytes()); - bytes.extend_from_slice(u64_key); - bytes.extend_from_slice(&10u32.to_le_bytes()); - bytes.extend_from_slice(&987_654_321u64.to_le_bytes()); - - Self::from_bytes(test_name, &bytes) - } + let err = context.key_at(0).unwrap_err(); + + assert_eq!(std::mem::discriminant(&err), utf8_error_disc()); } - impl Drop for SyntheticGgufFile { - fn drop(&mut self) { - std::fs::remove_file(&self.path) - .unwrap_or_else(|error| panic!("failed to remove synthetic GGUF: {error}")); - } + #[test] + fn read_tensor_f32_returns_shape_and_values() { + let fixture = SyntheticGgufBuilder::default() + .tensor(b"synthetic.bias", &[2], GGML_TYPE_F16, vec![0; 4]) + .tensor( + b"synthetic.weight", + &[3, 2], + llama_cpp_bindings_sys::GGML_TYPE_F32, + f32_bytes(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]), + ) + .write("read_tensor_f32_returns_shape_and_values"); + let context = GgufContext::from_file(&fixture.path).unwrap(); + + let tensor = context.read_tensor_f32("synthetic.weight").unwrap(); + + assert_eq!(context.n_tensors(), 2); + assert_eq!(tensor.shape, [3, 2, 1, 1]); + assert_eq!(tensor.values, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]); } #[test] - fn val_i32_and_val_u64_round_trip_through_synthetic_fixture() { - let fixture = SyntheticGgufFile::new("val_i32_and_val_u64_round_trip"); + fn read_tensor_f32_rejects_a_missing_tensor() { + let context = GgufContext::from_file(fixture_path()).unwrap(); + let err = context.read_tensor_f32("synthetic.missing").unwrap_err(); - let context = GgufContext::from_file(&fixture.path).unwrap(); + assert_eq!(std::mem::discriminant(&err), tensor_not_found_disc()); + } - let i32_index = context.find_key("synthetic.i32_value").unwrap(); - assert_eq!(context.kv_type(i32_index), Some(GgufType::Int32)); - assert_eq!(context.val_i32(i32_index), -12345); + #[test] + fn read_tensor_f32_rejects_a_tensor_stored_in_another_type() { + let fixture = SyntheticGgufBuilder::default() + .tensor(b"synthetic.half", &[2], GGML_TYPE_F16, vec![0; 4]) + .write("read_tensor_f32_rejects_another_type"); + let context = GgufContext::from_file(&fixture.path).unwrap(); + let err = context.read_tensor_f32("synthetic.half").unwrap_err(); - let u64_index = context.find_key("synthetic.u64_value").unwrap(); - assert_eq!(context.kv_type(u64_index), Some(GgufType::Uint64)); - assert_eq!(context.val_u64(u64_index), 987_654_321); + assert_eq!(std::mem::discriminant(&err), tensor_is_not_f32_disc()); } #[test] - fn val_str_returns_utf8_error_for_non_utf8_value() { - let mut bytes: Vec = Vec::new(); - bytes.extend_from_slice(b"GGUF"); - bytes.extend_from_slice(&3u32.to_le_bytes()); - bytes.extend_from_slice(&0u64.to_le_bytes()); - bytes.extend_from_slice(&1u64.to_le_bytes()); - - let value_key = b"synthetic.str_value"; - bytes.extend_from_slice(&(value_key.len() as u64).to_le_bytes()); - bytes.extend_from_slice(value_key); - bytes.extend_from_slice(&8u32.to_le_bytes()); - let non_utf8_value: [u8; 2] = [0xFF, 0xFE]; - bytes.extend_from_slice(&(non_utf8_value.len() as u64).to_le_bytes()); - bytes.extend_from_slice(&non_utf8_value); - - let fixture = - SyntheticGgufFile::from_bytes("val_str_returns_utf8_error_for_non_utf8_value", &bytes); + fn read_tensor_f32_reports_a_file_truncated_after_opening() { + let fixture = SyntheticGgufBuilder::default() + .tensor( + b"synthetic.weight", + &[2], + llama_cpp_bindings_sys::GGML_TYPE_F32, + f32_bytes(&[1.0, 2.0]), + ) + .write("read_tensor_f32_reports_a_file_truncated_after_opening"); let context = GgufContext::from_file(&fixture.path).unwrap(); - let value_index = context.find_key("synthetic.str_value").unwrap(); - let err = context.val_str(value_index).unwrap_err(); + std::fs::write(&fixture.path, b"").unwrap(); - assert_eq!(std::mem::discriminant(&err), utf8_error_disc()); + let err = context.read_tensor_f32("synthetic.weight").unwrap_err(); + + assert_eq!(std::mem::discriminant(&err), tensor_read_failed_disc()); } #[test] - fn key_at_returns_utf8_error_for_non_utf8_key() { - let mut bytes: Vec = Vec::new(); - bytes.extend_from_slice(b"GGUF"); - bytes.extend_from_slice(&3u32.to_le_bytes()); - bytes.extend_from_slice(&0u64.to_le_bytes()); - bytes.extend_from_slice(&1u64.to_le_bytes()); - - let non_utf8_key: [u8; 2] = [0xFF, 0xFE]; - bytes.extend_from_slice(&(non_utf8_key.len() as u64).to_le_bytes()); - bytes.extend_from_slice(&non_utf8_key); - bytes.extend_from_slice(&5u32.to_le_bytes()); - bytes.extend_from_slice(&42i32.to_le_bytes()); - - let fixture = - SyntheticGgufFile::from_bytes("key_at_returns_utf8_error_for_non_utf8_key", &bytes); + fn read_tensor_f32_reports_a_file_removed_after_opening() { + let fixture = SyntheticGgufBuilder::default() + .tensor( + b"synthetic.weight", + &[2], + llama_cpp_bindings_sys::GGML_TYPE_F32, + f32_bytes(&[1.0, 2.0]), + ) + .write("read_tensor_f32_reports_a_file_removed_after_opening"); let context = GgufContext::from_file(&fixture.path).unwrap(); - let err = context.key_at(0).unwrap_err(); + std::fs::remove_file(&fixture.path).unwrap(); - assert_eq!(std::mem::discriminant(&err), utf8_error_disc()); + let result = context.read_tensor_f32("synthetic.weight"); + + std::fs::write(&fixture.path, b"").unwrap(); + + assert_eq!( + std::mem::discriminant(&result.unwrap_err()), + tensor_read_failed_disc() + ); + } + + #[test] + fn read_tensor_f32_with_null_byte_in_name_returns_error() { + let context = GgufContext::from_file(fixture_path()).unwrap(); + let err = context.read_tensor_f32("foo\0bar").unwrap_err(); + + assert_eq!(std::mem::discriminant(&err), nul_error_disc()); } } diff --git a/llama-cpp-bindings/src/gguf_context_error.rs b/llama-cpp-bindings/src/gguf_context_error.rs index 69523c9d..051d6320 100644 --- a/llama-cpp-bindings/src/gguf_context_error.rs +++ b/llama-cpp-bindings/src/gguf_context_error.rs @@ -1,6 +1,8 @@ use std::ffi::NulError; use std::path::PathBuf; +use crate::gguf_type::GgufType; + #[derive(Debug, thiserror::Error)] pub enum GgufContextError { #[error("Failed to initialize GGUF context from file: {0}")] @@ -9,12 +11,39 @@ pub enum GgufContextError { #[error("Key not found in GGUF context: {key}")] KeyNotFound { key: String }, + #[error("GGUF key id {key_id} is outside of the {n_kv} keys the file holds")] + KeyIdOutOfRange { key_id: i64, n_kv: i64 }, + #[error("null byte in string: {0}")] NulError(#[from] NulError), #[error("failed to convert path {0} to str")] PathToStrError(PathBuf), + #[error("GGUF tensor {name} holds ggml type {ggml_type}, not F32")] + TensorIsNotF32 { name: String, ggml_type: u32 }, + + #[error("Tensor not found in GGUF context: {name}")] + TensorNotFound { name: String }, + + #[error("Failed to read GGUF tensor {name} from {path}")] + TensorReadFailed { + name: String, + path: PathBuf, + #[source] + source: std::io::Error, + }, + + #[error("GGUF key id {key_id} holds a value of unknown type {raw_type}")] + UnknownValueType { key_id: i64, raw_type: u32 }, + #[error("GGUF value is not valid UTF-8: {0}")] Utf8Error(#[from] std::str::Utf8Error), + + #[error("GGUF key id {key_id} holds a {actual:?} value, not {expected:?}")] + ValueTypeMismatch { + key_id: i64, + expected: GgufType, + actual: GgufType, + }, } diff --git a/llama-cpp-bindings/src/gguf_tensor_f32.rs b/llama-cpp-bindings/src/gguf_tensor_f32.rs new file mode 100644 index 00000000..5e37de22 --- /dev/null +++ b/llama-cpp-bindings/src/gguf_tensor_f32.rs @@ -0,0 +1,5 @@ +#[derive(Clone, Debug, PartialEq)] +pub struct GgufTensorF32 { + pub shape: [i64; llama_cpp_bindings_sys::GGML_MAX_DIMS as usize], + pub values: Vec, +} diff --git a/llama-cpp-bindings/src/lib.rs b/llama-cpp-bindings/src/lib.rs index 8f32a372..387af605 100644 --- a/llama-cpp-bindings/src/lib.rs +++ b/llama-cpp-bindings/src/lib.rs @@ -17,6 +17,7 @@ pub mod generation_progress; pub mod ggml_time_us; pub mod gguf_context; pub mod gguf_context_error; +pub mod gguf_tensor_f32; pub mod gguf_type; pub mod grammar_matcher; pub mod ingest_outcome; diff --git a/llama-cpp-bindings/src/llama_token_attrs.rs b/llama-cpp-bindings/src/llama_token_attrs.rs index d5ecd6de..d37d297c 100644 --- a/llama-cpp-bindings/src/llama_token_attrs.rs +++ b/llama-cpp-bindings/src/llama_token_attrs.rs @@ -5,16 +5,6 @@ use enumflags2::BitFlags; use crate::llama_token_attr::LlamaTokenAttr; use crate::llama_token_attrs_from_int_error::LlamaTokenAttrsFromIntError; -#[cfg(target_env = "msvc")] -const fn llama_token_type_to_u32(value: llama_cpp_bindings_sys::llama_token_type) -> u32 { - value.cast_unsigned() -} - -#[cfg(not(target_env = "msvc"))] -const fn llama_token_type_to_u32(value: llama_cpp_bindings_sys::llama_token_type) -> u32 { - value -} - #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct LlamaTokenAttrs(pub BitFlags); @@ -36,11 +26,11 @@ impl TryFrom for LlamaTokenAttrs { type Error = LlamaTokenAttrsFromIntError; fn try_from(value: llama_cpp_bindings_sys::llama_vocab_type) -> Result { - Ok(Self( - BitFlags::from_bits(llama_token_type_to_u32(value)).map_err(|bit_flag_error| { + Ok(Self(BitFlags::from_bits(value).map_err( + |bit_flag_error| { LlamaTokenAttrsFromIntError::UnknownValue(bit_flag_error.invalid_bits()) - })?, - )) + }, + )?)) } } diff --git a/llama-cpp-bindings/src/load_backends_from_path.rs b/llama-cpp-bindings/src/load_backends_from_path.rs index c2434e99..2ce89ea1 100644 --- a/llama-cpp-bindings/src/load_backends_from_path.rs +++ b/llama-cpp-bindings/src/load_backends_from_path.rs @@ -27,7 +27,6 @@ mod tests { use std::path::PathBuf; #[test] - #[cfg(unix)] fn load_backends_from_path_returns_path_null_byte_for_embedded_null() { use std::ffi::CString; use std::ffi::OsStr; @@ -45,7 +44,6 @@ mod tests { } #[test] - #[cfg(unix)] fn load_backends_from_path_returns_path_not_utf8_for_invalid_utf8() { use std::ffi::OsStr; use std::os::unix::ffi::OsStrExt; diff --git a/llama-cpp-bindings/src/model.rs b/llama-cpp-bindings/src/model.rs index 7da01217..99f7a71a 100644 --- a/llama-cpp-bindings/src/model.rs +++ b/llama-cpp-bindings/src/model.rs @@ -9,6 +9,7 @@ pub mod llama_lora_adapter; pub mod llama_split_mode_parse_error; pub mod lora_adapter_scale; pub mod params; +pub mod parse_special_tokens; pub mod rope_type; pub mod split_mode; pub mod tokenizer_input; @@ -59,6 +60,7 @@ pub use llama_lazy_mode_parse_error::LlamaLazyModeParseError; pub use llama_load_mode::LlamaLoadMode; pub use llama_load_mode_parse_error::LlamaLoadModeParseError; pub use llama_lora_adapter::LlamaLoraAdapter; +pub use parse_special_tokens::ParseSpecialTokens; pub use rope_type::RopeType; pub use vocab_type::VocabType; pub use vocab_type_from_int_error::VocabTypeFromIntError; @@ -488,20 +490,20 @@ impl LlamaModel { /// /// - if [`str`] contains a null byte /// - if an integer conversion fails during tokenization - /// - /// - /// ```no_run - /// use llama_cpp_bindings::model::LlamaModel; - /// pub fn str_to_token( &self, str: &str, add_bos: AddBos, + parse_special_tokens: ParseSpecialTokens, ) -> Result, StringToTokenError> { let add_bos = match add_bos { AddBos::Always => true, AddBos::Never => false, }; + let parse_special = match parse_special_tokens { + ParseSpecialTokens::Always => true, + ParseSpecialTokens::Never => false, + }; let tokens_estimation = std::cmp::max(8, (str.len() / 2) + usize::from(add_bos)); let TokenizerInput { @@ -518,6 +520,7 @@ impl LlamaModel { tokens, n_tokens_max, add_bos, + parse_special, ) }) } @@ -1000,7 +1003,7 @@ impl LlamaModel { if marker.is_empty() { return Ok(None); } - let tokens = self.str_to_token(marker, AddBos::Never)?; + let tokens = self.str_to_token(marker, AddBos::Never, ParseSpecialTokens::Always)?; if tokens.is_empty() { Ok(None) } else { @@ -1588,6 +1591,7 @@ fn invoke_rs_tokenize( tokens: *mut llama_cpp_bindings_sys::llama_token, n_tokens_max: c_int, add_bos: bool, + parse_special: bool, ) -> Result { let mut out_count: i32 = 0; let mut out_error: *mut c_char = ptr::null_mut(); @@ -1599,7 +1603,7 @@ fn invoke_rs_tokenize( tokens, n_tokens_max, add_bos, - true, + parse_special, &raw mut out_count, &raw mut out_error, ) diff --git a/llama-cpp-bindings/src/model/params.rs b/llama-cpp-bindings/src/model/params.rs index 949475db..83b5d3f8 100644 --- a/llama-cpp-bindings/src/model/params.rs +++ b/llama-cpp-bindings/src/model/params.rs @@ -788,7 +788,6 @@ mod tests { } #[test] - #[cfg(not(target_os = "windows"))] fn append_kv_override_with_high_byte_returns_invalid_character_error() { use crate::model::params::param_override_value::ParamOverrideValue; @@ -809,7 +808,6 @@ mod tests { } #[test] - #[cfg(not(target_os = "windows"))] fn add_cpu_buft_override_with_high_byte_returns_invalid_character_error() { let key_bytes: &[u8] = b"\xff\0"; let key = std::ffi::CStr::from_bytes_with_nul(key_bytes).unwrap(); diff --git a/llama-cpp-bindings/src/model/parse_special_tokens.rs b/llama-cpp-bindings/src/model/parse_special_tokens.rs new file mode 100644 index 00000000..45630037 --- /dev/null +++ b/llama-cpp-bindings/src/model/parse_special_tokens.rs @@ -0,0 +1,5 @@ +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ParseSpecialTokens { + Always, + Never, +} diff --git a/llama-cpp-bindings/src/sampling.rs b/llama-cpp-bindings/src/sampling.rs index 56269c3a..d8671e9f 100644 --- a/llama-cpp-bindings/src/sampling.rs +++ b/llama-cpp-bindings/src/sampling.rs @@ -367,8 +367,6 @@ impl LlamaSampler { no_perf: bool, ) -> Result { unsafe { - llama_cpp_bindings_sys::ggml_time_init(); - let chain = llama_cpp_bindings_sys::llama_sampler_chain_init( llama_cpp_bindings_sys::llama_sampler_chain_params { no_perf }, ); diff --git a/llama-cpp-bindings/src/send_logs_to_log.rs b/llama-cpp-bindings/src/send_logs_to_log.rs index 80796b87..9c5fa7e1 100644 --- a/llama-cpp-bindings/src/send_logs_to_log.rs +++ b/llama-cpp-bindings/src/send_logs_to_log.rs @@ -39,16 +39,6 @@ impl LogSource { static LLAMA_SOURCE: OnceLock = OnceLock::new(); static GGML_SOURCE: OnceLock = OnceLock::new(); -#[cfg(target_env = "msvc")] -const fn ggml_level_to_u32(level: llama_cpp_bindings_sys::ggml_log_level) -> u32 { - level.cast_unsigned() -} - -#[cfg(not(target_env = "msvc"))] -const fn ggml_level_to_u32(level: llama_cpp_bindings_sys::ggml_log_level) -> u32 { - level -} - const fn ggml_level_to_incoming(raw: llama_cpp_bindings_sys::ggml_log_level) -> IncomingLogLevel { match raw { llama_cpp_bindings_sys::GGML_LOG_LEVEL_NONE => IncomingLogLevel::None, @@ -57,7 +47,7 @@ const fn ggml_level_to_incoming(raw: llama_cpp_bindings_sys::ggml_log_level) -> llama_cpp_bindings_sys::GGML_LOG_LEVEL_WARN => IncomingLogLevel::Warn, llama_cpp_bindings_sys::GGML_LOG_LEVEL_ERROR => IncomingLogLevel::Error, llama_cpp_bindings_sys::GGML_LOG_LEVEL_CONT => IncomingLogLevel::Cont, - other => IncomingLogLevel::Unknown(ggml_level_to_u32(other)), + other => IncomingLogLevel::Unknown(other), } }