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: 93 additions & 90 deletions crates/onnx-genai-engine/src/engine/load.rs
Original file line number Diff line number Diff line change
Expand Up @@ -386,23 +386,25 @@ impl Engine {
// block. Derive the model-IO hidden output from the speculative target so
// the native session materializes the seed; never override an explicit
// `model.io.hidden_output`, and leave non-MTP models untouched.
if let Some(spec) = metadata.speculative.as_ref() {
if spec.proposal_type == ProposalType::Mtp {
if let Some(target_hidden) = spec
.target_hidden_output
let mtp_target_hidden = metadata
.speculative
.as_ref()
.filter(|spec| spec.proposal_type == ProposalType::Mtp)
.and_then(|spec| {
spec.target_hidden_output
.as_ref()
.filter(|name| !name.is_empty())
.cloned()
{
if let Some(model) = metadata.model.as_mut() {
if let Some(io) = model.io.as_mut() {
if io.hidden_output.as_deref().unwrap_or_default().is_empty() {
io.hidden_output = Some(target_hidden);
}
}
}
}
}
});
// Resolved into an owned value first on purpose: the borrow of
// `metadata.speculative` must end before `metadata.model.as_mut()` below,
// so this cannot be folded into one `let` chain.
if let Some(target_hidden) = mtp_target_hidden
&& let Some(model) = metadata.model.as_mut()
&& let Some(io) = model.io.as_mut()
&& io.hidden_output.as_deref().unwrap_or_default().is_empty()
{
io.hidden_output = Some(target_hidden);
}
let tokenizer = {
let _span = onnx_genai_ort::prof_span!("engine.tokenizer_load");
Expand Down Expand Up @@ -951,8 +953,11 @@ impl Engine {
if let Some(trace) = trace {
native_session.set_trace_context(trace);
}
let (native_shared_kv_proposer, shared_kv_mode) =
load_native_shared_kv_proposer(&metadata, &model_directory.root, native_device.clone())?;
let (native_shared_kv_proposer, shared_kv_mode) = load_native_shared_kv_proposer(
&metadata,
&model_directory.root,
native_device.clone(),
)?;
let environment = {
let _span = onnx_genai_ort::prof_span!("engine.ort_environment");
Environment::new("onnx-genai-engine")
Expand Down Expand Up @@ -2113,79 +2118,78 @@ fn build_mtp_model_from_resolved(
&mtp_config.public_config.head_model,
session_options.clone(),
)
.map_err(|error| anyhow::anyhow!("Failed to load MTP head: {error}"))?;
let decode_options = onnx_genai_ort::MtpDecodeOptions {
kv_mode: mtp_config.public_config.kv_mode,
batch_size: 1,
hc_mult: mtp_config.hc_mult,
hidden_state_rank4: mtp_config.target_hidden_layout == MtpHiddenLayout::Bshc,
hidden_output: mtp_config.mtp_hidden_output.clone(),
state_output: mtp_config.mtp_state_output.clone(),
};
let head_signature = MtpDecodeSession::new(&head_session, decode_options)
.map_err(|error| anyhow::anyhow!("Failed to inspect MTP head: {error}"))?
.signature()
.clone();
if head_signature.hidden_size != mtp_config.public_config.hidden_size {
anyhow::bail!(
"MTP head hidden size {} does not match configured target hidden size {}",
head_signature.hidden_size,
mtp_config.public_config.hidden_size
);
}
let (embedder, lm_head) = match (&mtp_config.embedding_weights, &mtp_config.lm_head_weights)
{
(MtpWeightSource::File(embedding), MtpWeightSource::File(lm_head)) => (
MtpEmbedder::Linear(
LinearEmbedder::new(
read_f32_weights(embedding)?,
mtp_config.public_config.vocab_size,
mtp_config.public_config.hidden_size,
)
.map_err(|error| anyhow::anyhow!("Invalid MTP embedding weights: {error}"))?,
),
MtpLmHead::Linear(
LinearLmHead::new(
read_f32_weights(lm_head)?,
mtp_config.public_config.hidden_size,
mtp_config.public_config.vocab_size,
)
.map_err(|error| anyhow::anyhow!("Invalid MTP LM-head weights: {error}"))?,
),
.map_err(|error| anyhow::anyhow!("Failed to load MTP head: {error}"))?;
let decode_options = onnx_genai_ort::MtpDecodeOptions {
kv_mode: mtp_config.public_config.kv_mode,
batch_size: 1,
hc_mult: mtp_config.hc_mult,
hidden_state_rank4: mtp_config.target_hidden_layout == MtpHiddenLayout::Bshc,
hidden_output: mtp_config.mtp_hidden_output.clone(),
state_output: mtp_config.mtp_state_output.clone(),
};
let head_signature = MtpDecodeSession::new(&head_session, decode_options)
.map_err(|error| anyhow::anyhow!("Failed to inspect MTP head: {error}"))?
.signature()
.clone();
if head_signature.hidden_size != mtp_config.public_config.hidden_size {
anyhow::bail!(
"MTP head hidden size {} does not match configured target hidden size {}",
head_signature.hidden_size,
mtp_config.public_config.hidden_size
);
}
let (embedder, lm_head) = match (&mtp_config.embedding_weights, &mtp_config.lm_head_weights) {
(MtpWeightSource::File(embedding), MtpWeightSource::File(lm_head)) => (
MtpEmbedder::Linear(
LinearEmbedder::new(
read_f32_weights(embedding)?,
mtp_config.public_config.vocab_size,
mtp_config.public_config.hidden_size,
)
.map_err(|error| anyhow::anyhow!("Invalid MTP embedding weights: {error}"))?,
),
(
MtpWeightSource::TargetInitializer(embedding),
MtpWeightSource::TargetInitializer(lm_head),
) => {
let (embedder, lm_head, vocab_size) = load_target_initializer_adapters(
&model_directory.model_path,
embedding,
lm_head,
MtpLmHead::Linear(
LinearLmHead::new(
read_f32_weights(lm_head)?,
mtp_config.public_config.hidden_size,
draft_projection,
)?;
if vocab_size != mtp_config.public_config.vocab_size {
anyhow::bail!(
"MTP target initializer vocabulary {vocab_size} does not match configured vocabulary {}",
mtp_config.public_config.vocab_size
);
}
(embedder, lm_head)
}
_ => anyhow::bail!(
"MTP embedding_weights and lm_head_weights must both use files or both use target initializers"
mtp_config.public_config.vocab_size,
)
.map_err(|error| anyhow::anyhow!("Invalid MTP LM-head weights: {error}"))?,
),
};
Ok(MtpModel {
config: mtp_config.public_config.clone(),
runtime_config: mtp_config.clone(),
session: Arc::new(head_session),
embedder,
lm_head,
hidden_output: mtp_config.public_config.target_hidden_output.clone(),
kv_mode: mtp_config.public_config.kv_mode,
num_speculative_tokens: mtp_config.public_config.num_speculative_tokens,
})
),
(
MtpWeightSource::TargetInitializer(embedding),
MtpWeightSource::TargetInitializer(lm_head),
) => {
let (embedder, lm_head, vocab_size) = load_target_initializer_adapters(
&model_directory.model_path,
embedding,
lm_head,
mtp_config.public_config.hidden_size,
draft_projection,
)?;
if vocab_size != mtp_config.public_config.vocab_size {
anyhow::bail!(
"MTP target initializer vocabulary {vocab_size} does not match configured vocabulary {}",
mtp_config.public_config.vocab_size
);
}
(embedder, lm_head)
}
_ => anyhow::bail!(
"MTP embedding_weights and lm_head_weights must both use files or both use target initializers"
),
};
Ok(MtpModel {
config: mtp_config.public_config.clone(),
runtime_config: mtp_config.clone(),
session: Arc::new(head_session),
embedder,
lm_head,
hidden_output: mtp_config.public_config.target_hidden_output.clone(),
kv_mode: mtp_config.public_config.kv_mode,
num_speculative_tokens: mtp_config.public_config.num_speculative_tokens,
})
}

/// Resolve and build a native-backend MTP proposer from a model directory's
Expand Down Expand Up @@ -2213,10 +2217,9 @@ fn load_native_mtp_proposer(
// The native target exposes no ORT `Session` to interrogate for the target
// vocabulary (that is the point of the native EP), so the MTP metadata block
// must declare it explicitly.
let vocab_size = config
.vocab_size
.filter(|&value| value > 0)
.context("native MTP speculation requires `speculative.vocab_size` in inference metadata")?;
let vocab_size = config.vocab_size.filter(|&value| value > 0).context(
"native MTP speculation requires `speculative.vocab_size` in inference metadata",
)?;
let descriptor =
onnx_genai_metadata::resolve_speculator_config(&model_directory.root, config.clone());
let spec = match descriptor.proposer {
Expand Down
10 changes: 9 additions & 1 deletion crates/onnx-genai-engine/src/engine/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -60,9 +60,17 @@ pub use crate::config::{
};
pub use crate::connector_bridge::{ConnectorLookupOutcome, ConnectorStats};
pub(crate) use crate::speculative::{
LinearEmbedder, LinearLmHead, MtpEmbedder, MtpLmHead, MtpProposer, SpeculativeStats,
LinearEmbedder, LinearLmHead, MtpEmbedder, MtpLmHead, SpeculativeStats,
load_target_initializer_adapters,
};
// `MtpProposer` is reached from exactly one place in this module tree --
// `runtime.rs`'s `generate_native_cold_with_callback`, which is itself
// `#[cfg(feature = "native-backend")]`. Importing it unconditionally therefore
// makes it an unused import in the default feature set, which `-D warnings`
// rejects. Its siblings above stay ungated because each has uses that are not
// feature-dependent (load.rs, model.rs, runtime.rs).
#[cfg(feature = "native-backend")]
pub(crate) use crate::speculative::MtpProposer;

mod decode_backend;
mod governor;
Expand Down
48 changes: 33 additions & 15 deletions crates/onnx-genai-engine/src/speculative/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -339,6 +339,11 @@ impl LmHead for MtpLmHead {
/// decode step (during proposal), so it never affects CUDA-graph capture of the
/// target decode step.
#[derive(Debug, Clone, Copy)]
// Outside `native-backend` the only constructions left are in `#[cfg(test)]`
// (`load.rs`'s `load_native_mtp_proposer` is itself gated on that feature), so
// the lib target sees both variants as never constructed. Same reason the
// `index` field below is gated -- one level up.
#[cfg_attr(not(feature = "native-backend"), allow(dead_code))]
pub(crate) enum DraftProjectionDevice {
/// Project on the CPU int4 `MatMulNBits` kernel (used by unit tests and the
/// native CPU backend).
Expand Down Expand Up @@ -626,7 +631,7 @@ fn build_quantized_draft_lm_head(
let k_blocks = weight_dims[1];
let blob = weight_dims[2];
let k = hidden_size;
if k_blocks == 0 || k % k_blocks != 0 {
if k_blocks == 0 || !k.is_multiple_of(k_blocks) {
anyhow::bail!(
"quantised LM-head weight '{lm_head_name}' k_blocks {k_blocks} does not divide hidden \
size {k}"
Expand Down Expand Up @@ -687,8 +692,11 @@ fn build_quantized_draft_lm_head(
projection_graph
.opset_imports
.insert("com.microsoft".to_string(), 1);
let hidden_value =
projection_graph.create_named_value("draft_hidden", IrDataType::Float32, static_shape([1, k]));
let hidden_value = projection_graph.create_named_value(
"draft_hidden",
IrDataType::Float32,
static_shape([1, k]),
);
projection_graph.add_input(hidden_value);
let add_initializer = |graph: &mut onnx_runtime_ir::Graph, name: &str, weight: &WeightRef| {
let value = graph.create_named_value(
Expand All @@ -703,8 +711,11 @@ fn build_quantized_draft_lm_head(
let scales_value = add_initializer(&mut projection_graph, "draft_lm_head.scales", &scales_ref);
let mut inputs = vec![Some(hidden_value), Some(weight_value), Some(scales_value)];
if let Some(zero_points_ref) = &zero_points_ref {
let zero_points_value =
add_initializer(&mut projection_graph, "draft_lm_head.zero_points", zero_points_ref);
let zero_points_value = add_initializer(
&mut projection_graph,
"draft_lm_head.zero_points",
zero_points_ref,
);
inputs.push(Some(zero_points_value));
}
let logits_value = projection_graph.create_named_value(
Expand All @@ -730,9 +741,7 @@ fn build_quantized_draft_lm_head(
projection_graph.add_output(logits_value);

let provider: Arc<dyn onnx_runtime_ep_api::ExecutionProvider> = match projection {
DraftProjectionDevice::Cpu => {
Arc::new(onnx_runtime_ep_cpu::CpuExecutionProvider::new())
}
DraftProjectionDevice::Cpu => Arc::new(onnx_runtime_ep_cpu::CpuExecutionProvider::new()),
#[cfg(feature = "native-cuda")]
DraftProjectionDevice::Cuda { index } => Arc::new(
onnx_runtime_ep_cuda::CudaExecutionProvider::initialized(index)
Expand Down Expand Up @@ -2489,8 +2498,11 @@ mod tests {
let mut graph = Graph::default();
graph.opset_imports.insert("com.microsoft".to_string(), 1);
let activation = graph.create_named_value("A", DataType::Float32, static_shape([1, k]));
let weight_value =
graph.create_named_value("lm_head.weight", DataType::Uint8, static_shape([n, k_blocks, blob]));
let weight_value = graph.create_named_value(
"lm_head.weight",
DataType::Uint8,
static_shape([n, k_blocks, blob]),
);
graph.set_initializer(
weight_value,
WeightRef::Inline(TensorData::from_raw(
Expand All @@ -2499,8 +2511,11 @@ mod tests {
weight,
)),
);
let scales_value =
graph.create_named_value("lm_head.scales", DataType::Float32, static_shape([n, k_blocks]));
let scales_value = graph.create_named_value(
"lm_head.scales",
DataType::Float32,
static_shape([n, k_blocks]),
);
graph.set_initializer(
scales_value,
WeightRef::Inline(TensorData::from_raw(
Expand All @@ -2518,9 +2533,12 @@ mod tests {
vec![logits_value],
);
node.domain = "com.microsoft".to_string();
node.attributes.insert("K".to_string(), Attribute::Int(k as i64));
node.attributes.insert("N".to_string(), Attribute::Int(n as i64));
node.attributes.insert("bits".to_string(), Attribute::Int(4));
node.attributes
.insert("K".to_string(), Attribute::Int(k as i64));
node.attributes
.insert("N".to_string(), Attribute::Int(n as i64));
node.attributes
.insert("bits".to_string(), Attribute::Int(4));
node.attributes
.insert("block_size".to_string(), Attribute::Int(block_size));
graph.insert_node(node);
Expand Down
Loading