Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
183 changes: 101 additions & 82 deletions crates/inference/examples/bench_serve_prepare.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,13 +15,15 @@
//! `Qwen35Model::generate_streaming_with_cancel`. This route
//! has no admission gate.
//! lattice-metal `lattice serve` on a Q4 checkpoint (Metal worker).
//! Handler: the same `prepare_chat_request` and mapping as
//! `cpu`, with the 4096-token context the binary uses.
//! Handler: the worker factory's `PreparationHandle`
//! (`prepare_lattice`, the same `prepare_chat_request` bound
//! to the worker's tokenizer and 4096-token context), then
//! `lattice_generate_config`.
//! lattice-serve `lattice_serve` (Metal worker, Q4 or safetensors).
//! Handler: `serve::contract::normalize_request` with the
//! `lattice_serve` profile, `serve::prepare::build_cfg`, then
//! `serve::into_engine_chat_messages`. No render or tokenize
//! happens in this handler.
//! Handler: the handle's `normalize_standalone` with the
//! `lattice_serve` defaults, `standalone_generate_config`,
//! then `serve::into_engine_chat_messages`. No render or
//! tokenize happens in this handler.
//! gemma-cpu Gemma 4 E2B text on a safetensors checkpoint, CPU.
//! Preparation: `serve::prepare::prepare_gemma_chat_request`
//! (validate, the Gemma prompt adapter's defaults, render
Expand Down Expand Up @@ -273,24 +275,24 @@ fn lattice_prepare(
))
}

metal_only! {
/// The `lattice_serve` handler's preparation: normalize, `build_cfg`, engine
/// messages.
fn lattice_serve_prepare(
body: &[u8],
model_max_context: usize,
) -> Result<Result<EngineRequest, ApiError>, String> {
let req = parse_body(body)?;
Ok(normalize_request(
&req,
GenerationDefaults::standard(LATTICE_SERVE_DEFAULT_MAX_TOKENS),
ServeProfile::lattice_serve(MODEL_ID, model_max_context).with_vision_support(false),
)
.and_then(|validated| {
let cfg = build_cfg(&validated);
into_engine_chat_messages(validated.messages).map(|messages| (messages, cfg))
}))
}
/// The `lattice_serve` handler's preparation before the worker factory:
/// normalize, `build_cfg`, engine messages. Kept for the preparation tests
/// below; the Metal route now prepares through the worker's handle.
#[cfg_attr(not(test), allow(dead_code))]
fn lattice_serve_prepare(
body: &[u8],
model_max_context: usize,
) -> Result<Result<EngineRequest, ApiError>, String> {
let req = parse_body(body)?;
Ok(normalize_request(
&req,
GenerationDefaults::standard(LATTICE_SERVE_DEFAULT_MAX_TOKENS),
ServeProfile::lattice_serve(MODEL_ID, model_max_context).with_vision_support(false),
)
.and_then(|validated| {
let cfg = build_cfg(&validated);
into_engine_chat_messages(validated.messages).map(|messages| (messages, cfg))
}))
}

/// Gemma E2B text preparation, with the `lattice serve` defaults.
Expand Down Expand Up @@ -655,6 +657,7 @@ fn run_metal(
use lattice_inference::serve::metal_worker::{
ContextWindowPolicy, MetalWorker, VisionRuntime, WorkerEvent, WorkerMetadata,
};
use lattice_inference::serving_factory::ServingFactory;
use lattice_inference::tokenizer::bpe::BpeTokenizer;
use std::time::Instant;

Expand All @@ -674,46 +677,48 @@ fn run_metal(
};

let load = Instant::now();
let (owner, client, meta) = MetalWorker::spawn_with_vision(
move || {
let tokenizer_path = loader_dir.join("tokenizer.json");
// `lattice serve` allocates its fixed 4096-token context;
// `lattice_serve` requests the checkpoint's configured window.
let state = match (route, format) {
(Route::LatticeServe, ModelFormat::Safetensors) => {
let model = Qwen35Model::from_safetensors(&loader_dir)
.map_err(|e| format!("safetensors load failed: {e}"))?;
let cfg = model.config().clone();
let context = cfg.max_position_embeddings;
MetalQwen35State::new(model.weights(), &cfg, context)
.map_err(|e| format!("Metal init failed: {e}"))?
}
_ => {
let cfg = Qwen35Config::from_model_dir(&loader_dir)
.map_err(|e| format!("config.json load failed: {e}"))?;
let context = match route {
Route::LatticeServe => cfg.max_position_embeddings,
_ => LATTICE_METAL_MAX_CONTEXT,
};
MetalQwen35State::from_q4_dir(&loader_dir, &tokenizer_path, &cfg, context)
.map_err(|e| format!("Q4 model load failed: {e}"))?
}
};
let model_max_context = match route {
Route::LatticeServe => state.max_context(),
_ => LATTICE_METAL_MAX_CONTEXT,
};
Ok((
state,
worker_tokenizer,
WorkerMetadata {
format: format!("{format:?}"),
model_max_context,
context_window_policy: policy,
},
))
},
VisionRuntime::unsupported(),
let (owner, client, meta, preparation) = MetalWorker::spawn_with_vision(
ServingFactory::qwen_metal(
move || {
let tokenizer_path = loader_dir.join("tokenizer.json");
// `lattice serve` allocates its fixed 4096-token context;
// `lattice_serve` requests the checkpoint's configured window.
let state = match (route, format) {
(Route::LatticeServe, ModelFormat::Safetensors) => {
let model = Qwen35Model::from_safetensors(&loader_dir)
.map_err(|e| format!("safetensors load failed: {e}"))?;
let cfg = model.config().clone();
let context = cfg.max_position_embeddings;
MetalQwen35State::new(model.weights(), &cfg, context)
.map_err(|e| format!("Metal init failed: {e}"))?
}
_ => {
let cfg = Qwen35Config::from_model_dir(&loader_dir)
.map_err(|e| format!("config.json load failed: {e}"))?;
let context = match route {
Route::LatticeServe => cfg.max_position_embeddings,
_ => LATTICE_METAL_MAX_CONTEXT,
};
MetalQwen35State::from_q4_dir(&loader_dir, &tokenizer_path, &cfg, context)
.map_err(|e| format!("Q4 model load failed: {e}"))?
}
};
let model_max_context = match route {
Route::LatticeServe => state.max_context(),
_ => LATTICE_METAL_MAX_CONTEXT,
};
Ok((
state,
worker_tokenizer,
WorkerMetadata {
format: format!("{format:?}"),
model_max_context,
context_window_policy: policy,
},
))
},
VisionRuntime::unsupported(),
),
1,
ResidencyLimits::default(),
)
Expand All @@ -723,25 +728,39 @@ fn run_metal(
println!("LOAD route={} load_ms={:.3}", route.name(), ms(load));

let prepare = |body: &[u8]| -> Result<Result<EngineRequest, ApiError>, String> {
let req = parse_body(body)?;
match route {
Route::LatticeServe => lattice_serve_prepare(body, model_max_context),
_ => Ok(lattice_prepare(
body,
|p| tokenizer.tokenize(p).real_length,
LATTICE_METAL_MAX_CONTEXT,
)?
.map(|prepared| {
let cfg = lattice_gen_cfg(
prepared.max_tokens,
prepared.temperature,
prepared.top_p,
prepared.seed,
prepared.stop_strings.clone(),
prepared.reasoning_budget,
prepared.logprobs,
);
(prepared.messages, cfg)
})),
Route::LatticeServe => Ok(preparation
.normalize_standalone(
&req,
GenerationDefaults::standard(LATTICE_SERVE_DEFAULT_MAX_TOKENS),
MODEL_ID,
false,
)
.and_then(|validated| {
let cfg = preparation.standalone_generate_config(&validated);
into_engine_chat_messages(validated.messages).map(|messages| (messages, cfg))
})),
_ => Ok(preparation
.prepare_lattice(
&req,
MODEL_ID,
LATTICE_DEFAULT_MAX_TOKENS,
LATTICE_MAX_TOKENS_CAP,
false,
)
.map(|prepared| {
let cfg = preparation.lattice_generate_config(
prepared.max_tokens,
prepared.temperature,
prepared.top_p,
prepared.seed,
prepared.stop_strings.clone(),
prepared.reasoning_budget,
prepared.logprobs,
);
(prepared.messages, cfg)
})),
}
};

Expand Down
120 changes: 93 additions & 27 deletions crates/inference/src/bin/lattice/serve.rs
Original file line number Diff line number Diff line change
Expand Up @@ -215,6 +215,72 @@ pub enum ModelBackend {
}

impl ModelBackend {
#[allow(clippy::too_many_arguments)]
fn generation_config(
&self,
max_tokens: usize,
temperature: f32,
top_p: f32,
seed: Option<u64>,
stop_strings: Vec<String>,
reasoning_budget: Option<usize>,
logprobs: Option<usize>,
) -> lattice_inference::GenerateConfig {
#[cfg(feature = "metal-gpu")]
if let ModelBackend::Metal { handle, .. } = self
&& let Some(preparation) = handle.client.preparation()
{
return preparation.lattice_generate_config(
max_tokens,
temperature,
top_p,
seed,
stop_strings,
reasoning_budget,
logprobs,
);
}
lattice_gen_cfg(
max_tokens,
temperature,
top_p,
seed,
stop_strings,
reasoning_budget,
logprobs,
)
}

fn prepare_chat_request(
&self,
req: &ChatCompletionRequest,
model_id: &str,
default_max_tokens: usize,
max_tokens_cap: usize,
) -> Result<PreparedChatRequest, ApiError> {
#[cfg(feature = "metal-gpu")]
if let ModelBackend::Metal { handle, .. } = self
&& let Some(preparation) = handle.client.preparation()
{
return preparation.prepare_lattice(
req,
model_id,
default_max_tokens,
max_tokens_cap,
self.supports_vision(),
);
}
prepare_chat_request(
req,
model_id,
default_max_tokens,
max_tokens_cap,
self.supports_vision(),
|prompt| self.tokenize_len(prompt),
|| self.max_context(),
)
}

pub fn tokenize_len(&self, text: &str) -> usize {
match self {
ModelBackend::Cpu(m) => m.tokenizer().tokenize(text).pre_truncation_len,
Expand Down Expand Up @@ -275,6 +341,7 @@ impl ModelBackend {
use lattice_inference::serve::metal_worker::{
ContextWindowPolicy, MetalWorker, StartupError, VisionRuntime, WorkerMetadata,
};
use lattice_inference::serving_factory::ServingFactory;

let tokenizer_path = tokenizer_dir
.as_deref()
Expand Down Expand Up @@ -318,28 +385,30 @@ impl ModelBackend {
}
let model_dir_for_loader = model_dir.clone();
let tokenizer_path_for_loader = tokenizer_path.clone();
let (owner, client, _meta) = MetalWorker::spawn_with_vision(
move || {
let cfg = crate::chat::load_q4_config(&model_dir_for_loader)?;
let state =
lattice_inference::forward::metal_qwen35::MetalQwen35State::from_q4_dir(
&model_dir_for_loader,
&tokenizer_path_for_loader,
&cfg,
max_context,
)
.map_err(|e| format!("Q4 model load failed: {e}"))?;
Ok((
state,
tokenizer_for_worker,
WorkerMetadata {
format: "q4".to_string(),
model_max_context: max_context,
context_window_policy: ContextWindowPolicy::PromptAndMaxTokens,
},
))
},
vision_runtime,
let (owner, client, _meta, _preparation) = MetalWorker::spawn_with_vision(
ServingFactory::qwen_metal(
move || {
let cfg = crate::chat::load_q4_config(&model_dir_for_loader)?;
let state =
lattice_inference::forward::metal_qwen35::MetalQwen35State::from_q4_dir(
&model_dir_for_loader,
&tokenizer_path_for_loader,
&cfg,
max_context,
)
.map_err(|e| format!("Q4 model load failed: {e}"))?;
Ok((
state,
tokenizer_for_worker,
WorkerMetadata {
format: "q4".to_string(),
model_max_context: max_context,
context_window_policy: ContextWindowPolicy::PromptAndMaxTokens,
},
))
},
vision_runtime,
),
max_pending,
residency_limits,
)
Expand Down Expand Up @@ -1118,17 +1187,14 @@ async fn chat_completions_with_request(
reasoning_budget,
seed,
stream,
} = prepare_chat_request(
} = state.model.prepare_chat_request(
&req,
&state.model_id,
state.default_max_tokens,
state.max_tokens_cap,
state.model.supports_vision(),
|p| state.model.tokenize_len(p),
|| state.model.max_context(),
)?;

let gen_cfg = lattice_gen_cfg(
let gen_cfg = state.model.generation_config(
max_tokens,
temperature,
top_p,
Expand Down
Loading
Loading