1use 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#[derive(Args, Debug)]
15pub struct EmbedCommand {
16 #[arg(required = true)]
18 pub model: String,
19
20 #[arg(short, long)]
22 pub text: Option<String>,
23
24 #[arg(short, long)]
26 pub image: Option<String>,
27
28 #[arg(short, long, default_value = "json")]
30 pub format: String,
31
32 #[arg(short, long, default_value = "true")]
34 pub normalize: bool,
35
36 #[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 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 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 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 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 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
190pub 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 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 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 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 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 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 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 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}