From 8fe30ca4caae488c8586d35ec2979ac045f86eb3 Mon Sep 17 00:00:00 2001 From: HugoCasa Date: Fri, 20 Oct 2023 15:49:04 +0200 Subject: [PATCH] fix: cache embedding model in docker img (#2474) --- Dockerfile | 2 ++ backend/src/main.rs | 16 ++++++++++++ backend/windmill-api/src/embeddings.rs | 34 ++++++++++++++++++++------ backend/windmill-api/src/jobs.rs | 2 +- backend/windmill-api/src/lib.rs | 2 +- 5 files changed, 46 insertions(+), 10 deletions(-) diff --git a/Dockerfile b/Dockerfile index 114e9e3316..147f3e7900 100644 --- a/Dockerfile +++ b/Dockerfile @@ -191,6 +191,8 @@ RUN ln -s ${APP}/windmill /usr/local/bin/windmill WORKDIR ${APP} +RUN windmill cache + EXPOSE 8000 CMD ["windmill"] diff --git a/backend/src/main.rs b/backend/src/main.rs index ad74a75497..3263d892e0 100644 --- a/backend/src/main.rs +++ b/backend/src/main.rs @@ -69,6 +69,22 @@ async fn main() -> anyhow::Result<()> { #[cfg(feature = "flamegraph")] let _guard = windmill_common::tracing_init::setup_flamegraph(); + let cli_arg = std::env::args().nth(1).unwrap_or_default(); + + match cli_arg.as_str() { + "cache" => { + tracing::info!("Caching embedding model..."); + windmill_api::embeddings::ModelInstance::load_model_files().await?; + tracing::info!("Cached embedding model"); + return Ok(()); + } + "-v" | "--version" | "version" => { + println!("Windmill {}", GIT_VERSION); + return Ok(()); + } + _ => {} + } + let mode = std::env::var("MODE") .map(|x| x.to_lowercase()) .map(|x| { diff --git a/backend/windmill-api/src/embeddings.rs b/backend/windmill-api/src/embeddings.rs index 570aa358f9..e6500b7be7 100644 --- a/backend/windmill-api/src/embeddings.rs +++ b/backend/windmill-api/src/embeddings.rs @@ -1,4 +1,4 @@ -use std::{collections::HashMap, sync::Arc}; +use std::{collections::HashMap, path::PathBuf, sync::Arc}; use anyhow::{self, Error, Result}; use axum::{extract::Query, routing::get, Extension, Json, Router}; @@ -110,11 +110,8 @@ pub struct ModelInstance { } impl ModelInstance { - pub async fn new() -> Result { - let device = Device::Cpu; - let model_id = "thenlper/gte-small".to_string(); - - let repo = Repo::model(model_id); + pub async fn load_model_files() -> Result<(PathBuf, PathBuf, PathBuf)> { + let repo = Repo::model("thenlper/gte-small".to_string()); let cache = Cache::default().repo(repo.clone()); @@ -132,9 +129,28 @@ impl ModelInstance { .ok_or(Error::msg("could not get tokenizer.json"))?, cache .get("model.safetensors") - .or_else(|| api.get("model.safetensors").ok()) + .and_then(|p| { + tracing::info!("Found embedding model in cache"); + Some(p) + }) + .or_else(|| { + tracing::info!("Downloading embedding model..."); + api.get("model.safetensors").ok().and_then(|p| { + tracing::info!("Downloaded embedding model"); + Some(p) + }) + }) .ok_or(Error::msg("could not get model.safetensors"))?, ); + + Ok((config_filename, tokenizer_filename, weights_filename)) + } + + pub async fn new() -> Result { + tracing::info!("Loading embedding model..."); + let device = Device::Cpu; + let (config_filename, tokenizer_filename, weights_filename) = + Self::load_model_files().await?; let config = std::fs::read_to_string(config_filename)?; let config: Config = serde_json::from_str(&config)?; let tokenizer = Tokenizer::from( @@ -149,6 +165,7 @@ impl ModelInstance { let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights_filename], DTYPE, &device)? }; let model = BertModel::load(vb, &config)?; + tracing::info!("Loaded embedding model"); Ok(Self { model, tokenizer }) } @@ -439,6 +456,7 @@ pub fn global_service(db: &Pool) -> Router { if let Ok(model_instance) = model_instance { let model_instance = Arc::new(model_instance); loop { + tracing::info!("Creating embeddings DB..."); let new_embeddings_db = EmbeddingsDb::new(&db_clone, model_instance.clone()).await; if let Err(e) = new_embeddings_db.as_ref() { @@ -446,7 +464,7 @@ pub fn global_service(db: &Pool) -> Router { } else { let mut embeddings_db = embeddings_clone.write().await; *embeddings_db = new_embeddings_db.ok(); - tracing::info!("Loaded embeddings DB"); + tracing::info!("Created embeddings DB"); } tokio::time::sleep(std::time::Duration::from_secs(3600 * 24)).await; diff --git a/backend/windmill-api/src/jobs.rs b/backend/windmill-api/src/jobs.rs index b8eb317616..635f51106e 100644 --- a/backend/windmill-api/src/jobs.rs +++ b/backend/windmill-api/src/jobs.rs @@ -41,7 +41,7 @@ use windmill_common::{ db::UserDB, error::{self, to_anyhow, Error}, flow_status::{Approval, FlowStatus, FlowStatusModule}, - flows::{FlowValue, InputTransform}, + flows::FlowValue, jobs::{script_path_to_payload, JobKind, JobPayload, QueuedJob, RawCode}, oauth2::HmacSha256, scripts::{Script, ScriptHash, ScriptLang}, diff --git a/backend/windmill-api/src/lib.rs b/backend/windmill-api/src/lib.rs index 0e5132be43..e88685c681 100644 --- a/backend/windmill-api/src/lib.rs +++ b/backend/windmill-api/src/lib.rs @@ -48,7 +48,7 @@ mod configs; mod db; mod drafts; pub mod ee; -mod embeddings; +pub mod embeddings; mod favorite; mod flows; mod folders;