diff --git a/kernels/build.rs b/kernels/build.rs index 0ba235a3..9d7e352a 100644 --- a/kernels/build.rs +++ b/kernels/build.rs @@ -30,7 +30,7 @@ fn main() -> Result<()> { builder.build_lib(build_dir.join("libpagedattention.a")); let kernel_dir = PathBuf::from("../kernels/"); - let absolute_kernel_dir = std::fs::canonicalize(&kernel_dir).unwrap(); + let absolute_kernel_dir = std::fs::canonicalize(&kernel_dir)?; println!( "cargo:rustc-link-search=native={}", diff --git a/src/backend/custom_ops/sort.rs b/src/backend/custom_ops/sort.rs index 1e7653b7..e0d34c17 100644 --- a/src/backend/custom_ops/sort.rs +++ b/src/backend/custom_ops/sort.rs @@ -36,7 +36,7 @@ impl candle::CustomOp1 for ArgSort { let dev = storage.device(); let elem_count = layout.shape().elem_count(); let ncols = self.last_dim as i32; - let nrows = (elem_count as i32 / ncols) as i32; + let nrows = elem_count as i32 / ncols; let dst = unsafe { dev.alloc::(elem_count) }.w()?; use std::ffi::c_void; diff --git a/src/backend/gptq.rs b/src/backend/gptq.rs index bd95bf73..c940ef1f 100644 --- a/src/backend/gptq.rs +++ b/src/backend/gptq.rs @@ -78,7 +78,7 @@ impl GPTQMatMul { let qzeros_ = qzeros_.slice(qzeros_l.start_offset()..); *qzeros_.device_ptr() as *const c_void } else { - std::ptr::null() as *const c_void + std::ptr::null() }; let g_idx_ptr = if self.g_idx.is_some() { @@ -91,7 +91,7 @@ impl GPTQMatMul { let g_idx_ = g_idx_.slice(g_idx_l.start_offset()..); *g_idx_.device_ptr() as *const c_void } else { - std::ptr::null() as *const c_void + std::ptr::null() }; unsafe { @@ -123,7 +123,7 @@ impl GPTQMatMul { size_k as i32, size_n as i32, workspace_ptr, - self.group_size as i32, + self.group_size, *dev.cu_stream() as i64, ); } else { @@ -138,7 +138,7 @@ impl GPTQMatMul { size_k as i32, //k size_n as i32, //n workspace_ptr, - self.group_size as i32, + self.group_size, *dev.cu_stream() as i64, ); } @@ -155,7 +155,7 @@ impl GPTQMatMul { size_k as i32, size_n as i32, workspace_ptr, - self.group_size as i32, + self.group_size, *dev.cu_stream() as i64, ); } else { @@ -170,7 +170,7 @@ impl GPTQMatMul { size_k as i32, //k size_n as i32, //n workspace_ptr, - self.group_size as i32, + self.group_size, *dev.cu_stream() as i64, ); } @@ -282,14 +282,14 @@ impl MarlinRepack { //in_dim 4096, out_dim 1024 (/pack_factor) //ws shape [4096, 128] //out_shape [256, 2048] - out_shape[0] = (q_shape[0] / pack_factor / 2) as usize; - out_shape[1] = (q_shape[1] * pack_factor * 2) as usize; + out_shape[0] = q_shape[0] / pack_factor / 2; + out_shape[1] = q_shape[1] * pack_factor * 2; } else { //in_dim 4096 (/pack_factor), out_dim 1024 //ws shape [512, 1024] //out_shape [256, 2048] - out_shape[0] = (q_shape[0] / 2) as usize; - out_shape[1] = (q_shape[1] * 2) as usize; + out_shape[0] = q_shape[0] / 2; + out_shape[1] = q_shape[1] * 2; } let oshape: Shape = out_shape.into(); diff --git a/src/backend/paged_attention.rs b/src/backend/paged_attention.rs index f5c17ec7..a004b050 100644 --- a/src/backend/paged_attention.rs +++ b/src/backend/paged_attention.rs @@ -156,7 +156,7 @@ impl PagedAttention { let kv_head_stride = kc_l.stride()[1]; let partition_size = 512; - let max_num_partitions = (self.max_context_len + partition_size - 1) / partition_size; + let max_num_partitions = self.max_context_len.div_ceil(partition_size); let use_v1 = (max_num_partitions == 1 || num_seqs * num_heads > 512) && partition_size % block_size == 0; diff --git a/src/backend/progress.rs b/src/backend/progress.rs index a3479ad6..0be0fe03 100644 --- a/src/backend/progress.rs +++ b/src/backend/progress.rs @@ -54,13 +54,13 @@ impl Progress { let pb = m.add(ProgressBar::new(size as u64)); pb.set_style(sty.clone()); if n > 1 { - pb.set_message(format!("On Rank {} Device", i)); + pb.set_message(format!("On Rank {i} Device")); } bars.push(pb); } if n > 1 { - m.println(format!("Loading model in {} ranks!", n)).unwrap(); + m.println(format!("Loading model in {n} ranks!")).unwrap(); } Self { m, bars, size } } @@ -71,9 +71,9 @@ impl Progress { self.bars[idx].inc(progress as u64 - pos); if self.bars.len() > 1 { if progress >= self.size { - self.bars[idx].set_message(format!("On Rank {} Device Finished", idx)); + self.bars[idx].set_message(format!("On Rank {idx} Device Finished")); } else { - self.bars[idx].set_message(format!("On Rank {} Device", idx)); + self.bars[idx].set_message(format!("On Rank {idx} Device")); } } } @@ -84,7 +84,7 @@ impl Progress { let pos = self.bars[idx].position(); self.bars[idx].inc(self.size as u64 - pos); if self.bars.len() > 1 { - self.bars[idx].set_message(format!("On Rank {} Device Finished", idx)); + self.bars[idx].set_message(format!("On Rank {idx} Device Finished")); } } self.m.clear().unwrap(); @@ -134,7 +134,7 @@ pub async fn progress_worker( #[cfg(not(feature = "nccl"))] let progress_bar = Some(Progress::new(1, length)); - let _ = thread::sleep(time::Duration::from_millis(1000 as u64)); + let _ = thread::sleep(time::Duration::from_millis(1000_u64)); loop { { @@ -181,7 +181,7 @@ pub async fn progress_worker( } } - let _ = thread::sleep(time::Duration::from_millis(500 as u64)); + let _ = thread::sleep(time::Duration::from_millis(500_u64)); } }); handle diff --git a/src/lib.rs b/src/lib.rs index c255192c..c1ecd4ec 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -864,7 +864,7 @@ pub fn hub_load_local_safetensors( pub fn new_device(ordinal: usize) -> Result { if cuda_is_available() { use candle_core::CudaDevice; - let device = Device::Cuda(CudaDevice::new_with_stream(ordinal).unwrap()); + let device = Device::Cuda(CudaDevice::new_with_stream(ordinal)?); Ok(device) } else if metal_is_available() { Ok(Device::new_metal(ordinal)?) diff --git a/src/main.rs b/src/main.rs index ed7ee5ba..fc6bee57 100644 --- a/src/main.rs +++ b/src/main.rs @@ -123,7 +123,7 @@ fn get_cache_config( / config.num_hidden_layers / 2; CacheConfig { - block_size: block_size, + block_size, num_gpu_blocks: Some(num_gpu_blocks), num_cpu_blocks: Some(num_cpu_blocks), fully_init: true, @@ -152,8 +152,8 @@ fn config_log( LevelFilter::Trace, ]; let level = level.to_uppercase(); - for (i, name) in log_level_names.to_vec().into_iter().enumerate() { - if level.find(name).is_some() { + for (i, name) in log_level_names.iter().copied().enumerate() { + if level.contains(name) { cfg_filter = log_levels[i] } } @@ -207,9 +207,9 @@ async fn main() -> Result<(), APIError> { filenames: { let path = path.clone().unwrap_or("".to_string()); if Path::new(&path).join(file).exists() { - vec![Path::new(&path).join(file).into()] + vec![Path::new(&path).join(file)] } else { - panic!("Model file not found {}", file); + panic!("Model file not found {file}"); } }, }, @@ -350,7 +350,7 @@ async fn main() -> Result<(), APIError> { #[cfg(not(feature = "nccl"))] let (pipelines, global_rank) = { - let log_file = format!("candle-vllm.log"); + let log_file = "candle-vllm.log".to_string(); let _ = config_log(logger, args.log, log_file); ( loader @@ -361,7 +361,7 @@ async fn main() -> Result<(), APIError> { }; let (default_pipelines, pipeline_config) = match pipelines { - Err(e) => panic!("{:?}", e), + Err(e) => panic!("{e:?}"), Ok((p, c)) => (p, c), }; let mut config: Option = None; @@ -382,7 +382,7 @@ async fn main() -> Result<(), APIError> { &cfg, &cache_cfg, cache_cfg.dtype, - &pipeline.device(), + pipeline.device(), num_shards, ) .unwrap(); @@ -458,7 +458,7 @@ async fn main() -> Result<(), APIError> { .route("/v1/chat/completions", post(chat_completions)) .with_state(Arc::new(server_data)); - let listener = tokio::net::TcpListener::bind(format!("0.0.0.0:{}", port)) + let listener = tokio::net::TcpListener::bind(format!("0.0.0.0:{port}")) .await .map_err(|e| APIError::new(e.to_string()))?; axum::serve(listener, app) diff --git a/src/openai/conversation/default_conversation.rs b/src/openai/conversation/default_conversation.rs index ebc7e3b2..f0df31ba 100644 --- a/src/openai/conversation/default_conversation.rs +++ b/src/openai/conversation/default_conversation.rs @@ -161,7 +161,7 @@ impl Conversation for DefaultConversation { env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback); let template = self.chat_template.as_ref().unwrap(); let mut template = template.replace("[::-1]", "|reverse"); - if template.find("{{ meta }}").is_some() { + if template.contains("{{ meta }}") { template = template.replace("{%- set meta = message.get(\"metadata\", \"\") %}", ""); template = template.replace("{{ meta }}", ""); } @@ -198,11 +198,11 @@ impl Conversation for DefaultConversation { tracing::warn!("apply chat template failed {:?}", e); } //no chat template exists? using the built-in template - let system_prompt = if self.system_message.is_some() { - format!("<|system|>\n {}", self.system_message.clone().unwrap()) - } else { - "".to_string() - }; + let system_prompt = self + .system_message + .as_ref() + .map_or("".to_string(), |msg| format!("<|system|>\n {msg}")); + match self.sep_style { SeparatorStyle::AddColonSingle | SeparatorStyle::AddColonSpaceSingle @@ -441,7 +441,7 @@ impl Conversation for DefaultConversation { SeparatorStyle::GLM => { let mut accum = "[gMASK]".to_string(); accum += &system_prompt.clone(); - for (_, message) in self.messages.iter().enumerate() { + for message in self.messages.iter() { if message.role.clone() == self.roles.0 { //user message accum += &format!("<|user|>\n {}", message.content); diff --git a/src/openai/distributed.rs b/src/openai/distributed.rs index ebb78dd2..e9058683 100644 --- a/src/openai/distributed.rs +++ b/src/openai/distributed.rs @@ -418,7 +418,7 @@ impl ReplicatedLinear { } pub fn forward(&self, x: &Tensor) -> Result { - let mut xs = self.linear.forward(&x)?; + let mut xs = self.linear.forward(x)?; if let Some(bias) = &self.bias { xs = xs.broadcast_add(bias)?; } diff --git a/src/openai/logits_processor.rs b/src/openai/logits_processor.rs index b5d992cc..858ec276 100644 --- a/src/openai/logits_processor.rs +++ b/src/openai/logits_processor.rs @@ -47,17 +47,9 @@ impl LogitsProcessor { top_p: Option, ) -> Sampling { let temperature = temperature.and_then(|v| if v < 1e-7 { None } else { Some(v) }); - let top_k: Option = if top_k.is_some() && top_k.unwrap() > 0 { - Some(top_k.unwrap() as usize) - } else { - None - }; + let top_k: Option = top_k.filter(|&k| k > 0).map(|k| k as usize); - let temperature: Option = if temperature.is_some() && temperature.unwrap() > 0. { - Some(temperature.unwrap()) - } else { - None - }; + let temperature: Option = temperature.filter(|&t| t > 0.0); match (temperature, top_k, top_p) { (None, _, _) => Sampling::ArgMax, @@ -182,18 +174,16 @@ impl LogitsProcessor { Ok(prs) }; - let sampling = if sampling_params.is_some() { - let param = sampling_params.as_ref().unwrap(); - LogitsProcessor::get_strategy(param.temperature, param.top_k, param.top_p) - } else { - self.sampling.to_owned() - }; + let sampling = sampling_params.as_ref().map_or_else( + || self.sampling.to_owned(), + |param| LogitsProcessor::get_strategy(param.temperature, param.top_k, param.top_p), + ); + let next_tokens = match &sampling { Sampling::ArgMax => self.sample_argmax(&logits)?, Sampling::All { temperature } => { let prs = prs(*temperature as f64)?.to_vec2()?; (0..batch) - .into_iter() .map(|b| self.sample_multinomial(&prs[b]).unwrap()) .collect() } @@ -203,7 +193,6 @@ impl LogitsProcessor { // simply sample from the predicted probability distribution let prs = prs.to_vec2()?; (0..batch) - .into_iter() .map(|b| self.sample_multinomial(&prs[b]).unwrap()) .collect() } else { diff --git a/src/openai/models/gemma.rs b/src/openai/models/gemma.rs index cab1c6ee..00924075 100644 --- a/src/openai/models/gemma.rs +++ b/src/openai/models/gemma.rs @@ -138,26 +138,23 @@ impl RotaryEmbedding { let sin = self.sin.narrow(0, seqlen_offset[0], seq_len)?; let x_q = q.narrow(0, b, 1)?; let x_k = k.narrow(0, b, 1)?; - let q_embed = candle_nn::rotary_emb::rope(&x_q, &cos, &sin).unwrap(); - let k_embed = candle_nn::rotary_emb::rope(&x_k, &cos, &sin).unwrap(); + let q_embed = candle_nn::rotary_emb::rope(&x_q, &cos, &sin)?; + let k_embed = candle_nn::rotary_emb::rope(&x_k, &cos, &sin)?; q_embeds.push(q_embed); k_embeds.push(k_embed); } - Ok(( - Tensor::cat(&q_embeds, 0).unwrap(), - Tensor::cat(&k_embeds, 0).unwrap(), - )) + Ok((Tensor::cat(&q_embeds, 0)?, Tensor::cat(&k_embeds, 0)?)) } } -struct MLP { +struct Mlp { gate_proj: TensorParallelColumnLinear, up_proj: TensorParallelColumnLinear, down_proj: TensorParallelRowLinear, act_fn: candle_nn::Activation, } -impl MLP { +impl Mlp { fn new(cfg: &Config, vb: VarBuilder, comm: Rc) -> Result { let hidden_sz = cfg.hidden_size; let intermediate_sz = cfg.intermediate_size; @@ -197,7 +194,7 @@ impl MLP { } } -impl Module for MLP { +impl Module for Mlp { fn forward(&self, xs: &Tensor) -> Result { let lhs = self.act_fn.forward(&self.gate_proj.forward(xs)?)?; let rhs = self.up_proj.forward(xs)?; @@ -359,7 +356,7 @@ impl Attention { struct DecoderLayer { self_attn: Attention, - mlp: MLP, + mlp: Mlp, input_layernorm: RmsNorm, post_feedforward_layernorm: Option, pre_feedforward_layernorm: Option, @@ -374,7 +371,7 @@ impl DecoderLayer { comm: Rc, ) -> Result { let self_attn = Attention::new(rotary_emb, cfg, vb.pp("self_attn"), comm.clone())?; - let mlp = MLP::new(cfg, vb.pp("mlp"), comm.clone())?; + let mlp = Mlp::new(cfg, vb.pp("mlp"), comm.clone())?; let input_layernorm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?; diff --git a/src/openai/models/gemma3.rs b/src/openai/models/gemma3.rs index 751fb770..d0e3b6b6 100644 --- a/src/openai/models/gemma3.rs +++ b/src/openai/models/gemma3.rs @@ -105,20 +105,17 @@ impl Gemma3Config { kv_cache_dtype: DType, scfg: &SpecificConfig, ) -> Config { - let bos_token_id = if self.text_config.bos_token_id.is_some() { - self.text_config.bos_token_id.unwrap() - } else if self.bos_token_id.is_some() { - self.bos_token_id.unwrap() - } else { - super::TokenID(Either::Left(Some(2))) - }; - let eos_token_id = if self.text_config.eos_token_id.is_some() { - self.text_config.eos_token_id.unwrap() - } else if self.eos_token_id.is_some() { - self.eos_token_id.unwrap() - } else { - super::TokenID(Either::Left(Some(1))) - }; + let bos_token_id = self + .text_config + .bos_token_id + .or(self.bos_token_id) + .unwrap_or(super::TokenID(Either::Left(Some(2)))); + + let eos_token_id = self + .text_config + .eos_token_id + .or(self.eos_token_id) + .unwrap_or(super::TokenID(Either::Left(Some(1)))); let ropescaling = if self.text_config.rope_scaling.is_some() { let mut ropescaling = HashMap::::new(); @@ -194,11 +191,10 @@ impl RotaryEmbedding { cfg: &Config, dev: &Device, ) -> Result<(Tensor, Tensor)> { - let rope_freq = if local_sliding_window.is_some() && cfg.rope_local_base_freq.is_some() { - cfg.rope_local_base_freq.unwrap() - } else { - cfg.rope_theta - }; + let rope_freq = local_sliding_window + .and(cfg.rope_local_base_freq) + .unwrap_or(cfg.rope_theta); + let dim = cfg .head_dim .unwrap_or(cfg.hidden_size / cfg.num_attention_heads); @@ -283,26 +279,23 @@ impl RotaryEmbedding { let x_q = q.narrow(0, b, 1)?; let x_k = k.narrow(0, b, 1)?; - let q_embed = candle_nn::rotary_emb::rope(&x_q, &cos, &sin).unwrap(); - let k_embed = candle_nn::rotary_emb::rope(&x_k, &cos, &sin).unwrap(); + let q_embed = candle_nn::rotary_emb::rope(&x_q, &cos, &sin)?; + let k_embed = candle_nn::rotary_emb::rope(&x_k, &cos, &sin)?; q_embeds.push(q_embed); k_embeds.push(k_embed); } - Ok(( - Tensor::cat(&q_embeds, 0).unwrap(), - Tensor::cat(&k_embeds, 0).unwrap(), - )) + Ok((Tensor::cat(&q_embeds, 0)?, Tensor::cat(&k_embeds, 0)?)) } } -struct MLP { +struct Mlp { gate_proj: TensorParallelColumnLinear, up_proj: TensorParallelColumnLinear, down_proj: TensorParallelRowLinear, act_fn: candle_nn::Activation, } -impl MLP { +impl Mlp { fn new(cfg: &Config, vb: VarBuilder, comm: Rc) -> Result { let hidden_sz = cfg.hidden_size; let intermediate_sz = cfg.intermediate_size; @@ -342,7 +335,7 @@ impl MLP { } } -impl Module for MLP { +impl Module for Mlp { fn forward(&self, xs: &Tensor) -> Result { let lhs = self.act_fn.forward(&self.gate_proj.forward(xs)?)?; let rhs = self.up_proj.forward(xs)?; @@ -506,7 +499,7 @@ impl Attention { struct DecoderLayer { self_attn: Attention, - mlp: MLP, + mlp: Mlp, input_layernorm: RmsNorm, post_feedforward_layernorm: RmsNorm, pre_feedforward_layernorm: RmsNorm, @@ -529,7 +522,7 @@ impl DecoderLayer { rotary_emb.clone(), sliding_window, )?; - let mlp = MLP::new(cfg, vb.pp("mlp"), comm.clone())?; + let mlp = Mlp::new(cfg, vb.pp("mlp"), comm.clone())?; let input_layernorm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?; diff --git a/src/openai/models/glm4.rs b/src/openai/models/glm4.rs index cc2b842f..87213fec 100644 --- a/src/openai/models/glm4.rs +++ b/src/openai/models/glm4.rs @@ -169,7 +169,7 @@ impl RotaryEmbedding { .unsqueeze(0)? .contiguous()?; let xs_pass = xs.i((b, .., .., self.rotary_dim..))?.unsqueeze(0)?; - let xs_rot = candle_nn::rotary_emb::rope_i(&xs_rot, &cos, &sin).unwrap(); + let xs_rot = candle_nn::rotary_emb::rope_i(&xs_rot, &cos, &sin)?; let embed = Tensor::cat(&[&xs_rot, &xs_pass], D::Minus1)?.contiguous()?; embeds.push(embed); } @@ -417,7 +417,7 @@ impl DecoderLayer { input_metadata: &InputMetadata, ) -> Result { let residual = xs; - let hidden_states = self.input_layernorm.forward(&xs)?; + let hidden_states = self.input_layernorm.forward(xs)?; let hidden_states = self.self_attn.forward( &hidden_states, attention_mask, diff --git a/src/openai/models/linear.rs b/src/openai/models/linear.rs index 45897cfe..665f8398 100644 --- a/src/openai/models/linear.rs +++ b/src/openai/models/linear.rs @@ -183,20 +183,11 @@ pub fn qlinear( ) -> Result { match quant_config { Some(cfg) => { - let marlin_compatible = if (cfg.quant_method != "gptq" && cfg.quant_method != "awq") - || (cfg.bits != 4 && cfg.bits != 8) - { - false - } else { - true - }; - let marlin_format = if cfg.checkpoint_format.is_some() - && cfg.checkpoint_format.as_ref().unwrap() == "marlin" - { - true - } else { - false - }; + let marlin_compatible = (cfg.quant_method != "gptq" && cfg.quant_method != "awq") + || (cfg.bits != 4 && cfg.bits != 8); + let marlin_format = cfg.checkpoint_format.is_some() + && cfg.checkpoint_format.as_ref().unwrap() == "marlin"; + let ws = vb.get_with_hints_dtype( if cfg.quant_method == "gptq" { //quantized gptq (k/pack_factor, n) format @@ -293,9 +284,9 @@ pub fn qlinear( if (cfg.sym.is_some() && !cfg.sym.unwrap()) || cfg.bits != 4 - || (cfg.group_size != 64 && cfg.group_size != 128 && cfg.group_size != -1) + || !matches!(cfg.group_size, 64 | 128 | -1) || (cfg.desc_act.is_some() - && cfg.desc_act.unwrap() == true + && cfg.desc_act.unwrap() && cfg.quant_method == "gptq") { //only model with 4-bit and desc_act==false can be repacked to marlin format @@ -362,8 +353,8 @@ pub fn qlinear( let scales = if marlin_compatible { marlin_permute_scales( &scales, - in_dim_partition as usize, - out_dim_partition as usize, + in_dim_partition, + out_dim_partition, cfg.group_size, cfg.bits as u32, )? @@ -517,7 +508,7 @@ impl QLinear { ); QLinear::from_linear( linear, - cfg.group_size as i32, + cfg.group_size, cfg.bits as i32, cfg.quant_method == "awq", ) @@ -661,7 +652,7 @@ impl QLinear { } else { let wdim = w.dims()[w.dims().len() - 1]; x.reshape((bsize * seq_len, dim1, dim2))? - .matmul(&w)? + .matmul(w)? .reshape((bsize, seq_len, dim1, wdim))? } } @@ -672,11 +663,11 @@ impl QLinear { } else { let wdim = w.dims()[w.dims().len() - 1]; x.reshape((bsize * seq_len, dim))? - .matmul(&w)? + .matmul(w)? .reshape((bsize, seq_len, wdim))? } } - _ => x.matmul(&w)?, + _ => x.matmul(w)?, }; // let x = x.to_dtype(DType::F16)?; if let Some(bias) = &self.bias { @@ -715,7 +706,7 @@ impl Module for QLinear { o.reshape((bsize, seq_len, dim1, ()))? } [_, _, _] => gptq_matmul( - &x, + x, qw, scale, qzeros, @@ -810,14 +801,14 @@ pub fn linear_x( dtype: DType, ) -> Result { if let Some(quantized_type) = quant { - let ln = qlinear(in_dim, out_dim, vb, shard, quant_config, true, dtype).unwrap(); + let ln = qlinear(in_dim, out_dim, vb, shard, quant_config, true, dtype)?; Ok(LinearX(Either::Right(QLinear::from_linear_x( ln, quantized_type.clone(), quant_config, )))) } else { - let ln = linear(in_dim, out_dim, vb, shard).unwrap(); + let ln = linear(in_dim, out_dim, vb, shard)?; Ok(LinearX(Either::Left(ln))) } } @@ -839,14 +830,10 @@ pub fn linear_no_bias_x( out_dim, vb, shard( - if shards.world_size < 2 { + if shards.world_size < 2 || shards.dim == 1 { 0 } else { - if shards.dim == 1 { - 0 - } else { - 1 - } + 1 }, shards.rank, shards.world_size, diff --git a/src/openai/models/llama.rs b/src/openai/models/llama.rs index 18625c6c..9bbc61b7 100644 --- a/src/openai/models/llama.rs +++ b/src/openai/models/llama.rs @@ -447,7 +447,7 @@ impl Llama { let blocks: Vec<_> = (0..cfg.num_hidden_layers) .map(|i| { let b = Block::load( - vb.pp(&format!("model.layers.{i}")), + vb.pp(format!("model.layers.{i}")), cfg, dtype, device, diff --git a/src/openai/models/mistral.rs b/src/openai/models/mistral.rs index 626e04b5..8cd2c20b 100644 --- a/src/openai/models/mistral.rs +++ b/src/openai/models/mistral.rs @@ -117,26 +117,23 @@ impl RotaryEmbedding { let sin = self.sin.narrow(0, seqlen_offset[0], seq_len)?; let x_q = q.narrow(0, b, 1)?; let x_k = k.narrow(0, b, 1)?; - let q_embed = candle_nn::rotary_emb::rope(&x_q, &cos, &sin).unwrap(); - let k_embed = candle_nn::rotary_emb::rope(&x_k, &cos, &sin).unwrap(); + let q_embed = candle_nn::rotary_emb::rope(&x_q, &cos, &sin)?; + let k_embed = candle_nn::rotary_emb::rope(&x_k, &cos, &sin)?; q_embeds.push(q_embed); k_embeds.push(k_embed); } - Ok(( - Tensor::cat(&q_embeds, 0).unwrap(), - Tensor::cat(&k_embeds, 0).unwrap(), - )) + Ok((Tensor::cat(&q_embeds, 0)?, Tensor::cat(&k_embeds, 0)?)) } } -struct MLP { +struct Mlp { gate_proj: TensorParallelColumnLinear, up_proj: TensorParallelColumnLinear, down_proj: TensorParallelRowLinear, act_fn: Activation, } -impl MLP { +impl Mlp { fn new(cfg: &Config, vb: VarBuilder, comm: Rc) -> Result { let hidden_sz = cfg.hidden_size; let intermediate_sz = cfg.intermediate_size; @@ -176,7 +173,7 @@ impl MLP { } } -impl Module for MLP { +impl Module for Mlp { fn forward(&self, xs: &Tensor) -> Result { let lhs = self.act_fn.forward(&self.gate_proj.forward(xs)?)?; let rhs = self.up_proj.forward(xs)?; @@ -330,7 +327,7 @@ impl Attention { struct DecoderLayer { self_attn: Attention, - mlp: MLP, + mlp: Mlp, input_layernorm: RmsNorm, post_attention_layernorm: RmsNorm, } @@ -343,7 +340,7 @@ impl DecoderLayer { comm: Rc, ) -> Result { let self_attn = Attention::new(rotary_emb, cfg, vb.pp("self_attn"), comm.clone())?; - let mlp = MLP::new(cfg, vb.pp("mlp"), comm.clone())?; + let mlp = Mlp::new(cfg, vb.pp("mlp"), comm.clone())?; let input_layernorm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?; let post_attention_layernorm = rms_norm( diff --git a/src/openai/models/mod.rs b/src/openai/models/mod.rs index e6f89652..a2f9b84b 100644 --- a/src/openai/models/mod.rs +++ b/src/openai/models/mod.rs @@ -242,9 +242,7 @@ pub fn get_attention_casual_mask( .unwrap() .to_dtype(dtype) .ok(), - _ => { - return None; - } + _ => None, } } @@ -295,8 +293,8 @@ impl NaiveAttention { } let mut cache = self.kv_cache.borrow_mut(); let (k, v) = match &mut *cache { - KvCache::Normal(c) => c.append(&k, &v)?, - KvCache::Rotating(c) => c.append(&k, &v)?, + KvCache::Normal(c) => c.append(k, v)?, + KvCache::Rotating(c) => c.append(k, v)?, }; let k = candle_transformers::utils::repeat_kv(k, self.num_kv_groups)?.contiguous()?; diff --git a/src/openai/models/phi2.rs b/src/openai/models/phi2.rs index 17385aaa..861078c9 100644 --- a/src/openai/models/phi2.rs +++ b/src/openai/models/phi2.rs @@ -89,7 +89,7 @@ struct RotaryEmbedding { impl RotaryEmbedding { fn new(cfg: &Config, _dtype: DType, dev: &Device) -> Result { let head_dim = cfg.hidden_size / cfg.num_attention_heads; - let dim = (cfg.partial_rotary_factor.unwrap() * head_dim as f32) as usize; + let dim = (cfg.partial_rotary_factor.unwrap_or(1.0) * head_dim as f32) as usize; let inv_freq: Vec<_> = (0..dim) .step_by(2) .map(|i| (1f64 / cfg.rope_theta.powf(i as f64 / dim as f64)) as f32) @@ -123,13 +123,13 @@ impl RotaryEmbedding { } } -struct MLP { +struct Mlp { fc1: TensorParallelColumnLinear, fc2: TensorParallelRowLinear, act: Activation, } -impl MLP { +impl Mlp { fn new(cfg: &Config, vb: VarBuilder, comm: Rc) -> Result { let fc1 = TensorParallelColumnLinear::load_with_hints( cfg.hidden_size, @@ -159,7 +159,7 @@ impl MLP { } } -impl Module for MLP { +impl Module for Mlp { fn forward(&self, xs: &Tensor) -> Result { self.fc2.forward(&self.act.forward(&self.fc1.forward(xs)?)?) } @@ -326,14 +326,14 @@ impl Attention { struct DecoderLayer { self_attn: Attention, - mlp: MLP, + mlp: Mlp, input_layernorm: LayerNorm, } impl DecoderLayer { fn new(cfg: &Config, vb: VarBuilder, comm: Rc) -> Result { let self_attn = Attention::new(cfg, vb.pp("self_attn"), comm.clone())?; - let mlp = MLP::new(cfg, vb.pp("mlp"), comm.clone())?; + let mlp = Mlp::new(cfg, vb.pp("mlp"), comm.clone())?; let input_layernorm = layer_norm( cfg.hidden_size, cfg.rms_norm_eps, diff --git a/src/openai/models/phi3.rs b/src/openai/models/phi3.rs index 717059d0..1b5b657b 100644 --- a/src/openai/models/phi3.rs +++ b/src/openai/models/phi3.rs @@ -230,10 +230,7 @@ impl RotaryEmbedding { q_embeds.push(q_embed); k_embeds.push(k_embed); } - Ok(( - Tensor::cat(&q_embeds, 0).unwrap(), - Tensor::cat(&k_embeds, 0).unwrap(), - )) + Ok((Tensor::cat(&q_embeds, 0)?, Tensor::cat(&k_embeds, 0)?)) } } @@ -407,7 +404,7 @@ impl Mlp { impl Module for Mlp { fn forward(&self, xs: &Tensor) -> Result { - let up_states = self.gate_up_proj.forward(&xs)?; + let up_states = self.gate_up_proj.forward(xs)?; let gate = up_states.narrow(D::Minus1, 0, self.i_size)?; let up_states = up_states.narrow(D::Minus1, self.i_size, self.i_size)?; let up_states = (up_states * gate.apply(&self.act_fn))?; diff --git a/src/openai/models/quantized_glm4.rs b/src/openai/models/quantized_glm4.rs index 549feddc..cea6c143 100644 --- a/src/openai/models/quantized_glm4.rs +++ b/src/openai/models/quantized_glm4.rs @@ -172,7 +172,7 @@ impl GGUFGLM4 { use_flash_attn: false, bos_token_id: super::TokenID(Either::Left(None)), eos_token_id: super::TokenID(Either::Left(None)), - max_seq_len: max_seq_len, + max_seq_len, sliding_window: None, sliding_window_pattern: None, hidden_act: None, @@ -349,10 +349,10 @@ impl GGUFGLM4 { head_dim, rotary_emb: rotary_emb.clone(), attn: PagedAttention::new( - head_count as usize, + head_count, head_dim, 1. / ((head_dim as f32).sqrt()), - Some(head_count_kv as usize), + Some(head_count_kv), None, device.clone(), None, diff --git a/src/openai/models/quantized_llama.rs b/src/openai/models/quantized_llama.rs index e44014a3..fe969985 100644 --- a/src/openai/models/quantized_llama.rs +++ b/src/openai/models/quantized_llama.rs @@ -145,10 +145,7 @@ impl LayerWeights { q_embeds.push(q_embed); k_embeds.push(k_embed); } - Ok(( - Tensor::cat(&q_embeds, 0).unwrap(), - Tensor::cat(&k_embeds, 0).unwrap(), - )) + Ok((Tensor::cat(&q_embeds, 0)?, Tensor::cat(&k_embeds, 0)?)) } fn forward_attn( @@ -267,7 +264,7 @@ impl GGUFLLaMa { use_flash_attn: false, bos_token_id: super::TokenID(Either::Left(Some(128256))), eos_token_id: super::TokenID(Either::Left(Some(128257))), - max_seq_len: max_seq_len, + max_seq_len, sliding_window: None, sliding_window_pattern: None, hidden_act: None, @@ -295,13 +292,8 @@ impl GGUFLLaMa { s_cfg: SpecificConfig, ) -> Result { let head_dim = (ct.hparams.n_embd / ct.hparams.n_head) as usize; - let (cos, sin) = precomput_freqs_cis( - head_dim, - 10000., - MAX_SEQ_LEN as usize, - &ct.device, - DType::F32, - )?; + let (cos, sin) = + precomput_freqs_cis(head_dim, 10000., MAX_SEQ_LEN, &ct.device, DType::F32)?; let tok_embeddings = ct.remove("tok_embeddings.weight")?; let tok_embeddings = tok_embeddings.dequantize(&ct.device)?; let norm = RmsNorm::from_qtensor(ct.remove("norm.weight")?, 1e-5)?; @@ -407,16 +399,13 @@ impl GGUFLLaMa { let embedding_length = md_get("llama.embedding_length")?.to_u32()? as usize; // let rope_dim = md_get("llama.rope.dimension_count")?.to_u32()? as usize; let context_length = md_get("llama.context_length")?.to_u32(); - let context_length = if context_length.is_ok() { - context_length.unwrap() as usize - } else { - MAX_SEQ_LEN as usize - }; + let context_length = context_length.map_or(MAX_SEQ_LEN, |v| v as usize); + let head_dim = md_get("llama.attention.key_length"); let head_dim = if head_dim.is_ok() { head_dim.unwrap().to_u32()? as usize } else { - (embedding_length / head_count) as usize + embedding_length / head_count }; // Strangely this value is generally 1e-6 in GGUF file but used to be 1e-5 by default. diff --git a/src/openai/models/quantized_phi3.rs b/src/openai/models/quantized_phi3.rs index 0abd596b..c05f9808 100644 --- a/src/openai/models/quantized_phi3.rs +++ b/src/openai/models/quantized_phi3.rs @@ -92,10 +92,7 @@ impl LayerWeights { q_embeds.push(q_embed); k_embeds.push(k_embed); } - Ok(( - Tensor::cat(&q_embeds, 0).unwrap(), - Tensor::cat(&k_embeds, 0).unwrap(), - )) + Ok((Tensor::cat(&q_embeds, 0)?, Tensor::cat(&k_embeds, 0)?)) } fn forward_attn( @@ -219,7 +216,7 @@ impl GGUFPhi3 { use_flash_attn: false, bos_token_id: super::TokenID(Either::Left(Some(1))), eos_token_id: super::TokenID(Either::Left(Some(2))), - max_seq_len: max_seq_len, + max_seq_len, sliding_window: None, sliding_window_pattern: None, hidden_act: None, @@ -277,13 +274,13 @@ impl GGUFPhi3 { let tok_embeddings = ct.tensor(reader, "token_embd.weight", device)?; let tok_embeddings = tok_embeddings.dequantize(device)?; let output_norm = rms_norm(ct.tensor(reader, "output_norm.weight", device)?, rms_eps)?; - let output = QLinear::new(&ct, reader, "output", device)?; + let output = QLinear::new(ct, reader, "output", device)?; let mut layers = Vec::with_capacity(block_count); for layer_idx in 0..block_count { let prefix = format!("blk.{layer_idx}"); - let ffn_up = QLinear::new(&ct, reader, &format!("{prefix}.ffn_up"), device)?; - let ffn_down = QLinear::new(&ct, reader, &format!("{prefix}.ffn_down"), device)?; + let ffn_up = QLinear::new(ct, reader, &format!("{prefix}.ffn_up"), device)?; + let ffn_down = QLinear::new(ct, reader, &format!("{prefix}.ffn_down"), device)?; let mlp = Mlp { ffn_up, ffn_down, @@ -298,8 +295,8 @@ impl GGUFPhi3 { rms_eps, )?; layers.push(LayerWeights { - attn_qkv: QLinear::new(&ct, reader, &format!("{prefix}.attn_qkv"), device)?, - attn_output: QLinear::new(&ct, reader, &format!("{prefix}.attn_output"), device)?, + attn_qkv: QLinear::new(ct, reader, &format!("{prefix}.attn_qkv"), device)?, + attn_output: QLinear::new(ct, reader, &format!("{prefix}.attn_output"), device)?, attn_norm, ffn_norm, mlp, diff --git a/src/openai/models/quantized_qwen.rs b/src/openai/models/quantized_qwen.rs index aca44523..46eb4ae9 100644 --- a/src/openai/models/quantized_qwen.rs +++ b/src/openai/models/quantized_qwen.rs @@ -68,10 +68,7 @@ impl LayerWeights { q_embeds.push(q_embed); k_embeds.push(k_embed); } - Ok(( - Tensor::cat(&q_embeds, 0).unwrap(), - Tensor::cat(&k_embeds, 0).unwrap(), - )) + Ok((Tensor::cat(&q_embeds, 0)?, Tensor::cat(&k_embeds, 0)?)) } fn forward_attn( @@ -125,19 +122,23 @@ impl LayerWeights { (q.contiguous()?, k.contiguous()?, v.contiguous()?) }; - let (q, k) = if self.q_norm.is_some() && self.k_norm.is_some() { - //Per‑head RMSNorm in qwen3 + let (q, k) = if let (Some(q_norm), Some(k_norm)) = (&self.q_norm, &self.k_norm) { + // Per‑head RMSNorm in qwen3 let q_flat = q.flatten(0, 2)?; // (B*H, L, D) -> (BHL, D) after transpose later let k_flat = k.flatten(0, 2)?; - //q_norm and k_norm weights stored in f32 format in qwen3 gguf - let q_flat = self.q_norm.as_ref().unwrap().forward(&q_flat)?; - let k_flat = self.k_norm.as_ref().unwrap().forward(&k_flat)?; + + // q_norm and k_norm weights stored in f32 format in qwen3 gguf + let q_flat = q_norm.forward(&q_flat)?; + let k_flat = k_norm.forward(&k_flat)?; + let q = q_flat.reshape((b_sz, self.n_head, seq_len, self.head_dim))?; let k = k_flat.reshape((b_sz, self.n_kv_head, seq_len, self.head_dim))?; + (q, k) } else { (q, k) }; + let (q, k) = self.apply_rotary_emb(&q, &k, input_positions)?; let (q, k, v) = ( q.to_dtype(self.dtype)?, @@ -222,7 +223,7 @@ impl GGUFQWen { use_flash_attn: false, bos_token_id: super::TokenID(Either::Left(Some(151644))), eos_token_id: super::TokenID(Either::Left(Some(151645))), - max_seq_len: max_seq_len, + max_seq_len, sliding_window: None, sliding_window_pattern: None, hidden_act: None, @@ -271,26 +272,25 @@ impl GGUFQWen { let version = if qwen3 { 3 } else { 2 }; let head_count = - md_get(format!("qwen{}.attention.head_count", version).as_str())?.to_u32()? as usize; + md_get(format!("qwen{version}.attention.head_count").as_str())?.to_u32()? as usize; let head_count_kv = - md_get(format!("qwen{}.attention.head_count_kv", version).as_str())?.to_u32()? as usize; + md_get(format!("qwen{version}.attention.head_count_kv").as_str())?.to_u32()? as usize; - let head_dim = md_get(format!("qwen{}.attention.key_length", version).as_str()); + let head_dim = md_get(format!("qwen{version}.attention.key_length").as_str()); let head_dim = if head_dim.is_ok() { Some(head_dim.unwrap().to_u32()? as usize) } else { None }; let embedding_length = - md_get(format!("qwen{}.embedding_length", version).as_str())?.to_u32()? as usize; + md_get(format!("qwen{version}.embedding_length").as_str())?.to_u32()? as usize; let context_length = - md_get(format!("qwen{}.context_length", version).as_str())?.to_u32()? as usize; - let block_count = - md_get(format!("qwen{}.block_count", version).as_str())?.to_u32()? as usize; + md_get(format!("qwen{version}.context_length").as_str())?.to_u32()? as usize; + let block_count = md_get(format!("qwen{version}.block_count").as_str())?.to_u32()? as usize; let rms_norm_eps = - md_get(format!("qwen{}.attention.layer_norm_rms_epsilon", version).as_str())? - .to_f32()? as f64; - let rope_freq_base = md_get(format!("qwen{}.rope.freq_base", version).as_str()) + md_get(format!("qwen{version}.attention.layer_norm_rms_epsilon").as_str())?.to_f32()? + as f64; + let rope_freq_base = md_get(format!("qwen{version}.rope.freq_base").as_str()) .and_then(|m| m.to_f32()) .unwrap_or(10000f32); @@ -406,10 +406,10 @@ impl GGUFQWen { n_kv_head: head_count_kv, head_dim, attn: PagedAttention::new( - head_count as usize, + head_count, head_dim, 1. / ((head_dim as f32).sqrt()), - Some(head_count_kv as usize), + Some(head_count_kv), None, device.clone(), None, diff --git a/src/openai/models/qwen.rs b/src/openai/models/qwen.rs index 5e65d9e5..140844d4 100644 --- a/src/openai/models/qwen.rs +++ b/src/openai/models/qwen.rs @@ -44,12 +44,11 @@ impl QwenConfig { kv_cache_dtype: DType, scfg: &SpecificConfig, ) -> Config { - let sliding_window = - if self.use_sliding_window.is_some() && self.use_sliding_window.unwrap() { - self.sliding_window - } else { - None - }; + let sliding_window = if self.use_sliding_window.unwrap_or(false) { + self.sliding_window + } else { + None + }; Config { hidden_size: self.hidden_size, head_dim: Some( @@ -68,7 +67,7 @@ impl QwenConfig { bos_token_id: self.bos_token_id, eos_token_id: self.eos_token_id, max_seq_len: self.max_position_embeddings, - sliding_window: sliding_window, + sliding_window, sliding_window_pattern: None, hidden_act: Some(self.hidden_act), tie_word_embeddings: self.tie_word_embeddings, @@ -131,26 +130,23 @@ impl RotaryEmbedding { let sin = self.sin.narrow(0, seqlen_offset[0], seq_len)?; let x_q = q.narrow(0, b, 1)?; let x_k = k.narrow(0, b, 1)?; - let q_embed = candle_nn::rotary_emb::rope(&x_q, &cos, &sin).unwrap(); - let k_embed = candle_nn::rotary_emb::rope(&x_k, &cos, &sin).unwrap(); + let q_embed = candle_nn::rotary_emb::rope(&x_q, &cos, &sin)?; + let k_embed = candle_nn::rotary_emb::rope(&x_k, &cos, &sin)?; q_embeds.push(q_embed); k_embeds.push(k_embed); } - Ok(( - Tensor::cat(&q_embeds, 0).unwrap(), - Tensor::cat(&k_embeds, 0).unwrap(), - )) + Ok((Tensor::cat(&q_embeds, 0)?, Tensor::cat(&k_embeds, 0)?)) } } -struct MLP { +struct Mlp { gate_proj: TensorParallelColumnLinear, up_proj: TensorParallelColumnLinear, down_proj: TensorParallelRowLinear, act_fn: candle_nn::Activation, } -impl MLP { +impl Mlp { fn new(cfg: &Config, vb: VarBuilder, comm: Rc) -> Result { let hidden_sz = cfg.hidden_size; let intermediate_sz = cfg.intermediate_size; @@ -191,7 +187,7 @@ impl MLP { } } -impl Module for MLP { +impl Module for Mlp { fn forward(&self, xs: &Tensor) -> Result { let lhs = self.act_fn.forward(&self.gate_proj.forward(xs)?)?; let rhs = self.up_proj.forward(xs)?; @@ -375,7 +371,7 @@ impl Attention { struct DecoderLayer { self_attn: Attention, - mlp: MLP, + mlp: Mlp, input_layernorm: RmsNorm, post_attention_layernorm: RmsNorm, } @@ -389,7 +385,7 @@ impl DecoderLayer { comm: Rc, ) -> Result { let self_attn = Attention::new(qwen3, rotary_emb, cfg, vb.pp("self_attn"), comm.clone())?; - let mlp = MLP::new(cfg, vb.pp("mlp"), comm.clone())?; + let mlp = Mlp::new(cfg, vb.pp("mlp"), comm.clone())?; let input_layernorm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?; let post_attention_layernorm = rms_norm( diff --git a/src/openai/models/yi.rs b/src/openai/models/yi.rs index c5c1a797..a87975e1 100644 --- a/src/openai/models/yi.rs +++ b/src/openai/models/yi.rs @@ -117,26 +117,23 @@ impl RotaryEmbedding { let sin = self.sin.narrow(0, seqlen_offset[0], seq_len)?; let x_q = q.narrow(0, b, 1)?; let x_k = k.narrow(0, b, 1)?; - let q_embed = candle_nn::rotary_emb::rope(&x_q, &cos, &sin).unwrap(); - let k_embed = candle_nn::rotary_emb::rope(&x_k, &cos, &sin).unwrap(); + let q_embed = candle_nn::rotary_emb::rope(&x_q, &cos, &sin)?; + let k_embed = candle_nn::rotary_emb::rope(&x_k, &cos, &sin)?; q_embeds.push(q_embed); k_embeds.push(k_embed); } - Ok(( - Tensor::cat(&q_embeds, 0).unwrap(), - Tensor::cat(&k_embeds, 0).unwrap(), - )) + Ok((Tensor::cat(&q_embeds, 0)?, Tensor::cat(&k_embeds, 0)?)) } } -struct MLP { +struct Mlp { gate_proj: TensorParallelColumnLinear, up_proj: TensorParallelColumnLinear, down_proj: TensorParallelRowLinear, act_fn: Activation, } -impl MLP { +impl Mlp { fn new(cfg: &Config, vb: VarBuilder, comm: Rc) -> Result { let hidden_sz = cfg.hidden_size; let intermediate_sz = cfg.intermediate_size; @@ -176,7 +173,7 @@ impl MLP { } } -impl Module for MLP { +impl Module for Mlp { fn forward(&self, xs: &Tensor) -> Result { let lhs = self.act_fn.forward(&self.gate_proj.forward(xs)?)?; let rhs = self.up_proj.forward(xs)?; @@ -331,7 +328,7 @@ impl Attention { struct DecoderLayer { self_attn: Attention, - mlp: MLP, + mlp: Mlp, ln1: RmsNorm, ln2: RmsNorm, } @@ -344,7 +341,7 @@ impl DecoderLayer { comm: Rc, ) -> Result { let self_attn = Attention::new(rotary_emb, cfg, vb.pp("self_attn"), comm.clone())?; - let mlp = MLP::new(cfg, vb.pp("mlp"), comm.clone())?; + let mlp = Mlp::new(cfg, vb.pp("mlp"), comm.clone())?; let ln1 = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?; let ln2 = rms_norm( cfg.hidden_size, diff --git a/src/openai/openai_server.rs b/src/openai/openai_server.rs index 8b464a4c..81539ac4 100644 --- a/src/openai/openai_server.rs +++ b/src/openai/openai_server.rs @@ -18,19 +18,6 @@ use tokio::sync::Notify; use tokio::time::Duration; use tracing::debug; use uuid::Uuid; -// fn verify_model(data: &OpenAIServerData<'_>, model_name: &String) -> Result<(), APIError> { -// let current_name = { -// let model = data.model.lock().unwrap(); -// model.get_pipeline().name().to_string() -// }; -// if ¤t_name != model_name { -// Err(APIError::new(format!( -// "Model name `{model_name}` is invalid." -// ))) -// } else { -// Ok(()) -// } -// } // Get prompt, roles async fn get_gen_prompt( @@ -38,11 +25,10 @@ async fn get_gen_prompt( request: &ChatCompletionRequest, ) -> Result { let mut model = data.model.write(); - let conversation = model + let pipeline = model .get_mut_pipeline(0) - .unwrap() - .0 - .get_conversation(data.record_conversation); + .ok_or(APIError::new("Missing pipeline".to_string()))?; + let conversation = pipeline.0.get_conversation(data.record_conversation); match &request.messages { Messages::Literal(msg) => { @@ -83,9 +69,10 @@ async fn check_length( ) -> Result { let token_ids = { let model = data.model.read(); - model + let pipeline = model .get_pipeline(0) - .unwrap() + .ok_or(APIError::new("Missing pipeline".to_string()))?; + pipeline .0 .tokenizer() .encode_fast(prompt, false) @@ -123,12 +110,6 @@ pub async fn chat_completions( State(data): State>, request: Json, ) -> ChatResponder { - // let model_name = &request.model; - // let res = verify_model(&data, model_name); - // if res.is_err() { - // return Either::Left(Err(res.err().unwrap())); - // } - #[cfg(feature = "nccl")] use crate::openai::communicator::DaemonManager; #[cfg(feature = "nccl")] @@ -146,23 +127,21 @@ pub async fn chat_completions( )); } - let prompt = get_gen_prompt(&data, &request).await; - if prompt.is_err() { - return ChatResponder::ValidationError(prompt.err().unwrap()); - } - let prompt = prompt.unwrap(); + let prompt = match get_gen_prompt(&data, &request).await { + Ok(p) => p, + Err(e) => return ChatResponder::ValidationError(e), + }; - let token_ids = check_length(&request, prompt.clone(), &data).await; - if token_ids.is_err() { - return ChatResponder::ValidationError(token_ids.err().unwrap()); - } - let token_ids: Encoding = token_ids.unwrap(); + let token_ids: Encoding = match check_length(&request, prompt.clone(), &data).await { + Ok(ids) => ids, + Err(e) => return ChatResponder::ValidationError(e), + }; debug!("\n\n\nPrompt {:?}", prompt); let request_id = format!("cmpl-{}", Uuid::new_v4()); - let sampling_params = SamplingParams::new( + let sampling_params = match SamplingParams::new( request.n.unwrap_or(1), request.best_of, request.presence_penalty.unwrap_or(0.0), @@ -186,11 +165,10 @@ pub async fn chat_completions( None, request.skip_special_tokens.unwrap_or(true), request.thinking.or(data.pipeline_config.thinking), - ); - if sampling_params.is_err() { - return ChatResponder::ValidationError(sampling_params.err().unwrap()); - } - let sampling_params = sampling_params.unwrap(); + ) { + Ok(params) => params, + Err(e) => return ChatResponder::ValidationError(e), + }; let (response_tx, rx) = flume::unbounded(); tracing::info!("{:?}", sampling_params); @@ -252,8 +230,7 @@ pub async fn chat_completions( let model = data_clone.model.read(); if !model.completion_records.contains_key(&request_id_clone) { return ChatResponder::ModelError(APIError::from(format!( - "Unable to generate response for request {}", - request_id_clone + "Unable to generate response for request {request_id_clone}" ))); } diff --git a/src/openai/pipelines/llm_engine.rs b/src/openai/pipelines/llm_engine.rs index 84b33f8d..05eb4077 100644 --- a/src/openai/pipelines/llm_engine.rs +++ b/src/openai/pipelines/llm_engine.rs @@ -153,7 +153,7 @@ impl LLMEngine { continue; } let result = &results[0]; - if results.len() == 0 || result.len() == 0 { + if results.is_empty() || result.is_empty() { continue; } @@ -469,7 +469,7 @@ impl LLMEngine { let x = e.sequence_groups.read(); x.clone() }; - if scheduled.len() == 0 { + if scheduled.is_empty() { continue; //data not ready } @@ -506,7 +506,7 @@ impl LLMEngine { false }; #[cfg(not(feature = "nccl"))] - let do_sample = if rank == 0 { true } else { false }; + let do_sample = rank == 0; let optional_results = if do_sample { let sample = { @@ -601,7 +601,7 @@ impl LLMEngine { group.arrival_time, Some(logprobs.bytes.clone()), None, - &pipeline, + pipeline, ); let ret = sender.send(ChatResponse::Chunk(chunk)); if ret.is_err() { @@ -624,7 +624,7 @@ impl LLMEngine { group.arrival_time, None, Some(finish_reason.clone()), - &pipeline, + pipeline, ); let ret = sender.send(ChatResponse::Chunk(chunk)); if ret.is_err() { @@ -745,8 +745,8 @@ impl LLMEngine { e.completion_records .insert(request_id.to_string(), responses[request_id].clone()); let notify = e.sync_notifies.get(request_id); - if notify.is_some() { - notify.unwrap().as_ref().unwrap().notify_one(); + if let Some(Some(notify)) = notify { + notify.notify_one(); } } } @@ -1028,7 +1028,7 @@ impl LLMEngine { seq_id: usize, group_id: usize, prompt: &Encoding, - request_id: &String, + request_id: &str, created: SystemTime, sampling_params: &SamplingParams, use_logprobs: bool, @@ -1044,17 +1044,16 @@ impl LLMEngine { seq_id, self.cache_config.block_size, )))); - let seq_group = SequenceGroup::new( + SequenceGroup::new( &[seq], get_created_time_secs(), group_id, - request_id.clone(), + request_id.to_owned(), created, sampling_params.clone(), use_logprobs, sender, - ); - seq_group + ) } pub fn add_request( @@ -1069,14 +1068,13 @@ impl LLMEngine { ) { let prompt_len = prompt.get_ids().len(); let sync_notify = sync_notify.clone(); - if sync_notify.is_some() { - self.sync_notifies - .insert(request_id.clone(), Some(sync_notify.unwrap())); + if let Some(sync) = sync_notify { + self.sync_notifies.insert(request_id.clone(), Some(sync)); } + let sender_clone = sender.clone(); - if sender_clone.is_some() { - self.senders - .insert(request_id.clone(), Some(sender_clone.unwrap())); + if let Some(sender) = sender_clone { + self.senders.insert(request_id.clone(), Some(sender)); } #[cfg(feature = "nccl")] diff --git a/src/openai/pipelines/pipeline.rs b/src/openai/pipelines/pipeline.rs index d385de85..bdf64664 100644 --- a/src/openai/pipelines/pipeline.rs +++ b/src/openai/pipelines/pipeline.rs @@ -220,11 +220,7 @@ impl DefaultLoader { ) -> Result<(Vec>, PipelineConfig), APIError> { let specific_args = self.config.clone(); let reporter = Arc::new(RwLock::new(ProgressReporter::new(local_rank.unwrap_or(0)))); - let num_subprogress = if local_world_size.is_none() { - 0 - } else { - local_world_size.unwrap() - 1 - }; + let num_subprogress = local_world_size.map_or(0, |n| n - 1); let (models, devices, config, sep_style) = if quant.is_some() && matches!(quant.as_ref().unwrap().as_str(), "ggml" | "gguf") @@ -238,7 +234,7 @@ impl DefaultLoader { ); let s_cfg = specific_args.clone(); let nlayers = { - let mut file = try_api!(std::fs::File::open(&path.clone())); + let mut file = try_api!(std::fs::File::open(path.clone())); let content = try_api!( gguf_file::Content::read(&mut file).map_err(|e| e.with_path(path.clone())) ); @@ -253,13 +249,9 @@ impl DefaultLoader { }; nlayers.unwrap() }; - let handle = progress_worker( - Some(num_subprogress as usize), - nlayers, - Arc::clone(&reporter), - ) - .await; - let mut file = try_api!(std::fs::File::open(&path.clone())); + let handle = + progress_worker(Some(num_subprogress), nlayers, Arc::clone(&reporter)).await; + let mut file = try_api!(std::fs::File::open(path.clone())); let content = try_api!( gguf_file::Content::read(&mut file).map_err(|e| e.with_path(path.clone())) ); @@ -405,7 +397,7 @@ impl DefaultLoader { info!("Loading {} model.", self.name); let handle = progress_worker( - Some(num_subprogress as usize), + Some(num_subprogress), config.num_hidden_layers, Arc::clone(&reporter), ) @@ -698,7 +690,7 @@ impl DefaultLoader { let tokenizer_cfg: Option = std::fs::read_to_string(tokenizer_cfg_file).ok(); let cfg_tokenizer: TokenizerConfig = - serde_json::from_str(&tokenizer_cfg.unwrap().as_str()).unwrap(); + serde_json::from_str(tokenizer_cfg.unwrap().as_str()).unwrap(); let bos = if cfg_tokenizer.bos_token.is_some() { match cfg_tokenizer.bos_token.unwrap() { BosEosToken(Either::Left(Some(id))) => Some(id), @@ -726,9 +718,8 @@ impl DefaultLoader { } else if quant.is_some() && matches!(quant.as_ref().unwrap().as_str(), "ggml" | "gguf") { use crate::backend::gguf::{get_gguf_info, Content, GGUFInfo}; let filename = paths.get_weight_filenames()[0].clone(); - let mut readers = Vec::new(); - readers.push(std::fs::File::open(filename).unwrap()); - let mut readers = readers.iter_mut().collect::>(); + let mut reader = std::fs::File::open(filename).unwrap(); + let mut readers = vec![&mut reader]; let content = Content::from_readers(&mut readers).unwrap(); let GGUFInfo { tokenizer, @@ -789,22 +780,23 @@ impl DefaultLoader { }; } - if chat_template.is_some() { - if chat_template.as_ref().unwrap().find("<|eom_id|>").is_some() { + if let Some(template) = chat_template.as_ref() { + if template.contains("<|eom_id|>") { tracing::warn!("custom stop token <|eom_id|> in chat template"); - stop_token_ids.push(128008) + stop_token_ids.push(128008); } - if chat_template.as_ref().unwrap().find("<|eot_id|>").is_some() { + if template.contains("<|eot_id|>") { tracing::warn!("custom stop token <|eot_id|> in chat template"); - stop_token_ids.push(128009) + stop_token_ids.push(128009); } - if chat_template.as_ref().unwrap().find("<|end|>").is_some() { + if template.contains("<|end|>") { tracing::warn!("custom stop token <|end|> in chat template"); if let Some(token) = tokenizer.get_vocab(true).get("<|end|>").copied() { - stop_token_ids.push(token) - }; + stop_token_ids.push(token); + } } } + if stop_token_ids.is_empty() { //if no eos_token defined in the config, use default if let Some(token) = tokenizer.get_vocab(true).get("<|endoftext|>").copied() { @@ -868,49 +860,49 @@ impl DefaultPipeline { match &self.model { LLMModel::Llama(llama) => llama - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::Phi2(phi) => phi - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::Phi3(phi) => phi - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::Qwen(qwen) => qwen - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::Gemma(gemma) => gemma - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::Gemma3(gemma3) => gemma3 - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::Mistral(mistral) => mistral - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::Yi(yi) => yi - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::StableLM(stablelm) => stablelm - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::GLM4(glm4) => glm4 - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::DeepSeek(deepseek) => deepseek - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::Phi3GGUF(phi3) => phi3 - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::LlamaGGUF(llama) => llama - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::QWenGGUF(qwen) => qwen - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), LLMModel::GLM4GGUF(glm4) => glm4 - .forward(&input_tokens, input_positions, kv_cache, &input_metadata) + .forward(&input_tokens, input_positions, kv_cache, input_metadata) .map_err(APIError::from), } } @@ -982,13 +974,13 @@ impl DefaultPipeline { let logits = if panalties.iter().any(|&v| v != 1.0 && v != 0.) { self.logits_processor - .apply_batch_repeat_penalty(&logits, panalties, reference_tokens) + .apply_batch_repeat_penalty(logits, panalties, reference_tokens) .unwrap() } else { logits.to_owned() }; - let group_ids: Vec = groups.into_iter().map(|group| group.group_id).collect(); + let group_ids: Vec = groups.iter().map(|group| group.group_id).collect(); let param = &groups[0].sampling_params; let sampling_params = if param.temperature.is_some() && (param.top_k.is_some() || param.top_p.is_some()) { @@ -1031,13 +1023,8 @@ impl DefaultPipeline { } } - let custom_stop_token_match = if custom_stop_tokens[i].len() > 0 - && custom_stop_tokens[i].contains(&text.trim().to_string()) - { - true - } else { - false - }; + let custom_stop_token_match = !custom_stop_tokens[i].is_empty() + && custom_stop_tokens[i].contains(&text.trim().to_string()); if tokens_generated[i] < 0 { Right("length".to_string()) diff --git a/src/openai/responses.rs b/src/openai/responses.rs index fe63de12..cb9a7a6c 100644 --- a/src/openai/responses.rs +++ b/src/openai/responses.rs @@ -6,7 +6,7 @@ use axum::response::{IntoResponse, Sse}; use derive_more::{Display, Error}; use serde::{Deserialize, Serialize}; #[derive(Debug, Display, Error, Serialize)] -#[display(fmt = "Error: {}", data)] +#[display(fmt = "Error: {data}")] pub struct APIError { data: String, } diff --git a/src/openai/sampling_params.rs b/src/openai/sampling_params.rs index 5d9985bc..85f8f8f4 100644 --- a/src/openai/sampling_params.rs +++ b/src/openai/sampling_params.rs @@ -152,44 +152,6 @@ impl SamplingParams { Ok(this) } - // pub fn get_logits_processor<'a>( - // &self, - // seed: u64, - // tokenizer: &'a Tokenizer, - // top_n_logprobs: usize, - // ) -> LogitsProcessor<'a> { - // if self.top_k == -1 && self.top_p == 1. { - // // Greedy - // LogitsProcessor::new( - // seed, - // Some(self.temperature.into()), - // SamplingMethod::Multinomial, - // top_n_logprobs, - // tokenizer, - // ) - // } else if self.top_k > 0 && self.top_p == 1. { - // // Top-k - // LogitsProcessor::new( - // seed, - // Some(self.temperature.into()), - // SamplingMethod::TopK(self.top_k.try_into().unwrap()), - // top_n_logprobs, - // tokenizer, - // ) - // } else if self.top_k == -1 && self.top_p != 1. { - // // Top-p - // LogitsProcessor::new( - // seed, - // Some(self.temperature.into()), - // SamplingMethod::TopP(self.top_p.into()), - // top_n_logprobs, - // tokenizer, - // ) - // } else { - // unreachable!() - // } - // } - fn verify_args(&self) -> Result<(), APIError> { if self.n < 1 { return Err(APIError::new(format!( @@ -249,17 +211,21 @@ impl SamplingParams { self.best_of ))); } - if self.temperature.is_some() && self.temperature.unwrap() > SAMPLING_EPS { + + if self.temperature.is_some_and(|t| t > SAMPLING_EPS) { return Err(APIError::new_str( "temperature must be 0 when using beam search", )); } - if self.top_p.is_some() && self.top_p.unwrap() < 1.0f32 - SAMPLING_EPS { + + if self.top_p.is_some_and(|p| p < 1.0 - SAMPLING_EPS) { return Err(APIError::new_str("top_p must be 1 when using beam search")); } - if self.top_k.is_some() && self.top_k.unwrap() != -1 { + + if self.top_k.is_some_and(|k| k != -1) { return Err(APIError::new_str("top_k must be -1 when using beam search")); } + Ok(()) } @@ -282,16 +248,19 @@ impl SamplingParams { self.best_of ))); } - if self.top_p.is_some() && self.top_p.unwrap() < 1.0f32 - SAMPLING_EPS { + + if self.top_p.is_some_and(|p| p < 1.0 - SAMPLING_EPS) { return Err(APIError::new_str( "top_p must be 1 when using greedy sampling (no temperature specified).", )); } - if self.top_k.is_some() && self.top_k.unwrap() != -1 { + + if self.top_k.is_some_and(|k| k != -1) { return Err(APIError::new_str( "top_k must be -1 when using greedy sampling (no temperature specified).", )); } + Ok(()) } } diff --git a/src/paged_attention/mod.rs b/src/paged_attention/mod.rs index 14de382f..48bcf8c9 100644 --- a/src/paged_attention/mod.rs +++ b/src/paged_attention/mod.rs @@ -152,7 +152,7 @@ impl PagedAttention { value.clone() }; - let num_chunks = (seq_len + chunk_size - 1) / chunk_size; + let num_chunks = seq_len.div_ceil(chunk_size); for c in 0..num_chunks { let offset = c * chunk_size; @@ -239,11 +239,11 @@ impl PagedAttention { // input_metadata: metadata for paged attention. // // alibi_slopes: shape = [num_heads] - let max_context_len = if self.sliding_window.is_some() { - self.sliding_window.unwrap() - } else { - input_metadata.max_context_len.unwrap() - }; + let max_context_len = self + .sliding_window + .or(input_metadata.max_context_len) + .expect("max_context_len must be set"); + paged_attention( &query, key_cache.as_ref().unwrap(),