Skip to content
Closed
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
61 changes: 0 additions & 61 deletions crates/skippy-prompt/src/prompt_cli/args.rs
Original file line number Diff line number Diff line change
Expand Up @@ -30,19 +30,10 @@ impl From<ReplLoadMode> for RuntimeLoadMode {
}
}

#[derive(Debug, Clone, Copy, PartialEq, Eq, ValueEnum)]
#[value(rename_all = "kebab-case")]
pub enum NgramProposalMode {
TransitionPool,
HistoryMatch,
}

#[derive(Parser)]
pub struct PromptArgs {
#[arg(long, default_value = "target/debug/metrics-server")]
pub metrics_server_bin: PathBuf,
#[arg(long, default_value = "target/debug/ngram-pool-server")]
pub ngram_pool_server_bin: PathBuf,
#[arg(long, default_value = "target/debug/skippy-server")]
pub stage_server_bin: PathBuf,
#[arg(long, default_value = "target/debug/skippy-model-package")]
Expand Down Expand Up @@ -95,34 +86,10 @@ pub struct PromptArgs {
pub speculative_window: usize,
#[arg(long)]
pub adaptive_speculative_window: bool,
#[arg(long)]
pub ngram_speculative: bool,
#[arg(long, value_enum, default_value = "transition-pool")]
pub ngram_proposal_mode: NgramProposalMode,
#[arg(long, default_value_t = 24)]
pub spec_ngram_size_n: usize,
#[arg(long, default_value_t = 1)]
pub ngram_history_min_hits: u32,
#[arg(long, default_value_t = 12)]
pub draft_min: usize,
#[arg(long, default_value_t = 48)]
pub draft_max: usize,
#[arg(long, default_value_t = DEFAULT_MIN_WINNER_COUNT)]
pub ngram_min_winner_count: u32,
#[arg(long, default_value_t = DEFAULT_MIN_CONFIDENCE)]
pub ngram_min_confidence: f32,
#[arg(long, default_value_t = DEFAULT_MIN_MARGIN)]
pub ngram_min_margin: u32,
#[arg(long, default_value_t = DEFAULT_CONFIDENCE_STEP)]
pub ngram_confidence_step: f32,
#[arg(long, default_value_t = DEFAULT_CONFIDENCE_STEP_TOKENS)]
pub ngram_confidence_step_tokens: usize,
#[arg(long, default_value_t = DEFAULT_MAX_CONFIDENCE)]
pub ngram_max_confidence: f32,
#[arg(long, default_value_t = DEFAULT_COUNT_STEP_TOKENS)]
pub ngram_count_step_tokens: usize,
#[arg(long, default_value_t = DEFAULT_MARGIN_STEP_TOKENS)]
pub ngram_margin_step_tokens: usize,
#[arg(long, default_value = "lookup-record")]
pub kv_mode: String,
#[arg(long, default_value_t = 512)]
Expand Down Expand Up @@ -155,8 +122,6 @@ pub struct PromptArgs {
pub history_path: Option<PathBuf>,
#[arg(long)]
pub session_id: Option<String>,
#[arg(long)]
pub ngram_pool_uds_path: Option<PathBuf>,
#[arg(long, default_value_t = 80)]
pub log_tail_lines: usize,
#[arg(long)]
Expand Down Expand Up @@ -210,34 +175,10 @@ pub struct BinaryReplArgs {
pub speculative_window: usize,
#[arg(long)]
pub adaptive_speculative_window: bool,
#[arg(long)]
pub ngram_speculative: bool,
#[arg(long, value_enum, default_value = "transition-pool")]
pub ngram_proposal_mode: NgramProposalMode,
#[arg(long, default_value_t = 24)]
pub spec_ngram_size_n: usize,
#[arg(long, default_value_t = 1)]
pub ngram_history_min_hits: u32,
#[arg(long, default_value_t = 12)]
pub draft_min: usize,
#[arg(long, default_value_t = 48)]
pub draft_max: usize,
#[arg(long, default_value_t = DEFAULT_MIN_WINNER_COUNT)]
pub ngram_min_winner_count: u32,
#[arg(long, default_value_t = DEFAULT_MIN_CONFIDENCE)]
pub ngram_min_confidence: f32,
#[arg(long, default_value_t = DEFAULT_MIN_MARGIN)]
pub ngram_min_margin: u32,
#[arg(long, default_value_t = DEFAULT_CONFIDENCE_STEP)]
pub ngram_confidence_step: f32,
#[arg(long, default_value_t = DEFAULT_CONFIDENCE_STEP_TOKENS)]
pub ngram_confidence_step_tokens: usize,
#[arg(long, default_value_t = DEFAULT_MAX_CONFIDENCE)]
pub ngram_max_confidence: f32,
#[arg(long, default_value_t = DEFAULT_COUNT_STEP_TOKENS)]
pub ngram_count_step_tokens: usize,
#[arg(long, default_value_t = DEFAULT_MARGIN_STEP_TOKENS)]
pub ngram_margin_step_tokens: usize,
#[arg(long, default_value_t = 60)]
pub startup_timeout_secs: u64,
#[arg(long, default_value_t = 30)]
Expand All @@ -247,8 +188,6 @@ pub struct BinaryReplArgs {
#[arg(long)]
pub session_id: Option<String>,
#[arg(long)]
pub ngram_pool_uds_path: Option<PathBuf>,
#[arg(long)]
pub native_logs: bool,
#[arg(
long,
Expand Down
24 changes: 0 additions & 24 deletions crates/skippy-prompt/src/prompt_cli/binary_repl.rs
Original file line number Diff line number Diff line change
Expand Up @@ -141,27 +141,6 @@ pub fn binary_repl(args: BinaryReplArgs) -> Result<()> {
let default_session_id = args.session_id.clone().unwrap_or_else(default_session_id);
let default_wire_session_id = stable_wire_id(&[default_session_id.as_bytes()]);
eprintln!("session_id={default_session_id} wire_session_id={default_wire_session_id}");
let mut ngram = if args.ngram_speculative {
eprintln!(
"ngram speculative enabled: mode={:?} n={} history_min_hits={} draft_min={} draft_max={} min_count={} min_confidence={:.2} min_margin={} confidence_step={:.2}/{} max_confidence={:.2} count_step={} margin_step={}",
args.ngram_proposal_mode,
args.spec_ngram_size_n,
args.ngram_history_min_hits,
args.draft_min,
args.draft_max,
args.ngram_min_winner_count,
args.ngram_min_confidence,
args.ngram_min_margin,
args.ngram_confidence_step,
args.ngram_confidence_step_tokens,
args.ngram_max_confidence,
args.ngram_count_step_tokens,
args.ngram_margin_step_tokens
);
Some(NgramSource::open(&args, &default_session_id)?)
} else {
None
};
let interrupt = install_prompt_interrupt_handler()?;
let mut history = PromptHistory::load(args.history_path.as_deref())?;
let mut prompt_input = prompt_input(&history)?;
Expand Down Expand Up @@ -198,7 +177,6 @@ pub fn binary_repl(args: BinaryReplArgs) -> Result<()> {
tokenizer: &tokenizer,
chat_template_model: chat_template_model.as_ref(),
draft: draft.as_mut(),
ngram: ngram.as_mut(),
interrupt: &interrupt,
wire_dtype,
session_id: &prompt_session_id,
Expand Down Expand Up @@ -263,7 +241,6 @@ pub fn binary_repl(args: BinaryReplArgs) -> Result<()> {
tokenizer: &tokenizer,
chat_template_model: chat_template_model.as_ref(),
draft: draft.as_mut(),
ngram: ngram.as_mut(),
interrupt: &interrupt,
wire_dtype,
session_id: &default_session_id,
Expand All @@ -283,7 +260,6 @@ pub fn binary_repl(args: BinaryReplArgs) -> Result<()> {
tokenizer: &tokenizer,
chat_template_model: chat_template_model.as_ref(),
draft: draft.as_mut(),
ngram: ngram.as_mut(),
interrupt: &interrupt,
wire_dtype,
session_id: &default_session_id,
Expand Down
25 changes: 0 additions & 25 deletions crates/skippy-prompt/src/prompt_cli/draft.rs
Original file line number Diff line number Diff line change
@@ -1,28 +1,3 @@
struct NgramSource;

impl NgramSource {
fn open(_args: &BinaryReplArgs, _session_id: &str) -> Result<Self> {
bail!("ngram speculative prompting is not imported into mesh-llm")
}

fn observe_sequence(&mut self, _session_id: &str, _tokens: &[i32]) -> Result<()> {
Ok(())
}

fn observe_accepted(&mut self, _session_id: &str, _context_tokens: &[i32]) -> Result<()> {
Ok(())
}

fn propose(
&mut self,
_session_id: &str,
_context_tokens: &[i32],
_remaining: usize,
) -> Result<Vec<i32>> {
Ok(Vec::new())
}
}

struct DraftRunner {
path: PathBuf,
window: usize,
Expand Down
29 changes: 5 additions & 24 deletions crates/skippy-prompt/src/prompt_cli/generation.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@ struct PromptRun<'a> {
tokenizer: &'a StageModel,
chat_template_model: Option<&'a StageModel>,
draft: Option<&'a mut DraftRunner>,
ngram: Option<&'a mut NgramSource>,
interrupt: &'a Arc<PromptInterruptState>,
wire_dtype: skippy_protocol::binary::WireActivationDType,
session_id: &'a str,
Expand All @@ -19,7 +18,6 @@ fn run_prompt(run: PromptRun<'_>) -> Result<()> {
tokenizer,
chat_template_model,
mut draft,
mut ngram,
interrupt,
wire_dtype,
session_id,
Expand Down Expand Up @@ -280,16 +278,13 @@ fn run_prompt(run: PromptRun<'_>) -> Result<()> {
if let Some(draft) = draft.as_deref_mut() {
draft.reset_to_context(&context_tokens)?;
}
if let Some(ngram) = ngram.as_deref_mut() {
ngram.observe_sequence(session_id, &context_tokens)?;
}
let max_speculative_window = args.speculative_window.max(1);
let mut adaptive_window = if args.adaptive_speculative_window {
max_speculative_window.min(4)
} else {
max_speculative_window
};
if draft.is_some() || ngram.is_some() {
if draft.is_some() {
speculative_stats.adaptive_window_max = max_speculative_window;
speculative_stats.adaptive_window_start = adaptive_window;
speculative_stats.adaptive_window_enabled = args.adaptive_speculative_window;
Expand All @@ -305,19 +300,11 @@ fn run_prompt(run: PromptRun<'_>) -> Result<()> {

let remaining = max_new_tokens - generated.len();
let proposal_limit = remaining.min(adaptive_window);
let draft_tokens = match ngram.as_deref_mut() {
Some(ngram) => ngram.propose(session_id, &context_tokens, proposal_limit)?,
None => Vec::new(),
};
let draft_tokens = if draft_tokens.is_empty() {
match draft.as_deref_mut() {
Some(draft) if draft.window > 0 => {
draft.propose(current, proposal_limit.min(draft.window))?
}
_ => Vec::new(),
let draft_tokens = match draft.as_deref_mut() {
Some(draft) if draft.window > 0 => {
draft.propose(current, proposal_limit.min(draft.window))?
}
} else {
draft_tokens
_ => Vec::new(),
};

if draft_tokens.is_empty() {
Expand All @@ -340,9 +327,6 @@ fn run_prompt(run: PromptRun<'_>) -> Result<()> {
current = reply.predicted;
generated.push(current);
context_tokens.push(current);
if let Some(ngram) = ngram.as_deref_mut() {
ngram.observe_accepted(session_id, &context_tokens)?;
}
first_time_to_token_ms.get_or_insert_with(|| elapsed_ms(wall_started));
if tokenizer.token_is_eog(current)? {
generation_reached_eog = true;
Expand Down Expand Up @@ -483,9 +467,6 @@ fn run_prompt(run: PromptRun<'_>) -> Result<()> {
current = predicted;
generated.push(current);
context_tokens.push(current);
if let Some(ngram) = ngram.as_deref_mut() {
ngram.observe_accepted(session_id, &context_tokens)?;
}
if tokenizer.token_is_eog(current)? {
reached_eog = true;
generation_reached_eog = true;
Expand Down
37 changes: 0 additions & 37 deletions crates/skippy-prompt/src/prompt_cli/launch.rs
Original file line number Diff line number Diff line change
Expand Up @@ -243,30 +243,6 @@ fn prompt_repl_launch(args: PromptArgs) -> Result<()> {
);
children.push(ChildGuard::spawn(metrics)?);

let ngram_pool_uds_path = if args.ngram_speculative {
let socket_path = args
.ngram_pool_uds_path
.clone()
.unwrap_or_else(|| run_dir.join("ngram-pool.sock"));
let mut ngram_pool = Command::new(&args.ngram_pool_server_bin);
ngram_pool.args(["serve", "--uds-path", path_str(&socket_path)?]);
configure_process_log(&mut ngram_pool, &run_dir.join("ngram-pool-server.log"))?;
eprintln!(
"launch: starting ngram-pool-server socket={} log={}",
socket_path.display(),
run_dir.join("ngram-pool-server.log").display()
);
children.push(ChildGuard::spawn(ngram_pool)?);
eprintln!(
"launch: waiting for ngram pool socket {}",
socket_path.display()
);
wait_for_socket(&socket_path, args.startup_timeout_secs)?;
Some(socket_path)
} else {
None
};

if remote {
rsync_remote_stage_inputs(&args, &stages, &model_package_dir, hf_package_ref)?;
eprintln!(
Expand Down Expand Up @@ -356,28 +332,15 @@ fn prompt_repl_launch(args: PromptArgs) -> Result<()> {
draft_model_path: args.draft_model_path,
speculative_window: args.speculative_window,
adaptive_speculative_window: args.adaptive_speculative_window,
ngram_speculative: args.ngram_speculative,
ngram_proposal_mode: args.ngram_proposal_mode,
spec_ngram_size_n: args.spec_ngram_size_n,
ngram_history_min_hits: args.ngram_history_min_hits,
draft_min: args.draft_min,
draft_max: args.draft_max,
ngram_min_winner_count: args.ngram_min_winner_count,
ngram_min_confidence: args.ngram_min_confidence,
ngram_min_margin: args.ngram_min_margin,
ngram_confidence_step: args.ngram_confidence_step,
ngram_confidence_step_tokens: args.ngram_confidence_step_tokens,
ngram_max_confidence: args.ngram_max_confidence,
ngram_count_step_tokens: args.ngram_count_step_tokens,
ngram_margin_step_tokens: args.ngram_margin_step_tokens,
startup_timeout_secs: args.startup_timeout_secs,
decode_timeout_secs: args.decode_timeout_secs,
history_path: Some(
args.history_path
.unwrap_or_else(|| args.run_root.join("prompt-history.txt")),
),
session_id: args.session_id,
ngram_pool_uds_path,
native_logs: args.native_logs,
raw_prompt: args.raw_prompt,
no_think: args.no_think,
Expand Down
8 changes: 0 additions & 8 deletions crates/skippy-prompt/src/prompt_cli/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -41,14 +41,6 @@ use skippy_topology::{
dense_attention_layers, infer_family_capability, plan_contiguous_with_splits,
};

const DEFAULT_MIN_WINNER_COUNT: u32 = 2;
const DEFAULT_MIN_CONFIDENCE: f32 = 0.55;
const DEFAULT_MIN_MARGIN: u32 = 1;
const DEFAULT_CONFIDENCE_STEP: f32 = 0.0;
const DEFAULT_CONFIDENCE_STEP_TOKENS: usize = usize::MAX;
const DEFAULT_MAX_CONFIDENCE: f32 = 0.95;
const DEFAULT_COUNT_STEP_TOKENS: usize = usize::MAX;
const DEFAULT_MARGIN_STEP_TOKENS: usize = usize::MAX;
const DEFAULT_MESH_CTX_SIZE: u32 = 4096;
const DEFAULT_MESH_PROMPT_MAX_NEW_TOKENS: usize = 0;
const PROMPT_EXACT_PREFIX_RESTORE_MIN_TOKENS: usize = 512;
Expand Down
13 changes: 0 additions & 13 deletions crates/skippy-prompt/src/prompt_cli/topology.rs
Original file line number Diff line number Diff line change
Expand Up @@ -153,19 +153,6 @@ fn even_stage_ranges(stage_count: usize, layer_end: u32) -> Result<Vec<(u32, u32
Ok(ranges)
}

fn wait_for_socket(socket_path: &Path, timeout_secs: u64) -> Result<()> {
let deadline = Instant::now() + Duration::from_secs(timeout_secs.max(1));
loop {
if socket_path.exists() {
return Ok(());
}
if Instant::now() >= deadline {
bail!("timed out waiting for socket {}", socket_path.display());
}
std::thread::sleep(std::time::Duration::from_millis(100));
}
}

fn write_json(path: &Path, value: &serde_json::Value) -> Result<()> {
fs::write(path, serde_json::to_vec_pretty(value)?)
.with_context(|| format!("write {}", path.display()))
Expand Down
Loading
Loading