feat: first cut at actual token generation and model loading
This commit is contained in:
297
src/app.rs
297
src/app.rs
@@ -3,6 +3,8 @@ mod view;
|
||||
pub(crate) use view::app_theme;
|
||||
|
||||
use crate::database::{AppPreferences, Database, ProjectWithSessions};
|
||||
#[cfg(target_os = "macos")]
|
||||
use crate::engine::{ChatTurn, Generator};
|
||||
use crate::model::{self, DownloadOutcome, DownloadProgress, ManagedArtifactId, ModelChoice};
|
||||
use crate::settings::{
|
||||
DiagnosticPreferences, ExecutionPreferences, GIB, GenerationPreferences, ReasoningMode,
|
||||
@@ -398,9 +400,46 @@ pub(crate) struct App {
|
||||
pending_project_path: Option<PathBuf>,
|
||||
project_name_input: String,
|
||||
model_download: ModelDownload,
|
||||
pub(super) composer: String,
|
||||
pub(super) conversation: Vec<ChatMessage>,
|
||||
pub(super) generating: bool,
|
||||
#[cfg(target_os = "macos")]
|
||||
generation_worker: Option<GenerationWorker>,
|
||||
error: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(super) struct ChatMessage {
|
||||
pub(super) user: bool,
|
||||
reasoning: bool,
|
||||
pub(super) content: String,
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
struct GenerationWorker {
|
||||
commands: mpsc::Sender<GenerationCommand>,
|
||||
events: mpsc::Receiver<GenerationEvent>,
|
||||
cancel: Option<Arc<AtomicBool>>,
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
enum GenerationCommand {
|
||||
Generate {
|
||||
engine: crate::settings::EngineSettings,
|
||||
turn: crate::settings::TurnSettings,
|
||||
messages: Vec<ChatTurn>,
|
||||
idle_timeout: Duration,
|
||||
cancel: Arc<AtomicBool>,
|
||||
},
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
enum GenerationEvent {
|
||||
Loading,
|
||||
Chunk(String),
|
||||
Finished(Result<(), String>),
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(super) enum ModelDownload {
|
||||
Idle,
|
||||
@@ -479,6 +518,10 @@ pub(crate) enum Message {
|
||||
CancelDeleteArtifact,
|
||||
StopModelDownload,
|
||||
DownloadProgressTick,
|
||||
ComposerChanged(String),
|
||||
SubmitPrompt,
|
||||
StopGeneration,
|
||||
GenerationTick,
|
||||
ChooseProjectFolder,
|
||||
ProjectFolderPicked(Option<PathBuf>),
|
||||
ProjectNameChanged(String),
|
||||
@@ -519,6 +562,11 @@ impl App {
|
||||
pending_project_path: None,
|
||||
project_name_input: String::new(),
|
||||
model_download: ModelDownload::Idle,
|
||||
composer: String::new(),
|
||||
conversation: Vec::new(),
|
||||
generating: false,
|
||||
#[cfg(target_os = "macos")]
|
||||
generation_worker: None,
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
@@ -550,6 +598,11 @@ impl App {
|
||||
pending_project_path: None,
|
||||
project_name_input: String::new(),
|
||||
model_download: ModelDownload::Idle,
|
||||
composer: String::new(),
|
||||
conversation: Vec::new(),
|
||||
generating: false,
|
||||
#[cfg(target_os = "macos")]
|
||||
generation_worker: None,
|
||||
error: Some(format!("Could not open the project database: {error}")),
|
||||
}
|
||||
}
|
||||
@@ -816,6 +869,19 @@ impl App {
|
||||
}
|
||||
}
|
||||
Message::DownloadProgressTick => self.update_download_progress(),
|
||||
Message::ComposerChanged(value) => self.composer = value,
|
||||
Message::SubmitPrompt => self.start_generation(),
|
||||
Message::StopGeneration => {
|
||||
#[cfg(target_os = "macos")]
|
||||
if let Some(cancel) = self
|
||||
.generation_worker
|
||||
.as_ref()
|
||||
.and_then(|worker| worker.cancel.as_ref())
|
||||
{
|
||||
cancel.store(true, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
Message::GenerationTick => self.poll_generation(),
|
||||
Message::ChooseProjectFolder => {
|
||||
self.choosing_folder = true;
|
||||
return Task::perform(
|
||||
@@ -843,11 +909,21 @@ impl App {
|
||||
self.error = None;
|
||||
}
|
||||
Message::SelectProject(project_id) => {
|
||||
if self.generating {
|
||||
self.error =
|
||||
Some("Stop the active generation before changing sessions.".into());
|
||||
return Task::none();
|
||||
}
|
||||
self.selected_project = Some(project_id);
|
||||
self.selected_session = None;
|
||||
self.error = None;
|
||||
}
|
||||
Message::DeleteProject(project_id) => {
|
||||
if self.generating && self.selected_project == Some(project_id) {
|
||||
self.error =
|
||||
Some("Stop the active generation before deleting its project.".into());
|
||||
return Task::none();
|
||||
}
|
||||
if let Some(database) = &mut self.database {
|
||||
match database.delete_project(project_id) {
|
||||
Ok(()) => {
|
||||
@@ -863,11 +939,25 @@ impl App {
|
||||
}
|
||||
Message::CreateSession => self.create_session(),
|
||||
Message::SelectSession(project_id, session_id) => {
|
||||
if self.generating && self.selected_session != Some(session_id) {
|
||||
self.error =
|
||||
Some("Stop the active generation before changing sessions.".into());
|
||||
return Task::none();
|
||||
}
|
||||
if self.selected_session != Some(session_id) {
|
||||
self.conversation.clear();
|
||||
self.composer.clear();
|
||||
}
|
||||
self.selected_project = Some(project_id);
|
||||
self.selected_session = Some(session_id);
|
||||
self.error = None;
|
||||
}
|
||||
Message::DeleteSession(session_id) => {
|
||||
if self.generating && self.selected_session == Some(session_id) {
|
||||
self.error =
|
||||
Some("Stop the active generation before deleting its session.".into());
|
||||
return Task::none();
|
||||
}
|
||||
if let Some(database) = &mut self.database {
|
||||
match database.delete_session(session_id) {
|
||||
Ok(()) => {
|
||||
@@ -905,6 +995,11 @@ impl App {
|
||||
iced::time::every(Duration::from_secs(1)).map(|_| Message::DownloadProgressTick),
|
||||
);
|
||||
}
|
||||
if self.generating {
|
||||
subscriptions.push(
|
||||
iced::time::every(Duration::from_millis(50)).map(|_| Message::GenerationTick),
|
||||
);
|
||||
}
|
||||
Subscription::batch(subscriptions)
|
||||
}
|
||||
|
||||
@@ -1108,6 +1203,10 @@ impl App {
|
||||
}
|
||||
|
||||
fn create_session(&mut self) {
|
||||
if self.generating {
|
||||
self.error = Some("Stop the active generation before creating a session.".into());
|
||||
return;
|
||||
}
|
||||
let Some(project_id) = self.selected_project else {
|
||||
self.error = Some("Select a project first.".into());
|
||||
return;
|
||||
@@ -1205,6 +1304,196 @@ impl App {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn start_generation(&mut self) {
|
||||
if self.generating || self.selected_session.is_none() {
|
||||
return;
|
||||
}
|
||||
let prompt = self.composer.trim().to_owned();
|
||||
if prompt.is_empty() {
|
||||
return;
|
||||
}
|
||||
let model = match ModelChoice::from_id(&self.preferences.selected_model) {
|
||||
Some(model) => model,
|
||||
None => {
|
||||
self.error = Some("The selected model is not supported.".into());
|
||||
return;
|
||||
}
|
||||
};
|
||||
let effective = self.preferences.generation().and_then(|generation| {
|
||||
self.preferences.runtime().and_then(|runtime| {
|
||||
crate::settings::effective_settings(model, &generation, &runtime, &models_path())
|
||||
})
|
||||
});
|
||||
let effective = match effective {
|
||||
Ok(settings) => settings,
|
||||
Err(error) => {
|
||||
self.error = Some(error);
|
||||
return;
|
||||
}
|
||||
};
|
||||
let assistant_reasoning = effective.turn.reasoning_mode != ReasoningMode::Direct;
|
||||
#[cfg(target_os = "macos")]
|
||||
let mut messages = self
|
||||
.conversation
|
||||
.iter()
|
||||
.map(|message| ChatTurn {
|
||||
user: message.user,
|
||||
reasoning: message.reasoning,
|
||||
content: message.content.clone(),
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
#[cfg(target_os = "macos")]
|
||||
messages.push(ChatTurn {
|
||||
user: true,
|
||||
reasoning: false,
|
||||
content: prompt.clone(),
|
||||
});
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
{
|
||||
if self.generation_worker.is_none() {
|
||||
match spawn_generation_worker() {
|
||||
Ok(worker) => self.generation_worker = Some(worker),
|
||||
Err(error) => {
|
||||
self.error = Some(error);
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
let cancel = Arc::new(AtomicBool::new(false));
|
||||
let idle_timeout =
|
||||
Duration::from_secs(self.preferences.idle_timeout_minutes.max(1) as u64 * 60);
|
||||
let command = GenerationCommand::Generate {
|
||||
engine: effective.engine,
|
||||
turn: effective.turn,
|
||||
messages,
|
||||
idle_timeout,
|
||||
cancel: Arc::clone(&cancel),
|
||||
};
|
||||
let worker = self.generation_worker.as_mut().expect("worker was created");
|
||||
if worker.commands.send(command).is_err() {
|
||||
self.generation_worker = None;
|
||||
self.error = Some("The local generation worker stopped unexpectedly.".into());
|
||||
return;
|
||||
}
|
||||
worker.cancel = Some(cancel);
|
||||
}
|
||||
#[cfg(not(target_os = "macos"))]
|
||||
{
|
||||
let _ = effective;
|
||||
self.error = Some("Local Metal generation requires macOS.".into());
|
||||
return;
|
||||
}
|
||||
self.composer.clear();
|
||||
self.conversation.push(ChatMessage {
|
||||
user: true,
|
||||
reasoning: false,
|
||||
content: prompt,
|
||||
});
|
||||
self.conversation.push(ChatMessage {
|
||||
user: false,
|
||||
reasoning: assistant_reasoning,
|
||||
content: String::new(),
|
||||
});
|
||||
self.generating = true;
|
||||
self.error = None;
|
||||
}
|
||||
|
||||
fn poll_generation(&mut self) {
|
||||
#[cfg(target_os = "macos")]
|
||||
let Some(worker) = &mut self.generation_worker else {
|
||||
self.generating = false;
|
||||
return;
|
||||
};
|
||||
#[cfg(target_os = "macos")]
|
||||
loop {
|
||||
match worker.events.try_recv() {
|
||||
Ok(GenerationEvent::Loading) => {}
|
||||
Ok(GenerationEvent::Chunk(chunk)) => {
|
||||
if let Some(message) = self.conversation.last_mut()
|
||||
&& !message.user
|
||||
{
|
||||
message.content.push_str(&chunk);
|
||||
}
|
||||
}
|
||||
Ok(GenerationEvent::Finished(result)) => {
|
||||
self.generating = false;
|
||||
worker.cancel = None;
|
||||
if let Err(error) = result {
|
||||
self.error = Some(error);
|
||||
}
|
||||
break;
|
||||
}
|
||||
Err(TryRecvError::Empty) => break,
|
||||
Err(TryRecvError::Disconnected) => {
|
||||
self.generating = false;
|
||||
self.generation_worker = None;
|
||||
self.error = Some("The local generation worker stopped unexpectedly.".into());
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
fn spawn_generation_worker() -> Result<GenerationWorker, String> {
|
||||
let (command_sender, command_receiver) = mpsc::channel();
|
||||
let (event_sender, event_receiver) = mpsc::channel();
|
||||
thread::Builder::new()
|
||||
.name("local-generation".into())
|
||||
.spawn(move || {
|
||||
let mut loaded = None::<(crate::settings::EngineSettings, Generator)>;
|
||||
let mut last_used = Instant::now();
|
||||
let mut idle_timeout = Duration::from_secs(15 * 60);
|
||||
loop {
|
||||
match command_receiver.recv_timeout(Duration::from_secs(1)) {
|
||||
Ok(GenerationCommand::Generate {
|
||||
engine,
|
||||
turn,
|
||||
messages,
|
||||
idle_timeout: requested_timeout,
|
||||
cancel,
|
||||
}) => {
|
||||
idle_timeout = requested_timeout;
|
||||
if loaded
|
||||
.as_ref()
|
||||
.is_none_or(|(current, _)| current != &engine)
|
||||
{
|
||||
let _ = event_sender.send(GenerationEvent::Loading);
|
||||
loaded = match Generator::open(&engine) {
|
||||
Ok(generator) => Some((engine.clone(), generator)),
|
||||
Err(error) => {
|
||||
let _ =
|
||||
event_sender.send(GenerationEvent::Finished(Err(error)));
|
||||
None
|
||||
}
|
||||
};
|
||||
}
|
||||
if let Some((_, generator)) = &mut loaded {
|
||||
let result = generator.generate(&messages, &turn, &cancel, |chunk| {
|
||||
let _ = event_sender.send(GenerationEvent::Chunk(chunk));
|
||||
});
|
||||
let _ = event_sender.send(GenerationEvent::Finished(result));
|
||||
last_used = Instant::now();
|
||||
}
|
||||
}
|
||||
Err(mpsc::RecvTimeoutError::Timeout) => {
|
||||
if loaded.is_some() && last_used.elapsed() >= idle_timeout {
|
||||
loaded = None;
|
||||
}
|
||||
}
|
||||
Err(mpsc::RecvTimeoutError::Disconnected) => break,
|
||||
}
|
||||
}
|
||||
})
|
||||
.map_err(|error| format!("Could not start local generation: {error}"))?;
|
||||
Ok(GenerationWorker {
|
||||
commands: command_sender,
|
||||
events: event_receiver,
|
||||
cancel: None,
|
||||
})
|
||||
}
|
||||
|
||||
impl Drop for App {
|
||||
@@ -1212,6 +1501,14 @@ impl Drop for App {
|
||||
if let ModelDownload::Active(download) = &self.model_download {
|
||||
download.cancel.store(true, Ordering::Relaxed);
|
||||
}
|
||||
#[cfg(target_os = "macos")]
|
||||
if let Some(cancel) = self
|
||||
.generation_worker
|
||||
.as_ref()
|
||||
.and_then(|worker| worker.cancel.as_ref())
|
||||
{
|
||||
cancel.store(true, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -224,18 +224,51 @@ impl App {
|
||||
.align_y(Alignment::Center);
|
||||
|
||||
let body: Element<'_, Message> = if let Some(session) = self.selected_session(item) {
|
||||
let mut messages = column![].spacing(12);
|
||||
if self.conversation.is_empty() {
|
||||
messages = messages.push(
|
||||
column![
|
||||
text(&session.title).size(26),
|
||||
text("Run DeepSeek locally with the Rust Metal engine.").size(14),
|
||||
]
|
||||
.spacing(8),
|
||||
);
|
||||
} else {
|
||||
for message in &self.conversation {
|
||||
let label = if message.user { "You" } else { "DS4" };
|
||||
let content = if !message.user && message.content.is_empty() && self.generating
|
||||
{
|
||||
"Loading model…"
|
||||
} else {
|
||||
&message.content
|
||||
};
|
||||
messages = messages.push(
|
||||
container(column![text(label).size(11), text(content).size(14)].spacing(5))
|
||||
.padding(14)
|
||||
.width(Length::Fill)
|
||||
.style(overview_style),
|
||||
);
|
||||
}
|
||||
}
|
||||
let composer = text_input("Ask DS4Server anything…", &self.composer)
|
||||
.on_input(Message::ComposerChanged)
|
||||
.on_submit(Message::SubmitPrompt)
|
||||
.padding(12)
|
||||
.size(14);
|
||||
let action = if self.generating {
|
||||
action_button(text("Stop").size(12)).on_press(Message::StopGeneration)
|
||||
} else if self.composer.trim().is_empty() {
|
||||
action_button(icon(ICON_SEND, 18)).padding(8)
|
||||
} else {
|
||||
action_button(icon(ICON_SEND, 18))
|
||||
.padding(8)
|
||||
.on_press(Message::SubmitPrompt)
|
||||
};
|
||||
let conversation = column![
|
||||
Space::with_height(Length::Fill),
|
||||
column![
|
||||
text(&session.title).size(26),
|
||||
text("This session is ready for the agent runtime.").size(14),
|
||||
]
|
||||
.spacing(8),
|
||||
Space::with_height(Length::Fill),
|
||||
scrollable(messages).height(Length::Fill),
|
||||
container(
|
||||
column![
|
||||
text("Ask DS4Server anything…"),
|
||||
Space::with_height(28),
|
||||
composer,
|
||||
row![
|
||||
icon(ICON_PAPERCLIP, 19),
|
||||
Space::with_width(Length::Fill),
|
||||
@@ -246,7 +279,7 @@ impl App {
|
||||
.to_string(),
|
||||
)
|
||||
.size(12),
|
||||
action_button(icon(ICON_SEND, 18)).padding(8),
|
||||
action,
|
||||
]
|
||||
.align_y(Alignment::Center),
|
||||
]
|
||||
@@ -658,7 +691,8 @@ impl App {
|
||||
let panel = container(
|
||||
column![
|
||||
header,
|
||||
scrollable(container(fields).padding(iced::Padding::ZERO.right(18))),
|
||||
scrollable(container(fields).padding(iced::Padding::ZERO.right(18)))
|
||||
.height(Length::Fill),
|
||||
footer
|
||||
]
|
||||
.spacing(16),
|
||||
|
||||
267
src/engine.rs
267
src/engine.rs
@@ -1,12 +1,19 @@
|
||||
mod gguf;
|
||||
#[cfg(target_os = "macos")]
|
||||
mod metal;
|
||||
mod tokenizer;
|
||||
|
||||
use crate::model::ModelChoice;
|
||||
use crate::settings::TurnSettings;
|
||||
use crate::settings::{EngineSettings, ReasoningMode};
|
||||
use gguf::{F16, F32, Gguf, I32, IQ2_XXS, Q2_K, Q4_0, Q4_K, Q5_K, Q6_K, Q8_0, Tensor, Value};
|
||||
use std::path::Path;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use tokenizer::Tokenizer;
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
pub(crate) use metal::configure_sources as configure_metal_sources;
|
||||
|
||||
const DENSE: &[u32] = &[Q8_0, Q4_K, Q4_0];
|
||||
const ROUTED: &[u32] = &[Q8_0, IQ2_XXS, Q2_K, Q4_K, Q5_K, Q6_K];
|
||||
const PLAIN: &[u32] = &[F16, F32];
|
||||
@@ -236,6 +243,16 @@ impl Model {
|
||||
self.tokenizer.encode_chat(system, prompt, reasoning)
|
||||
}
|
||||
|
||||
fn render_conversation(
|
||||
&self,
|
||||
system: &str,
|
||||
messages: &[ChatTurn],
|
||||
reasoning: ReasoningMode,
|
||||
) -> Vec<i32> {
|
||||
self.tokenizer
|
||||
.encode_conversation(system, messages, reasoning)
|
||||
}
|
||||
|
||||
pub(crate) fn token_bytes(&self, token: i32) -> Option<Vec<u8>> {
|
||||
self.tokenizer.token_bytes(token)
|
||||
}
|
||||
@@ -253,6 +270,230 @@ impl Model {
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
pub(crate) struct Generator {
|
||||
executor: metal::Executor,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct ChatTurn {
|
||||
pub(crate) user: bool,
|
||||
pub(crate) reasoning: bool,
|
||||
pub(crate) content: String,
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
impl Generator {
|
||||
pub(crate) fn open(settings: &EngineSettings) -> Result<Self, String> {
|
||||
if settings.model != ModelChoice::DeepSeekV4Flash {
|
||||
return Err("local generation currently supports DeepSeek V4 Flash only".into());
|
||||
}
|
||||
if settings.speculative.dspark || settings.ssd.enabled || settings.steering.file.is_some() {
|
||||
return Err(
|
||||
"DSpark, SSD streaming, and steering are not yet available in the Rust executor"
|
||||
.into(),
|
||||
);
|
||||
}
|
||||
let model = Model::open(settings)?;
|
||||
let executor = metal::Executor::open(
|
||||
model,
|
||||
settings.context_tokens.max(1) as u32,
|
||||
settings.execution.quality,
|
||||
)?;
|
||||
Ok(Self { executor })
|
||||
}
|
||||
|
||||
pub(crate) fn generate(
|
||||
&mut self,
|
||||
messages: &[ChatTurn],
|
||||
settings: &TurnSettings,
|
||||
cancelled: &AtomicBool,
|
||||
mut emit: impl FnMut(String),
|
||||
) -> Result<(), String> {
|
||||
self.executor.reset()?;
|
||||
let tokens = self.executor.model().render_conversation(
|
||||
&settings.system_prompt,
|
||||
messages,
|
||||
settings.reasoning_mode,
|
||||
);
|
||||
if tokens.is_empty() {
|
||||
return Err("the rendered prompt is empty".into());
|
||||
}
|
||||
let max_context = self.executor.context() as usize;
|
||||
if tokens.len() >= max_context {
|
||||
return Err(format!(
|
||||
"the conversation uses {} tokens; the current Rust attention port supports fewer than {max_context}",
|
||||
tokens.len()
|
||||
));
|
||||
}
|
||||
let mut rng = Rng::new(settings.seed.unwrap_or(0x4453_3453_4552_5645));
|
||||
for token in tokens {
|
||||
if cancelled.load(Ordering::Relaxed) {
|
||||
return Ok(());
|
||||
}
|
||||
self.executor.eval(token)?;
|
||||
}
|
||||
for _ in 0..settings
|
||||
.max_generated_tokens
|
||||
.max(0)
|
||||
.min((max_context - self.executor.position() as usize) as i32)
|
||||
{
|
||||
if cancelled.load(Ordering::Relaxed) {
|
||||
return Ok(());
|
||||
}
|
||||
let token = sample(
|
||||
self.executor.logits(),
|
||||
settings.temperature,
|
||||
settings.top_p,
|
||||
settings.min_p,
|
||||
&mut rng,
|
||||
);
|
||||
if self.executor.model().is_stop_token(token) {
|
||||
return Ok(());
|
||||
}
|
||||
if let Some(bytes) = self.executor.model().token_bytes(token) {
|
||||
emit(String::from_utf8_lossy(&bytes).into_owned());
|
||||
}
|
||||
self.executor.eval(token)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
fn sample(logits: &[f32], temperature: f32, top_p: f32, min_p: f32, rng: &mut Rng) -> i32 {
|
||||
if temperature <= 0.0 {
|
||||
return logits
|
||||
.iter()
|
||||
.enumerate()
|
||||
.max_by(|a, b| a.1.total_cmp(b.1))
|
||||
.map_or(0, |(index, _)| index as i32);
|
||||
}
|
||||
let maximum = logits
|
||||
.iter()
|
||||
.copied()
|
||||
.filter(|value| value.is_finite())
|
||||
.fold(f32::NEG_INFINITY, f32::max);
|
||||
if !maximum.is_finite() {
|
||||
return 0;
|
||||
}
|
||||
let top_p = if top_p <= 0.0 || top_p > 1.0 {
|
||||
1.0
|
||||
} else {
|
||||
top_p
|
||||
};
|
||||
let min_p = min_p.max(0.0);
|
||||
let mut probabilities: Vec<(usize, f32)> = logits
|
||||
.iter()
|
||||
.enumerate()
|
||||
.filter(|(_, logit)| logit.is_finite())
|
||||
.map(|(index, logit)| (index, ((*logit - maximum) / temperature).exp()))
|
||||
.filter(|(_, probability)| *probability >= min_p)
|
||||
.collect();
|
||||
if probabilities.is_empty() {
|
||||
return logits
|
||||
.iter()
|
||||
.enumerate()
|
||||
.max_by(|a, b| a.1.total_cmp(b.1))
|
||||
.map_or(0, |(index, _)| index as i32);
|
||||
}
|
||||
if top_p < 1.0 {
|
||||
probabilities.sort_unstable_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
|
||||
let total: f32 = logits
|
||||
.iter()
|
||||
.filter(|logit| logit.is_finite())
|
||||
.map(|logit| ((*logit - maximum) / temperature).exp())
|
||||
.sum();
|
||||
let mut kept = 0.0;
|
||||
let count = probabilities
|
||||
.iter()
|
||||
.position(|(_, probability)| {
|
||||
kept += *probability;
|
||||
kept / total >= top_p
|
||||
})
|
||||
.map_or(probabilities.len(), |index| index + 1);
|
||||
probabilities.truncate(count);
|
||||
}
|
||||
let kept_total: f32 = probabilities.iter().map(|(_, p)| p).sum();
|
||||
let mut choice = rng.unit() * kept_total;
|
||||
for (token, probability) in &probabilities {
|
||||
choice -= probability;
|
||||
if choice <= 0.0 {
|
||||
return *token as i32;
|
||||
}
|
||||
}
|
||||
probabilities.last().map_or(0, |(token, _)| *token as i32)
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
struct Rng(u64);
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
impl Rng {
|
||||
fn new(seed: u64) -> Self {
|
||||
Self(seed.max(1))
|
||||
}
|
||||
|
||||
fn unit(&mut self) -> f32 {
|
||||
let mut value = self.0;
|
||||
if value == 0 {
|
||||
value = 0x9e37_79b9_7f4a_7c15;
|
||||
}
|
||||
value ^= value >> 12;
|
||||
value ^= value << 25;
|
||||
value ^= value >> 27;
|
||||
self.0 = value;
|
||||
let value = value.wrapping_mul(0x2545_f491_4f6c_dd1d);
|
||||
((value >> 40) & 0xff_ffff) as f32 / 16_777_216.0
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(all(test, target_os = "macos"))]
|
||||
mod sampling_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn zero_temperature_is_greedy() {
|
||||
let mut rng = Rng::new(1);
|
||||
assert_eq!(sample(&[1.0, 4.0, 2.0], 0.0, 1.0, 0.0, &mut rng), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires the 80 GiB Flash checkpoint and Apple Metal"]
|
||||
fn metal_executes_real_flash_token() {
|
||||
configure_metal_sources().unwrap();
|
||||
let path = Path::new(env!("CARGO_MANIFEST_DIR")).join(
|
||||
"../ds4/gguf/DeepSeek-V4-Flash-IQ2XXS-w2Q2K-AProjQ8-SExpQ8-OutQ8-chat-v2-imatrix.gguf",
|
||||
);
|
||||
let model = Model::open_main(&path, ModelChoice::DeepSeekV4Flash).unwrap();
|
||||
let tokens = model.render_prompt(
|
||||
"You are a helpful assistant",
|
||||
"Hello",
|
||||
ReasoningMode::Direct,
|
||||
);
|
||||
assert_eq!(tokens.len(), 10);
|
||||
let mut executor = metal::Executor::open(model, 128, false).unwrap();
|
||||
for token in tokens {
|
||||
executor.eval(token).unwrap();
|
||||
}
|
||||
assert!(executor.logits().iter().all(|logit| logit.is_finite()));
|
||||
let argmax = executor
|
||||
.logits()
|
||||
.iter()
|
||||
.enumerate()
|
||||
.max_by(|a, b| a.1.total_cmp(b.1))
|
||||
.unwrap();
|
||||
eprintln!(
|
||||
"Rust logits: argmax={} value={} logit0={}",
|
||||
argmax.0,
|
||||
argmax.1,
|
||||
executor.logits()[0]
|
||||
);
|
||||
assert_eq!(argmax.0, 19_923);
|
||||
assert!((executor.logits()[0] - -7.675_424).abs() < 0.1);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn validate_model_artifact(
|
||||
path: &Path,
|
||||
expected: ModelChoice,
|
||||
@@ -1144,6 +1385,32 @@ mod tests {
|
||||
model.render_prompt("", "Hello", ReasoningMode::Direct),
|
||||
[0, 128_803, 19_923, 128_804, 128_822]
|
||||
);
|
||||
assert_eq!(
|
||||
model.render_conversation(
|
||||
"",
|
||||
&[
|
||||
ChatTurn {
|
||||
user: true,
|
||||
reasoning: false,
|
||||
content: "Hello".into(),
|
||||
},
|
||||
ChatTurn {
|
||||
user: false,
|
||||
reasoning: false,
|
||||
content: "Hello".into(),
|
||||
},
|
||||
ChatTurn {
|
||||
user: true,
|
||||
reasoning: false,
|
||||
content: "Hello".into(),
|
||||
},
|
||||
],
|
||||
ReasoningMode::Direct,
|
||||
),
|
||||
[
|
||||
0, 128_803, 19_923, 128_804, 128_822, 19_923, 128_803, 19_923, 128_804, 128_822,
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -170,6 +170,10 @@ impl Gguf {
|
||||
self.map.len() as u64
|
||||
}
|
||||
|
||||
pub(super) fn map_ptr(&self) -> *const u8 {
|
||||
self.map.as_ptr()
|
||||
}
|
||||
|
||||
pub(super) fn tensor(&self, name: &str) -> Result<&Tensor, String> {
|
||||
self.tensors
|
||||
.get(name)
|
||||
|
||||
1729
src/engine/metal.rs
Normal file
1729
src/engine/metal.rs
Normal file
File diff suppressed because it is too large
Load Diff
@@ -1,5 +1,5 @@
|
||||
use super::ModelFamily;
|
||||
use super::gguf::Gguf;
|
||||
use super::{ChatTurn, ModelFamily};
|
||||
use crate::settings::ReasoningMode;
|
||||
use std::collections::HashMap;
|
||||
|
||||
@@ -131,6 +131,23 @@ impl Tokenizer {
|
||||
system_prompt: &str,
|
||||
prompt: &str,
|
||||
reasoning: ReasoningMode,
|
||||
) -> Vec<i32> {
|
||||
self.encode_conversation(
|
||||
system_prompt,
|
||||
&[ChatTurn {
|
||||
user: true,
|
||||
reasoning: false,
|
||||
content: prompt.to_owned(),
|
||||
}],
|
||||
reasoning,
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn encode_conversation(
|
||||
&self,
|
||||
system_prompt: &str,
|
||||
messages: &[ChatTurn],
|
||||
reasoning: ReasoningMode,
|
||||
) -> Vec<i32> {
|
||||
let mut output = vec![self.bos];
|
||||
if self.family == ModelFamily::Glm && self.sop >= 0 {
|
||||
@@ -156,8 +173,23 @@ impl Tokenizer {
|
||||
}
|
||||
output.extend(self.tokenize(system_prompt));
|
||||
}
|
||||
output.push(self.user);
|
||||
output.extend(self.tokenize(prompt));
|
||||
for message in messages {
|
||||
output.push(if message.user {
|
||||
self.user
|
||||
} else {
|
||||
self.assistant
|
||||
});
|
||||
if !message.user {
|
||||
if message.reasoning {
|
||||
output.push(self.think_start);
|
||||
} else if self.family == ModelFamily::Glm {
|
||||
output.extend([self.think_start, self.think_end]);
|
||||
} else {
|
||||
output.push(self.think_end);
|
||||
}
|
||||
}
|
||||
output.extend(self.tokenize(&message.content));
|
||||
}
|
||||
output.push(self.assistant);
|
||||
if reasoning != ReasoningMode::Direct {
|
||||
output.push(self.think_start);
|
||||
|
||||
@@ -11,6 +11,10 @@ use app::{App, Message, app_icon, app_theme};
|
||||
use iced::{Size, window};
|
||||
|
||||
fn main() -> iced::Result {
|
||||
#[cfg(target_os = "macos")]
|
||||
if let Err(error) = engine::configure_metal_sources() {
|
||||
eprintln!("DS4Server: {error}");
|
||||
}
|
||||
iced::daemon(App::title, App::update, App::view)
|
||||
.subscription(App::subscription)
|
||||
.theme(|_, _| app_theme())
|
||||
|
||||
@@ -668,7 +668,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_ds4_cli_option_is_mapped_or_explicitly_not_a_preference() {
|
||||
fn captured_ds4_cli_options_are_uniquely_classified() {
|
||||
const PREFERENCES: &[&str] = &[
|
||||
"--ctx",
|
||||
"--dir-steering-attn",
|
||||
@@ -753,32 +753,16 @@ mod tests {
|
||||
"--tensor-parallel-token-prefill",
|
||||
"--transport",
|
||||
];
|
||||
let source = [
|
||||
include_str!("../../ds4/ds4_cli.c"),
|
||||
include_str!("../../ds4/ds4_tp.c"),
|
||||
include_str!("../../ds4/ds4_distributed.c"),
|
||||
]
|
||||
.join("\n");
|
||||
let actual = option_branches(&source);
|
||||
let classified = PREFERENCES
|
||||
.iter()
|
||||
.chain(NON_PREFERENCES)
|
||||
.copied()
|
||||
.collect::<BTreeSet<_>>();
|
||||
assert_eq!(actual, classified);
|
||||
assert_eq!(classified.len(), PREFERENCES.len() + NON_PREFERENCES.len());
|
||||
assert!(
|
||||
PREFERENCES
|
||||
.iter()
|
||||
.all(|option| !NON_PREFERENCES.contains(option))
|
||||
);
|
||||
}
|
||||
|
||||
fn option_branches(source: &str) -> BTreeSet<&str> {
|
||||
source
|
||||
.split("!strcmp(arg, \"")
|
||||
.skip(1)
|
||||
.filter_map(|tail| tail.split('"').next())
|
||||
.filter(|option| option.starts_with("--"))
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user