diff options
| author | sigoden <sigoden@gmail.com> | 2024-12-04 21:03:59 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-12-04 21:03:59 +0800 |
| commit | 3a3388375be05758d5f5574cce1d0f8eb7bdd604 (patch) | |
| tree | cc8e19e218c20f28a04feca33850fce4bb2cbe98 /src/rag/mod.rs | |
| parent | 7d42fe9429f75d195f865b07cef10d040d5397f2 (diff) | |
| download | aichat-3a3388375be05758d5f5574cce1d0f8eb7bdd604.tar.gz | |
refactor: improve retrieve model (#1036)
- check the model type while retrieve model
- select chat/reranker model even if it is missed in client models
- find predefined-models for openai-compatible client with startsWith
- remove client::ApiType
Diffstat (limited to 'src/rag/mod.rs')
| -rw-r--r-- | src/rag/mod.rs | 11 |
1 files changed, 7 insertions, 4 deletions
diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 03b8f7e..ceb83f6 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -110,7 +110,8 @@ impl Rag { pub fn create(config: &GlobalConfig, name: &str, path: &Path, data: RagData) -> Result<Self> { let hnsw = data.build_hnsw(); let bm25 = data.build_bm25(); - let embedding_model = Model::retrieve_embedding(&config.read(), &data.embedding_model)?; + let embedding_model = + Model::retrieve_model(&config.read(), &data.embedding_model, ModelType::Embedding)?; let rag = Rag { config: config.clone(), name: name.to_string(), @@ -164,14 +165,15 @@ impl Rag { value } None => { - let models = list_embedding_models(&config.read()); + let models = list_models(&config.read(), ModelType::Embedding); if models.is_empty() { bail!("No available embedding model"); } select_embedding_model(&models)? } }; - let embedding_model = Model::retrieve_embedding(&config.read(), &embedding_model_id)?; + let embedding_model = + Model::retrieve_model(&config.read(), &embedding_model_id, ModelType::Embedding)?; let chunk_size = match chunk_size { Some(value) => { @@ -516,7 +518,8 @@ impl Rag { let ids = match rerank_model { Some(model_id) => { - let model = Model::retrieve_reranker(&self.config.read(), model_id)?; + let model = + Model::retrieve_model(&self.config.read(), model_id, ModelType::Reranker)?; let client = init_client(&self.config, Some(model))?; let ids: IndexSet<DocumentId> = [vector_search_ids, keyword_search_ids] .concat() |
