diff --git a/src/exec/mod.rs b/src/exec/mod.rs index 191024a..0e59094 100644 --- a/src/exec/mod.rs +++ b/src/exec/mod.rs @@ -8,6 +8,7 @@ pub mod fun_asr_nano; pub mod glm_asr_nano; pub mod glm_ocr; pub mod hunyuan_ocr; +pub mod lfm2; pub mod minicpm4; pub mod paddleocr_vl; pub mod qwen2_5vl; @@ -18,7 +19,6 @@ pub mod qwen3vl; pub mod rmbg2_0; pub mod voxcpm; pub mod voxcpm1_5; -pub mod lfm2; use anyhow::Result; diff --git a/src/models/lfm2/config.rs b/src/models/lfm2/config.rs index 8ec24c1..2d97e31 100644 --- a/src/models/lfm2/config.rs +++ b/src/models/lfm2/config.rs @@ -1,5 +1,5 @@ +use anyhow::{Result, anyhow}; use serde::{Deserialize, Serialize}; -use anyhow::{anyhow, Result}; #[derive(Debug, PartialEq, Deserialize, Serialize)] pub struct Lfm2Config { pub architectures: Vec, @@ -43,7 +43,6 @@ pub struct Lfm2Config { pub tie_embedding: Option, } - impl Lfm2Config { pub fn full_attn_idx2layer_type(&mut self) { if self.layer_types.is_none() @@ -75,7 +74,9 @@ impl Lfm2Config { } Ok(layer_types) } else { - Err(anyhow!("layer_types full_attn_idxs cannot be none at the same time")) + Err(anyhow!( + "layer_types full_attn_idxs cannot be none at the same time" + )) } } } @@ -84,5 +85,5 @@ impl Lfm2Config { pub struct Lfm2GenerateConfig { pub bos_token_id: u32, pub eos_token_id: u32, - pub pad_token_id: u32 -} \ No newline at end of file + pub pad_token_id: u32, +} diff --git a/src/models/lfm2/generate.rs b/src/models/lfm2/generate.rs index 1c58aff..72ab927 100644 --- a/src/models/lfm2/generate.rs +++ b/src/models/lfm2/generate.rs @@ -1,15 +1,18 @@ +use crate::utils::build_completion_chunk_response; use crate::{ chat_template::ChatTemplate, - models::{GenerateModel, lfm2::{ - config::{Lfm2Config, Lfm2GenerateConfig}, - model::Lfm2Model, - }}, + models::{ + GenerateModel, + lfm2::{ + config::{Lfm2Config, Lfm2GenerateConfig}, + model::Lfm2Model, + }, + }, tokenizer::TokenizerModel, utils::{ build_completion_response, find_type_files, get_device, get_dtype, get_logit_processor, }, }; -use crate::utils::build_completion_chunk_response; use aha_openai_dive::v1::resources::chat::{ChatCompletionParameters, ChatCompletionResponse}; use anyhow::Result; use candle_core::{DType, Device, Tensor}; @@ -59,8 +62,6 @@ impl<'a> Lfm2GenerateModel<'a> { model_name, }) } - - } impl<'a> GenerateModel for Lfm2GenerateModel<'a> { @@ -103,16 +104,20 @@ impl<'a> GenerateModel for Lfm2GenerateModel<'a> { } fn generate_stream( - &mut self, - mes: ChatCompletionParameters, - ) -> Result< - Box< - dyn rocket::futures::Stream> - + Send - + Unpin - + '_, - >, - > { + &mut self, + mes: ChatCompletionParameters, + ) -> Result< + Box< + dyn rocket::futures::Stream< + Item = Result< + aha_openai_dive::v1::resources::chat::ChatCompletionChunkResponse, + anyhow::Error, + >, + > + Send + + Unpin + + '_, + >, + > { let mes_render = self.chat_template.apply_chat_template(&mes)?; let mut logits = get_logit_processor( mes.temperature, @@ -160,4 +165,4 @@ impl<'a> GenerateModel for Lfm2GenerateModel<'a> { }; Ok(Box::new(Box::pin(stream))) } -} \ No newline at end of file +} diff --git a/src/models/lfm2/mod.rs b/src/models/lfm2/mod.rs index de0a4df..7fca417 100644 --- a/src/models/lfm2/mod.rs +++ b/src/models/lfm2/mod.rs @@ -1,3 +1,3 @@ pub mod config; +pub mod generate; pub mod model; -pub mod generate; \ No newline at end of file diff --git a/src/models/lfm2/model.rs b/src/models/lfm2/model.rs index db7d4d1..8dd13c7 100644 --- a/src/models/lfm2/model.rs +++ b/src/models/lfm2/model.rs @@ -70,7 +70,7 @@ impl Lfm2ShortConv { bx.narrow(D::Minus1, pad_num.unsigned_abs(), self.l_cache)? }; self.cache = Some(conv_state); - let bx = bx.pad_with_zeros(D::Minus1, self.l_cache-1, self.l_cache-1)?; + let bx = bx.pad_with_zeros(D::Minus1, self.l_cache - 1, self.l_cache - 1)?; let bx = conv1d_depthwise(&bx, self.conv.weight(), self.conv.bias())?; bx.narrow(D::Minus1, 0, seq_len)? } else { @@ -145,9 +145,8 @@ impl Lfm2DecoderLayer { let intermediate_size = if config.block_auto_adjust_ff_dim { let inter_size = 2 * config.block_ff_dim / 3; let inter_size = (config.block_ffn_dim_multiplier * inter_size as f64) as usize; - let inter_size = config.block_multiple_of - * ((inter_size + config.block_multiple_of - 1) / config.block_multiple_of); - inter_size + config.block_multiple_of + * ((inter_size + config.block_multiple_of - 1) / config.block_multiple_of) } else { config.block_ff_dim }; @@ -291,9 +290,7 @@ impl Lfm2Model { ); match linear { Ok(linear) => linear, - Err(_) => { - Linear::new(model.embed_tokens.embeddings().clone(), None) - } + Err(_) => Linear::new(model.embed_tokens.embeddings().clone(), None), } }; Ok(Self { model, lm_head }) diff --git a/src/models/mod.rs b/src/models/mod.rs index c9ce259..625784e 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -7,6 +7,7 @@ pub mod fun_asr_nano; pub mod glm_asr_nano; pub mod glm_ocr; pub mod hunyuan_ocr; +pub mod lfm2; pub mod mask_gct; pub mod minicpm4; pub mod paddleocr_vl; @@ -19,7 +20,6 @@ pub mod qwen3vl; pub mod rmbg2_0; pub mod voxcpm; pub mod w2v_bert_2_0; -pub mod lfm2; use aha_openai_dive::v1::resources::chat::{ ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, @@ -28,7 +28,15 @@ use anyhow::{Result, anyhow}; use rocket::futures::Stream; use crate::models::{ - deepseek_ocr::generate::DeepseekOCRGenerateModel, fun_asr_nano::generate::FunAsrNanoGenerateModel, glm_asr_nano::generate::GlmAsrNanoGenerateModel, glm_ocr::generate::GlmOcrGenerateModel, hunyuan_ocr::generate::HunyuanOCRGenerateModel, lfm2::generate::Lfm2GenerateModel, minicpm4::generate::MiniCPMGenerateModel, paddleocr_vl::generate::PaddleOCRVLGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel, qwen3::generate::Qwen3GenerateModel, qwen3_5::generate::Qwen3_5GenerateModel, qwen3_asr::generate::Qwen3AsrGenerateModel, qwen3vl::generate::Qwen3VLGenerateModel, rmbg2_0::generate::RMBG2_0Model, voxcpm::generate::VoxCPMGenerate + deepseek_ocr::generate::DeepseekOCRGenerateModel, + fun_asr_nano::generate::FunAsrNanoGenerateModel, + glm_asr_nano::generate::GlmAsrNanoGenerateModel, glm_ocr::generate::GlmOcrGenerateModel, + hunyuan_ocr::generate::HunyuanOCRGenerateModel, lfm2::generate::Lfm2GenerateModel, + minicpm4::generate::MiniCPMGenerateModel, paddleocr_vl::generate::PaddleOCRVLGenerateModel, + qwen2_5vl::generate::Qwen2_5VLGenerateModel, qwen3::generate::Qwen3GenerateModel, + qwen3_5::generate::Qwen3_5GenerateModel, qwen3_asr::generate::Qwen3AsrGenerateModel, + qwen3vl::generate::Qwen3VLGenerateModel, rmbg2_0::generate::RMBG2_0Model, + voxcpm::generate::VoxCPMGenerate, }; #[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)] @@ -130,7 +138,8 @@ impl WhichModel { pub fn model_type(self) -> &'static str { match self { // LLM models - WhichModel::MiniCPM4_0_5B | WhichModel::Qwen3_0_6B + WhichModel::MiniCPM4_0_5B + | WhichModel::Qwen3_0_6B | WhichModel::LFM2_1_2B | WhichModel::LFM2_5_1_2BInstruct => "llm", WhichModel::Qwen2_5vl3B diff --git a/tests/config_tests.rs b/tests/config_tests.rs index 94409c8..e78645f 100644 --- a/tests/config_tests.rs +++ b/tests/config_tests.rs @@ -1,5 +1,8 @@ use aha::models::{ - deepseek_ocr::config::DeepseekOCRConfig, hunyuan_ocr::config::HunYuanVLConfig, lfm2::config::Lfm2Config, minicpm4::config::MiniCPM4Config, paddleocr_vl::config::PaddleOCRVLConfig, qwen2_5vl::config::Qwen2_5VLConfig, qwen3vl::config::Qwen3VLConfig, voxcpm::config::VoxCPMConfig + deepseek_ocr::config::DeepseekOCRConfig, hunyuan_ocr::config::HunYuanVLConfig, + lfm2::config::Lfm2Config, minicpm4::config::MiniCPM4Config, + paddleocr_vl::config::PaddleOCRVLConfig, qwen2_5vl::config::Qwen2_5VLConfig, + qwen3vl::config::Qwen3VLConfig, voxcpm::config::VoxCPMConfig, }; use anyhow::Result; diff --git a/tests/messy_test.rs b/tests/messy_test.rs index e80cccc..9740826 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -43,7 +43,7 @@ fn messy_test() -> Result<()> { // aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; // let model_path = format!("{}/deepseek-ai/DeepSeek-OCR-2/", save_dir); // let stem = std::path::Path::new(&model_path) - // .file_name() + // .file_name() // .and_then(|s| s.to_str()) // .unwrap_or("qwen3.5"); // println!("stem: {:?}", stem); diff --git a/tests/test_lfm2.rs b/tests/test_lfm2.rs index 60e2c9d..aea053a 100644 --- a/tests/test_lfm2.rs +++ b/tests/test_lfm2.rs @@ -1,7 +1,10 @@ -use std::{pin::pin, time::Instant}; +use aha::{ + chat::ChatCompletionParameters, + models::{GenerateModel, lfm2::generate::Lfm2GenerateModel}, +}; use anyhow::Result; -use aha::{chat::ChatCompletionParameters, models::{GenerateModel, lfm2::generate::Lfm2GenerateModel}}; use rocket::futures::StreamExt; +use std::{pin::pin, time::Instant}; #[test] fn lfm2_generate() -> Result<()> { @@ -44,7 +47,6 @@ fn lfm2_generate() -> Result<()> { Ok(()) } - #[tokio::test] async fn lfm2_stream() -> Result<()> { // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_lfm2 lfm2_stream -r -- --nocapture @@ -88,4 +90,4 @@ async fn lfm2_stream() -> Result<()> { println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) -} \ No newline at end of file +}