diff --git a/sgl-model-gateway/bindings/golang/Cargo.toml b/sgl-model-gateway/bindings/golang/Cargo.toml index 13c60bdba..76a448b7c 100644 --- a/sgl-model-gateway/bindings/golang/Cargo.toml +++ b/sgl-model-gateway/bindings/golang/Cargo.toml @@ -17,6 +17,7 @@ uuid = { version = "1.10", features = ["v4", "serde"] } once_cell = "1.21.3" futures-util = "0.3" tracing = "0.1" +libc = "0.2.179" [dependencies.sgl-model-gateway] path = "../.." diff --git a/sgl-model-gateway/bindings/golang/src/client.rs b/sgl-model-gateway/bindings/golang/src/client.rs index 124fca88a..2ac444d20 100644 --- a/sgl-model-gateway/bindings/golang/src/client.rs +++ b/sgl-model-gateway/bindings/golang/src/client.rs @@ -160,7 +160,7 @@ pub unsafe extern "C" fn sgl_client_chat_completion_stream( }; // Tokenize - let token_ids = match tokenizer.encode(&processed_messages.text) { + let token_ids = match tokenizer.encode(&processed_messages.text, false) { Ok(encoding) => encoding.token_ids().to_vec(), Err(e) => { set_error_message(error_out, &format!("Failed to tokenize: {}", e)); diff --git a/sgl-model-gateway/bindings/golang/src/preprocessor.rs b/sgl-model-gateway/bindings/golang/src/preprocessor.rs index c959beca5..56d0475b3 100644 --- a/sgl-model-gateway/bindings/golang/src/preprocessor.rs +++ b/sgl-model-gateway/bindings/golang/src/preprocessor.rs @@ -115,7 +115,7 @@ pub unsafe extern "C" fn sgl_preprocess_chat_request( }; // Tokenize the processed text - let encoding = match tokenizer.encode(&processed_messages.text) { + let encoding = match tokenizer.encode(&processed_messages.text, false) { Ok(enc) => enc, Err(e) => { set_error_message(error_out, &format!("Tokenization failed: {}", e)); @@ -267,7 +267,7 @@ pub unsafe extern "C" fn sgl_preprocess_chat_request_with_tokenizer( }; // Tokenize the processed text - let encoding = match tokenizer.encode(&processed_messages.text) { + let encoding = match tokenizer.encode(&processed_messages.text, false) { Ok(enc) => enc, Err(e) => { set_error_message(error_out, &format!("Tokenization failed: {}", e)); diff --git a/sgl-model-gateway/bindings/golang/src/tokenizer.rs b/sgl-model-gateway/bindings/golang/src/tokenizer.rs index 65cea6d2e..ace39d7bd 100644 --- a/sgl-model-gateway/bindings/golang/src/tokenizer.rs +++ b/sgl-model-gateway/bindings/golang/src/tokenizer.rs @@ -15,6 +15,11 @@ use smg::tokenizer::{ use super::error::{SglErrorCode, set_error_message, clear_error_message}; +#[cfg(target_os = "macos")] +type BooleanT = libc::boolean_t; +#[cfg(not(target_os = "macos"))] +type BooleanT = libc::c_int; + /// Opaque handle for a tokenizer instance #[repr(C)] pub struct TokenizerHandle { @@ -69,6 +74,7 @@ pub unsafe extern "C" fn sgl_tokenizer_create_from_file( /// # Arguments /// * `handle` - Tokenizer handle (must not be null) /// * `text` - Input text (null-terminated C string) +/// * `add_special_tokens` - Whether to add special tokens /// * `token_ids_out` - Pointer to receive array of token IDs (must be freed with sgl_free_token_ids) /// * `token_count_out` - Pointer to receive token count /// * `error_out` - Optional pointer to receive error message @@ -82,6 +88,7 @@ pub unsafe extern "C" fn sgl_tokenizer_create_from_file( pub unsafe extern "C" fn sgl_tokenizer_encode( handle: *mut TokenizerHandle, text: *const c_char, + add_special_tokens: BooleanT, token_ids_out: *mut *mut u32, token_count_out: *mut usize, error_out: *mut *mut c_char, @@ -99,8 +106,10 @@ pub unsafe extern "C" fn sgl_tokenizer_encode( } }; + let add_special_tokens_bool = add_special_tokens != 0; + let tokenizer = &(*handle).tokenizer; - match tokenizer.encode(text_str) { + match tokenizer.encode(text_str, add_special_tokens_bool) { Ok(encoding) => { let token_ids = encoding.token_ids(); let count = token_ids.len();