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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions docs/speculative-acceptance.md
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,7 @@ real and provable.
| Path | `temperature == 0` | `temperature > 0` default | `temperature > 0` opt-in |
|------|--------------------|---------------------------|--------------------------|
| `SpeculativeGenerator` (classic; `mlxcel generate --draft-model`) | greedy argmax (lossless) | **sampler-match** (lossless) | modified rejection sampling (lossless, acceptance-optimal) |
| `PromptLookupGenerator` (`mlxcel generate --prompt-lookup`) | greedy argmax (lossless) | **sampler-match** (lossless, and acceptance-optimal here) | none needed |
| Gemma 4 MTP round loop | argmax (lossless here) | argmax-against-argmax (**biased**) | not wired |
| DFlash round loop (Qwen 3.5 DFlash drafter) | argmax (lossless here) | argmax-against-argmax (**biased**) | not wired |
| DFlash round loop (LFM2 DSpark drafter) | argmax (lossless here, probe-gated) | declines to classic decode | not wired |
Expand All @@ -125,6 +126,13 @@ Those are the paths where this feature would buy correctness rather than only
acceptance rate, and where the acceptance trade against argmax is a real
decision rather than a free win.

Prompt lookup proposes deterministically, so its `q` is one-hot on the drafted
token `d`. The two acceptance probabilities then coincide at `p(d)`
(`sum_x p(x) q(x) = sum_x min(p(x), q(x)) = p(d)`), and on a rejection both
rules emit a draw from `p` conditioned on not being `d`. Sampler-match is
therefore already the acceptance-optimal rule for that path, and there is
nothing to opt into.

## RNG dependency

Modified rejection sampling draws randomness the previous rules did not:
Expand Down
106 changes: 104 additions & 2 deletions src/commands/generate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -444,7 +444,7 @@ fn validate_pipeline_parallel_args(args: &GenerateArgs) -> Result<()> {
// Single-adapter only; multi-adapter stacking and runtime hot-swap
// remain out of scope for v1.
ensure!(
args.model.draft_model.is_none(),
args.model.draft_model.is_none() && !args.prompt_lookup.prompt_lookup,
"CLI pipeline parallelism does not support speculative decoding yet"
);
ensure!(
Expand Down Expand Up @@ -1730,7 +1730,17 @@ pub(super) fn run_generation_mode(
// models whose suppressed set is empty.
token_bias.suppress_tokens(&model.output_suppressed_token_ids());

let output = if let Some(ref draft_model_path) = args.model.draft_model {
let output = if args.prompt_lookup.prompt_lookup {
run_prompt_lookup(
model,
args,
prompt_tokens,
sampling_config,
vlm_embeddings.is_some(),
kv_cache_mode,
token_bias,
)?
} else if let Some(ref draft_model_path) = args.model.draft_model {
// resolve the effective DrafterKind from
// (a) the explicit `--draft-kind` CLI flag, OR
// (b) the drafter's `config.json::model_type` auto-detection.
Expand Down Expand Up @@ -1962,6 +1972,90 @@ pub(super) fn run_generation_mode(
Ok(output)
}

/// Refuse a `--prompt-lookup` run that cannot start, before the checkpoint is
/// resolved or loaded: an invalid lookup configuration, or multimodal input,
/// whose embeddings the generator cannot take. A no-op without the flag.
pub(super) fn validate_prompt_lookup_args(args: &GenerateArgs) -> Result<()> {
if !args.prompt_lookup.prompt_lookup {
return Ok(());
}
args.prompt_lookup
.config()
.validate()
.map_err(|err| anyhow!("--prompt-lookup: {err}"))?;
ensure!(
args.generation.image.is_empty()
&& args.generation.audio.is_none()
&& args.generation.video.is_empty(),
"--prompt-lookup supports text-only prompts; drop --image/--audio/--video"
);
Ok(())
}

/// Offline prompt-lookup speculative decoding (`--prompt-lookup`).
///
/// Refuses the model families whose state cannot be rewound after a rejected
/// block, and multimodal prompts, whose embeddings the generator cannot take.
/// Runs one short warmup generation first, as `generate_standard` does, so the
/// timed run does not pay for kernel compilation, then one forward at every
/// verify width so the widths the warmup's own proposals missed are compiled
/// too.
///
/// The returned stats time the whole call, prefill included, exactly as
/// `generate_standard` does, so the printed tok/s compares like for like.
fn run_prompt_lookup(
model: &mlxcel::LoadedModel,
args: &GenerateArgs,
prompt_tokens: &[i32],
sampling_config: &SamplingConfig,
has_multimodal_input: bool,
kv_cache_mode: KVCacheMode,
token_bias: TokenBiasMap,
) -> Result<(Vec<i32>, GenerationStats)> {
ensure!(
!has_multimodal_input,
"--prompt-lookup supports text-only prompts; drop --image/--audio/--video"
);
if let Some(reason) = mlxcel::prompt_lookup_unsupported_reason(model) {
anyhow::bail!("--prompt-lookup cannot run on this model: {reason}");
}
let config = args.prompt_lookup.config();
config
.validate()
.map_err(|err| anyhow!("--prompt-lookup: {err}"))?;

let mut generator = mlxcel::PromptLookupGenerator::new(config)
.with_kv_cache_mode(kv_cache_mode)
.with_token_bias(token_bias);
let warmup_tokens = args.generation.max_tokens.min(16);
let _ = generator.generate(model, prompt_tokens, warmup_tokens, sampling_config);
// The warmup generation compiles only the verify widths its own proposals
// used; a reply with nothing to copy uses none, and the timed run would
// then pay each width's first-use cost mid-decode.
generator.warm_up_verify_widths(model, prompt_tokens);
let start_time = Instant::now();
let (tokens, measured) = generator.generate(
model,
prompt_tokens,
args.generation.max_tokens,
sampling_config,
);
let total_time = start_time.elapsed();
// Diagnostics go to stderr on their own line: stdout still holds the
// echoed prompt without a trailing newline, and the reply follows it.
eprintln!();
eprintln!("{}", generator.stats().summary_line(tokens.len()));
// `--profile` prints the prefill / decode split the generator measured,
// as the plain path does; otherwise the one-line rate covers the whole
// call, prefill included, like `generate_standard`.
let stats = if args.generation.profile {
measured
} else {
generation_stats_from_duration(prompt_tokens.len(), tokens.len(), total_time)
};
Ok((tokens, stats))
}

/// Routing gate for the offline MTP speculative path (issue #166).
///
/// Returns `true` only when the operator explicitly passed `--draft-kind mtp`
Expand Down Expand Up @@ -2558,6 +2652,10 @@ pub(crate) fn run_generate(mut args: GenerateArgs) -> Result<()> {
if args.generation.layout_detections.is_some() {
args.generation.prompt = Some(String::new());
} else {
ensure!(
!args.prompt_lookup.prompt_lookup,
"--prompt-lookup requires a one-shot -p/--prompt run (not interactive chat)"
);
let opts = chat_options_from_args(&args)?;
return crate::commands::run_chat(opts);
}
Expand Down Expand Up @@ -2590,6 +2688,10 @@ fn run_generate_once(mut args: GenerateArgs) -> Result<()> {
#[cfg(feature = "surgery")]
install_surgery_pipeline_from_cli(&args)?;

// Model-independent too, so a bad `--prompt-lookup` combination fails
// before the resolver below can auto-download the checkpoint.
validate_prompt_lookup_args(&args)?;

// Resolve `-m` into a concrete model directory (epic #92, issue #94)
// before any consumer reads it. An existing path is used verbatim
// (byte-identical to the pre-#94 local-path behavior); an `owner/name`
Expand Down
46 changes: 45 additions & 1 deletion src/commands/generate_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ use super::{
generated_suffix, generation_stats_from_duration, memory_preflight_ctx_len,
reject_dflash_drafter_offline, resolve_cli_pipeline_assignments, resolve_cli_prompt,
should_route_offline_mtp, strip_trailing_eos, validate_muse_glimmer_cli_unsupported_options,
validate_pipeline_parallel_args, validate_tensor_parallel_args,
validate_pipeline_parallel_args, validate_prompt_lookup_args, validate_tensor_parallel_args,
validate_xla_cli_image_cardinality, validate_xla_output_audio,
};
use mlxcel::server::chat_template::{ChatMessage, ChatTemplateProcessor, flatten_template_text};
Expand Down Expand Up @@ -638,6 +638,7 @@ fn sample_generate_args(model_path: PathBuf) -> crate::GenerateArgs {
},
lang_bias: mlxcel::lang_bias::LangBiasCliArgs::default(),
speculative: mlxcel::cli::speculative_args::SpeculativeArgs::default(),
prompt_lookup: crate::PromptLookupOptions::default(),
// Default to None so existing tests stay on
// the bit-exact baseline load path; tests that need surgery
// override this field explicitly.
Expand Down Expand Up @@ -964,6 +965,9 @@ fn validate_pipeline_parallel_args_rejects_incompatible_modes() {
args.model.draft_model = Some(PathBuf::from("draft"));
assert!(validate_pipeline_parallel_args(&args).is_err());
args.model.draft_model = None;
args.prompt_lookup.prompt_lookup = true;
assert!(validate_pipeline_parallel_args(&args).is_err());
args.prompt_lookup.prompt_lookup = false;

// Tensor parallelism + PP is now accepted (2D PP × TP composition landed). Positive coverage for the 2D path lives in
// `validate_pipeline_parallel_args_accepts_2d_pp_tp` below.
Expand Down Expand Up @@ -1417,3 +1421,43 @@ fn cli_video_frames_are_private_and_removed_with_the_run() {
}
assert!(!dir.exists(), "the frame directory outlived the run");
}

#[test]
fn validate_prompt_lookup_args_is_a_no_op_without_the_flag() {
let mut args = sample_generate_args(temp_model_dir("pl-disabled"));
args.generation.image = vec![PathBuf::from("page.png")];
args.prompt_lookup.prompt_lookup_ngram_min = 0;
assert!(validate_prompt_lookup_args(&args).is_ok());
fs::remove_dir_all(args.model.model).unwrap();
}

#[test]
fn validate_prompt_lookup_args_refuses_before_the_model_loads() {
let mut args = sample_generate_args(temp_model_dir("pl-refusals"));
args.prompt_lookup.prompt_lookup = true;
assert!(validate_prompt_lookup_args(&args).is_ok());

args.prompt_lookup.prompt_lookup_ngram_min = 0;
let err = validate_prompt_lookup_args(&args).unwrap_err().to_string();
assert!(err.contains("ngram-min"), "{err}");
args.prompt_lookup.prompt_lookup_ngram_min = 2;

args.prompt_lookup.prompt_lookup_ngram_max =
mlxcel_core::speculative::prompt_lookup::NGRAM_MAX_LIMIT + 1;
let err = validate_prompt_lookup_args(&args).unwrap_err().to_string();
assert!(err.contains("at most"), "{err}");
args.prompt_lookup.prompt_lookup_ngram_max = 3;

for media in ["image", "audio", "video"] {
let mut args = sample_generate_args(args.model.model.clone());
args.prompt_lookup.prompt_lookup = true;
match media {
"image" => args.generation.image = vec![PathBuf::from("page.png")],
"audio" => args.generation.audio = Some(PathBuf::from("clip.wav")),
_ => args.generation.video = vec![PathBuf::from("clip.mp4")],
}
let err = validate_prompt_lookup_args(&args).unwrap_err().to_string();
assert!(err.contains("text-only"), "{media}: {err}");
}
fs::remove_dir_all(args.model.model).unwrap();
}
1 change: 1 addition & 0 deletions src/commands/run.rs
Original file line number Diff line number Diff line change
Expand Up @@ -143,6 +143,7 @@ impl RunArgs {
tensor_parallel: crate::TensorParallelOptions::default(),
lang_bias: mlxcel::lang_bias::LangBiasCliArgs::default(),
speculative: mlxcel::cli::speculative_args::SpeculativeArgs::default(),
prompt_lookup: crate::PromptLookupOptions::default(),
#[cfg(feature = "surgery")]
surgery: None,
}
Expand Down
4 changes: 4 additions & 0 deletions src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,10 @@ pub use mlxcel_core::generate::{
SamplingConfig,
};
pub use mlxcel_core::speculative::SpeculativeGenerator;
pub use mlxcel_core::speculative::prompt_lookup::{
PromptLookupConfig, PromptLookupGenerator, prompt_lookup_unsupported_reason,
supports_prompt_lookup,
};
#[cfg(feature = "xla-diagnostics")]
pub use multimodal::host_preprocessor::LlavaHostReferenceCapture;
#[cfg(feature = "xla-iree")]
Expand Down
104 changes: 65 additions & 39 deletions src/lib/mlxcel-core/src/generate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -343,6 +343,69 @@ fn chunked_prefill_last_logits<M: LanguageModel + ?Sized>(
logits.expect("chunked_prefill_last_logits requires a non-empty prompt")
}

/// Prefill the whole prompt and return the `[1, 1, vocab]` logits of its last
/// position, the way [`CxxGenerator::generate_streaming`] does.
///
/// Picks cache-level chunked prefill when `MLXCEL_PREFILL_CHUNK` applies,
/// tile-aligned padded prefill on M5+ hardware, and one `forward_last_logits`
/// over the prompt otherwise. Speculative generators that must stay
/// byte-identical to plain decoding at greedy sampling call this instead of
/// their own prefill: splitting the prompt differently changes the fp16
/// rounding of the cached keys and values, and with it the first token.
///
/// Used by: `CxxGenerator::generate_streaming`, `PromptLookupGenerator`
pub(crate) fn prefill_prompt_last_logits<M: LanguageModel>(
model: &M,
caches: &mut [KVCache],
prompt_tokens: &[i32],
) -> UniquePtr<MlxArray> {
let actual_len = prompt_tokens.len();
let prefill_chunk = effective_prefill_chunk(
prefill_chunk_len(),
model.supports_chunked_prefill(),
actual_len,
);
if let Some(chunk) = prefill_chunk {
// Cache-level chunked prefill (MLXCEL_PREFILL_CHUNK, default 2048).
chunked_prefill_last_logits(model, caches, prompt_tokens, chunk)
} else if should_align_prefill() && model.supports_padded_prefill() {
let padded_len = align_to_na_tile(actual_len);
let (padded_tokens, mask_opt) = pad_tokens_for_prefill(
prompt_tokens,
padded_len,
model.supports_maskless_padded_prefill(),
);
let input = ffi::from_slice_i32(&padded_tokens, &[1, padded_len as i32]);
// Last *real* token position; `forward_last_logits` slices there,
// replacing the previous forward + `logits_at_position` pair.
let raw_logits = model.forward_last_logits(
&input,
caches,
mask_opt.as_ref().map(|m| m.as_ref().unwrap()),
actual_len.saturating_sub(1),
);
// Trim padding positions from all KV caches so decode uses the
// correct cache offset (actual_len, not padded_len).
if padded_len > actual_len {
trim_caches_to_actual_len(caches, actual_len, padded_len);
model.trim_internal_caches((padded_len - actual_len) as i32);
}
raw_logits
} else {
let input = ffi::from_slice_i32(prompt_tokens, &[1, actual_len as i32]);
model.forward_last_logits(&input, caches, None, actual_len.saturating_sub(1))
}
}

/// Per-layer KV cache modes for `n_layers` caches under the nominal `mode`,
/// with the Boundary-V upgrade (`MLXCEL_KV_BOUNDARY_V_LAYERS`) applied.
///
/// Used by: `CxxGenerator`, `PromptLookupGenerator`
pub(crate) fn resolve_kv_cache_layer_modes(mode: KVCacheMode, n_layers: usize) -> Vec<KVCacheMode> {
let requested = crate::cache::turbo::boundary_v_layers_from_env();
crate::cache::turbo::resolve_layer_modes(mode, n_layers, requested)
}

/// Trait for language models that can be used for generation
pub trait LanguageModel {
/// Forward pass through the model
Expand Down Expand Up @@ -1662,9 +1725,7 @@ impl CxxGenerator {
/// including `reset_with_model` boundary cases.
///
fn resolved_layer_modes_for(&self, n_layers: usize) -> Vec<KVCacheMode> {
let nominal = self.kv_cache_mode;
let requested = crate::cache::turbo::boundary_v_layers_from_env();
crate::cache::turbo::resolve_layer_modes(nominal, n_layers, requested)
resolve_kv_cache_layer_modes(self.kv_cache_mode, n_layers)
}

fn apply_kv_cache_mode_with_boundary_policy(&mut self) -> Vec<KVCacheMode> {
Expand Down Expand Up @@ -1756,42 +1817,7 @@ impl CxxGenerator {
// Prefill: process all prompt tokens at once.
// On M5+ hardware pad the sequence to a 32-token tile boundary for
// optimal Neural Accelerator throughput.
let actual_len = prompt_tokens.len();
let prefill_chunk = effective_prefill_chunk(
prefill_chunk_len(),
model.supports_chunked_prefill(),
actual_len,
);
let logits = if let Some(chunk) = prefill_chunk {
// Cache-level chunked prefill (MLXCEL_PREFILL_CHUNK, default 2048).
chunked_prefill_last_logits(model, &mut self.caches, prompt_tokens, chunk)
} else if should_align_prefill() && model.supports_padded_prefill() {
let padded_len = align_to_na_tile(actual_len);
let (padded_tokens, mask_opt) = pad_tokens_for_prefill(
prompt_tokens,
padded_len,
model.supports_maskless_padded_prefill(),
);
let input = ffi::from_slice_i32(&padded_tokens, &[1, padded_len as i32]);
// Last *real* token position; `forward_last_logits` slices there,
// replacing the previous forward + `logits_at_position` pair.
let raw_logits = model.forward_last_logits(
&input,
&mut self.caches,
mask_opt.as_ref().map(|m| m.as_ref().unwrap()),
actual_len.saturating_sub(1),
);
// Trim padding positions from all KV caches so decode uses the
// correct cache offset (actual_len, not padded_len).
if padded_len > actual_len {
trim_caches_to_actual_len(&mut self.caches, actual_len, padded_len);
model.trim_internal_caches((padded_len - actual_len) as i32);
}
raw_logits
} else {
let input = ffi::from_slice_i32(prompt_tokens, &[1, actual_len as i32]);
model.forward_last_logits(&input, &mut self.caches, None, actual_len.saturating_sub(1))
};
let logits = prefill_prompt_last_logits(model, &mut self.caches, prompt_tokens);

if trace_dtype {
ffi::eval(&logits);
Expand Down
4 changes: 4 additions & 0 deletions src/lib/mlxcel-core/src/speculative/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -57,8 +57,12 @@
//! `trim_caches`). See [`mtp::MtpGenerator`].
//! - [`stochastic_accept`] — the distribution-preserving acceptance rule and
//! its residual resample, shared by every speculative verify path.
//! - [`prompt_lookup`] — drafter-free speculation that proposes the tokens
//! following an earlier occurrence of the sequence's tail. See
//! [`prompt_lookup::PromptLookupGenerator`].

pub mod mtp;
pub mod prompt_lookup;
pub mod stochastic_accept;

use crate::cache::can_trim_prompt_cache;
Expand Down
Loading
Loading