[model-gateway] fix tokenizer encode in golang bindings (#16482)
This commit is contained in:
@@ -17,6 +17,7 @@ uuid = { version = "1.10", features = ["v4", "serde"] }
|
|||||||
once_cell = "1.21.3"
|
once_cell = "1.21.3"
|
||||||
futures-util = "0.3"
|
futures-util = "0.3"
|
||||||
tracing = "0.1"
|
tracing = "0.1"
|
||||||
|
libc = "0.2.179"
|
||||||
|
|
||||||
[dependencies.sgl-model-gateway]
|
[dependencies.sgl-model-gateway]
|
||||||
path = "../.."
|
path = "../.."
|
||||||
|
|||||||
@@ -160,7 +160,7 @@ pub unsafe extern "C" fn sgl_client_chat_completion_stream(
|
|||||||
};
|
};
|
||||||
|
|
||||||
// Tokenize
|
// 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(),
|
Ok(encoding) => encoding.token_ids().to_vec(),
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
set_error_message(error_out, &format!("Failed to tokenize: {}", e));
|
set_error_message(error_out, &format!("Failed to tokenize: {}", e));
|
||||||
|
|||||||
@@ -115,7 +115,7 @@ pub unsafe extern "C" fn sgl_preprocess_chat_request(
|
|||||||
};
|
};
|
||||||
|
|
||||||
// Tokenize the processed text
|
// Tokenize the processed text
|
||||||
let encoding = match tokenizer.encode(&processed_messages.text) {
|
let encoding = match tokenizer.encode(&processed_messages.text, false) {
|
||||||
Ok(enc) => enc,
|
Ok(enc) => enc,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
set_error_message(error_out, &format!("Tokenization failed: {}", 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
|
// Tokenize the processed text
|
||||||
let encoding = match tokenizer.encode(&processed_messages.text) {
|
let encoding = match tokenizer.encode(&processed_messages.text, false) {
|
||||||
Ok(enc) => enc,
|
Ok(enc) => enc,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
set_error_message(error_out, &format!("Tokenization failed: {}", e));
|
set_error_message(error_out, &format!("Tokenization failed: {}", e));
|
||||||
|
|||||||
@@ -15,6 +15,11 @@ use smg::tokenizer::{
|
|||||||
|
|
||||||
use super::error::{SglErrorCode, set_error_message, clear_error_message};
|
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
|
/// Opaque handle for a tokenizer instance
|
||||||
#[repr(C)]
|
#[repr(C)]
|
||||||
pub struct TokenizerHandle {
|
pub struct TokenizerHandle {
|
||||||
@@ -69,6 +74,7 @@ pub unsafe extern "C" fn sgl_tokenizer_create_from_file(
|
|||||||
/// # Arguments
|
/// # Arguments
|
||||||
/// * `handle` - Tokenizer handle (must not be null)
|
/// * `handle` - Tokenizer handle (must not be null)
|
||||||
/// * `text` - Input text (null-terminated C string)
|
/// * `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_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
|
/// * `token_count_out` - Pointer to receive token count
|
||||||
/// * `error_out` - Optional pointer to receive error message
|
/// * `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(
|
pub unsafe extern "C" fn sgl_tokenizer_encode(
|
||||||
handle: *mut TokenizerHandle,
|
handle: *mut TokenizerHandle,
|
||||||
text: *const c_char,
|
text: *const c_char,
|
||||||
|
add_special_tokens: BooleanT,
|
||||||
token_ids_out: *mut *mut u32,
|
token_ids_out: *mut *mut u32,
|
||||||
token_count_out: *mut usize,
|
token_count_out: *mut usize,
|
||||||
error_out: *mut *mut c_char,
|
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;
|
let tokenizer = &(*handle).tokenizer;
|
||||||
match tokenizer.encode(text_str) {
|
match tokenizer.encode(text_str, add_special_tokens_bool) {
|
||||||
Ok(encoding) => {
|
Ok(encoding) => {
|
||||||
let token_ids = encoding.token_ids();
|
let token_ids = encoding.token_ids();
|
||||||
let count = token_ids.len();
|
let count = token_ids.len();
|
||||||
|
|||||||
Reference in New Issue
Block a user