use crate::decode::{
DecodeState, clone_value, extract_next_token_logits_with_io, is_present_output,
run_decode_step_with_extra,
};
use crate::decode_loop::{DecodeLoopBackend, DecodeLoopState, run_decode_loop};
use crate::engine::{
Engine, EngineConfig, model_requires_native_backend, requested_decode_backend,
resolved_host_ram_budget,
};
use crate::kv_bridge::infer_kv_model_info;
use crate::logits::TokenId;
use crate::processors::build_processor_chain;
use crate::{
EngineDecodeBackend, GeneratePrompt, GenerateRequest, GenerateResult, GenerateTokenCallback,
};
use anyhow::Context;
use onnx_genai_metadata::{
AbsentInputKind, DataflowEdge, PhaseRunOn, PipelineSpec, PipelineStrategy,
PipelineStrategyKind, PipelineVisionConfig, SchedulerSpec, TensorDimension,
};
use onnx_genai_ort::{
DataType, PipelineModelDirectory, PipelineModels, Session, SessionOptions, Tokenizer, Value,
};
use std::collections::{BTreeSet, HashMap};
use std::path::Path;
use std::sync::{Arc, Mutex};
pub type PipelineTensors = HashMap<String, Value>;
pub struct PipelineSynthesis {
pub generation: GenerateResult,
pub tensors: PipelineTensors,
}
#[derive(Debug, Clone, Default)]
pub struct IterativeOverrides {
pub num_steps: Option<usize>,
pub guidance_scale: Option<f32>,
pub start_step: Option<usize>,
}
pub struct PipelineGenerateRequest {
pub request: GenerateRequest,
pub inputs: PipelineTensors,
pub present: BTreeSet<String>,
pub num_image_tiles: Option<usize>,
pub iterative_overrides: IterativeOverrides,
}
impl PipelineGenerateRequest {
pub fn new(request: GenerateRequest) -> Self {
Self {
request,
inputs: HashMap::new(),
present: BTreeSet::new(),
num_image_tiles: None,
iterative_overrides: IterativeOverrides::default(),
}
}
pub fn with_input(mut self, endpoint: impl Into<String>, value: Value) -> Self {
self.inputs.insert(endpoint.into(), value);
self
}
pub fn with_presence(mut self, key: impl Into<String>) -> Self {
self.present.insert(key.into());
self
}
pub fn with_present_keys(mut self, keys: impl IntoIterator<Item = impl Into<String>>) -> Self {
self.present.extend(keys.into_iter().map(Into::into));
self
}
pub fn with_image_tile_count(mut self, num_image_tiles: usize) -> Self {
self.num_image_tiles = Some(num_image_tiles);
self
}
pub fn with_iterative_overrides(mut self, overrides: IterativeOverrides) -> Self {
self.iterative_overrides = overrides;
self
}
}
impl From<GenerateRequest> for PipelineGenerateRequest {
fn from(request: GenerateRequest) -> Self {
Self::new(request)
}
}
pub struct PipelineEngine {
models: PipelineModels,
plan: PipelinePlan,
decoder_state: Option<DecodeState>,
tokenizer_component: String,
fixed_state_budget_bytes: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum PipelineBackend {
Ort,
Native,
}
fn resolve_auto_pipeline_backend(
directory: &PipelineModelDirectory,
) -> anyhow::Result<PipelineBackend> {
for model_path in directory.model_paths.values() {
if model_requires_native_backend(model_path)? {
return Ok(PipelineBackend::Native);
}
}
Ok(PipelineBackend::Ort)
}
#[cfg(not(feature = "native-backend"))]
fn native_backend_not_compiled_error() -> anyhow::Error {
anyhow::anyhow!(
"the native backend was requested for a pipeline model, but this build of \
onnx-genai-engine was compiled without the 'native-backend' feature. Rebuild with \
`--features native-backend` (and `cuda` for GPU) to run pipelines natively, or set \
decode_backend = EngineDecodeBackend::Ort (or ONNX_GENAI_BACKEND=ort) to use ONNX \
Runtime."
)
}
#[cfg(feature = "native-backend")]
fn build_native_pipeline_and_report_gap(
directory: &PipelineModelDirectory,
config: &EngineConfig,
) -> anyhow::Error {
let components = match build_native_pipeline_components(directory, config) {
Ok(components) => components,
Err(err) => return err,
};
let component_list = components.keys().cloned().collect::<Vec<_>>().join(", ");
anyhow::anyhow!(
"native pipeline decode is not yet implemented. All {} pipeline component(s) loaded \
successfully on the native backend and expose their graph I/O through the \
backend-neutral component-session interface (components: {}), so backend selection and \
construction are backend-neutral. Native target decode now accepts metadata-declared \
token or embedding sequence inputs plus arbitrary named routed step tensors. The \
remaining GAP 3 work is replacing `DecodeState` and `PipelineDecodeLoopBackend` ownership \
of ORT `Value`/`Session` with backend-neutral tensors/component sessions, then invoking \
the native target step with those routed tensors. To run this pipeline today, set \
decode_backend = EngineDecodeBackend::Ort (or \
ONNX_GENAI_BACKEND=ort).",
components.len(),
component_list,
)
}
#[cfg(feature = "native-backend")]
fn build_native_pipeline_components(
directory: &PipelineModelDirectory,
config: &EngineConfig,
) -> anyhow::Result<
std::collections::BTreeMap<String, Box<dyn onnx_genai_metadata::ComponentSession>>,
> {
use crate::native_component::NativeComponentSession;
use onnx_genai_metadata::ComponentSession;
let device = crate::engine::resolve_native_decode_device(
config.native_device,
&SessionOptions::default(),
)?;
let mut components: std::collections::BTreeMap<String, Box<dyn ComponentSession>> =
std::collections::BTreeMap::new();
for (name, path) in &directory.model_paths {
let session = NativeComponentSession::load(path, device).with_context(|| {
format!("failed to construct pipeline component '{name}' on the native backend")
})?;
components.insert(name.clone(), Box::new(session));
}
Ok(components)
}
impl Engine {
pub fn from_pipeline_dir(
pipeline_dir: &Path,
config: EngineConfig,
) -> anyhow::Result<PipelineEngine> {
PipelineEngine::from_dir_with_config(pipeline_dir, config)
}
pub fn from_pipeline_dir_with_schedulers(
pipeline_dir: &Path,
config: EngineConfig,
schedulers: &SchedulerRegistry,
) -> anyhow::Result<PipelineEngine> {
PipelineEngine::from_dir_with_schedulers(pipeline_dir, config, schedulers)
}
}
impl PipelineEngine {
pub fn from_dir(pipeline_dir: &Path) -> anyhow::Result<Self> {
Self::from_dir_with_config(pipeline_dir, EngineConfig::default())
}
pub fn from_dir_with_config(pipeline_dir: &Path, config: EngineConfig) -> anyhow::Result<Self> {
Self::from_dir_with_schedulers(pipeline_dir, config, &SchedulerRegistry::builtin())
}
pub fn from_dir_with_schedulers(
pipeline_dir: &Path,
config: EngineConfig,
schedulers: &SchedulerRegistry,
) -> anyhow::Result<Self> {
let decode_backend = requested_decode_backend(config.decode_backend)?;
let backend = match decode_backend {
EngineDecodeBackend::Ort => PipelineBackend::Ort,
EngineDecodeBackend::Native => PipelineBackend::Native,
EngineDecodeBackend::Auto => {
let directory = PipelineModelDirectory::load(pipeline_dir)
.map_err(|e| anyhow::anyhow!("Failed to resolve pipeline models: {}", e))?;
resolve_auto_pipeline_backend(&directory)?
}
};
if backend == PipelineBackend::Native {
#[cfg(not(feature = "native-backend"))]
{
return Err(native_backend_not_compiled_error());
}
#[cfg(feature = "native-backend")]
{
let directory = PipelineModelDirectory::load(pipeline_dir)
.map_err(|e| anyhow::anyhow!("Failed to resolve pipeline models: {}", e))?;
return Err(build_native_pipeline_and_report_gap(&directory, &config));
}
}
let models = PipelineModels::load_with_options(pipeline_dir, SessionOptions::default())
.map_err(|e| anyhow::anyhow!("Failed to load pipeline models: {}", e))?;
let plan = PipelinePlan::from_spec(&models.directory.spec, schedulers)?;
let (decoder_state, tokenizer_component, fixed_state_budget_bytes) = match &plan {
PipelinePlan::Autoregressive(ar) => {
let decoder = models
.session(&ar.decoder)
.with_context(|| format!("pipeline decoder '{}' was not loaded", ar.decoder))?;
let kv_model =
infer_kv_model_info(decoder, config.page_size, config.kv_cache_dtype)?;
let fixed_state_budget_bytes =
resolved_host_ram_budget(&config, kv_model.as_ref())?;
let decoder_io = models
.directory
.spec
.models
.get(&ar.decoder)
.and_then(|component| component.io.as_ref());
let positions = models.directory.spec.positions.as_ref();
(
Some(DecodeState::new_with_io_positions_and_state_budget(
decoder,
decoder_io,
positions,
fixed_state_budget_bytes,
)?),
ar.decoder.clone(),
fixed_state_budget_bytes,
)
}
PipelinePlan::NestedAutoregressive(nested) => (None, nested.outer.clone(), 0),
PipelinePlan::SinglePass(sp) => (None, sp.model.clone(), 0),
PipelinePlan::Iterative(it) => (None, it.denoiser.clone(), 0),
PipelinePlan::Composite(c) => (
None,
c.stages
.last()
.map(|stage| match &stage.kind {
CompositeStageKind::SinglePass { model } => model.clone(),
})
.unwrap_or_default(),
0,
),
};
Ok(Self {
models,
plan,
decoder_state,
tokenizer_component,
fixed_state_budget_bytes,
})
}
pub fn generate(&mut self, request: GenerateRequest) -> anyhow::Result<GenerateResult> {
self.generate_with_pipeline_request(request.into())
}
pub fn generate_with_pipeline_request(
&mut self,
pipeline_request: PipelineGenerateRequest,
) -> anyhow::Result<GenerateResult> {
self.generate_with_callback(pipeline_request, None)
}
pub fn generate_with_callback(
&mut self,
pipeline_request: PipelineGenerateRequest,
callback: Option<&mut GenerateTokenCallback<'_>>,
) -> anyhow::Result<GenerateResult> {
if matches!(self.plan, PipelinePlan::NestedAutoregressive(_)) {
return self
.run_nested_autoregressive(pipeline_request)
.map(|(result, _pool)| result);
}
self.run_autoregressive(pipeline_request, callback)
.map(|(result, _pool)| result)
}
pub fn synthesize(
&mut self,
pipeline_request: PipelineGenerateRequest,
) -> anyhow::Result<PipelineSynthesis> {
let present = pipeline_request.present.clone();
if matches!(self.plan, PipelinePlan::NestedAutoregressive(_)) {
return self.synthesize_nested(pipeline_request);
}
let (generation, mut tensors) = self.run_autoregressive(pipeline_request, None)?;
let ar = self.plan.autoregressive_plan()?.clone();
let codes: Vec<i64> = generation.token_ids.iter().map(|&t| i64::from(t)).collect();
let codes_endpoint = format!("{}.output_ids", ar.decoder);
let codes_value =
Value::from_slice_i64(&codes, &[1, codes.len() as i64]).with_context(|| {
format!("failed to build generated-codes tensor '{codes_endpoint}'")
})?;
tensors.insert(codes_endpoint, codes_value);
self.run_prompt_phase_components(
&ar.post_decode_components,
&mut tensors,
"postlogue",
&present,
None,
)?;
Ok(PipelineSynthesis {
generation,
tensors,
})
}
fn run_autoregressive(
&mut self,
pipeline_request: PipelineGenerateRequest,
callback: Option<&mut GenerateTokenCallback<'_>>,
) -> anyhow::Result<(GenerateResult, PipelineTensors)> {
let ar = self
.plan
.autoregressive_plan()
.context(
"generate() requires an autoregressive pipeline; use run_pipeline() for \
single-pass or iterative (diffusion) pipelines",
)?
.clone();
let present = pipeline_request.present.clone();
self.ensure_component_present(&ar.decoder, &present, "autoregressive decoder")?;
let mut options = pipeline_request.request.options.clone();
options.validate()?;
if options.eos_token_id.is_none() {
options.eos_token_id = self.tokenizer()?.eos_token_id();
}
let prompt_tokens = tokenize_with(self.tokenizer()?, &pipeline_request.request.prompt)?;
if prompt_tokens.is_empty() {
anyhow::bail!("prompt must contain at least one token");
}
if pipeline_request.num_image_tiles == Some(0) {
anyhow::bail!("image tile count must be greater than zero");
}
let prompt_tokens = expand_image_placeholders_count_based(
prompt_tokens,
pipeline_request.num_image_tiles,
self.models.directory.spec.vision.as_ref(),
)?;
let mut tensors = self.prepare_request_tensors(pipeline_request.inputs, &present)?;
self.seed_prompt_token_inputs(&ar.prompt_components, &prompt_tokens, &mut tensors)?;
self.run_prompt_phase_components(
&ar.prompt_components,
&mut tensors,
"prologue",
&present,
None,
)?;
let decoder_in_edges = self.decoder_in_edges(&ar.decoder, &present, &tensors)?;
let step_bindings = self.build_step_bindings(&ar.step_components, &present)?;
let chain = build_processor_chain(&options, Some(self.tokenizer()?))?;
self.decoder_state = Some({
let decoder = self
.models
.session(&ar.decoder)
.with_context(|| format!("pipeline decoder '{}' was not loaded", ar.decoder))?;
let decoder_io = self
.models
.directory
.spec
.models
.get(&ar.decoder)
.and_then(|component| component.io.as_ref());
DecodeState::new_with_io_positions_and_state_budget(
decoder,
decoder_io,
self.models.directory.spec.positions.as_ref(),
self.fixed_state_budget_bytes,
)?
});
let cross_kv_pairs = self
.decoder_state
.as_ref()
.expect("autoregressive pipeline has decode state")
.io
.cross_kv_pairs
.clone();
let static_cross_kv = self.static_cross_kv_bindings(&cross_kv_pairs, &tensors)?;
let decoder = self
.models
.session(&ar.decoder)
.with_context(|| format!("pipeline decoder '{}' was not loaded", ar.decoder))?;
let step_components = step_bindings
.into_iter()
.map(|binding| {
let session = self.models.session(&binding.component).with_context(|| {
format!(
"pipeline every_step component '{}' was not loaded",
binding.component
)
})?;
Ok((binding, session))
})
.collect::<anyhow::Result<Vec<_>>>()?;
let tokenizer = self
.models
.tokenizer_for(&self.tokenizer_component)
.with_context(|| {
format!("no tokenizer available for '{}'", self.tokenizer_component)
})?;
let mut backend = PipelineDecodeLoopBackend {
decoder,
decoder_state: self
.decoder_state
.as_mut()
.expect("autoregressive pipeline has decode state"),
pool: &mut tensors,
step_components,
decoder_in_edges,
static_cross_kv,
context_tokens: prompt_tokens,
prompt_len: 0,
generated_count: 0,
};
backend.prompt_len = backend.context_tokens.len();
let mut loop_state = DecodeLoopState::new(0, options.seed, options.top_logprobs);
let result = run_decode_loop(
&mut backend,
&mut loop_state,
&options,
&chain,
tokenizer,
None,
callback,
)?;
Ok((result, tensors))
}
fn synthesize_nested(
&mut self,
pipeline_request: PipelineGenerateRequest,
) -> anyhow::Result<PipelineSynthesis> {
let present = pipeline_request.present.clone();
let post_decode_components = match &self.plan {
PipelinePlan::NestedAutoregressive(plan) => plan.post_decode_components.clone(),
_ => anyhow::bail!("internal error: synthesize_nested on a non-nested plan"),
};
let (generation, mut tensors) = self.run_nested_autoregressive(pipeline_request)?;
self.run_prompt_phase_components(
&post_decode_components,
&mut tensors,
"postlogue",
&present,
None,
)?;
Ok(PipelineSynthesis {
generation,
tensors,
})
}
fn run_nested_autoregressive(
&mut self,
pipeline_request: PipelineGenerateRequest,
) -> anyhow::Result<(GenerateResult, PipelineTensors)> {
let plan = match &self.plan {
PipelinePlan::NestedAutoregressive(plan) => plan.clone(),
_ => anyhow::bail!(
"synthesize()/generate() on a nested pipeline requires a nested_autoregressive plan"
),
};
let present = pipeline_request.present.clone();
self.ensure_component_present(&plan.outer, &present, "nested outer decoder")?;
self.ensure_component_present(&plan.inner, &present, "nested inner decoder")?;
let options = pipeline_request.request.options.clone();
options.validate()?;
let prompt_tokens = tokenize_with(self.tokenizer()?, &pipeline_request.request.prompt)?;
if prompt_tokens.is_empty() {
anyhow::bail!("prompt must contain at least one token");
}
let mut tensors = self.prepare_request_tensors(pipeline_request.inputs, &present)?;
self.seed_prompt_token_inputs(&plan.prompt_components, &prompt_tokens, &mut tensors)?;
if let Some(prefill) = plan
.prefill_embedder
.as_ref()
.filter(|binding| self.plan.component_is_present(&binding.component, &present))
{
let endpoint = format!("{}.{}", prefill.component, prefill.prompt_input);
let routed = plan.dataflow.iter().any(|edge| edge.to == endpoint);
if !routed && !tensors.contains_key(&endpoint) {
let ids: Vec<i64> = prompt_tokens.iter().map(|&t| i64::from(t)).collect();
let value = Value::from_slice_i64(&ids, &[1, ids.len() as i64])?;
tensors.insert(endpoint, value);
}
}
self.run_prompt_phase_components(
&plan.prompt_components,
&mut tensors,
"prologue",
&present,
None,
)?;
let outer_extra_exclude = plan.pre_embedder.as_ref().map(|p| p.outer_input.as_str());
let outer_extras =
self.decoder_extra_inputs(&plan.outer, &tensors, outer_extra_exclude, &present)?;
let inner_extras = self.decoder_extra_inputs(
&plan.inner,
&tensors,
Some(&plan.inner_embeds_input),
&present,
)?;
let outer_session = self
.models
.session(&plan.outer)
.with_context(|| format!("nested outer decoder '{}' was not loaded", plan.outer))?;
let inner_session = self
.models
.session(&plan.inner)
.with_context(|| format!("nested inner decoder '{}' was not loaded", plan.inner))?;
let pre_embed = match plan
.pre_embedder
.as_ref()
.filter(|binding| self.plan.component_is_present(&binding.component, &present))
{
Some(binding) => {
let session = self.models.session(&binding.component).with_context(|| {
format!("nested pre_embedder '{}' was not loaded", binding.component)
})?;
let frame_codes_input = binding.frame_codes_input.clone();
if !session
.inputs()
.iter()
.any(|info| info.name == frame_codes_input)
{
anyhow::bail!(
"nested pre_embedder '{}' has no declared frame_codes input '{}'",
binding.component,
frame_codes_input
);
}
let text_embed_input = binding.text_embed_input.clone();
if let Some(name) = &text_embed_input
&& !session.inputs().iter().any(|info| &info.name == name)
{
anyhow::bail!(
"nested pre_embedder '{}' has no declared text_embed input '{}'",
binding.component,
name
);
}
if !session
.output_names()
.iter()
.any(|name| name == &binding.output_port)
{
anyhow::bail!(
"nested pre_embedder '{}' has no declared output port '{}'",
binding.component,
binding.output_port
);
}
let hidden = outer_session
.inputs()
.iter()
.find(|info| info.name == binding.outer_input)
.and_then(|info| info.shape.last().copied())
.filter(|dim| *dim > 0)
.or_else(|| {
session
.outputs()
.iter()
.find(|info| info.name == binding.output_port)
.and_then(|info| info.shape.last().copied())
.filter(|dim| *dim > 0)
})
.map(|dim| dim as usize)
.with_context(|| {
format!(
"could not determine hidden size for nested pre_embedder '{}' \
(outer '{}' input '{}' has no static last dim)",
binding.component, plan.outer, binding.outer_input
)
})?;
Some(ResolvedPreEmbedder {
session,
outer_input: binding.outer_input.clone(),
output_port: binding.output_port.clone(),
frame_codes_input,
text_embed_input,
hidden,
})
}
None => None,
};
let prefill = match plan
.prefill_embedder
.as_ref()
.filter(|binding| self.plan.component_is_present(&binding.component, &present))
{
Some(binding) => {
let component = binding.component.as_str();
let pre = pre_embed.as_ref().with_context(|| {
format!(
"nested prefill_embedder '{component}' requires a pre_embedder to be set"
)
})?;
let _ = self.models.session(component).with_context(|| {
format!("nested prefill_embedder '{component}' was not loaded")
})?;
let prefill_name = binding.prefill_output.as_str();
let trailing_name = binding.trailing_output.as_str();
let prefill_value = tensors
.get(&format!("{component}.{prefill_name}"))
.with_context(|| {
format!(
"nested prefill_embedder '{component}' produced no pooled \
'{prefill_name}' output (did it run in the prompt phase?)"
)
})?;
let prefill_len = match prefill_value.shape() {
[1, p, _] if *p > 0 => *p as usize,
other => anyhow::bail!(
"nested prefill_embedder '{component}' '{prefill_name}' must be \
[1, prefill_len, hidden]; got {other:?}"
),
};
let prefill_embeds = clone_value(prefill_value)?;
let trailing_value = tensors
.get(&format!("{component}.{trailing_name}"))
.with_context(|| {
format!(
"nested prefill_embedder '{component}' produced no pooled \
'{trailing_name}' output (did it run in the prompt phase?)"
)
})?;
let trailing_len = match trailing_value.shape() {
[1, t, h] if *h as usize == pre.hidden => *t as usize,
other => anyhow::bail!(
"nested prefill_embedder '{component}' '{trailing_name}' must be \
[1, trailing_len, {}]; got {other:?}",
pre.hidden
),
};
let trailing = trailing_value.to_vec_f32_lossy().map_err(|e| {
anyhow::anyhow!("failed to read trailing_text_embeds tensor: {e}")
})?;
Some(ResolvedPrefill {
prefill_embeds,
prefill_len,
trailing,
trailing_len,
hidden: pre.hidden,
})
}
None => None,
};
let inner_embed_output = inner_session
.output_names()
.iter()
.find(|name| {
let lower = name.to_ascii_lowercase();
!lower.contains("logits") && !is_present_output(name)
})
.cloned()
.with_context(|| {
format!(
"nested inner decoder '{}' must expose a per-code embedding output (a \
non-logits, non-KV output) to thread across inner steps",
plan.inner
)
})?;
let mut outer_state = DecodeState::new(outer_session)?;
let mut codes: Vec<i64> = Vec::with_capacity(plan.max_frames * plan.num_code_groups);
let mut outer_input_tokens = prompt_tokens.clone();
let mut outer_past_len = 0usize;
let mut prev_frame_codes: Option<Vec<i64>> = None;
for _frame in 0..plan.max_frames {
let outer_outputs = if let Some(pre) = &pre_embed {
let (inputs_embeds, positions) =
if let Some(prefill) = prefill.as_ref().filter(|_| _frame == 0) {
(clone_value(&prefill.prefill_embeds)?, prefill.prefill_len)
} else {
let frame_codes = prev_frame_codes
.clone()
.unwrap_or_else(|| vec![0i64; plan.num_code_groups]);
let text_embed = match prefill.as_ref() {
Some(prefill) => {
let idx = _frame - 1;
let hidden = prefill.hidden;
let slice = if idx < prefill.trailing_len {
prefill.trailing[idx * hidden..(idx + 1) * hidden].to_vec()
} else {
vec![0.0f32; hidden]
};
Some(slice)
}
None => None,
};
(
run_pre_embedder(pre, &frame_codes, text_embed.as_deref())?,
1,
)
};
let mut step_extras = Vec::with_capacity(outer_extras.len() + 1);
for (name, value) in &outer_extras {
step_extras.push((name.clone(), clone_value(value)?));
}
step_extras.push((pre.outer_input.clone(), inputs_embeds));
let position_tokens = vec![0u32; positions];
let outputs = run_decode_step_with_extra(
outer_session,
&mut outer_state,
&position_tokens,
outer_past_len,
&step_extras,
)?;
outer_past_len += positions;
outputs
} else {
let outputs = run_decode_step_with_extra(
outer_session,
&mut outer_state,
&outer_input_tokens,
outer_past_len,
&outer_extras,
)?;
outer_past_len += outer_input_tokens.len();
outputs
};
let outer_logits = named_output(outer_session, &outer_outputs, "logits", true)?;
let outer_token = argmax_last_row(outer_logits)?;
let hidden = named_output(
outer_session,
&outer_outputs,
&plan.outer_hidden_output,
false,
)?;
let seed = last_position_hidden(hidden)?;
outer_input_tokens = vec![u32::try_from(outer_token).unwrap_or(0)];
let mut inner_state = DecodeState::new(inner_session)?;
let mut inner_embeds = seed;
let mut frame_inner_codes: Vec<i64> = Vec::with_capacity(plan.num_code_groups);
for step in 0..plan.num_code_groups {
let mut step_extras = Vec::with_capacity(inner_extras.len() + 1);
for (name, value) in &inner_extras {
step_extras.push((name.clone(), clone_value(value)?));
}
step_extras.push((plan.inner_embeds_input.clone(), inner_embeds));
let inner_outputs = run_decode_step_with_extra(
inner_session,
&mut inner_state,
&[0],
step,
&step_extras,
)?;
let inner_logits = named_output(inner_session, &inner_outputs, "logits", true)?;
let inner_code = argmax_last_row(inner_logits)?;
codes.push(inner_code);
frame_inner_codes.push(inner_code);
inner_embeds = clone_value(named_output(
inner_session,
&inner_outputs,
&inner_embed_output,
false,
)?)?;
}
if pre_embed.is_some() {
let mut tuple = Vec::with_capacity(plan.num_code_groups);
tuple.push(outer_token);
tuple.extend_from_slice(&frame_inner_codes[1..]);
prev_frame_codes = Some(tuple);
}
}
let codes_endpoint = format!("{}.output_codes", plan.outer);
let codes_value = Value::from_slice_i64(
&codes,
&[1, plan.max_frames as i64, plan.num_code_groups as i64],
)
.with_context(|| format!("failed to build generated-codes tensor '{codes_endpoint}'"))?;
tensors.insert(codes_endpoint, codes_value);
let token_ids: Vec<TokenId> = codes
.iter()
.map(|&c| u32::try_from(c).unwrap_or(0))
.collect();
let result = GenerateResult {
text: String::new(),
token_ids,
finish_reason: crate::FinishReason::MaxTokens,
prefix_cache_hit_len: 0,
logprobs: None,
};
Ok((result, tensors))
}
pub fn spec(&self) -> &PipelineSpec {
&self.models.directory.spec
}
pub fn diffusion_init_noise_sigma(&self) -> Option<f32> {
match &self.plan {
PipelinePlan::Iterative(iterative) => iterative
.scheduler
.as_ref()
.map(|scheduler| scheduler.init_noise_sigma()),
_ => None,
}
}
pub fn run_pipeline(
&mut self,
request: PipelineGenerateRequest,
) -> anyhow::Result<PipelineTensors> {
match &self.plan {
PipelinePlan::Iterative(_) => self.run_iterative(request),
PipelinePlan::SinglePass(_) => self.run_single_pass(request),
PipelinePlan::Composite(_) => self.run_composite(request),
PipelinePlan::Autoregressive(_) => anyhow::bail!(
"run_pipeline() runs single-pass or iterative pipelines; use generate() for \
autoregressive text pipelines"
),
PipelinePlan::NestedAutoregressive(_) => anyhow::bail!(
"run_pipeline() runs single-pass or iterative pipelines; use synthesize() for \
a nested-autoregressive (multi-decoder TTS) pipeline"
),
}
}
fn run_iterative(&self, request: PipelineGenerateRequest) -> anyhow::Result<PipelineTensors> {
let PipelinePlan::Iterative(plan) = &self.plan else {
anyhow::bail!("internal error: run_iterative on a non-iterative plan");
};
let present = request.present.clone();
let overrides = &request.iterative_overrides;
let num_steps = overrides.num_steps.unwrap_or(plan.num_steps);
let start_step = overrides.start_step.unwrap_or(plan.start_step);
if num_steps == 0 {
anyhow::bail!("iterative override num_steps must be >= 1");
}
if start_step >= num_steps {
anyhow::bail!(
"iterative override start_step ({start_step}) must be < num_steps ({num_steps})"
);
}
let rebuilt_scheduler = if num_steps != plan.num_steps {
if plan.timesteps.is_some() {
anyhow::bail!(
"cannot override num_steps for a pipeline with an explicit timestep schedule"
);
}
match &plan.scheduler_spec {
Some(spec) => Some(plan.scheduler_registry.build(spec, num_steps)?),
None => None,
}
} else {
None
};
let scheduler = rebuilt_scheduler.as_ref().or(plan.scheduler.as_ref());
let guidance = overrides
.guidance_scale
.or(plan.guidance_scale)
.filter(|s| *s != 1.0);
let mut constants = self.prepare_request_tensors(request.inputs, &present)?;
let mut stage_timings: Vec<serde_json::Value> = Vec::new();
self.run_prompt_phase_components(
&plan.prompt_components,
&mut constants,
"encode",
&present,
Some(&mut stage_timings),
)?;
if !self.plan.component_is_present(&plan.denoiser, &present) {
self.run_prompt_phase_components(
&plan.final_components,
&mut constants,
"decode",
&present,
Some(&mut stage_timings),
)?;
dump_stage_timings(&stage_timings);
return Ok(constants);
}
let denoiser = self
.models
.session(&plan.denoiser)
.with_context(|| format!("pipeline denoiser '{}' was not loaded", plan.denoiser))?;
let cfg_uncond: Vec<(String, Value)> = if guidance.is_some() {
if let Some(primary) = plan.cfg_conditioning_input.clone() {
let mut overrides: Vec<(String, Value)> = Vec::new();
let mut seen: BTreeSet<String> = BTreeSet::new();
for info in denoiser.inputs() {
let port = info.name.as_str();
let uncond_endpoint = format!("{}.{}.uncond", plan.denoiser, port);
if let Some(u) = constants.get(&uncond_endpoint) {
overrides.push((port.to_string(), clone_value(u)?));
seen.insert(port.to_string());
}
}
if !seen.contains(&primary) {
let cond_endpoint = format!("{}.{}", plan.denoiser, primary);
let cond = constants
.get(&cond_endpoint)
.or_else(|| {
plan.dataflow
.iter()
.find(|e| e.to == cond_endpoint)
.and_then(|e| constants.get(&e.from))
})
.with_context(|| format!("cfg conditioning '{cond_endpoint}' not found"))?;
overrides.push((
primary.clone(),
Value::from_slice_f32(&vec![0.0f32; cond.numel()], cond.shape())?,
));
}
overrides
} else {
Vec::new()
}
} else {
Vec::new()
};
let mut carried: HashMap<String, Value> = HashMap::new();
let mut last_outputs: HashMap<String, Value> = HashMap::new();
if let Some(scheduler) = scheduler {
scheduler.reset();
}
let scheduler_timesteps: Option<Vec<f32>> = if plan.timesteps.is_some() {
None
} else {
scheduler.and_then(|scheduler| scheduler.timesteps())
};
let denoise_start = std::time::Instant::now();
for step in start_step..num_steps {
let step_start = std::time::Instant::now();
let is_first = step == start_step;
let timestep = plan
.timesteps
.as_ref()
.or(scheduler_timesteps.as_ref())
.and_then(|ts| ts.get(step).copied())
.unwrap_or(step as f32);
let mut raw_samples: HashMap<String, Value> = HashMap::new();
for (_, in_port) in &plan.loop_edges {
let raw = if is_first {
let endpoint = format!("{}.{}", plan.denoiser, in_port);
constants.get(&endpoint).with_context(|| {
format!("missing iterative pipeline seed '{endpoint}' at start step")
})?
} else {
carried.get(in_port).with_context(|| {
format!(
"loop-carried input '{}.{in_port}' was not produced",
plan.denoiser
)
})?
};
raw_samples.insert(in_port.clone(), clone_value(raw)?);
}
let mut scaled_inputs: HashMap<String, Value> = HashMap::new();
if let Some(scheduler) = scheduler {
for (_, in_port) in &plan.loop_edges {
let raw = &raw_samples[in_port];
if let Some(scaled) = scheduler.scale_input(step, num_steps, raw)? {
scaled_inputs.insert(in_port.clone(), scaled);
}
}
}
let scale_overrides: Vec<(&str, &Value)> = scaled_inputs
.iter()
.map(|(port, value)| (port.as_str(), value))
.collect();
let cond_out = self.run_denoiser_pass(
denoiser,
plan,
start_step,
&constants,
&carried,
step,
timestep,
&scale_overrides,
)?;
let out_map = if let Some(scale) = guidance {
let mut cfg_overrides = scale_overrides.clone();
for (port, value) in &cfg_uncond {
cfg_overrides.retain(|(p, _)| *p != port.as_str());
cfg_overrides.push((port.as_str(), value));
}
let mut prompt_masked_inputs: Vec<(String, Value)> = Vec::new();
if let Some(scheduler) = scheduler {
for (_, in_port) in &plan.loop_edges {
let raw = &raw_samples[in_port];
if let Some(uncond_sample) = scheduler.cfg_uncond_sample(raw)? {
prompt_masked_inputs.push((in_port.clone(), uncond_sample));
}
}
}
for (port, value) in &prompt_masked_inputs {
cfg_overrides.retain(|(p, _)| *p != port.as_str());
cfg_overrides.push((port.as_str(), value));
}
let uncond_out = self.run_denoiser_pass(
denoiser,
plan,
start_step,
&constants,
&carried,
step,
timestep,
&cfg_overrides,
)?;
let mut combined: HashMap<String, Value> = HashMap::new();
for (port, cond_value) in &cond_out {
let uncond_value = uncond_out.get(port).with_context(|| {
format!(
"unconditional pass did not produce '{}.{port}'",
plan.denoiser
)
})?;
let cond_v = cond_value.to_vec_f32_lossy()?;
let uncond_v = uncond_value.to_vec_f32_lossy()?;
let guided: Vec<f32> = uncond_v
.iter()
.zip(&cond_v)
.map(|(u, c)| u + scale * (c - u))
.collect();
combined.insert(
port.clone(),
Value::from_slice_f32(&guided, cond_value.shape())?,
);
}
combined
} else {
cond_out
};
for (out_port, in_port) in &plan.loop_edges {
let model_output = out_map.get(out_port).with_context(|| {
format!(
"denoiser did not produce loop output '{}.{out_port}'",
plan.denoiser
)
})?;
let next = if let Some(scheduler) = scheduler {
let sample = raw_samples.get(in_port).with_context(|| {
format!(
"missing loop-carried sample for '{}.{in_port}'",
plan.denoiser
)
})?;
if scheduler.needs_noise() {
let noise =
self.step_noise(plan, num_steps, &constants, in_port, step, sample)?;
scheduler.step_with_noise(
step,
num_steps,
sample,
model_output,
Some(&noise),
)?
} else {
scheduler.step(step, num_steps, sample, model_output)?
}
} else {
clone_value(model_output)?
};
dump_iterative_step(
&plan.denoiser,
in_port,
step,
&next,
step_start.elapsed().as_secs_f64() * 1e3,
);
carried.insert(in_port.clone(), next);
}
last_outputs = out_map;
}
let denoise_ms = denoise_start.elapsed().as_secs_f64() * 1e3;
stage_timings.push(serde_json::json!({
"component": plan.denoiser,
"phase": "denoise",
"ms": denoise_ms,
"steps": num_steps - start_step,
}));
let mut tensors = constants;
for (out_port, value) in last_outputs {
tensors.insert(format!("{}.{}", plan.denoiser, out_port), value);
}
for (in_port, value) in carried {
tensors.insert(format!("{}.{}", plan.denoiser, in_port), value);
}
self.run_prompt_phase_components(
&plan.final_components,
&mut tensors,
"decode",
&present,
Some(&mut stage_timings),
)?;
dump_stage_timings(&stage_timings);
Ok(tensors)
}
#[allow(clippy::too_many_arguments)]
fn run_denoiser_pass(
&self,
denoiser: &Session,
plan: &IterativePlan,
start_step: usize,
constants: &PipelineTensors,
carried: &HashMap<String, Value>,
step: usize,
timestep: f32,
overrides: &[(&str, &Value)],
) -> anyhow::Result<HashMap<String, Value>> {
let mut inputs: Vec<(String, Value)> = Vec::new();
for info in denoiser.inputs() {
let port = info.name.as_str();
let endpoint = format!("{}.{}", plan.denoiser, port);
if let Some((_, over_value)) = overrides.iter().find(|(p, _)| *p == port) {
inputs.push((
port.to_string(),
coerce_value_to_dtype(over_value, info.dtype)?,
));
continue;
}
if plan.timestep_input.as_deref() == Some(port) {
let ts = match info.dtype {
DataType::Int64 => Value::from_vec_i64(vec![timestep as i64], &[1])?,
_ => Value::from_slice_f32(&[timestep], &[1])?,
};
inputs.push((port.to_string(), ts));
continue;
}
let is_loop = plan.loop_edges.iter().any(|(_, in_port)| in_port == port);
let value = if is_loop {
if step == start_step {
constants.get(&endpoint).with_context(|| {
format!("missing iterative pipeline seed '{endpoint}' at start step")
})?
} else {
carried.get(port).with_context(|| {
format!("loop-carried input '{endpoint}' was not produced")
})?
}
} else {
let routed = plan
.dataflow
.iter()
.find(|edge| edge.to == endpoint)
.and_then(|edge| constants.get(&edge.from));
constants
.get(&endpoint)
.or(routed)
.with_context(|| format!("missing pipeline input '{endpoint}'"))?
};
inputs.push((port.to_string(), coerce_value_to_dtype(value, info.dtype)?));
}
let refs = inputs
.iter()
.map(|(name, value)| (name.as_str(), value))
.collect::<Vec<_>>();
let outputs = denoiser.run(&refs).map_err(|e| {
anyhow::anyhow!(
"ORT denoiser '{}' failed at step {step}: {e}",
plan.denoiser
)
})?;
let mut out_map: HashMap<String, Value> = HashMap::new();
for (name, value) in denoiser.output_names().iter().zip(outputs) {
out_map.insert(name.clone(), value);
}
Ok(out_map)
}
fn step_noise(
&self,
plan: &IterativePlan,
num_steps: usize,
constants: &PipelineTensors,
in_port: &str,
step: usize,
sample: &Value,
) -> anyhow::Result<Value> {
let endpoint = format!("{}.{}.noise", plan.denoiser, in_port);
let all = constants.get(&endpoint).with_context(|| {
format!(
"ancestral scheduler requires per-step noise tensor '{endpoint}' \
shaped [num_steps, ...]"
)
})?;
let elem: usize = sample.shape().iter().map(|&d| d as usize).product();
let data = all.to_vec_f32_lossy()?;
let want = num_steps * elem;
if data.len() != want {
anyhow::bail!(
"noise tensor '{endpoint}' has {} elements but expected {want} \
({num_steps} steps x {elem})",
data.len(),
);
}
let slice = &data[step * elem..(step + 1) * elem];
Value::from_slice_f32(slice, sample.shape()).map_err(Into::into)
}
fn run_composite(&self, request: PipelineGenerateRequest) -> anyhow::Result<PipelineTensors> {
let PipelinePlan::Composite(plan) = &self.plan else {
anyhow::bail!("internal error: run_composite on a non-composite plan");
};
let present = request.present;
let mut tensors = self.prepare_request_tensors(request.inputs, &present)?;
for stage in &plan.stages {
match &stage.kind {
CompositeStageKind::SinglePass { model } => {
self.run_prompt_phase_components(
std::slice::from_ref(model),
&mut tensors,
&stage.name,
&present,
None,
)?;
}
}
}
Ok(tensors)
}
fn run_single_pass(&self, request: PipelineGenerateRequest) -> anyhow::Result<PipelineTensors> {
let PipelinePlan::SinglePass(plan) = &self.plan else {
anyhow::bail!("internal error: run_single_pass on a non-single-pass plan");
};
let present = request.present;
let mut tensors = self.prepare_request_tensors(request.inputs, &present)?;
self.run_prompt_phase_components(
&plan.prompt_components,
&mut tensors,
"prologue",
&present,
None,
)?;
if !self.plan.component_is_present(&plan.model, &present) {
return Ok(tensors);
}
let model = self
.models
.session(&plan.model)
.with_context(|| format!("pipeline model '{}' was not loaded", plan.model))?;
let inputs = self.component_inputs(&plan.model, model, &tensors, &present)?;
let refs = inputs
.iter()
.map(|(name, value)| (name.as_str(), value))
.collect::<Vec<_>>();
let outputs = model
.run(&refs)
.map_err(|e| anyhow::anyhow!("ORT pipeline model '{}' failed: {e}", plan.model))?;
for (name, value) in model.output_names().iter().zip(outputs) {
tensors.insert(format!("{}.{}", plan.model, name), value);
}
Ok(tensors)
}
fn tokenizer(&self) -> anyhow::Result<&Tokenizer> {
self.models
.tokenizer_for(&self.tokenizer_component)
.with_context(|| format!("no tokenizer available for '{}'", self.tokenizer_component))
}
fn prepare_request_tensors(
&self,
inputs: PipelineTensors,
present: &BTreeSet<String>,
) -> anyhow::Result<PipelineTensors> {
if present.iter().any(String::is_empty) {
anyhow::bail!("pipeline request presence keys must be non-empty");
}
let mut dimensions = HashMap::<String, i64>::new();
for (component, model) in &self.models.directory.spec.models {
let Some(io) = model.io.as_ref() else {
continue;
};
let session = self
.models
.session(component)
.with_context(|| format!("pipeline component '{component}' was not loaded"))?;
for (port, optional) in &io.optional_inputs {
let endpoint = format!("{component}.{port}");
let route = self.plan.dataflow().iter().find(|edge| edge.to == endpoint);
let supplied_endpoint = inputs
.get(&endpoint)
.map(|value| (endpoint.as_str(), value));
let supplied_route = route.and_then(|edge| {
inputs
.get(&edge.from)
.map(|value| (edge.from.as_str(), value))
});
let supplied = supplied_endpoint.or(supplied_route);
let is_present = present.contains(&optional.presence);
if !is_present {
if let Some((supplied_name, _)) = supplied {
anyhow::bail!(
"pipeline input '{supplied_name}' is associated with presence key '{}' \
but that key was declared absent",
optional.presence
);
}
} else if supplied.is_none() {
let active_route = route.is_some_and(|edge| {
endpoint_component(&edge.from).is_some_and(|producer| {
self.plan.component_is_present(producer, present)
})
});
if !active_route {
anyhow::bail!(
"missing optional-but-present pipeline input '{endpoint}' for presence \
key '{}': supply the destination endpoint or an active routed source",
optional.presence
);
}
}
let info = session
.inputs()
.iter()
.find(|info| info.name == *port)
.with_context(|| {
format!(
"optional pipeline input '{endpoint}' is not exposed by its ONNX graph"
)
})?;
if info.shape.len() != optional.absent.shape.len() {
anyhow::bail!(
"invalid fallback for optional pipeline input '{endpoint}': declared rank {} \
does not match graph rank {}",
optional.absent.shape.len(),
info.shape.len()
);
}
for (index, dimension) in optional.absent.shape.iter().enumerate() {
let TensorDimension::Symbol(symbol) = dimension else {
continue;
};
if info.shape[index] >= 0 {
bind_dimension(&mut dimensions, symbol, info.shape[index], &endpoint)?;
}
if let Some((_, value)) = supplied {
if value.shape().len() != optional.absent.shape.len() {
anyhow::bail!(
"pipeline input '{endpoint}' has rank {}, expected {} from its \
optional-input contract",
value.shape().len(),
optional.absent.shape.len()
);
}
bind_dimension(&mut dimensions, symbol, value.shape()[index], &endpoint)?;
}
}
}
}
let mut tensors = inputs;
for (component, model) in &self.models.directory.spec.models {
let Some(io) = model.io.as_ref() else {
continue;
};
let session = self
.models
.session(component)
.with_context(|| format!("pipeline component '{component}' was not loaded"))?;
for (port, optional) in &io.optional_inputs {
if present.contains(&optional.presence) {
continue;
}
let endpoint = format!("{component}.{port}");
if tensors.contains_key(&endpoint) {
continue;
}
let info = session
.inputs()
.iter()
.find(|info| info.name == *port)
.with_context(|| {
format!(
"optional pipeline input '{endpoint}' is not exposed by its ONNX graph"
)
})?;
let shape = optional
.absent
.shape
.iter()
.map(|dimension| match dimension {
TensorDimension::Fixed(value) => Ok(*value),
TensorDimension::Symbol(symbol) => {
dimensions.get(symbol).copied().with_context(|| {
format!(
"unresolved fallback shape symbol '{symbol}' for optional \
pipeline input '{endpoint}'"
)
})
}
})
.collect::<anyhow::Result<Vec<_>>>()?;
let value = match optional.absent.kind {
AbsentInputKind::Zeros => zero_value(&shape, info.dtype).with_context(|| {
format!(
"invalid fallback for optional pipeline input '{endpoint}' with dtype \
{:?} and shape {shape:?}",
info.dtype
)
})?,
};
tensors.insert(endpoint, value);
}
}
Ok(tensors)
}
fn ensure_component_present(
&self,
component: &str,
present: &BTreeSet<String>,
role: &str,
) -> anyhow::Result<()> {
if let Some(key) = self.plan.presence_condition(component)
&& !present.contains(key)
{
anyhow::bail!(
"{role} '{component}' is gated by absent presence key '{key}' and cannot execute"
);
}
Ok(())
}
fn missing_input_error(
&self,
component: &str,
port: &str,
present: &BTreeSet<String>,
) -> anyhow::Error {
let endpoint = format!("{component}.{port}");
let optional = self
.models
.directory
.spec
.models
.get(component)
.and_then(|model| model.io.as_ref())
.and_then(|io| io.optional_inputs.get(port));
match optional {
Some(optional) if present.contains(&optional.presence) => anyhow::anyhow!(
"missing optional-but-present pipeline input '{endpoint}' for presence key '{}'",
optional.presence
),
Some(optional) => anyhow::anyhow!(
"missing or invalid fallback for absent optional pipeline input '{endpoint}' \
(presence key '{}')",
optional.presence
),
None => anyhow::anyhow!("missing required pipeline input '{endpoint}'"),
}
}
fn run_prompt_phase_components(
&self,
components: &[String],
tensors: &mut PipelineTensors,
phase: &str,
present: &BTreeSet<String>,
mut timings: Option<&mut Vec<serde_json::Value>>,
) -> anyhow::Result<()> {
for component in components {
if !self.plan.component_is_present(component, present) {
continue;
}
let session = self
.models
.session(component)
.with_context(|| format!("pipeline component '{component}' was not loaded"))?;
let inputs = self.component_inputs(component, session, tensors, present)?;
let refs = inputs
.iter()
.map(|(name, value)| (name.as_str(), value))
.collect::<Vec<_>>();
let started = std::time::Instant::now();
let outputs = session
.run(&refs)
.map_err(|e| anyhow::anyhow!("ORT pipeline component '{component}' failed: {e}"))?;
if let Some(sink) = timings.as_deref_mut() {
sink.push(serde_json::json!({
"component": component,
"phase": phase,
"ms": started.elapsed().as_secs_f64() * 1e3,
}));
}
for (name, value) in session.output_names().iter().zip(outputs) {
tensors.insert(format!("{component}.{name}"), value);
}
}
Ok(())
}
fn component_inputs(
&self,
component: &str,
session: &Session,
tensors: &PipelineTensors,
present: &BTreeSet<String>,
) -> anyhow::Result<Vec<(String, Value)>> {
let mut inputs = Vec::new();
for info in session.inputs() {
let endpoint = format!("{component}.{}", info.name);
let routed = self
.plan
.dataflow()
.iter()
.find(|edge| {
edge.to == endpoint
&& endpoint_component(&edge.from)
.is_none_or(|source| self.plan.component_is_present(source, present))
})
.and_then(|edge| tensors.get(&edge.from));
let value = tensors
.get(&endpoint)
.or(routed)
.ok_or_else(|| self.missing_input_error(component, &info.name, present))?;
inputs.push((info.name.clone(), coerce_value_to_dtype(value, info.dtype)?));
}
Ok(inputs)
}
fn decoder_extra_inputs(
&self,
decoder: &str,
tensors: &PipelineTensors,
exclude_input: Option<&str>,
present: &BTreeSet<String>,
) -> anyhow::Result<Vec<(String, Value)>> {
let mut extras = Vec::new();
let mut bound = BTreeSet::new();
for edge in self
.plan
.edges_to_component(decoder)
.filter(|edge| endpoint_component(&edge.from).is_some_and(|from| from != decoder))
.filter(|edge| {
endpoint_component(&edge.from)
.is_none_or(|source| self.plan.component_is_present(source, present))
})
{
let (_, input) = parse_endpoint(&edge.to)?;
if exclude_input == Some(input) {
continue;
}
let value = tensors
.get(&edge.to)
.or_else(|| tensors.get(&edge.from))
.with_context(|| {
format!(
"missing pipeline tensor '{}' and routed source '{}'",
edge.to, edge.from
)
})?;
extras.push((input.to_string(), clone_value(value)?));
bound.insert(input.to_string());
}
if let Some(optional_inputs) = self
.models
.directory
.spec
.models
.get(decoder)
.and_then(|model| model.io.as_ref())
.map(|io| &io.optional_inputs)
{
let session = self
.models
.session(decoder)
.with_context(|| format!("pipeline decoder '{decoder}' was not loaded"))?;
for port in optional_inputs.keys() {
if exclude_input == Some(port.as_str()) || bound.contains(port) {
continue;
}
let endpoint = format!("{decoder}.{port}");
let value = tensors
.get(&endpoint)
.ok_or_else(|| self.missing_input_error(decoder, port, present))?;
let dtype = session
.inputs()
.iter()
.find(|info| info.name == *port)
.with_context(|| {
format!("optional pipeline input '{endpoint}' is not exposed by its graph")
})?
.dtype;
extras.push((port.clone(), coerce_value_to_dtype(value, dtype)?));
}
}
Ok(extras)
}
fn seed_prompt_token_inputs(
&self,
components: &[String],
prompt_tokens: &[TokenId],
tensors: &mut PipelineTensors,
) -> anyhow::Result<()> {
for component in components {
let Some(token_input) = self
.models
.directory
.spec
.models
.get(component)
.and_then(|model| model.io.as_ref())
.and_then(|io| io.token_input.as_deref())
else {
continue;
};
let endpoint = format!("{component}.{token_input}");
let routed = self.plan.dataflow().iter().any(|edge| edge.to == endpoint);
if routed || tensors.contains_key(&endpoint) {
continue;
}
let ids: Vec<i64> = prompt_tokens.iter().map(|&t| i64::from(t)).collect();
let value = Value::from_slice_i64(&ids, &[1, ids.len() as i64])?;
tensors.insert(endpoint, value);
}
Ok(())
}
fn static_cross_kv_bindings(
&self,
cross_kv_pairs: &[(String, String)],
tensors: &PipelineTensors,
) -> anyhow::Result<Vec<(String, Arc<Value>)>> {
let mut bindings = Vec::with_capacity(cross_kv_pairs.len());
for (decoder_input, encoder_output) in cross_kv_pairs {
let suffix = format!(".{encoder_output}");
let mut matches = tensors
.iter()
.filter(|(key, _)| key.ends_with(&suffix) || key.as_str() == encoder_output);
let (_, value) = matches.next().with_context(|| {
format!(
"encoder-decoder cross-attention: no pooled encoder output '{encoder_output}' \
to bind decoder input '{decoder_input}'; the encoder prologue must run and \
publish it before decode"
)
})?;
if matches.next().is_some() {
anyhow::bail!(
"encoder-decoder cross-attention: multiple pooled tensors match encoder output \
'{encoder_output}' for decoder input '{decoder_input}'; the producing \
component is ambiguous"
);
}
#[allow(clippy::arc_with_non_send_sync)]
let shared = Arc::new(clone_value(value)?);
bindings.push((decoder_input.clone(), shared));
}
Ok(bindings)
}
fn decoder_in_edges(
&self,
decoder: &str,
present: &BTreeSet<String>,
tensors: &PipelineTensors,
) -> anyhow::Result<Vec<(String, String)>> {
let mut edges = Vec::new();
let mut bound = BTreeSet::new();
for edge in self
.plan
.edges_to_component(decoder)
.filter(|edge| endpoint_component(&edge.from).is_some_and(|from| from != decoder))
.filter(|edge| {
endpoint_component(&edge.from)
.is_none_or(|source| self.plan.component_is_present(source, present))
})
{
let (_, input) = parse_endpoint(&edge.to)?;
let source = if tensors.contains_key(&edge.to) {
edge.to.clone()
} else {
edge.from.clone()
};
edges.push((source, input.to_string()));
bound.insert(input.to_string());
}
if let Some(io) = self
.models
.directory
.spec
.models
.get(decoder)
.and_then(|model| model.io.as_ref())
{
for port in io.optional_inputs.keys() {
if bound.contains(port) {
continue;
}
let endpoint = format!("{decoder}.{port}");
if tensors.contains_key(&endpoint) {
edges.push((endpoint, port.clone()));
}
}
}
Ok(edges)
}
fn build_step_bindings(
&self,
step_components: &[String],
present: &BTreeSet<String>,
) -> anyhow::Result<Vec<StepComponentBinding>> {
let mut bindings = Vec::with_capacity(step_components.len());
for component in step_components {
if !self.plan.component_is_present(component, present) {
continue;
}
let session = self.models.session(component).with_context(|| {
format!("pipeline every_step component '{component}' was not loaded")
})?;
let token_input = self
.models
.directory
.spec
.models
.get(component)
.and_then(|spec| spec.io.as_ref())
.and_then(|io| io.token_input.clone());
if let Some(port) = &token_input
&& !session.inputs().iter().any(|info| &info.name == port)
{
anyhow::bail!(
"every_step component '{component}' declares io.token_input '{port}' but \
the graph does not expose it; graph inputs: {:?}",
session.input_names()
);
}
let mut routed_inputs = Vec::new();
for info in session.inputs() {
if token_input.as_deref() == Some(info.name.as_str()) {
continue;
}
let endpoint = format!("{component}.{}", info.name);
let routed_from = self
.plan
.dataflow()
.iter()
.find(|edge| {
edge.to == endpoint
&& endpoint_component(&edge.from).is_none_or(|source| {
self.plan.component_is_present(source, present)
})
})
.map(|edge| edge.from.clone());
routed_inputs.push(StepComponentInput {
port: info.name.clone(),
endpoint,
routed_from,
dtype: info.dtype,
missing_message: self
.missing_input_error(component, &info.name, present)
.to_string(),
});
}
bindings.push(StepComponentBinding {
component: component.clone(),
token_input,
routed_inputs,
});
}
Ok(bindings)
}
}
fn tokenize_with(tokenizer: &Tokenizer, prompt: &GeneratePrompt) -> anyhow::Result<Vec<TokenId>> {
match prompt {
GeneratePrompt::TokenIds(tokens) => Ok(tokens.clone()),
GeneratePrompt::Text(text) => tokenizer
.encode(text)
.map_err(|e| anyhow::anyhow!("Failed to tokenize prompt: {}", e)),
}
}
fn expand_image_placeholders_count_based(
prompt_tokens: Vec<TokenId>,
num_image_tiles: Option<usize>,
vision: Option<&PipelineVisionConfig>,
) -> anyhow::Result<Vec<TokenId>> {
let num_tiles = match num_image_tiles {
None => return Ok(prompt_tokens),
Some(n) => n,
};
let (placeholder_i64, tokens_per_tile) = match vision {
Some(v) => match (v.image_placeholder_token_id, v.tokens_per_tile) {
(Some(id), Some(tpt)) => (id, tpt),
_ => anyhow::bail!(
"image tile count supplied but pipeline metadata vision contract is incomplete: \
both image_placeholder_token_id and tokens_per_tile must be set"
),
},
None => anyhow::bail!(
"image tile count supplied but pipeline metadata declares no vision section; \
add pipeline.vision with image_placeholder_token_id and tokens_per_tile"
),
};
if tokens_per_tile == 0 {
anyhow::bail!("pipeline metadata tokens_per_tile is 0; must be at least 1");
}
let placeholder_id: TokenId = u32::try_from(placeholder_i64).with_context(|| {
format!("image_placeholder_token_id {placeholder_i64} is out of range for token ID (u32)")
})?;
let placeholder_count = prompt_tokens
.iter()
.filter(|&&t| t == placeholder_id)
.count();
if placeholder_count == 0 {
anyhow::bail!(
"num_image_tiles supplied but prompt contains no image placeholder token \
(id={placeholder_id}); the prompt must contain exactly one placeholder"
);
}
if placeholder_count > 1 {
anyhow::bail!(
"multi-image count-based expansion is not supported: found {placeholder_count} image \
placeholders (id={placeholder_id}) but only an aggregate tile count is available; \
supply a single image or thread per-image tile counts"
);
}
let expansion: usize = tokens_per_tile.checked_mul(num_tiles).context(
"image token expansion overflow: tokens_per_tile * num_image_tiles is too large",
)?;
let new_len = prompt_tokens
.len()
.checked_sub(1)
.and_then(|base| base.checked_add(expansion))
.context("expanded prompt token sequence length overflows")?;
let mut expanded = Vec::new();
expanded
.try_reserve_exact(new_len)
.context("failed to allocate expanded prompt token sequence")?;
for token in prompt_tokens {
if token == placeholder_id {
for _ in 0..expansion {
expanded.push(placeholder_id);
}
} else {
expanded.push(token);
}
}
if expanded.is_empty() {
anyhow::bail!(
"image placeholder expansion produced an empty token sequence; \
check that num_image_tiles > 0 and the prompt contains non-placeholder tokens"
);
}
Ok(expanded)
}
struct StepComponentBinding {
component: String,
token_input: Option<String>,
routed_inputs: Vec<StepComponentInput>,
}
struct StepComponentInput {
port: String,
endpoint: String,
routed_from: Option<String>,
dtype: DataType,
missing_message: String,
}
struct PipelineDecodeLoopBackend<'a> {
decoder: &'a Session,
decoder_state: &'a mut DecodeState,
pool: &'a mut PipelineTensors,
step_components: Vec<(StepComponentBinding, &'a Session)>,
decoder_in_edges: Vec<(String, String)>,
static_cross_kv: Vec<(String, Arc<Value>)>,
context_tokens: Vec<TokenId>,
prompt_len: usize,
generated_count: usize,
}
impl PipelineDecodeLoopBackend<'_> {
fn run_step_components(&mut self, seed: &[TokenId]) -> anyhow::Result<()> {
if self.step_components.is_empty() {
return Ok(());
}
let ids: Vec<i64> = seed.iter().map(|&t| i64::from(t)).collect();
let seq = ids.len() as i64;
for (binding, session) in &self.step_components {
let mut inputs: Vec<(String, Value)> =
Vec::with_capacity(binding.routed_inputs.len() + 1);
for routed in &binding.routed_inputs {
let value = self
.pool
.get(&routed.endpoint)
.or_else(|| {
routed
.routed_from
.as_deref()
.and_then(|from| self.pool.get(from))
})
.with_context(|| routed.missing_message.clone())?;
inputs.push((
routed.port.clone(),
coerce_value_to_dtype(value, routed.dtype)?,
));
}
if let Some(port) = &binding.token_input {
inputs.push((port.clone(), Value::from_slice_i64(&ids, &[1, seq])?));
}
let refs = inputs
.iter()
.map(|(name, value)| (name.as_str(), value))
.collect::<Vec<_>>();
let outputs = session.run(&refs).map_err(|e| {
anyhow::anyhow!(
"ORT every_step component '{}' failed: {e}",
binding.component
)
})?;
for (name, value) in session.output_names().iter().zip(outputs) {
self.pool
.insert(format!("{}.{}", binding.component, name), value);
}
}
Ok(())
}
fn decoder_extras(&self) -> anyhow::Result<Vec<(String, Value)>> {
let mut extras = Vec::with_capacity(self.decoder_in_edges.len() + self.static_cross_kv.len());
for (from, port) in &self.decoder_in_edges {
let value = self.pool.get(from).with_context(|| {
format!("missing routed pipeline tensor '{from}' for decoder input '{port}'")
})?;
extras.push((port.clone(), clone_value(value)?));
}
for (port, value) in &self.static_cross_kv {
let aliased = Value::alias_with_shape(Arc::clone(value), value.shape())?;
extras.push((port.clone(), aliased));
}
Ok(extras)
}
}
impl DecodeLoopBackend for PipelineDecodeLoopBackend<'_> {
fn context_len(&self) -> usize {
self.context_tokens.len()
}
fn processor_prompt_tokens(&self) -> &[TokenId] {
&self.context_tokens
}
fn next_logits(&mut self) -> anyhow::Result<Vec<f32>> {
let past_len = if self.decoder_state.use_kv {
self.context_tokens
.len()
.saturating_sub(if self.generated_count == 0 {
self.prompt_len
} else {
1
})
} else {
0
};
let input_tokens = if self.decoder_state.use_kv && self.generated_count > 0 {
self.context_tokens[self.context_tokens.len() - 1..].to_vec()
} else {
self.context_tokens.clone()
};
self.run_step_components(&input_tokens)?;
let extras = self.decoder_extras()?;
let outputs = run_decode_step_with_extra(
self.decoder,
self.decoder_state,
&input_tokens,
past_len,
&extras,
)?;
extract_next_token_logits_with_io(
self.decoder,
outputs,
self.decoder_state.io.logits_output.as_deref(),
)
}
fn commit_token(&mut self, token_id: TokenId) -> anyhow::Result<()> {
self.context_tokens.push(token_id);
self.generated_count += 1;
Ok(())
}
}
#[derive(Debug, Clone)]
#[allow(clippy::large_enum_variant)] enum PipelinePlan {
Autoregressive(AutoregressivePlan),
NestedAutoregressive(NestedAutoregressivePlan),
SinglePass(SinglePassPlan),
Iterative(Box<IterativePlan>),
Composite(CompositePlan),
}
#[derive(Debug, Clone)]
struct CompositePlan {
stages: Vec<CompositeStage>,
dataflow: Vec<DataflowEdge>,
presence_conditions: HashMap<String, String>,
}
#[derive(Debug, Clone)]
struct CompositeStage {
name: String,
kind: CompositeStageKind,
}
#[derive(Debug, Clone)]
enum CompositeStageKind {
SinglePass { model: String },
}
#[derive(Debug, Clone)]
struct AutoregressivePlan {
decoder: String,
prompt_components: Vec<String>,
step_components: Vec<String>,
post_decode_components: Vec<String>,
dataflow: Vec<DataflowEdge>,
presence_conditions: HashMap<String, String>,
}
#[derive(Debug, Clone)]
struct NestedAutoregressivePlan {
outer: String,
inner: String,
num_code_groups: usize,
max_frames: usize,
outer_hidden_output: String,
inner_embeds_input: String,
prompt_components: Vec<String>,
post_decode_components: Vec<String>,
pre_embedder: Option<PreEmbedderBinding>,
prefill_embedder: Option<PrefillEmbedderBinding>,
dataflow: Vec<DataflowEdge>,
presence_conditions: HashMap<String, String>,
}
#[derive(Debug, Clone)]
struct PreEmbedderBinding {
component: String,
outer_input: String,
output_port: String,
frame_codes_input: String,
text_embed_input: Option<String>,
}
#[derive(Debug, Clone)]
struct PrefillEmbedderBinding {
component: String,
prompt_input: String,
prefill_output: String,
trailing_output: String,
}
#[derive(Debug, Clone)]
struct SinglePassPlan {
model: String,
prompt_components: Vec<String>,
dataflow: Vec<DataflowEdge>,
presence_conditions: HashMap<String, String>,
}
fn bind_dimension(
dimensions: &mut HashMap<String, i64>,
symbol: &str,
value: i64,
endpoint: &str,
) -> anyhow::Result<()> {
if value < 0 {
anyhow::bail!(
"cannot resolve fallback shape symbol '{symbol}' for '{endpoint}' from dynamic \
dimension {value}"
);
}
if let Some(previous) = dimensions.insert(symbol.to_string(), value)
&& previous != value
{
anyhow::bail!(
"conflicting values for fallback shape symbol '{symbol}': {previous} and {value} \
while resolving '{endpoint}'"
);
}
Ok(())
}
fn zero_value(shape: &[i64], dtype: DataType) -> anyhow::Result<Value> {
let numel = shape.iter().try_fold(1usize, |count, &dimension| {
let dimension = usize::try_from(dimension)
.map_err(|_| anyhow::anyhow!("negative tensor dimension {dimension}"))?;
count
.checked_mul(dimension)
.context("fallback tensor element count overflow")
})?;
match dtype {
DataType::Float32 | DataType::Float16 | DataType::BFloat16 => {
Value::from_f32_slice_as(&vec![0.0; numel], shape, dtype).map_err(Into::into)
}
DataType::Int64 => Value::from_slice_i64(&vec![0; numel], shape).map_err(Into::into),
other => anyhow::bail!(
"zero fallback materialization does not support graph input dtype {other:?}"
),
}
}
fn coerce_value_to_dtype(value: &Value, target: DataType) -> anyhow::Result<Value> {
if value.dtype() == target {
return clone_value(value);
}
match (value.dtype(), target) {
(
DataType::Float32 | DataType::Float16 | DataType::BFloat16,
DataType::Float32 | DataType::Float16 | DataType::BFloat16,
) => {
let data = value.to_vec_f32_lossy()?;
Value::from_f32_slice_as(&data, value.shape(), target)
.map_err(|e| anyhow::anyhow!("failed to coerce value to {target:?}: {e}"))
}
_ => clone_value(value),
}
}
fn dump_iterative_step(denoiser: &str, port: &str, step: usize, value: &Value, step_ms: f64) {
let Ok(dir) = std::env::var("ONNX_GENAI_STEP_DUMP_DIR") else {
return;
};
let shape: Vec<i64> = value.shape().to_vec();
let payload = match value.dtype() {
DataType::Int64 | DataType::Int32 | DataType::Int16 | DataType::Int8 => value
.to_vec_i64()
.ok()
.map(|data| serde_json::json!({"dtype": "i64", "shape": shape, "data": data, "step_ms": step_ms})),
_ => value
.to_vec_f32()
.ok()
.map(|data| serde_json::json!({"dtype": "f32", "shape": shape, "data": data, "step_ms": step_ms})),
};
if let Some(payload) = payload {
let path =
std::path::Path::new(&dir).join(format!("step_{step:04}_{denoiser}_{port}.json"));
let _ = std::fs::write(path, payload.to_string());
}
}
fn dump_stage_timings(stages: &[serde_json::Value]) {
let Ok(dir) = std::env::var("ONNX_GENAI_STEP_DUMP_DIR") else {
return;
};
let path = std::path::Path::new(&dir).join("stages.json");
let _ = std::fs::write(path, serde_json::json!({ "stages": stages }).to_string());
}
#[derive(Debug, Clone)]
struct IterativePlan {
denoiser: String,
num_steps: usize,
guidance_scale: Option<f32>,
prompt_components: Vec<String>,
final_components: Vec<String>,
loop_edges: Vec<(String, String)>,
timestep_input: Option<String>,
start_step: usize,
timesteps: Option<Vec<f32>>,
scheduler: Option<Arc<dyn Scheduler>>,
cfg_conditioning_input: Option<String>,
dataflow: Vec<DataflowEdge>,
scheduler_spec: Option<SchedulerSpec>,
scheduler_registry: SchedulerRegistry,
presence_conditions: HashMap<String, String>,
}
pub trait Scheduler: Send + Sync + std::fmt::Debug {
fn step(
&self,
step: usize,
num_steps: usize,
sample: &Value,
model_output: &Value,
) -> anyhow::Result<Value>;
fn reset(&self) {}
fn needs_noise(&self) -> bool {
false
}
fn step_with_noise(
&self,
step: usize,
num_steps: usize,
sample: &Value,
model_output: &Value,
_noise: Option<&Value>,
) -> anyhow::Result<Value> {
self.step(step, num_steps, sample, model_output)
}
fn scale_input(
&self,
_step: usize,
_num_steps: usize,
_sample: &Value,
) -> anyhow::Result<Option<Value>> {
Ok(None)
}
fn init_noise_sigma(&self) -> f32 {
1.0
}
fn timesteps(&self) -> Option<Vec<f32>> {
None
}
fn cfg_uncond_sample(&self, _sample: &Value) -> anyhow::Result<Option<Value>> {
Ok(None)
}
}
pub type SchedulerFactory =
Arc<dyn Fn(&SchedulerSpec, usize) -> anyhow::Result<Arc<dyn Scheduler>> + Send + Sync>;
#[derive(Clone)]
pub struct SchedulerRegistry {
factories: HashMap<String, SchedulerFactory>,
}
impl std::fmt::Debug for SchedulerRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SchedulerRegistry")
.field("kinds", &self.factories.keys().collect::<Vec<_>>())
.finish()
}
}
impl SchedulerRegistry {
pub fn builtin() -> Self {
let mut factories: HashMap<String, SchedulerFactory> = HashMap::new();
factories.insert(
"ddim".to_string(),
Arc::new(|cfg: &SchedulerSpec, num_steps: usize| {
let prediction = PredictionType::parse(cfg.prediction_type.as_deref())?;
let sched = DdimSchedule::with_schedule(
cfg.num_train_timesteps.unwrap_or(1000),
cfg.beta_start.unwrap_or(0.00085),
cfg.beta_end.unwrap_or(0.012),
cfg.beta_schedule.as_deref().unwrap_or("linear"),
num_steps,
)?
.with_prediction(prediction);
Ok(Arc::new(sched) as Arc<dyn Scheduler>)
}),
);
factories.insert(
"euler".to_string(),
Arc::new(|cfg: &SchedulerSpec, num_steps: usize| {
let prediction = PredictionType::parse(cfg.prediction_type.as_deref())?;
let sched = EulerSchedule::with_schedule(
cfg.num_train_timesteps.unwrap_or(1000),
cfg.beta_start.unwrap_or(0.00085),
cfg.beta_end.unwrap_or(0.012),
cfg.beta_schedule.as_deref().unwrap_or("scaled_linear"),
num_steps,
sigma_spacing(cfg)?,
)?
.with_prediction(prediction);
Ok(Arc::new(sched) as Arc<dyn Scheduler>)
}),
);
factories.insert(
"euler_ancestral".to_string(),
Arc::new(|cfg: &SchedulerSpec, num_steps: usize| {
let prediction = PredictionType::parse(cfg.prediction_type.as_deref())?;
let sched = EulerAncestral::with_schedule(
cfg.num_train_timesteps.unwrap_or(1000),
cfg.beta_start.unwrap_or(0.00085),
cfg.beta_end.unwrap_or(0.012),
cfg.beta_schedule.as_deref().unwrap_or("scaled_linear"),
num_steps,
sigma_spacing(cfg)?,
)?
.with_prediction(prediction);
Ok(Arc::new(sched) as Arc<dyn Scheduler>)
}),
);
factories.insert(
"dpmpp_2m".to_string(),
Arc::new(|cfg: &SchedulerSpec, num_steps: usize| {
let prediction = PredictionType::parse(cfg.prediction_type.as_deref())?;
let sched = Dpmpp2m::with_schedule(
cfg.num_train_timesteps.unwrap_or(1000),
cfg.beta_start.unwrap_or(0.00085),
cfg.beta_end.unwrap_or(0.012),
cfg.beta_schedule.as_deref().unwrap_or("scaled_linear"),
num_steps,
sigma_spacing(cfg)?,
)?
.with_prediction(prediction);
Ok(Arc::new(sched) as Arc<dyn Scheduler>)
}),
);
factories.insert(
"masked_diffusion".to_string(),
Arc::new(|cfg: &SchedulerSpec, _num_steps: usize| {
let mask_token_id = cfg
.mask_token_id
.context("masked_diffusion scheduler requires 'mask_token_id'")?;
let temperature = cfg.temperature.unwrap_or(0.0);
if temperature < 0.0 {
anyhow::bail!("masked_diffusion temperature must be >= 0");
}
if let Some(block_length) = cfg.block_length
&& block_length == 0
{
anyhow::bail!("masked_diffusion block_length must be >= 1");
}
let remasking = match cfg.remasking.as_deref() {
None | Some("low_confidence") => Remasking::LowConfidence,
Some("random") => Remasking::Random,
Some(other) => anyhow::bail!(
"masked_diffusion remasking must be 'low_confidence' or 'random', \
got '{other}'"
),
};
Ok(Arc::new(MaskedDiffusion {
mask_token_id,
temperature,
block_length: cfg.block_length,
remasking,
generation_start: Mutex::new(None),
}) as Arc<dyn Scheduler>)
}),
);
Self { factories }
}
pub fn register(&mut self, kind: impl Into<String>, factory: SchedulerFactory) {
self.factories.insert(kind.into(), factory);
}
fn build(&self, spec: &SchedulerSpec, num_steps: usize) -> anyhow::Result<Arc<dyn Scheduler>> {
let factory = self.factories.get(&spec.kind).with_context(|| {
format!(
"unknown scheduler kind '{}' (registered: {:?})",
spec.kind,
self.factories.keys().collect::<Vec<_>>()
)
})?;
factory(spec, num_steps)
}
}
impl Default for SchedulerRegistry {
fn default() -> Self {
Self::builtin()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Remasking {
LowConfidence,
Random,
}
#[derive(Debug)]
struct MaskedDiffusion {
mask_token_id: i64,
temperature: f32,
block_length: Option<usize>,
remasking: Remasking,
generation_start: Mutex<Option<Vec<usize>>>,
}
impl MaskedDiffusion {
fn ensure_generation_start(&self, tokens: &[i64], batch: usize, sequence_length: usize) {
let mut guard = self.generation_start.lock().unwrap();
if guard.is_some() {
return;
}
let mut starts = Vec::with_capacity(batch);
for row_index in 0..batch {
let start = row_index * sequence_length;
let first_mask = tokens[start..start + sequence_length]
.iter()
.position(|&token| token == self.mask_token_id)
.unwrap_or(sequence_length);
starts.push(first_mask);
}
*guard = Some(starts);
}
fn predict_row(&self, row: &[f32], gumbel: &[f64]) -> (i64, f32) {
let max_logit = row.iter().copied().fold(f32::MIN, f32::max);
let sum_exp: f32 = row.iter().map(|&x| (x - max_logit).exp()).sum();
let mut best_index = 0usize;
let mut best_score = f32::MIN;
for (j, &logit) in row.iter().enumerate() {
let score = if self.temperature > 0.0 {
let u = gumbel[j];
logit - self.temperature * (-u.ln()).ln() as f32
} else {
logit
};
if score > best_score {
best_score = score;
best_index = j;
}
}
let confidence = (row[best_index] - max_logit).exp() / sum_exp;
(best_index as i64, confidence)
}
fn sample_token(&self, row: &[f32], step: usize, position: usize, vocab: usize) -> i64 {
let gumbel = if self.temperature > 0.0 {
gumbel_uniforms(step, position, vocab)
} else {
Vec::new()
};
let mut best_index: Option<usize> = None;
let mut best_score = f32::MIN;
for (j, &logit) in row.iter().enumerate() {
if j as i64 == self.mask_token_id {
continue;
}
let score = if self.temperature > 0.0 {
logit - self.temperature * (-gumbel[j].ln()).ln() as f32
} else {
logit
};
if best_index.is_none() || score > best_score {
best_score = score;
best_index = Some(j);
}
}
best_index.unwrap_or(0) as i64
}
}
impl Scheduler for MaskedDiffusion {
fn reset(&self) {
*self.generation_start.lock().unwrap() = None;
}
fn cfg_uncond_sample(&self, sample: &Value) -> anyhow::Result<Option<Value>> {
let shape = sample.shape().to_vec();
let tokens = sample.to_vec_i64()?;
let count = tokens.len();
let sequence_length = *shape.last().unwrap_or(&(count as i64)) as usize;
if sequence_length == 0 {
return Ok(None);
}
let batch = count.checked_div(sequence_length).unwrap_or(0).max(1);
self.ensure_generation_start(&tokens, batch, sequence_length);
let generation_start = self.generation_start.lock().unwrap().clone().unwrap();
let mut output = tokens;
for (row_index, &prompt_length) in generation_start.iter().enumerate() {
let row_start = row_index * sequence_length;
for offset in 0..prompt_length.min(sequence_length) {
output[row_start + offset] = self.mask_token_id;
}
}
Value::from_slice_i64(&output, &shape)
.map(Some)
.map_err(Into::into)
}
fn step(
&self,
step: usize,
num_steps: usize,
tokens: &Value,
logits: &Value,
) -> anyhow::Result<Value> {
let token_shape = tokens.shape().to_vec();
let tokens = tokens.to_vec_i64()?;
let sequence_count = tokens.len();
let logit_shape = logits.shape();
let vocab = *logit_shape
.last()
.context("masked_diffusion logits must be rank >= 1")? as usize;
if vocab == 0 || sequence_count == 0 || logits.numel() != sequence_count * vocab {
anyhow::bail!(
"masked_diffusion shape mismatch: tokens {token_shape:?}, logits {logit_shape:?}"
);
}
let sequence_length = *token_shape.last().unwrap_or(&(sequence_count as i64)) as usize;
let batch = sequence_count
.checked_div(sequence_length)
.unwrap_or(0)
.max(1);
self.ensure_generation_start(&tokens, batch, sequence_length);
let generation_start = self.generation_start.lock().unwrap().clone().unwrap();
let all_logits = logits.to_vec_f32()?;
let mut output = tokens.clone();
for (row_index, &prompt_length) in generation_start.iter().enumerate() {
let row_start = row_index * sequence_length;
let generation_length = sequence_length.saturating_sub(prompt_length);
if generation_length == 0 {
continue;
}
let block_length = self
.block_length
.unwrap_or(generation_length)
.min(generation_length)
.max(1);
if !generation_length.is_multiple_of(block_length) {
anyhow::bail!(
"masked_diffusion: generation length {generation_length} is not divisible \
by block_length {block_length}"
);
}
let num_blocks = generation_length / block_length;
if !num_steps.is_multiple_of(num_blocks) {
anyhow::bail!(
"masked_diffusion: num_steps {num_steps} is not divisible by num_blocks \
{num_blocks} (generation_length {generation_length} / block_length \
{block_length})"
);
}
let steps_per_block = num_steps / num_blocks;
let block_index = (step / steps_per_block).min(num_blocks - 1);
let step_in_block = step % steps_per_block;
let block_start = prompt_length + block_index * block_length;
let block_end = (block_start + block_length).min(sequence_length);
let remaining_steps_in_block = steps_per_block - step_in_block;
match self.remasking {
Remasking::LowConfidence => {
let mut candidates: Vec<(usize, i64, f32)> = Vec::new();
for offset in block_start..block_end {
let position = row_start + offset;
if tokens[position] != self.mask_token_id {
continue;
}
let logit_row = &all_logits[position * vocab..(position + 1) * vocab];
let gumbel = if self.temperature > 0.0 {
gumbel_uniforms(step, position, vocab)
} else {
Vec::new()
};
let (predicted, confidence) = self.predict_row(logit_row, &gumbel);
candidates.push((position, predicted, confidence));
}
if candidates.is_empty() {
continue;
}
candidates
.sort_by(|a, b| b.2.partial_cmp(&a.2).unwrap_or(std::cmp::Ordering::Equal));
let commit = candidates.len().div_ceil(remaining_steps_in_block);
for &(position, predicted, _) in candidates.iter().take(commit) {
output[position] = predicted;
}
}
Remasking::Random => {
let last_step_in_block = remaining_steps_in_block <= 1;
let unmask_prob = 1.0f64 / remaining_steps_in_block as f64;
for offset in block_start..block_end {
let position = row_start + offset;
if tokens[position] != self.mask_token_id {
continue;
}
if last_step_in_block || unmask_uniform(step, position) < unmask_prob {
let logit_row = &all_logits[position * vocab..(position + 1) * vocab];
output[position] = self.sample_token(logit_row, step, position, vocab);
}
}
}
}
}
Value::from_slice_i64(&output, &token_shape).map_err(Into::into)
}
}
fn gumbel_uniforms(step: usize, position: usize, vocab: usize) -> Vec<f64> {
use rand::{Rng, SeedableRng};
let seed = (step as u64)
.wrapping_mul(0x9E37_79B9_7F4A_7C15)
.wrapping_add(position as u64);
let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
(0..vocab)
.map(|_| rng.random::<f64>().clamp(1e-9, 1.0 - 1e-9))
.collect()
}
fn unmask_uniform(step: usize, position: usize) -> f64 {
use rand::{Rng, SeedableRng};
let seed = (step as u64)
.wrapping_mul(0x2545_F491_4F6C_DD1D)
.wrapping_add((position as u64).wrapping_mul(0xD1B5_4A32_D192_ED03))
.wrapping_add(0xA076_1D64_78BD_642F);
let mut rng = rand::rngs::StdRng::seed_from_u64(seed);
rng.random::<f64>()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum PredictionType {
Epsilon,
VPrediction,
Sample,
}
impl PredictionType {
fn parse(value: Option<&str>) -> anyhow::Result<Self> {
match value.unwrap_or("epsilon") {
"epsilon" => Ok(Self::Epsilon),
"v_prediction" => Ok(Self::VPrediction),
"sample" | "x0" => Ok(Self::Sample),
other => anyhow::bail!(
"unsupported prediction_type '{other}' (expected 'epsilon', 'v_prediction', or 'sample'/'x0')"
),
}
}
}
#[inline]
fn epsilon_from_model_output(
model_out: f32,
x_t: f32,
alpha_t: f32,
sigma_t: f32,
prediction: PredictionType,
) -> f32 {
match prediction {
PredictionType::Epsilon => model_out,
PredictionType::VPrediction => alpha_t * model_out + sigma_t * x_t,
PredictionType::Sample => (x_t - alpha_t * model_out) / sigma_t,
}
}
#[inline]
fn x0_from_model_output(
model_out: f32,
x_t: f32,
alpha_t: f32,
sigma_t: f32,
prediction: PredictionType,
) -> f32 {
match prediction {
PredictionType::Epsilon => (x_t - sigma_t * model_out) / alpha_t,
PredictionType::VPrediction => alpha_t * x_t - sigma_t * model_out,
PredictionType::Sample => model_out,
}
}
#[derive(Debug, Clone)]
struct DdimSchedule {
steps: Vec<(f32, f32)>,
timesteps: Vec<f32>,
prediction: PredictionType,
}
impl DdimSchedule {
fn with_schedule(
num_train_timesteps: usize,
beta_start: f32,
beta_end: f32,
beta_schedule: &str,
num_steps: usize,
) -> anyhow::Result<Self> {
if num_train_timesteps < 2 {
anyhow::bail!("scheduler num_train_timesteps must be >= 2");
}
if num_steps == 0 || num_steps > num_train_timesteps {
anyhow::bail!("scheduler num_steps ({num_steps}) must be in 1..={num_train_timesteps}");
}
let denom = (num_train_timesteps - 1) as f32;
let (lo, hi, square) = match beta_schedule {
"linear" => (beta_start, beta_end, false),
"scaled_linear" => (beta_start.sqrt(), beta_end.sqrt(), true),
other => anyhow::bail!(
"unsupported scheduler beta_schedule '{other}' (expected 'linear' or 'scaled_linear')"
),
};
let mut alpha_cumprod = Vec::with_capacity(num_train_timesteps);
let mut prod = 1.0f32;
for i in 0..num_train_timesteps {
let mut beta = lo + (hi - lo) * (i as f32) / denom;
if square {
beta *= beta;
}
prod *= 1.0 - beta;
alpha_cumprod.push(prod);
}
let step_ratio = num_train_timesteps / num_steps;
let ascending: Vec<usize> = (0..num_steps).map(|i| i * step_ratio).collect();
let mut steps = Vec::with_capacity(num_steps);
let mut timesteps = Vec::with_capacity(num_steps);
for k in 0..num_steps {
let t = ascending[num_steps - 1 - k];
timesteps.push(t as f32);
let a_t = alpha_cumprod[t];
let a_prev = if k + 1 < num_steps {
alpha_cumprod[ascending[num_steps - 1 - (k + 1)]]
} else {
1.0
};
steps.push((a_t, a_prev));
}
Ok(Self {
steps,
timesteps,
prediction: PredictionType::Epsilon,
})
}
fn with_prediction(mut self, prediction: PredictionType) -> Self {
self.prediction = prediction;
self
}
fn step(&self, k: usize, sample: &[f32], model_out: &[f32]) -> anyhow::Result<Vec<f32>> {
if sample.len() != model_out.len() {
anyhow::bail!(
"scheduler sample/model_output length mismatch: {} vs {}",
sample.len(),
model_out.len()
);
}
let (a_t, a_prev) = self.steps[k];
let sqrt_a_t = a_t.sqrt();
let sqrt_one_minus_a_t = (1.0 - a_t).sqrt();
let sqrt_a_prev = a_prev.sqrt();
let sqrt_one_minus_a_prev = (1.0 - a_prev).sqrt();
Ok(sample
.iter()
.zip(model_out)
.map(|(&x, &m)| {
let e = epsilon_from_model_output(
m,
x,
sqrt_a_t,
sqrt_one_minus_a_t,
self.prediction,
);
let x0_hat = (x - sqrt_one_minus_a_t * e) / sqrt_a_t;
sqrt_a_prev * x0_hat + sqrt_one_minus_a_prev * e
})
.collect())
}
}
impl Scheduler for DdimSchedule {
fn step(
&self,
step: usize,
_num_steps: usize,
sample: &Value,
model_output: &Value,
) -> anyhow::Result<Value> {
let shape = sample.shape().to_vec();
let stepped = DdimSchedule::step(
self,
step,
&sample.to_vec_f32_lossy()?,
&model_output.to_vec_f32_lossy()?,
)?;
Value::from_slice_f32(&stepped, &shape).map_err(Into::into)
}
fn timesteps(&self) -> Option<Vec<f32>> {
Some(self.timesteps.clone())
}
}
#[derive(Debug, Clone)]
struct EulerSchedule {
sigmas: Vec<f32>,
timesteps: Vec<f32>,
prediction: PredictionType,
}
impl EulerSchedule {
fn with_schedule(
num_train_timesteps: usize,
beta_start: f32,
beta_end: f32,
beta_schedule: &str,
num_steps: usize,
spacing: &str,
) -> anyhow::Result<Self> {
if num_train_timesteps < 2 {
anyhow::bail!("scheduler num_train_timesteps must be >= 2");
}
if num_steps == 0 || num_steps > num_train_timesteps {
anyhow::bail!("scheduler num_steps ({num_steps}) must be in 1..={num_train_timesteps}");
}
if let Some(sigmas) = spacing_sigmas(
spacing,
num_train_timesteps,
beta_start,
beta_end,
beta_schedule,
num_steps,
)? {
let train = training_sigmas(num_train_timesteps, beta_start, beta_end, beta_schedule)?;
let timesteps = sigmas[..num_steps]
.iter()
.map(|&s| sigma_to_t(&train, s))
.collect();
return Ok(Self {
sigmas,
timesteps,
prediction: PredictionType::Epsilon,
});
}
let denom = (num_train_timesteps - 1) as f32;
let (lo, hi, square) = match beta_schedule {
"linear" => (beta_start, beta_end, false),
"scaled_linear" => (beta_start.sqrt(), beta_end.sqrt(), true),
other => anyhow::bail!(
"unsupported scheduler beta_schedule '{other}' (expected 'linear' or 'scaled_linear')"
),
};
let mut train_sigmas = Vec::with_capacity(num_train_timesteps);
let mut prod = 1.0f32;
for i in 0..num_train_timesteps {
let mut beta = lo + (hi - lo) * (i as f32) / denom;
if square {
beta *= beta;
}
prod *= 1.0 - beta;
train_sigmas.push(((1.0 - prod) / prod).sqrt());
}
let ts_denom = if num_steps > 1 {
(num_steps - 1) as f32
} else {
1.0
};
let interp = |t: f32| -> f32 {
let low = t.floor().max(0.0) as usize;
let high = (low + 1).min(num_train_timesteps - 1);
let frac = t - low as f32;
train_sigmas[low] * (1.0 - frac) + train_sigmas[high] * frac
};
let mut sigmas = Vec::with_capacity(num_steps + 1);
let mut timesteps = Vec::with_capacity(num_steps);
for k in 0..num_steps {
let idx = num_steps - 1 - k;
let t = idx as f32 * denom / ts_denom;
timesteps.push(t);
sigmas.push(interp(t));
}
sigmas.push(0.0);
Ok(Self {
sigmas,
timesteps,
prediction: PredictionType::Epsilon,
})
}
fn with_prediction(mut self, prediction: PredictionType) -> Self {
self.prediction = prediction;
self
}
fn scale(&self, step: usize, sample: &[f32]) -> Vec<f32> {
let factor = (self.sigmas[step] * self.sigmas[step] + 1.0).sqrt();
sample.iter().map(|&x| x / factor).collect()
}
fn step_vec(&self, step: usize, sample: &[f32], model_out: &[f32]) -> anyhow::Result<Vec<f32>> {
if sample.len() != model_out.len() {
anyhow::bail!(
"scheduler sample/model_output length mismatch: {} vs {}",
sample.len(),
model_out.len()
);
}
let sigma = self.sigmas[step];
let alpha_t = 1.0 / (sigma * sigma + 1.0).sqrt();
let sigma_t = sigma * alpha_t;
let dt = self.sigmas[step + 1] - self.sigmas[step];
Ok(sample
.iter()
.zip(model_out)
.map(|(&x, &m)| {
let e = epsilon_from_model_output(m, alpha_t * x, alpha_t, sigma_t, self.prediction);
x + e * dt
})
.collect())
}
}
impl Scheduler for EulerSchedule {
fn step(
&self,
step: usize,
_num_steps: usize,
sample: &Value,
model_output: &Value,
) -> anyhow::Result<Value> {
let shape = sample.shape().to_vec();
let stepped = self.step_vec(
step,
&sample.to_vec_f32_lossy()?,
&model_output.to_vec_f32_lossy()?,
)?;
Value::from_slice_f32(&stepped, &shape).map_err(Into::into)
}
fn scale_input(
&self,
step: usize,
_num_steps: usize,
sample: &Value,
) -> anyhow::Result<Option<Value>> {
let shape = sample.shape().to_vec();
let scaled = self.scale(step, &sample.to_vec_f32_lossy()?);
Ok(Some(Value::from_slice_f32(&scaled, &shape)?))
}
fn init_noise_sigma(&self) -> f32 {
self.sigmas[0]
}
fn timesteps(&self) -> Option<Vec<f32>> {
Some(self.timesteps.clone())
}
}
#[derive(Debug, Clone)]
struct EulerAncestral {
sigmas: Vec<f32>,
timesteps: Vec<f32>,
prediction: PredictionType,
}
impl EulerAncestral {
fn with_schedule(
num_train_timesteps: usize,
beta_start: f32,
beta_end: f32,
beta_schedule: &str,
num_steps: usize,
spacing: &str,
) -> anyhow::Result<Self> {
let euler = EulerSchedule::with_schedule(
num_train_timesteps,
beta_start,
beta_end,
beta_schedule,
num_steps,
spacing,
)?;
Ok(Self {
sigmas: euler.sigmas,
timesteps: euler.timesteps,
prediction: PredictionType::Epsilon,
})
}
fn with_prediction(mut self, prediction: PredictionType) -> Self {
self.prediction = prediction;
self
}
}
impl Scheduler for EulerAncestral {
fn step(
&self,
_step: usize,
_num_steps: usize,
_sample: &Value,
_model_output: &Value,
) -> anyhow::Result<Value> {
anyhow::bail!("euler_ancestral is stochastic; the loop must call step_with_noise")
}
fn needs_noise(&self) -> bool {
true
}
fn step_with_noise(
&self,
step: usize,
_num_steps: usize,
sample: &Value,
model_output: &Value,
noise: Option<&Value>,
) -> anyhow::Result<Value> {
let shape = sample.shape().to_vec();
let x = sample.to_vec_f32_lossy()?;
let model_out = model_output.to_vec_f32_lossy()?;
let sigma_from = self.sigmas[step];
let sigma_to = self.sigmas[step + 1];
let sigma_up = (sigma_to * sigma_to * (sigma_from * sigma_from - sigma_to * sigma_to)
/ (sigma_from * sigma_from))
.max(0.0)
.sqrt();
let sigma_down = (sigma_to * sigma_to - sigma_up * sigma_up).max(0.0).sqrt();
let dt = sigma_down - sigma_from;
let alpha_t = 1.0 / (sigma_from * sigma_from + 1.0).sqrt();
let sigma_t = sigma_from * alpha_t;
let noise = noise
.context("euler_ancestral requires per-step noise")?
.to_vec_f32_lossy()?;
if noise.len() != x.len() {
anyhow::bail!(
"euler_ancestral noise length {} != sample {}",
noise.len(),
x.len()
);
}
let out: Vec<f32> = (0..x.len())
.map(|i| {
let e = epsilon_from_model_output(
model_out[i],
alpha_t * x[i],
alpha_t,
sigma_t,
self.prediction,
);
x[i] + e * dt + noise[i] * sigma_up
})
.collect();
Value::from_slice_f32(&out, &shape).map_err(Into::into)
}
fn scale_input(
&self,
step: usize,
_num_steps: usize,
sample: &Value,
) -> anyhow::Result<Option<Value>> {
let factor = (self.sigmas[step] * self.sigmas[step] + 1.0).sqrt();
let scaled: Vec<f32> = sample
.to_vec_f32_lossy()?
.iter()
.map(|&x| x / factor)
.collect();
Ok(Some(Value::from_slice_f32(&scaled, sample.shape())?))
}
fn init_noise_sigma(&self) -> f32 {
self.sigmas[0]
}
fn timesteps(&self) -> Option<Vec<f32>> {
Some(self.timesteps.clone())
}
}
#[derive(Debug)]
struct Dpmpp2m {
sigmas: Vec<f32>,
timesteps: Vec<f32>,
prev_x0: Mutex<Option<Vec<f32>>>,
prediction: PredictionType,
}
impl Dpmpp2m {
fn with_schedule(
num_train_timesteps: usize,
beta_start: f32,
beta_end: f32,
beta_schedule: &str,
num_steps: usize,
spacing: &str,
) -> anyhow::Result<Self> {
if num_train_timesteps < 2 {
anyhow::bail!("scheduler num_train_timesteps must be >= 2");
}
if num_steps == 0 || num_steps > num_train_timesteps {
anyhow::bail!("scheduler num_steps ({num_steps}) must be in 1..={num_train_timesteps}");
}
if let Some(sigmas) = spacing_sigmas(
spacing,
num_train_timesteps,
beta_start,
beta_end,
beta_schedule,
num_steps,
)? {
let train = training_sigmas(num_train_timesteps, beta_start, beta_end, beta_schedule)?;
let timesteps = sigmas[..num_steps]
.iter()
.map(|&s| sigma_to_t(&train, s))
.collect();
return Ok(Self {
sigmas,
timesteps,
prev_x0: Mutex::new(None),
prediction: PredictionType::Epsilon,
});
}
let denom = (num_train_timesteps - 1) as f32;
let (lo, hi, square) = match beta_schedule {
"linear" => (beta_start, beta_end, false),
"scaled_linear" => (beta_start.sqrt(), beta_end.sqrt(), true),
other => anyhow::bail!(
"unsupported scheduler beta_schedule '{other}' (expected 'linear' or 'scaled_linear')"
),
};
let mut train = Vec::with_capacity(num_train_timesteps);
let mut prod = 1.0f32;
for i in 0..num_train_timesteps {
let mut beta = lo + (hi - lo) * (i as f32) / denom;
if square {
beta *= beta;
}
prod *= 1.0 - beta;
train.push(((1.0 - prod) / prod).sqrt());
}
let mut ts_int: Vec<usize> = (0..=num_steps)
.map(|j| (j as f32 * denom / num_steps as f32).round_ties_even() as usize)
.collect();
ts_int.reverse();
ts_int.pop();
let timesteps: Vec<f32> = ts_int.iter().map(|&t| t as f32).collect();
let mut sigmas: Vec<f32> = ts_int
.iter()
.map(|&t| train[t.min(num_train_timesteps - 1)])
.collect();
sigmas.push(0.0);
Ok(Self {
sigmas,
timesteps,
prev_x0: Mutex::new(None),
prediction: PredictionType::Epsilon,
})
}
fn with_prediction(mut self, prediction: PredictionType) -> Self {
self.prediction = prediction;
self
}
}
fn dpm_alpha_sigma(sigma: f32) -> (f32, f32) {
let alpha_t = 1.0 / (sigma * sigma + 1.0).sqrt();
(alpha_t, sigma * alpha_t)
}
fn training_sigmas(
num_train_timesteps: usize,
beta_start: f32,
beta_end: f32,
beta_schedule: &str,
) -> anyhow::Result<Vec<f32>> {
let denom = (num_train_timesteps - 1) as f32;
let (lo, hi, square) = match beta_schedule {
"linear" => (beta_start, beta_end, false),
"scaled_linear" => (beta_start.sqrt(), beta_end.sqrt(), true),
other => anyhow::bail!(
"unsupported scheduler beta_schedule '{other}' (expected 'linear' or 'scaled_linear')"
),
};
let mut out = Vec::with_capacity(num_train_timesteps);
let mut prod = 1.0f32;
for i in 0..num_train_timesteps {
let mut beta = lo + (hi - lo) * (i as f32) / denom;
if square {
beta *= beta;
}
prod *= 1.0 - beta;
out.push(((1.0 - prod) / prod).sqrt());
}
Ok(out)
}
fn sigma_to_t(train: &[f32], sigma: f32) -> f32 {
let log_sigma = sigma.max(1e-10).ln();
let count = train
.iter()
.filter(|&&s| s.max(1e-10).ln() <= log_sigma)
.count();
let low_idx = count.saturating_sub(1).min(train.len().saturating_sub(2));
let high_idx = low_idx + 1;
let low = train[low_idx].max(1e-10).ln();
let high = train[high_idx].max(1e-10).ln();
let weight = ((low - log_sigma) / (low - high)).clamp(0.0, 1.0);
(1.0 - weight) * low_idx as f32 + weight * high_idx as f32
}
fn karras_sigmas(
num_train_timesteps: usize,
beta_start: f32,
beta_end: f32,
beta_schedule: &str,
num_steps: usize,
) -> anyhow::Result<Vec<f32>> {
const RHO: f32 = 7.0;
let train = training_sigmas(num_train_timesteps, beta_start, beta_end, beta_schedule)?;
let sigma_min = train[0];
let sigma_max = train[num_train_timesteps - 1];
let min_inv = sigma_min.powf(1.0 / RHO);
let max_inv = sigma_max.powf(1.0 / RHO);
let mut sigmas = Vec::with_capacity(num_steps + 1);
for k in 0..num_steps {
let ramp = if num_steps > 1 {
k as f32 / (num_steps - 1) as f32
} else {
0.0
};
sigmas.push((max_inv + ramp * (min_inv - max_inv)).powf(RHO));
}
sigmas.push(0.0);
Ok(sigmas)
}
fn exponential_sigmas(
num_train_timesteps: usize,
beta_start: f32,
beta_end: f32,
beta_schedule: &str,
num_steps: usize,
) -> anyhow::Result<Vec<f32>> {
let train = training_sigmas(num_train_timesteps, beta_start, beta_end, beta_schedule)?;
let log_min = train[0].ln();
let log_max = train[num_train_timesteps - 1].ln();
let mut sigmas = Vec::with_capacity(num_steps + 1);
for k in 0..num_steps {
let ramp = if num_steps > 1 {
k as f32 / (num_steps - 1) as f32
} else {
0.0
};
sigmas.push((log_max + ramp * (log_min - log_max)).exp());
}
sigmas.push(0.0);
Ok(sigmas)
}
fn sigma_spacing(cfg: &SchedulerSpec) -> anyhow::Result<&'static str> {
let karras = cfg.use_karras_sigmas.unwrap_or(false);
let exponential = cfg.use_exponential_sigmas.unwrap_or(false);
if karras && exponential {
anyhow::bail!("scheduler cannot set both use_karras_sigmas and use_exponential_sigmas");
}
Ok(if karras {
"karras"
} else if exponential {
"exponential"
} else {
"linspace"
})
}
fn spacing_sigmas(
spacing: &str,
num_train_timesteps: usize,
beta_start: f32,
beta_end: f32,
beta_schedule: &str,
num_steps: usize,
) -> anyhow::Result<Option<Vec<f32>>> {
match spacing {
"karras" => Ok(Some(karras_sigmas(
num_train_timesteps,
beta_start,
beta_end,
beta_schedule,
num_steps,
)?)),
"exponential" => Ok(Some(exponential_sigmas(
num_train_timesteps,
beta_start,
beta_end,
beta_schedule,
num_steps,
)?)),
"linspace" | "" => Ok(None),
other => anyhow::bail!("unsupported sigma spacing '{other}' (karras/exponential/linspace)"),
}
}
impl Scheduler for Dpmpp2m {
fn step(
&self,
step: usize,
num_steps: usize,
sample: &Value,
model_output: &Value,
) -> anyhow::Result<Value> {
let shape = sample.shape().to_vec();
let x = sample.to_vec_f32_lossy()?;
let model_out = model_output.to_vec_f32_lossy()?;
if x.len() != model_out.len() {
anyhow::bail!(
"dpm++ sample/model_output length mismatch: {} vs {}",
x.len(),
model_out.len()
);
}
let sigma = self.sigmas[step];
let (alpha_t0, sigma_t0) = dpm_alpha_sigma(sigma);
let x0: Vec<f32> = x
.iter()
.zip(&model_out)
.map(|(&xi, &mi)| x0_from_model_output(mi, xi, alpha_t0, sigma_t0, self.prediction))
.collect();
let s_next = self.sigmas[step + 1];
let (a_t, sig_t) = dpm_alpha_sigma(s_next);
let (a_s0, sig_s0) = dpm_alpha_sigma(sigma);
let lam_t = a_t.ln() - sig_t.ln(); let lam_s0 = a_s0.ln() - sig_s0.ln();
let h = lam_t - lam_s0;
let neg_expm1 = (-h).exp() - 1.0;
let mut prev = self
.prev_x0
.lock()
.map_err(|_| anyhow::anyhow!("dpm++ scheduler state poisoned"))?;
let lower_order_final = step + 1 == num_steps && (num_steps < 15 || s_next <= 0.0);
let first_order = lower_order_final || prev.is_none();
let out: Vec<f32> = if first_order {
x.iter()
.zip(&x0)
.map(|(&xi, &d0)| (sig_t / sig_s0) * xi - a_t * neg_expm1 * d0)
.collect()
} else {
let prev_x0 = prev.as_ref().unwrap();
let s_prev = self.sigmas[step - 1];
let (a_s1, sig_s1) = dpm_alpha_sigma(s_prev);
let lam_s1 = a_s1.ln() - sig_s1.ln();
let h0 = lam_s0 - lam_s1;
let r0 = h0 / h;
x.iter()
.enumerate()
.map(|(i, &xi)| {
let d0 = x0[i];
let d1 = (1.0 / r0) * (x0[i] - prev_x0[i]);
(sig_t / sig_s0) * xi - a_t * neg_expm1 * d0 - 0.5 * a_t * neg_expm1 * d1
})
.collect()
};
*prev = Some(x0);
drop(prev);
Value::from_slice_f32(&out, &shape).map_err(Into::into)
}
fn reset(&self) {
if let Ok(mut prev) = self.prev_x0.lock() {
*prev = None;
}
}
fn timesteps(&self) -> Option<Vec<f32>> {
Some(self.timesteps.clone())
}
}
impl PipelinePlan {
fn from_spec(spec: &PipelineSpec, schedulers: &SchedulerRegistry) -> anyhow::Result<Self> {
if let Some(stage) = nested_autoregressive_strategy(&spec.strategy) {
return Self::nested_autoregressive(spec, stage);
}
if let Some(decoder) = autoregressive_decoder(&spec.strategy) {
return Self::autoregressive(spec, decoder);
}
match spec.strategy.kind {
PipelineStrategyKind::SinglePass => Self::single_pass(spec),
PipelineStrategyKind::Iterative => Self::iterative(spec, schedulers),
PipelineStrategyKind::Composite => Self::composite(spec),
PipelineStrategyKind::Autoregressive => {
anyhow::bail!("autoregressive strategy is missing its 'decoder' component")
}
PipelineStrategyKind::NestedAutoregressive => {
anyhow::bail!(
"nested_autoregressive strategy is missing its 'outer'/'inner' decoders"
)
}
PipelineStrategyKind::Other(ref value) => {
anyhow::bail!("unsupported pipeline strategy kind '{value}'")
}
}
}
fn autoregressive(spec: &PipelineSpec, decoder: String) -> anyhow::Result<Self> {
if !spec.models.contains_key(&decoder) {
anyhow::bail!("pipeline decoder '{decoder}' is not declared in models");
}
let prompt_components = prompt_phase_components(spec, &decoder)?;
let step_components = step_phase_components(spec, &decoder)?;
let post_decode_components = post_decode_components(spec, &decoder)?;
Ok(Self::Autoregressive(AutoregressivePlan {
decoder,
prompt_components,
step_components,
post_decode_components,
dataflow: spec.dataflow.clone(),
presence_conditions: presence_conditions(spec),
}))
}
fn nested_autoregressive(
spec: &PipelineSpec,
nested: &PipelineStrategy,
) -> anyhow::Result<Self> {
let outer = nested
.outer
.clone()
.context("nested_autoregressive strategy is missing its 'outer' decoder")?;
let inner = nested
.inner
.clone()
.context("nested_autoregressive strategy is missing its 'inner' decoder")?;
if !spec.models.contains_key(&outer) {
anyhow::bail!(
"nested_autoregressive outer decoder '{outer}' is not declared in models"
);
}
if !spec.models.contains_key(&inner) {
anyhow::bail!(
"nested_autoregressive inner decoder '{inner}' is not declared in models"
);
}
if outer == inner {
anyhow::bail!(
"nested_autoregressive 'outer' and 'inner' must be distinct decoders (both '{outer}')"
);
}
let num_code_groups = nested
.num_code_groups
.context("nested_autoregressive strategy is missing 'num_code_groups'")?;
if num_code_groups == 0 {
anyhow::bail!("nested_autoregressive 'num_code_groups' must be greater than zero");
}
let max_frames = nested
.max_tokens
.context("nested_autoregressive strategy is missing 'max_tokens' (max audio frames)")?;
if max_frames == 0 {
anyhow::bail!(
"nested_autoregressive 'max_tokens' (max frames) must be greater than zero"
);
}
let inner_embeds_endpoint_edge = spec
.dataflow
.iter()
.find(|edge| {
endpoint_component(&edge.to) == Some(inner.as_str())
&& endpoint_component(&edge.from) == Some(outer.as_str())
})
.with_context(|| {
format!(
"nested_autoregressive needs a per-frame hidden binding: a dataflow edge \
'{outer}.last_hidden_state -> {inner}.inputs_embeds'"
)
})?;
let (_, outer_hidden_output) = parse_endpoint(&inner_embeds_endpoint_edge.from)?;
let (_, inner_embeds_input) = parse_endpoint(&inner_embeds_endpoint_edge.to)?;
let outer_hidden_output = outer_hidden_output.to_string();
let inner_embeds_input = inner_embeds_input.to_string();
let pre_embedder = match nested.pre_embedder.as_ref() {
Some(spec_pre) => {
let name = spec_pre.component.as_str();
if !spec.models.contains_key(name) {
anyhow::bail!(
"nested_autoregressive pre_embedder '{name}' is not declared in models"
);
}
if name == outer || name == inner {
anyhow::bail!(
"nested_autoregressive pre_embedder '{name}' must be distinct from the \
outer/inner decoders"
);
}
let edge = spec
.dataflow
.iter()
.find(|edge| {
endpoint_component(&edge.from) == Some(name)
&& endpoint_component(&edge.to) == Some(outer.as_str())
})
.with_context(|| {
format!(
"nested_autoregressive pre_embedder '{name}' needs a per-step feed: a \
dataflow edge '{name}.<output> -> {outer}.inputs_embeds'"
)
})?;
let (_, outer_input) = parse_endpoint(&edge.to)?;
let (_, output_port) = parse_endpoint(&edge.from)?;
Some(PreEmbedderBinding {
component: name.to_string(),
outer_input: outer_input.to_string(),
output_port: output_port.to_string(),
frame_codes_input: spec_pre.frame_codes_input.clone(),
text_embed_input: spec_pre.text_embed_input.clone(),
})
}
None => None,
};
let pre_embedder_component = pre_embedder.as_ref().map(|p| p.component.clone());
let prefill_embedder = match nested.prefill_embedder.as_ref() {
Some(spec_prefill) => {
let name = spec_prefill.component.as_str();
if !spec.models.contains_key(name) {
anyhow::bail!(
"nested_autoregressive prefill_embedder '{name}' is not declared in models"
);
}
if name == outer || name == inner {
anyhow::bail!(
"nested_autoregressive prefill_embedder '{name}' must be distinct from \
the outer/inner decoders"
);
}
if pre_embedder_component.as_deref() == Some(name) {
anyhow::bail!(
"nested_autoregressive prefill_embedder '{name}' must be distinct from \
the pre_embedder"
);
}
if pre_embedder.is_none() {
anyhow::bail!(
"nested_autoregressive prefill_embedder '{name}' requires a 'pre_embedder' \
(frames >= 1 thread its trailing-text vectors through the pre-embedder)"
);
}
Some(PrefillEmbedderBinding {
component: name.to_string(),
prompt_input: spec_prefill.prompt_input.clone(),
prefill_output: spec_prefill.prefill_output.clone(),
trailing_output: spec_prefill.trailing_output.clone(),
})
}
None => None,
};
let mut prompt_components = Vec::new();
let mut post_decode_components = Vec::new();
for component in topological_components(spec)? {
if component == outer || component == inner {
continue;
}
if pre_embedder_component.as_deref() == Some(component.as_str()) {
continue;
}
match component_phase(spec, &component, &outer) {
PhaseRunOn::PromptOnly => prompt_components.push(component),
PhaseRunOn::FinalOnly => post_decode_components.push(component),
PhaseRunOn::OnDemand => {}
PhaseRunOn::EveryStep => anyhow::bail!(
"nested_autoregressive component '{component}' declares run_on: every_step, \
but only the outer/inner decoders may run inside the nested loop"
),
PhaseRunOn::Other(value) => anyhow::bail!(
"unsupported phase '{value}' for pipeline component '{component}'"
),
}
}
Ok(Self::NestedAutoregressive(NestedAutoregressivePlan {
outer,
inner,
num_code_groups,
max_frames,
outer_hidden_output,
inner_embeds_input,
prompt_components,
post_decode_components,
pre_embedder,
prefill_embedder,
dataflow: spec.dataflow.clone(),
presence_conditions: presence_conditions(spec),
}))
}
fn single_pass(spec: &PipelineSpec) -> anyhow::Result<Self> {
let model = spec
.strategy
.model
.clone()
.context("single_pass strategy is missing its 'model' component")?;
if !spec.models.contains_key(&model) {
anyhow::bail!("pipeline model '{model}' is not declared in models");
}
let mut prompt_components = Vec::new();
for component in topological_components(spec)? {
if component == model {
continue;
}
match component_phase(spec, &component, &model) {
PhaseRunOn::PromptOnly => prompt_components.push(component),
PhaseRunOn::OnDemand => {}
PhaseRunOn::EveryStep | PhaseRunOn::FinalOnly => anyhow::bail!(
"component '{component}' declares a run_on phase unsupported by a single_pass \
pipeline (only prompt_only / on_demand components are allowed)"
),
PhaseRunOn::Other(value) => anyhow::bail!(
"unsupported phase '{value}' for pipeline component '{component}'"
),
}
}
Ok(Self::SinglePass(SinglePassPlan {
model,
prompt_components,
dataflow: spec.dataflow.clone(),
presence_conditions: presence_conditions(spec),
}))
}
fn composite(spec: &PipelineSpec) -> anyhow::Result<Self> {
if spec.strategy.stages.is_empty() {
anyhow::bail!("composite pipeline strategy declares no stages");
}
let mut stages = Vec::with_capacity(spec.strategy.stages.len());
let mut seen_names = BTreeSet::new();
for stage in &spec.strategy.stages {
if !seen_names.insert(stage.name.clone()) {
anyhow::bail!("composite stage name '{}' is not unique", stage.name);
}
let kind = match stage.strategy.kind {
PipelineStrategyKind::SinglePass => {
let model = stage.strategy.model.clone().with_context(|| {
format!(
"composite stage '{}' (single_pass) is missing 'model'",
stage.name
)
})?;
if !spec.models.contains_key(&model) {
anyhow::bail!(
"composite stage '{}' model '{model}' is not declared in models",
stage.name
);
}
CompositeStageKind::SinglePass { model }
}
PipelineStrategyKind::Iterative => anyhow::bail!(
"composite iterative stage '{}' is not yet supported (single-pass stages only)",
stage.name
),
PipelineStrategyKind::Autoregressive
| PipelineStrategyKind::Composite
| PipelineStrategyKind::NestedAutoregressive => {
anyhow::bail!(
"composite stage '{}' has an unsupported nested strategy kind for a \
non-autoregressive composite",
stage.name
)
}
PipelineStrategyKind::Other(ref value) => anyhow::bail!(
"composite stage '{}' has unsupported strategy kind '{value}'",
stage.name
),
};
stages.push(CompositeStage {
name: stage.name.clone(),
kind,
});
}
Ok(Self::Composite(CompositePlan {
stages,
dataflow: spec.dataflow.clone(),
presence_conditions: presence_conditions(spec),
}))
}
fn iterative(spec: &PipelineSpec, schedulers: &SchedulerRegistry) -> anyhow::Result<Self> {
let denoiser = spec
.strategy
.denoiser
.clone()
.context("iterative strategy is missing its 'denoiser' component")?;
if !spec.models.contains_key(&denoiser) {
anyhow::bail!("pipeline denoiser '{denoiser}' is not declared in models");
}
let num_steps = spec
.strategy
.num_steps
.context("iterative strategy is missing 'num_steps'")?;
if num_steps == 0 {
anyhow::bail!("iterative strategy 'num_steps' must be greater than zero");
}
let start_step = spec.strategy.start_step.unwrap_or(0);
if start_step >= num_steps {
anyhow::bail!(
"iterative strategy 'start_step' ({start_step}) must be less than 'num_steps' ({num_steps})"
);
}
let guidance_active = spec.strategy.guidance_scale.is_some_and(|s| s != 1.0);
let scheduler_supplies_uncond = spec
.strategy
.scheduler_config
.as_ref()
.is_some_and(|scheduler| scheduler.kind == "masked_diffusion");
if guidance_active
&& spec.strategy.cfg_conditioning_input.is_none()
&& !scheduler_supplies_uncond
{
anyhow::bail!(
"classifier-free guidance (guidance_scale != 1.0) requires \
'cfg_conditioning_input' naming the denoiser conditioning port to zero on the \
unconditional pass"
);
}
let mut loop_edges = Vec::new();
for edge in &spec.dataflow {
let (from_component, from_port) = parse_endpoint(&edge.from)?;
let (to_component, to_port) = parse_endpoint(&edge.to)?;
if from_component == denoiser && to_component == denoiser {
loop_edges.push((from_port.to_string(), to_port.to_string()));
}
}
if let Some(cfg_port) = &spec.strategy.cfg_conditioning_input
&& guidance_active
&& loop_edges.iter().any(|(_, in_port)| in_port == cfg_port)
{
anyhow::bail!(
"cfg_conditioning_input '{cfg_port}' must not also be a loop-carried input \
port: the unconditional conditioning override would clobber the loop sample"
);
}
let mut prompt_components = Vec::new();
let mut final_components = Vec::new();
for component in topological_components(spec)? {
if component == denoiser {
continue;
}
match component_phase(spec, &component, &denoiser) {
PhaseRunOn::PromptOnly => prompt_components.push(component),
PhaseRunOn::FinalOnly => final_components.push(component),
PhaseRunOn::OnDemand => {}
PhaseRunOn::EveryStep => anyhow::bail!(
"component '{component}' declares run_on: every_step, but running a \
non-denoiser component inside the iterative loop is not yet supported"
),
PhaseRunOn::Other(value) => anyhow::bail!(
"unsupported phase '{value}' for pipeline component '{component}'"
),
}
}
Ok(Self::Iterative(Box::new(IterativePlan {
denoiser,
num_steps,
guidance_scale: spec.strategy.guidance_scale,
prompt_components,
final_components,
loop_edges,
timestep_input: spec.strategy.timestep_input.clone(),
start_step,
timesteps: spec.strategy.timesteps.clone(),
scheduler: build_scheduler(
spec.strategy.scheduler_config.as_ref(),
num_steps,
schedulers,
)?,
cfg_conditioning_input: spec.strategy.cfg_conditioning_input.clone(),
dataflow: spec.dataflow.clone(),
scheduler_spec: spec.strategy.scheduler_config.clone(),
scheduler_registry: schedulers.clone(),
presence_conditions: presence_conditions(spec),
})))
}
fn autoregressive_plan(&self) -> anyhow::Result<&AutoregressivePlan> {
match self {
Self::Autoregressive(plan) => Ok(plan),
_ => anyhow::bail!("pipeline strategy is not autoregressive"),
}
}
fn dataflow(&self) -> &[DataflowEdge] {
match self {
Self::Autoregressive(plan) => &plan.dataflow,
Self::NestedAutoregressive(plan) => &plan.dataflow,
Self::SinglePass(plan) => &plan.dataflow,
Self::Iterative(plan) => &plan.dataflow,
Self::Composite(plan) => &plan.dataflow,
}
}
fn presence_condition(&self, component: &str) -> Option<&str> {
let conditions = match self {
Self::Autoregressive(plan) => &plan.presence_conditions,
Self::NestedAutoregressive(plan) => &plan.presence_conditions,
Self::SinglePass(plan) => &plan.presence_conditions,
Self::Iterative(plan) => &plan.presence_conditions,
Self::Composite(plan) => &plan.presence_conditions,
};
conditions.get(component).map(String::as_str)
}
fn component_is_present(&self, component: &str, present: &BTreeSet<String>) -> bool {
self.presence_condition(component)
.is_none_or(|key| present.contains(key))
}
fn edges_to_component<'a>(
&'a self,
component: &'a str,
) -> impl Iterator<Item = &'a DataflowEdge> + 'a {
self.dataflow()
.iter()
.filter(move |edge| endpoint_component(&edge.to) == Some(component))
}
}
fn presence_conditions(spec: &PipelineSpec) -> HashMap<String, String> {
spec.phases
.iter()
.filter_map(|(component, phase)| {
phase
.when_present
.as_ref()
.map(|key| (component.clone(), key.clone()))
})
.collect()
}
fn prompt_phase_components(spec: &PipelineSpec, primary: &str) -> anyhow::Result<Vec<String>> {
let mut prompt_components = Vec::new();
for component in topological_components(spec)? {
if component == primary {
continue;
}
match component_phase(spec, &component, primary) {
PhaseRunOn::PromptOnly => prompt_components.push(component),
PhaseRunOn::EveryStep | PhaseRunOn::OnDemand | PhaseRunOn::FinalOnly => {}
PhaseRunOn::Other(value) => {
anyhow::bail!("unsupported phase '{value}' for pipeline component '{component}'")
}
}
}
Ok(prompt_components)
}
fn step_phase_components(spec: &PipelineSpec, decoder: &str) -> anyhow::Result<Vec<String>> {
let mut step = Vec::new();
for component in topological_components(spec)? {
if component == decoder {
continue;
}
if let PhaseRunOn::EveryStep = component_phase(spec, &component, decoder) {
step.push(component);
}
}
Ok(step)
}
fn post_decode_components(spec: &PipelineSpec, decoder: &str) -> anyhow::Result<Vec<String>> {
let mut post = Vec::new();
for component in topological_components(spec)? {
if component == decoder {
continue;
}
if let PhaseRunOn::FinalOnly = component_phase(spec, &component, decoder) {
post.push(component);
}
}
Ok(post)
}
fn build_scheduler(
config: Option<&SchedulerSpec>,
num_steps: usize,
registry: &SchedulerRegistry,
) -> anyhow::Result<Option<Arc<dyn Scheduler>>> {
let Some(cfg) = config else {
return Ok(None);
};
Ok(Some(registry.build(cfg, num_steps)?))
}
fn autoregressive_decoder(strategy: &PipelineStrategy) -> Option<String> {
match strategy.kind {
PipelineStrategyKind::Autoregressive => strategy.decoder.clone(),
PipelineStrategyKind::Composite => strategy
.stages
.iter()
.find_map(|stage| autoregressive_decoder(&stage.strategy)),
PipelineStrategyKind::Iterative
| PipelineStrategyKind::SinglePass
| PipelineStrategyKind::NestedAutoregressive
| PipelineStrategyKind::Other(_) => None,
}
}
fn nested_autoregressive_strategy(strategy: &PipelineStrategy) -> Option<&PipelineStrategy> {
match strategy.kind {
PipelineStrategyKind::NestedAutoregressive => Some(strategy),
PipelineStrategyKind::Composite => strategy
.stages
.iter()
.find_map(|stage| nested_autoregressive_strategy(&stage.strategy)),
_ => None,
}
}
fn component_phase(spec: &PipelineSpec, component: &str, decoder: &str) -> PhaseRunOn {
spec.phases
.get(component)
.map(|phase| phase.run_on.clone())
.unwrap_or_else(|| {
if component == decoder {
PhaseRunOn::EveryStep
} else {
PhaseRunOn::PromptOnly
}
})
}
fn topological_components(spec: &PipelineSpec) -> anyhow::Result<Vec<String>> {
let mut remaining = spec.models.keys().cloned().collect::<BTreeSet<_>>();
let mut ordered = Vec::new();
while !remaining.is_empty() {
let ready = remaining
.iter()
.find(|component| {
spec.dataflow.iter().all(|edge| {
let to = endpoint_component(&edge.to);
let from = endpoint_component(&edge.from);
to != Some(component.as_str())
|| from == Some(component.as_str())
|| from.is_some_and(|f| !remaining.contains(f))
})
})
.cloned();
let Some(component) = ready else {
anyhow::bail!("pipeline dataflow contains a cycle");
};
remaining.remove(&component);
ordered.push(component);
}
Ok(ordered)
}
fn parse_endpoint(endpoint: &str) -> anyhow::Result<(&str, &str)> {
endpoint
.split_once('.')
.filter(|(component, port)| !component.is_empty() && !port.is_empty())
.with_context(|| format!("pipeline endpoint must be component.port: {endpoint}"))
}
fn endpoint_component(endpoint: &str) -> Option<&str> {
parse_endpoint(endpoint)
.ok()
.map(|(component, _)| component)
}
fn named_output<'a>(
session: &Session,
outputs: &'a [Value],
name: &str,
contains: bool,
) -> anyhow::Result<&'a Value> {
let index = session
.output_names()
.iter()
.position(|out| out == name)
.or_else(|| {
if contains {
let needle = name.to_ascii_lowercase();
session
.output_names()
.iter()
.position(|out| out.to_ascii_lowercase().contains(&needle))
} else {
None
}
})
.with_context(|| format!("model did not expose output '{name}'"))?;
outputs
.get(index)
.with_context(|| format!("output '{name}' index was out of range"))
}
fn argmax_last_row(logits: &Value) -> anyhow::Result<i64> {
let shape = logits.shape();
let data = logits
.to_vec_f32_lossy()
.map_err(|e| anyhow::anyhow!("failed to read logits tensor: {e}"))?;
let vocab = match shape {
[vocab] if *vocab > 0 => *vocab as usize,
[seq, vocab] if *seq > 0 && *vocab > 0 => *vocab as usize,
[batch, seq, vocab] if *batch == 1 && *seq > 0 && *vocab > 0 => *vocab as usize,
other => anyhow::bail!("unsupported logits tensor shape: {other:?}"),
};
let start = data.len() - vocab;
let row = &data[start..];
let mut best = 0usize;
for (i, &value) in row.iter().enumerate() {
if value > row[best] {
best = i;
}
}
Ok(best as i64)
}
fn last_position_hidden(hidden: &Value) -> anyhow::Result<Value> {
let shape = hidden.shape();
let data = hidden
.to_vec_f32_lossy()
.map_err(|e| anyhow::anyhow!("failed to read hidden-state tensor: {e}"))?;
let hidden_dim = match shape {
[h] if *h > 0 => *h as usize,
[seq, h] if *seq > 0 && *h > 0 => *h as usize,
[batch, seq, h] if *batch == 1 && *seq > 0 && *h > 0 => *h as usize,
other => anyhow::bail!("unsupported hidden-state tensor shape: {other:?}"),
};
let start = data.len() - hidden_dim;
Value::from_slice_f32(&data[start..], &[1, 1, hidden_dim as i64])
.map_err(|e| anyhow::anyhow!("failed to build inner seed embedding: {e}"))
}
struct ResolvedPreEmbedder<'a> {
session: &'a Session,
outer_input: String,
output_port: String,
frame_codes_input: String,
text_embed_input: Option<String>,
hidden: usize,
}
struct ResolvedPrefill {
prefill_embeds: Value,
prefill_len: usize,
trailing: Vec<f32>,
trailing_len: usize,
hidden: usize,
}
fn run_pre_embedder(
pre: &ResolvedPreEmbedder<'_>,
frame_codes: &[i64],
text_embed: Option<&[f32]>,
) -> anyhow::Result<Value> {
let mut inputs: Vec<(String, Value)> = Vec::with_capacity(2);
inputs.push((
pre.frame_codes_input.clone(),
Value::from_slice_i64(frame_codes, &[1, frame_codes.len() as i64])?,
));
if let Some(name) = &pre.text_embed_input {
let dtype = pre
.session
.inputs()
.iter()
.find(|info| &info.name == name)
.map(|info| info.dtype)
.unwrap_or(DataType::Float32);
let data = match text_embed {
Some(slice) => slice.to_vec(),
None => vec![0.0f32; pre.hidden],
};
inputs.push((
name.clone(),
Value::from_f32_slice_as(&data, &[1, 1, pre.hidden as i64], dtype)
.map_err(|e| anyhow::anyhow!("failed to build text_embed: {e}"))?,
));
}
let refs = inputs
.iter()
.map(|(name, value)| (name.as_str(), value))
.collect::<Vec<_>>();
let outputs = pre
.session
.run(&refs)
.map_err(|e| anyhow::anyhow!("ORT pre-embedder run failed: {e}"))?;
let index = pre
.session
.output_names()
.iter()
.position(|name| name == &pre.output_port)
.with_context(|| {
format!(
"pre-embedder has no declared output port '{}'",
pre.output_port
)
})?;
let value = outputs
.get(index)
.context("pre-embedder produced no output for its declared port")?;
clone_value(value)
}
#[cfg(test)]
mod tests {
use super::*;
use onnx_genai_metadata::{PhaseConfig, PipelineComponentSpec, PipelineStrategyStage};
use std::collections::BTreeMap;
#[test]
fn ddim_step_matches_hand_computed_closed_form() {
let sched = DdimSchedule::with_schedule(2, 0.5, 0.5, "linear", 1).expect("schedule builds");
let n0 = sched.step(0, &[1.0], &[0.0]).unwrap();
assert!((n0[0] - std::f32::consts::SQRT_2).abs() < 1e-5, "{}", n0[0]);
let n1 = sched.step(0, &[1.0], &[1.0]).unwrap();
assert!(
(n1[0] - (std::f32::consts::SQRT_2 - 1.0)).abs() < 1e-5,
"{}",
n1[0]
);
}
#[test]
fn prediction_type_parse_accepts_known_aliases() {
assert_eq!(
PredictionType::parse(None).unwrap(),
PredictionType::Epsilon
);
assert_eq!(
PredictionType::parse(Some("epsilon")).unwrap(),
PredictionType::Epsilon
);
assert_eq!(
PredictionType::parse(Some("v_prediction")).unwrap(),
PredictionType::VPrediction
);
assert_eq!(
PredictionType::parse(Some("sample")).unwrap(),
PredictionType::Sample
);
assert_eq!(
PredictionType::parse(Some("x0")).unwrap(),
PredictionType::Sample
);
assert!(PredictionType::parse(Some("nonsense")).is_err());
}
#[test]
fn model_output_conversion_matches_diffusers_formulas() {
let alpha_t = 0.5f32.sqrt();
let sigma_t = (1.0f32 - 0.5).sqrt();
let x_t = 0.7f32;
let model_out = 0.3f32;
assert!(
(epsilon_from_model_output(model_out, x_t, alpha_t, sigma_t, PredictionType::Epsilon)
- model_out)
.abs()
< 1e-6
);
let x0_eps = x0_from_model_output(model_out, x_t, alpha_t, sigma_t, PredictionType::Epsilon);
assert!((x0_eps - (x_t - sigma_t * model_out) / alpha_t).abs() < 1e-6);
let eps_v =
epsilon_from_model_output(model_out, x_t, alpha_t, sigma_t, PredictionType::VPrediction);
assert!((eps_v - (alpha_t * model_out + sigma_t * x_t)).abs() < 1e-6);
let x0_v =
x0_from_model_output(model_out, x_t, alpha_t, sigma_t, PredictionType::VPrediction);
assert!((x0_v - (alpha_t * x_t - sigma_t * model_out)).abs() < 1e-6);
assert!((alpha_t * x0_v + sigma_t * eps_v - x_t).abs() < 1e-6);
let x0_s = x0_from_model_output(model_out, x_t, alpha_t, sigma_t, PredictionType::Sample);
assert!((x0_s - model_out).abs() < 1e-6);
let eps_s =
epsilon_from_model_output(model_out, x_t, alpha_t, sigma_t, PredictionType::Sample);
assert!((eps_s - (x_t - alpha_t * model_out) / sigma_t).abs() < 1e-6);
assert!((alpha_t * x0_s + sigma_t * eps_s - x_t).abs() < 1e-6);
}
#[test]
fn model_output_conversion_endpoints() {
let (alpha_t, sigma_t) = (1.0f32, 0.0f32);
let x_t = 0.9f32;
let model_out = -0.4f32;
assert!(
(epsilon_from_model_output(model_out, x_t, alpha_t, sigma_t, PredictionType::VPrediction)
- model_out)
.abs()
< 1e-6
);
assert!(
(x0_from_model_output(model_out, x_t, alpha_t, sigma_t, PredictionType::VPrediction)
- x_t)
.abs()
< 1e-6
);
}
#[test]
fn v_prediction_epsilon_path_stays_byte_identical() {
let registry = SchedulerRegistry::builtin();
for kind in ["ddim", "euler", "dpmpp_2m"] {
let eps_spec = SchedulerSpec {
kind: kind.to_string(),
prediction_type: Some("epsilon".to_string()),
..SchedulerSpec::default()
};
let default_spec = SchedulerSpec {
kind: kind.to_string(),
prediction_type: None,
..SchedulerSpec::default()
};
let eps = registry.build(&eps_spec, 6).expect("epsilon builds");
let dflt = registry.build(&default_spec, 6).expect("default builds");
let sample = Value::from_slice_f32(&[0.1, -0.2, 0.3, 0.4], &[1, 4]).unwrap();
let model_out = Value::from_slice_f32(&[0.5, 0.6, -0.7, 0.8], &[1, 4]).unwrap();
eps.reset();
dflt.reset();
let a = eps.step(0, 6, &sample, &model_out).unwrap();
let b = dflt.step(0, 6, &sample, &model_out).unwrap();
assert_eq!(
a.to_vec_f32_lossy().unwrap(),
b.to_vec_f32_lossy().unwrap(),
"{kind}: explicit epsilon must equal the default"
);
}
}
#[test]
fn v_prediction_schedulers_construct_and_step_finite() {
let registry = SchedulerRegistry::builtin();
let num_steps = 4usize;
let shape = [1i64, 2, 2, 2];
let elems = 8usize;
for kind in ["ddim", "euler", "euler_ancestral", "dpmpp_2m"] {
for prediction in ["v_prediction", "x0", "sample"] {
let spec = SchedulerSpec {
kind: kind.to_string(),
prediction_type: Some(prediction.to_string()),
num_train_timesteps: Some(1000),
..SchedulerSpec::default()
};
let sched = registry
.build(&spec, num_steps)
.unwrap_or_else(|e| panic!("{kind}/{prediction} must construct: {e}"));
sched.reset();
let init = sched.init_noise_sigma();
let mut latent: Vec<f32> = (0..elems)
.map(|i| ((i as f32 * 0.37).sin()) * init)
.collect();
for step in 0..num_steps {
let sample = Value::from_slice_f32(&latent, &shape).unwrap();
let model_out: Vec<f32> = (0..elems)
.map(|i| ((step as f32 + 1.0) * 0.11 + i as f32 * 0.19).cos())
.collect();
let model_value = Value::from_slice_f32(&model_out, &shape).unwrap();
let noise = Value::from_slice_f32(
&(0..elems)
.map(|i| ((i as f32 + step as f32) * 0.53).sin())
.collect::<Vec<_>>(),
&shape,
)
.unwrap();
let next = sched
.step_with_noise(step, num_steps, &sample, &model_value, Some(&noise))
.unwrap_or_else(|e| panic!("{kind}/{prediction} step {step}: {e}"));
assert_eq!(next.shape(), shape, "{kind}/{prediction} preserves shape");
latent = next.to_vec_f32_lossy().unwrap();
for (i, v) in latent.iter().enumerate() {
assert!(
v.is_finite(),
"{kind}/{prediction} step {step} elem {i} non-finite: {v}"
);
}
}
}
}
}
#[test]
fn ddim_new_rejects_invalid_step_counts() {
assert!(DdimSchedule::with_schedule(1, 0.1, 0.2, "linear", 1).is_err()); assert!(DdimSchedule::with_schedule(4, 0.1, 0.2, "linear", 0).is_err()); assert!(DdimSchedule::with_schedule(4, 0.1, 0.2, "linear", 5).is_err()); }
#[test]
fn dpmpp_timesteps_match_diffusers_linspace() {
let num_train = 1000usize;
let num_steps = 25usize;
let sched =
Dpmpp2m::with_schedule(num_train, 0.00085, 0.012, "scaled_linear", num_steps, "")
.expect("schedule builds");
let timesteps = sched.timesteps().expect("dpm++ exposes timesteps");
let denom = (num_train - 1) as f32;
let mut expected: Vec<f32> = (0..=num_steps)
.map(|j| (j as f32 * denom / num_steps as f32).round_ties_even())
.collect();
expected.reverse();
expected.pop();
assert_eq!(timesteps.len(), num_steps);
assert!(
(timesteps[0] - 999.0).abs() < 1e-3,
"first timestep {}",
timesteps[0]
);
for (got, want) in timesteps.iter().zip(&expected) {
assert!((got - want).abs() < 1e-3, "timestep {got} != {want}");
}
}
#[test]
fn ddim_exposes_descending_integer_timesteps() {
let sched =
DdimSchedule::with_schedule(1000, 0.00085, 0.012, "scaled_linear", 4).expect("builds");
assert_eq!(sched.timesteps(), Some(vec![750.0, 500.0, 250.0, 0.0]));
}
#[test]
fn masked_diffusion_random_unmasks_all_and_never_emits_mask() {
let vocab = 5usize;
let mask_id = 4i64;
let seq = 6usize;
let prompt_len = 2usize;
let num_steps = 4usize;
let mut logits = vec![0f32; seq * vocab];
for pos in 0..seq {
logits[pos * vocab + mask_id as usize] = 100.0;
logits[pos * vocab + (pos % 4)] = 10.0;
}
let logits_value =
Value::from_slice_f32(&logits, &[1, seq as i64, vocab as i64]).expect("logits");
let sched = MaskedDiffusion {
mask_token_id: mask_id,
temperature: 0.0, block_length: None,
remasking: Remasking::Random,
generation_start: Mutex::new(None),
};
let seed = vec![1i64, 2, mask_id, mask_id, mask_id, mask_id];
let run = |sched: &MaskedDiffusion| -> Vec<i64> {
sched.reset();
let mut value = Value::from_slice_i64(&seed, &[1, seq as i64]).expect("seed");
for step in 0..num_steps {
value = sched
.step(step, num_steps, &value, &logits_value)
.expect("step");
}
value.to_vec_i64().expect("tokens")
};
let out = run(&sched);
assert_eq!(&out[..prompt_len], &[1, 2], "prompt prefix preserved");
for (pos, &tok) in out.iter().enumerate() {
assert_ne!(
tok, mask_id,
"position {pos} still masked / emitted the mask token"
);
}
for (offset, &token) in out[prompt_len..seq].iter().enumerate() {
let pos = prompt_len + offset;
assert_eq!(token, (pos % 4) as i64, "position {pos} token");
}
assert_eq!(run(&sched), out, "ancestral sampling is deterministic");
}
#[test]
fn masked_diffusion_rejects_unknown_remasking() {
let registry = SchedulerRegistry::default();
let spec = SchedulerSpec {
kind: "masked_diffusion".to_string(),
mask_token_id: Some(4),
remasking: Some("nonsense".to_string()),
..SchedulerSpec::default()
};
assert!(registry.build(&spec, 4).is_err());
}
#[test]
fn dpmpp_final_step_stays_finite_with_zero_final_sigma() {
let num_steps = 20usize;
let sched = Dpmpp2m::with_schedule(1000, 0.00085, 0.012, "scaled_linear", num_steps, "")
.expect("schedule builds");
sched.reset();
let mut sample = Value::from_slice_f32(&[1.0, -0.5, 0.25], &[3]).unwrap();
for step in 0..num_steps {
let eps = Value::from_slice_f32(&[0.3, -0.2, 0.1], &[3]).unwrap();
sample = sched.step(step, num_steps, &sample, &eps).unwrap();
}
assert!(
sample
.to_vec_f32()
.unwrap()
.iter()
.all(|value| value.is_finite()),
"final dpm++ sample must be finite"
);
}
fn component(role: &str) -> PipelineComponentSpec {
PipelineComponentSpec {
filename: format!("{role}.onnx"),
role: role.to_string(),
device_preference: None,
tokenizer: None,
io: None,
}
}
#[cfg(not(feature = "native-backend"))]
#[test]
fn explicit_native_backend_without_feature_reports_actionable_build_error() {
let error = PipelineEngine::from_dir_with_config(
Path::new("does-not-need-to-exist"),
EngineConfig {
decode_backend: EngineDecodeBackend::Native,
..EngineConfig::default()
},
)
.err()
.expect("native pipeline backend must report an actionable error");
let message = error.to_string();
assert!(
message.contains("without the 'native-backend' feature"),
"unexpected error: {message}"
);
assert!(!message.contains("native backend not supported for pipeline models"));
}
#[cfg(feature = "native-backend")]
#[test]
fn auto_backend_routes_native_only_pipeline_to_the_native_backend() -> anyhow::Result<()> {
use onnx_runtime_loader::proto::{
ModelProto,
onnx::{GraphProto, NodeProto, OperatorSetIdProto},
};
use prost::Message;
let root = Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../../target/test-fixtures/pipeline-native-backend-rejection");
std::fs::create_dir_all(&root)?;
let model = ModelProto {
opset_import: vec![OperatorSetIdProto {
domain: "pkg.nxrt".to_string(),
version: 1,
}],
graph: Some(GraphProto {
node: vec![NodeProto {
domain: "pkg.nxrt".to_string(),
op_type: "BlockQuantizedMatMul".to_string(),
..NodeProto::default()
}],
..GraphProto::default()
}),
..ModelProto::default()
};
std::fs::write(root.join("decoder.onnx"), model.encode_to_vec())?;
std::fs::write(
root.join("inference_metadata.yaml"),
r#"
pipeline:
models:
decoder:
filename: decoder.onnx
type: decoder
dataflow: []
strategy:
kind: autoregressive
decoder: decoder
"#,
)?;
let error = PipelineEngine::from_dir_with_config(&root, EngineConfig::default())
.err()
.expect("Auto must engage the native backend for native-only pipeline components");
let message = error.to_string();
assert!(
message.contains("native"),
"error should reference the native backend: {message}"
);
assert!(!message.contains("native backend not supported for pipeline models"));
Ok(())
}
#[test]
fn plan_routes_prompt_encoder_outputs_to_decoder_inputs() -> anyhow::Result<()> {
let spec = PipelineSpec {
models: BTreeMap::from([
("vision_encoder".to_string(), component("encoder")),
("decoder".to_string(), component("decoder")),
]),
dataflow: vec![DataflowEdge {
from: "vision_encoder.image_features".to_string(),
to: "decoder.encoder_hidden_states".to_string(),
dtype: Some("fp32".to_string()),
device_transfer: Some(false),
}],
strategy: PipelineStrategy {
kind: PipelineStrategyKind::Composite,
decoder: None,
max_tokens: None,
stop_conditions: None,
kv_cache: None,
speculative: None,
model: None,
batching: None,
denoiser: None,
scheduler: None,
num_steps: None,
timestep_input: None,
timesteps: None,
start_step: None,
scheduler_config: None,
cfg_conditioning_input: None,
guidance_scale: None,
state: None,
outer: None,
inner: None,
num_code_groups: None,
pre_embedder: None,
prefill_embedder: None,
stages: vec![
PipelineStrategyStage {
name: "encode".to_string(),
strategy: Box::new(PipelineStrategy {
kind: PipelineStrategyKind::SinglePass,
decoder: None,
max_tokens: None,
stop_conditions: None,
kv_cache: None,
speculative: None,
model: Some("vision_encoder".to_string()),
batching: None,
denoiser: None,
scheduler: None,
num_steps: None,
timestep_input: None,
timesteps: None,
start_step: None,
scheduler_config: None,
cfg_conditioning_input: None,
guidance_scale: None,
state: None,
outer: None,
inner: None,
num_code_groups: None,
pre_embedder: None,
prefill_embedder: None,
stages: vec![],
}),
run_on: Some(PhaseRunOn::PromptOnly),
},
PipelineStrategyStage {
name: "decode".to_string(),
strategy: Box::new(PipelineStrategy {
kind: PipelineStrategyKind::Autoregressive,
decoder: Some("decoder".to_string()),
max_tokens: None,
stop_conditions: None,
kv_cache: None,
speculative: None,
model: None,
batching: None,
denoiser: None,
scheduler: None,
num_steps: None,
timestep_input: None,
timesteps: None,
start_step: None,
scheduler_config: None,
cfg_conditioning_input: None,
guidance_scale: None,
state: None,
outer: None,
inner: None,
num_code_groups: None,
pre_embedder: None,
prefill_embedder: None,
stages: vec![],
}),
run_on: Some(PhaseRunOn::EveryStep),
},
],
},
phases: BTreeMap::from([
(
"vision_encoder".to_string(),
PhaseConfig {
run_on: PhaseRunOn::PromptOnly,
when_present: None,
},
),
(
"decoder".to_string(),
PhaseConfig {
run_on: PhaseRunOn::EveryStep,
when_present: None,
},
),
]),
vision: None,
positions: None,
};
let plan = PipelinePlan::from_spec(&spec, &SchedulerRegistry::builtin())?;
let ar = plan.autoregressive_plan()?;
assert_eq!(ar.prompt_components, ["vision_encoder"]);
assert_eq!(ar.decoder, "decoder");
let routed = plan.edges_to_component("decoder").collect::<Vec<_>>();
assert_eq!(routed.len(), 1);
assert_eq!(
parse_endpoint(&routed[0].to)?,
("decoder", "encoder_hidden_states")
);
assert_eq!(routed[0].from, "vision_encoder.image_features");
Ok(())
}
fn bare_strategy(kind: PipelineStrategyKind) -> PipelineStrategy {
PipelineStrategy {
kind,
decoder: None,
max_tokens: None,
stop_conditions: None,
kv_cache: None,
speculative: None,
model: None,
batching: None,
denoiser: None,
scheduler: None,
num_steps: None,
timestep_input: None,
timesteps: None,
start_step: None,
scheduler_config: None,
cfg_conditioning_input: None,
guidance_scale: None,
state: None,
outer: None,
inner: None,
num_code_groups: None,
pre_embedder: None,
prefill_embedder: None,
stages: vec![],
}
}
fn single_pass_stage(name: &str, model: &str) -> PipelineStrategyStage {
PipelineStrategyStage {
name: name.to_string(),
strategy: Box::new(PipelineStrategy {
model: Some(model.to_string()),
..bare_strategy(PipelineStrategyKind::SinglePass)
}),
run_on: None,
}
}
#[test]
fn plan_builds_composite_single_pass_stages() -> anyhow::Result<()> {
let spec = PipelineSpec {
models: BTreeMap::from([
("encoder".to_string(), component("encoder")),
("decoder".to_string(), component("decoder")),
]),
dataflow: vec![DataflowEdge {
from: "encoder.codes".to_string(),
to: "decoder.codes".to_string(),
dtype: Some("int64".to_string()),
device_transfer: Some(false),
}],
strategy: PipelineStrategy {
stages: vec![
single_pass_stage("encode", "encoder"),
single_pass_stage("decode", "decoder"),
],
..bare_strategy(PipelineStrategyKind::Composite)
},
phases: BTreeMap::new(),
vision: None,
positions: None,
};
let plan = PipelinePlan::from_spec(&spec, &SchedulerRegistry::builtin())?;
match &plan {
PipelinePlan::Composite(composite) => {
assert_eq!(composite.stages.len(), 2);
assert_eq!(composite.stages[0].name, "encode");
assert_eq!(composite.stages[1].name, "decode");
assert!(matches!(
&composite.stages[0].kind,
CompositeStageKind::SinglePass { model } if model == "encoder"
));
assert!(matches!(
&composite.stages[1].kind,
CompositeStageKind::SinglePass { model } if model == "decoder"
));
}
other => panic!("expected a Composite plan, got {other:?}"),
}
let routed = plan.edges_to_component("decoder").collect::<Vec<_>>();
assert_eq!(routed.len(), 1);
assert_eq!(routed[0].from, "encoder.codes");
Ok(())
}
#[test]
fn composite_iterative_stage_is_rejected_for_now() {
let spec = PipelineSpec {
models: BTreeMap::from([("encoder".to_string(), component("encoder"))]),
dataflow: vec![],
strategy: PipelineStrategy {
stages: vec![PipelineStrategyStage {
name: "loop".to_string(),
strategy: Box::new(PipelineStrategy {
denoiser: Some("encoder".to_string()),
..bare_strategy(PipelineStrategyKind::Iterative)
}),
run_on: None,
}],
..bare_strategy(PipelineStrategyKind::Composite)
},
phases: BTreeMap::new(),
vision: None,
positions: None,
};
let error = PipelinePlan::from_spec(&spec, &SchedulerRegistry::builtin()).unwrap_err();
assert!(
error.to_string().contains("iterative stage"),
"unexpected error: {error}"
);
}
fn vision_config(placeholder_id: i64, tpt: usize) -> PipelineVisionConfig {
PipelineVisionConfig {
image_placeholder_token_id: Some(placeholder_id),
tokens_per_tile: Some(tpt),
..Default::default()
}
}
#[test]
fn image_placeholder_expansion_replaces_tokens() {
let tokens: Vec<TokenId> = vec![1, 100, 2];
let cfg = vision_config(100, 3);
let expanded = expand_image_placeholders_count_based(tokens, Some(2), Some(&cfg)).unwrap();
assert_eq!(expanded, vec![1, 100, 100, 100, 100, 100, 100, 2]);
}
#[test]
fn image_placeholder_expansion_multiple_placeholders_errors() {
let tokens: Vec<TokenId> = vec![100, 5, 100];
let cfg = vision_config(100, 4);
let err = expand_image_placeholders_count_based(tokens, Some(1), Some(&cfg)).unwrap_err();
assert!(
err.to_string()
.contains("multi-image count-based expansion is not supported"),
"unexpected error: {err}"
);
}
#[test]
fn image_placeholder_expansion_none_tiles_is_noop() {
let tokens: Vec<TokenId> = vec![1, 100, 2];
let cfg = vision_config(100, 256);
let result =
expand_image_placeholders_count_based(tokens.clone(), None, Some(&cfg)).unwrap();
assert_eq!(result, tokens);
}
#[test]
fn image_placeholder_expansion_no_vision_config_with_tiles_errors() {
let tokens: Vec<TokenId> = vec![1, 100, 2];
let err = expand_image_placeholders_count_based(tokens, Some(1), None).unwrap_err();
assert!(err.to_string().contains("no vision section"));
}
#[test]
fn image_placeholder_expansion_incomplete_contract_errors() {
let tokens: Vec<TokenId> = vec![1, 100, 2];
let cfg = PipelineVisionConfig {
image_placeholder_token_id: Some(100),
tokens_per_tile: None,
..Default::default()
};
let err = expand_image_placeholders_count_based(tokens, Some(1), Some(&cfg)).unwrap_err();
assert!(err.to_string().contains("vision contract is incomplete"));
}
#[test]
fn image_placeholder_expansion_missing_placeholder_errors() {
let tokens: Vec<TokenId> = vec![1, 2, 3];
let cfg = vision_config(100, 4);
let err = expand_image_placeholders_count_based(tokens, Some(1), Some(&cfg)).unwrap_err();
assert!(err.to_string().contains("no image placeholder token"));
}
#[test]
fn image_placeholder_expansion_negative_id_errors() {
let tokens: Vec<TokenId> = vec![1, 2];
let cfg = PipelineVisionConfig {
image_placeholder_token_id: Some(-1),
tokens_per_tile: Some(4),
..Default::default()
};
let err = expand_image_placeholders_count_based(tokens, Some(1), Some(&cfg)).unwrap_err();
assert!(err.to_string().contains("out of range"));
}
#[test]
fn image_placeholder_expansion_tokens_per_tile_zero_errors() {
let tokens: Vec<TokenId> = vec![1, 100, 2];
let cfg = vision_config(100, 0);
let err = expand_image_placeholders_count_based(tokens, Some(1), Some(&cfg)).unwrap_err();
assert!(
err.to_string().contains("tokens_per_tile is 0"),
"unexpected error: {err}"
);
}
#[test]
fn image_placeholder_expansion_zero_tiles_produces_empty_errors() {
let tokens: Vec<TokenId> = vec![100];
let cfg = vision_config(100, 4);
let err = expand_image_placeholders_count_based(tokens, Some(0), Some(&cfg)).unwrap_err();
assert!(
err.to_string().contains("empty token sequence"),
"unexpected error: {err}"
);
}
}