mirror of
https://github.com/windmill-labs/windmill.git
synced 2026-08-19 08:01:25 +00:00
fix: cache embedding model in docker img (#2474)
This commit is contained in:
@@ -191,6 +191,8 @@ RUN ln -s ${APP}/windmill /usr/local/bin/windmill
|
||||
|
||||
WORKDIR ${APP}
|
||||
|
||||
RUN windmill cache
|
||||
|
||||
EXPOSE 8000
|
||||
|
||||
CMD ["windmill"]
|
||||
|
||||
@@ -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| {
|
||||
|
||||
@@ -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<Self> {
|
||||
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<Self> {
|
||||
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<Postgres>) -> 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<Postgres>) -> 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;
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -48,7 +48,7 @@ mod configs;
|
||||
mod db;
|
||||
mod drafts;
|
||||
pub mod ee;
|
||||
mod embeddings;
|
||||
pub mod embeddings;
|
||||
mod favorite;
|
||||
mod flows;
|
||||
mod folders;
|
||||
|
||||
Reference in New Issue
Block a user