Skip to main content

ferrum_cli/commands/
embed.rs

1//! Embed command - Generate embeddings using BERT models
2
3use crate::config::CliConfig;
4use clap::Args;
5use colored::Colorize;
6use ferrum_models::source::{ModelFormat, ResolvedModelSource};
7use ferrum_models::HfDownloader;
8use ferrum_models::{BertModelExecutor, ConfigManager};
9use ferrum_types::Result;
10use std::io::{self, BufRead};
11use std::path::PathBuf;
12
13/// Generate embeddings using BERT or CLIP models
14#[derive(Args, Debug)]
15pub struct EmbedCommand {
16    /// Model name (e.g., google-bert/bert-base-chinese, OFA-Sys/chinese-clip-vit-base-patch16)
17    #[arg(required = true)]
18    pub model: String,
19
20    /// Text to embed (if not provided, reads from stdin)
21    #[arg(short, long)]
22    pub text: Option<String>,
23
24    /// Image path to embed (CLIP models only)
25    #[arg(short, long)]
26    pub image: Option<String>,
27
28    /// Output format: json, csv, or raw
29    #[arg(short, long, default_value = "json")]
30    pub format: String,
31
32    /// Normalize embeddings to unit length
33    #[arg(short, long, default_value = "true")]
34    pub normalize: bool,
35
36    /// Backend: auto, cpu, cuda, metal (default: auto)
37    #[arg(short, long, default_value = "auto")]
38    pub backend: String,
39}
40
41pub async fn execute(cmd: EmbedCommand, config: CliConfig) -> Result<()> {
42    eprintln!("{}", format!("Loading {}...", cmd.model).dimmed());
43
44    // Resolve model path using same logic as list/run commands
45    let model_id = cmd.model.clone();
46    let cache_dir = get_hf_cache_dir(&config);
47
48    let source = match find_cached_model(&cache_dir, &model_id) {
49        Some(source) => source,
50        None => {
51            eprintln!(
52                "{} Model '{}' not found locally, downloading...",
53                "📥".cyan(),
54                model_id
55            );
56
57            let token = std::env::var("HF_TOKEN")
58                .or_else(|_| std::env::var("HUGGING_FACE_HUB_TOKEN"))
59                .ok();
60
61            let downloader = HfDownloader::new(cache_dir, token)?;
62            let snapshot_path = downloader.download(&model_id, None).await?;
63
64            let format = detect_format(&snapshot_path);
65            if format == ModelFormat::Unknown {
66                return Err(ferrum_types::FerrumError::model(
67                    "Downloaded model has unknown format",
68                ));
69            }
70
71            ResolvedModelSource {
72                original: model_id.clone(),
73                local_path: snapshot_path,
74                format,
75                from_cache: false,
76            }
77        }
78    };
79
80    let model_path = source.local_path.to_string_lossy().to_string();
81
82    // Load model definition to detect architecture
83    let mut config_manager = ConfigManager::new();
84    let model_def = config_manager.load_from_path(&source.local_path).await?;
85
86    let device = crate::commands::transcribe::select_candle_device(&cmd.backend)?;
87    eprintln!("{} {:?}", "Device:".dimmed(), &device);
88    let is_clip = model_def.architecture == ferrum_models::Architecture::Clip;
89
90    let mut all_embeddings: Vec<(String, Vec<f32>)> = Vec::new();
91
92    if is_clip {
93        // CLIP path: supports both text and image
94        let executor = ferrum_models::ClipModelExecutor::from_path(
95            &model_path,
96            device.clone(),
97            candle_core::DType::F32,
98        )?;
99        eprintln!("{}", "CLIP model loaded.".green());
100
101        if let Some(ref image_path) = cmd.image {
102            let embedding_tensor = executor.embed_image_path(image_path)?;
103            let embedding = tensor_to_vec(&embedding_tensor, cmd.normalize)?;
104            all_embeddings.push((format!("[image] {image_path}"), embedding));
105        }
106
107        let texts = collect_texts(&cmd)?;
108        if !texts.is_empty() {
109            let tokenizer = load_tokenizer(&source.local_path)?;
110
111            for text in &texts {
112                let encoding = tokenizer
113                    .encode(text.as_str(), true)
114                    .map_err(|e| ferrum_types::FerrumError::model(format!("Tokenize: {e}")))?;
115                let embedding_tensor = executor.embed_text(encoding.get_ids())?;
116                let embedding = tensor_to_vec(&embedding_tensor, cmd.normalize)?;
117                all_embeddings.push((text.clone(), embedding));
118            }
119        }
120
121        if all_embeddings.is_empty() {
122            eprintln!("{}", "No input provided. Use --text or --image.".yellow());
123            return Ok(());
124        }
125    } else {
126        // BERT path (existing)
127        let executor = BertModelExecutor::from_path(&model_path, &model_def, device).await?;
128        eprintln!("{}", "BERT model loaded.".green());
129
130        let tokenizer = load_tokenizer(&source.local_path)?;
131
132        let texts = collect_texts(&cmd)?;
133        if texts.is_empty() {
134            eprintln!("{}", "No text provided.".yellow());
135            return Ok(());
136        }
137
138        for text in &texts {
139            let encoding = tokenizer
140                .encode(text.as_str(), true)
141                .map_err(|e| ferrum_types::FerrumError::model(format!("Tokenize: {e}")))?;
142            let embedding_tensor = executor.get_embeddings(encoding.get_ids())?;
143            let embedding = tensor_to_vec(&embedding_tensor, cmd.normalize)?;
144            all_embeddings.push((text.clone(), embedding));
145        }
146    }
147
148    // Output embeddings
149    match cmd.format.as_str() {
150        "json" => {
151            let output: Vec<serde_json::Value> = all_embeddings
152                .iter()
153                .map(|(text, emb)| {
154                    serde_json::json!({
155                        "text": text,
156                        "embedding": emb,
157                        "dimensions": emb.len()
158                    })
159                })
160                .collect();
161            println!("{}", serde_json::to_string_pretty(&output).unwrap());
162        }
163        "csv" => {
164            if let Some((_, first_emb)) = all_embeddings.first() {
165                let header: Vec<String> =
166                    (0..first_emb.len()).map(|i| format!("dim_{}", i)).collect();
167                println!("text,{}", header.join(","));
168            }
169            for (text, emb) in &all_embeddings {
170                let emb_str: Vec<String> = emb.iter().map(|v| format!("{:.6}", v)).collect();
171                println!("\"{}\",{}", text.replace("\"", "\\\""), emb_str.join(","));
172            }
173        }
174        "raw" => {
175            for (text, emb) in &all_embeddings {
176                eprintln!("{}: {} dimensions", text.dimmed(), emb.len());
177                let preview: Vec<String> =
178                    emb.iter().take(5).map(|v| format!("{:.4}", v)).collect();
179                println!("[{}, ...]", preview.join(", "));
180            }
181        }
182        _ => {
183            eprintln!("{}", format!("Unknown format: {}", cmd.format).red());
184        }
185    }
186
187    Ok(())
188}
189
190/// Load tokenizer: try tokenizer.json first, fall back to vocab.txt (BERT-style).
191pub fn load_tokenizer(model_dir: &std::path::Path) -> Result<tokenizers::Tokenizer> {
192    let tokenizer_json = model_dir.join("tokenizer.json");
193    if tokenizer_json.exists() {
194        return tokenizers::Tokenizer::from_file(&tokenizer_json)
195            .map_err(|e| ferrum_types::FerrumError::model(format!("Load tokenizer.json: {e}")));
196    }
197
198    // Fall back to vocab.txt (Chinese-CLIP, older BERT models)
199    let vocab_txt = model_dir.join("vocab.txt");
200    if vocab_txt.exists() {
201        use tokenizers::models::wordpiece::WordPiece;
202        use tokenizers::processors::template::TemplateProcessing;
203        use tokenizers::Model;
204        let wp = WordPiece::from_file(vocab_txt.to_str().unwrap())
205            .unk_token("[UNK]".to_string())
206            .build()
207            .map_err(|e| ferrum_types::FerrumError::model(format!("Load vocab.txt: {e}")))?;
208
209        // Look up [CLS] and [SEP] IDs from vocab dynamically
210        let vocab = wp.get_vocab();
211        let cls_id = vocab.get("[CLS]").copied().unwrap_or(101);
212        let sep_id = vocab.get("[SEP]").copied().unwrap_or(102);
213
214        let mut tokenizer = tokenizers::Tokenizer::new(wp);
215        // BertPreTokenizer + Chinese char splitting (add spaces around CJK chars
216        // so WordPiece treats each character as a separate token, matching Python's
217        // BertTokenizer._tokenize_chinese_chars behavior)
218        use tokenizers::pre_tokenizers::sequence::Sequence;
219        use tokenizers::pre_tokenizers::unicode_scripts::UnicodeScripts;
220        tokenizer.with_pre_tokenizer(Some(Sequence::new(vec![
221            tokenizers::pre_tokenizers::PreTokenizerWrapper::UnicodeScripts(UnicodeScripts),
222            tokenizers::pre_tokenizers::PreTokenizerWrapper::BertPreTokenizer(
223                tokenizers::pre_tokenizers::bert::BertPreTokenizer,
224            ),
225        ])));
226        let template = TemplateProcessing::builder()
227            .try_single("[CLS] $A [SEP]")
228            .unwrap()
229            .special_tokens(vec![("[CLS]", cls_id), ("[SEP]", sep_id)])
230            .build()
231            .map_err(|e| ferrum_types::FerrumError::model(format!("Template: {e}")))?;
232        tokenizer.with_post_processor(Some(template));
233        return Ok(tokenizer);
234    }
235
236    Err(ferrum_types::FerrumError::model(
237        "No tokenizer.json or vocab.txt found in model directory",
238    ))
239}
240
241fn collect_texts(cmd: &EmbedCommand) -> Result<Vec<String>> {
242    if let Some(ref text) = cmd.text {
243        Ok(vec![text.clone()])
244    } else if cmd.image.is_some() {
245        // Image-only mode, no text needed
246        Ok(vec![])
247    } else {
248        eprintln!(
249            "{}",
250            "Reading text from stdin (one per line, Ctrl+D to finish):".dimmed()
251        );
252        let stdin = io::stdin();
253        Ok(stdin.lock().lines().filter_map(|l| l.ok()).collect())
254    }
255}
256
257fn tensor_to_vec(tensor: &candle_core::Tensor, normalize: bool) -> Result<Vec<f32>> {
258    let mut embedding = tensor
259        .flatten_all()
260        .map_err(|e| ferrum_types::FerrumError::model(format!("Flatten: {e}")))?
261        .to_vec1::<f32>()
262        .map_err(|e| ferrum_types::FerrumError::model(format!("to_vec1: {e}")))?;
263
264    if normalize {
265        let norm: f32 = embedding.iter().map(|x| x * x).sum::<f32>().sqrt();
266        if norm > 0.0 {
267            for v in &mut embedding {
268                *v /= norm;
269            }
270        }
271    }
272    Ok(embedding)
273}
274
275fn find_cached_model(cache_dir: &PathBuf, model_id: &str) -> Option<ResolvedModelSource> {
276    // HuggingFace cache structure: hub/models--Org--ModelName/snapshots/<hash>/
277    let hub_dir = cache_dir.join("hub");
278    let model_dir_name = format!("models--{}", model_id.replace("/", "--"));
279    let model_dir = hub_dir.join(&model_dir_name);
280
281    if model_dir.exists() {
282        // Find the latest snapshot
283        let snapshots_dir = model_dir.join("snapshots");
284        if snapshots_dir.exists() {
285            if let Ok(entries) = std::fs::read_dir(&snapshots_dir) {
286                for entry in entries.filter_map(|e| e.ok()) {
287                    let snapshot_path = entry.path();
288                    if snapshot_path.is_dir() && snapshot_path.join("config.json").exists() {
289                        let format = detect_format(&snapshot_path);
290                        if format != ModelFormat::Unknown {
291                            return Some(ResolvedModelSource {
292                                original: model_id.to_string(),
293                                local_path: snapshot_path,
294                                format,
295                                from_cache: true,
296                            });
297                        }
298                    }
299                }
300            }
301        }
302    }
303
304    // Also check direct path (for models downloaded to custom locations)
305    let direct = cache_dir.join(model_id);
306    if direct.exists() && direct.join("config.json").exists() {
307        let format = detect_format(&direct);
308        if format != ModelFormat::Unknown {
309            return Some(ResolvedModelSource {
310                original: model_id.to_string(),
311                local_path: direct,
312                format,
313                from_cache: true,
314            });
315        }
316    }
317
318    None
319}
320
321fn get_hf_cache_dir(config: &CliConfig) -> PathBuf {
322    if let Ok(hf_home) = std::env::var("HF_HOME") {
323        return PathBuf::from(hf_home);
324    }
325    let configured = shellexpand::tilde(&config.models.download.hf_cache_dir).to_string();
326    PathBuf::from(configured)
327}
328
329fn detect_format(path: &PathBuf) -> ModelFormat {
330    if path.join("model.safetensors").exists() {
331        ModelFormat::SafeTensors
332    } else if std::fs::read_dir(path)
333        .map(|d| {
334            d.filter_map(|e| e.ok())
335                .any(|e| e.path().extension().is_some_and(|ext| ext == "safetensors"))
336        })
337        .unwrap_or(false)
338    {
339        ModelFormat::SafeTensors
340    } else if path.join("pytorch_model.bin").exists() {
341        ModelFormat::PyTorchBin
342    } else {
343        ModelFormat::Unknown
344    }
345}