@@ -10,7 +10,7 @@ use anyhow::Context;
1010use async_trait:: async_trait;
1111use ndarray:: Array2 ;
1212use tokenizers:: Tokenizer ;
13- use tract_onnx:: prelude:: { Framework , Tensor , TValue , tvec } ;
13+ use tract_onnx:: prelude:: { tvec , Framework , TValue , Tensor } ;
1414
1515use super :: Embedder ;
1616
@@ -22,7 +22,11 @@ fn huggingface_cache_dir() -> PathBuf {
2222 . join ( "hub" )
2323}
2424
25- fn ensure_cached ( model_name : & str , onnx_path : & PathBuf , tokenizer_path : & PathBuf ) -> anyhow:: Result < ( ) > {
25+ fn ensure_cached (
26+ model_name : & str ,
27+ onnx_path : & PathBuf ,
28+ tokenizer_path : & PathBuf ,
29+ ) -> anyhow:: Result < ( ) > {
2630 if onnx_path. exists ( ) && tokenizer_path. exists ( ) {
2731 return Ok ( ( ) ) ;
2832 }
@@ -57,17 +61,19 @@ fn download_from_huggingface(repo: &str, path: &str, dest: &PathBuf) -> anyhow::
5761 . with_context ( || format ! ( "tract: download {} from HF: {}" , path, url) ) ?
5862 . error_for_status ( )
5963 . with_context ( || format ! ( "tract: HF download failed for {}: {}" , path, url) ) ?;
60- let bytes = response. bytes ( )
64+ let bytes = response
65+ . bytes ( )
6166 . with_context ( || format ! ( "tract: read body for {} from HF" , path) ) ?;
62- std:: fs:: write ( dest, bytes)
63- . with_context ( || format ! ( "tract: write {} to cache" , path) ) ?;
67+ std:: fs:: write ( dest, bytes) . with_context ( || format ! ( "tract: write {} to cache" , path) ) ?;
6468 Ok ( ( ) )
6569}
6670
6771fn normalize_l2 ( mut v : Vec < f32 > ) -> Vec < f32 > {
6872 let norm: f32 = v. iter ( ) . map ( |x| x * x) . sum :: < f32 > ( ) . sqrt ( ) ;
6973 if norm > 1e-9 {
70- for x in & mut v { * x /= norm; }
74+ for x in & mut v {
75+ * x /= norm;
76+ }
7177 }
7278 v
7379}
@@ -78,19 +84,25 @@ fn probe_dimension(model_path: &PathBuf) -> anyhow::Result<usize> {
7884 let inference_model = tract_onnx:: onnx ( )
7985 . model_for_path ( model_path)
8086 . context ( "tract: load for probe" ) ?;
81- let runnable = inference_model. into_runnable ( ) . context ( "tract: build runnable for probe" ) ?;
87+ let runnable = inference_model
88+ . into_runnable ( )
89+ . context ( "tract: build runnable for probe" ) ?;
8290
8391 // Token IDs as f32 (most ONNX models expect float inputs for token IDs)
8492 let input_ids: Vec < f32 > = vec ! [ 1.0 , 2.0 , 3.0 ] ;
8593 let input_tensor: Tensor = Array2 :: from_shape_vec ( ( 1 , 3 ) , input_ids)
8694 . map_err ( |e| anyhow:: anyhow!( "tract: probe shape: {}" , e) ) ?
8795 . into ( ) ;
8896 let input_tvalue: TValue = input_tensor. into ( ) ;
89- let result = runnable. run ( tvec ! ( input_tvalue) )
97+ let result = runnable
98+ . run ( tvec ! ( input_tvalue) )
9099 . map_err ( |e| anyhow:: anyhow!( "tract: probe run: {}" , e) ) ?;
91- let out = result. into_iter ( ) . next ( )
100+ let out = result
101+ . into_iter ( )
102+ . next ( )
92103 . ok_or_else ( || anyhow:: anyhow!( "tract: no probe output" ) ) ?;
93- let view = out. to_array_view :: < f32 > ( )
104+ let view = out
105+ . to_array_view :: < f32 > ( )
94106 . map_err ( |e| anyhow:: anyhow!( "tract: probe view: {}" , e) ) ?;
95107 let hidden = * view. shape ( ) . last ( ) . unwrap_or ( & 384 ) ;
96108 if hidden == 0 {
@@ -108,7 +120,10 @@ pub struct TractEmbedder {
108120}
109121
110122impl TractEmbedder {
111- pub fn with_model ( model_name : impl Into < String > , _cache_dir : Option < PathBuf > ) -> anyhow:: Result < Self > {
123+ pub fn with_model (
124+ model_name : impl Into < String > ,
125+ _cache_dir : Option < PathBuf > ,
126+ ) -> anyhow:: Result < Self > {
112127 let model_name_owned = model_name. into ( ) ;
113128
114129 let cache = huggingface_cache_dir ( ) ;
@@ -133,21 +148,32 @@ impl TractEmbedder {
133148
134149#[ async_trait]
135150impl Embedder for TractEmbedder {
136- fn dim ( & self ) -> usize { self . dim }
137- fn fingerprint ( & self ) -> & str { & self . fingerprint }
151+ fn dim ( & self ) -> usize {
152+ self . dim
153+ }
154+ fn fingerprint ( & self ) -> & str {
155+ & self . fingerprint
156+ }
138157
139158 async fn embed ( & self , text : & str ) -> anyhow:: Result < Vec < f32 > > {
140159 let mut out = self . embed_batch ( & [ text] ) . await ?;
141160 Ok ( out. pop ( ) . unwrap_or_default ( ) )
142161 }
143162
144163 async fn embed_batch ( & self , texts : & [ & str ] ) -> anyhow:: Result < Vec < Vec < f32 > > > {
145- if texts. is_empty ( ) { return Ok ( Vec :: new ( ) ) ; }
164+ if texts. is_empty ( ) {
165+ return Ok ( Vec :: new ( ) ) ;
166+ }
146167
147168 let tokenizer_path = self . tokenizer_path . clone ( ) ;
148169 let model_name = tokenizer_path
149170 . parent ( )
150- . map ( |p| p. file_name ( ) . unwrap_or_default ( ) . to_string_lossy ( ) . into_owned ( ) )
171+ . map ( |p| {
172+ p. file_name ( )
173+ . unwrap_or_default ( )
174+ . to_string_lossy ( )
175+ . into_owned ( )
176+ } )
151177 . unwrap_or_default ( ) ;
152178
153179 let owned: Vec < String > = texts. iter ( ) . map ( |s| ( * s) . to_owned ( ) ) . collect ( ) ;
@@ -185,11 +211,19 @@ fn run_embed_batch(
185211
186212 let encodings: Vec < _ > = texts
187213 . iter ( )
188- . map ( |s| tokenizer. encode ( s. as_str ( ) , true )
189- . map_err ( |e| anyhow:: anyhow!( "tract: tokenize: {}" , e) ) )
214+ . map ( |s| {
215+ tokenizer
216+ . encode ( s. as_str ( ) , true )
217+ . map_err ( |e| anyhow:: anyhow!( "tract: tokenize: {}" , e) )
218+ } )
190219 . collect :: < Result < Vec < _ > , _ > > ( ) ?;
191220
192- let max_len = encodings. iter ( ) . map ( |e| e. get_ids ( ) . len ( ) ) . max ( ) . unwrap_or ( 1 ) . min ( DEFAULT_MAX_LEN ) ;
221+ let max_len = encodings
222+ . iter ( )
223+ . map ( |e| e. get_ids ( ) . len ( ) )
224+ . max ( )
225+ . unwrap_or ( 1 )
226+ . min ( DEFAULT_MAX_LEN ) ;
193227 let batch_size = encodings. len ( ) ;
194228 let mut input_ids = vec ! [ 0f32 ; batch_size * max_len] ;
195229 let mut attention_mask = vec ! [ 0f32 ; batch_size * max_len] ;
@@ -206,23 +240,30 @@ fn run_embed_batch(
206240 let input_ids_tensor: Tensor = Array2 :: from_shape_vec ( ( batch_size, max_len) , input_ids)
207241 . map_err ( |e| anyhow:: anyhow!( "tract: input_ids: {}" , e) ) ?
208242 . into ( ) ;
209- let attention_mask_tensor: Tensor = Array2 :: from_shape_vec ( ( batch_size, max_len) , attention_mask)
210- . map_err ( |e| anyhow:: anyhow!( "tract: attention_mask: {}" , e) ) ?
211- . into ( ) ;
243+ let attention_mask_tensor: Tensor =
244+ Array2 :: from_shape_vec ( ( batch_size, max_len) , attention_mask)
245+ . map_err ( |e| anyhow:: anyhow!( "tract: attention_mask: {}" , e) ) ?
246+ . into ( ) ;
212247 let input_tvalue: TValue = input_ids_tensor. into ( ) ;
213248 let mask_tvalue: TValue = attention_mask_tensor. into ( ) ;
214249
215250 let result = runnable. run ( tvec ! ( input_tvalue, mask_tvalue) ) ?;
216251
217- let output = result. into_iter ( ) . next ( )
252+ let output = result
253+ . into_iter ( )
254+ . next ( )
218255 . ok_or_else ( || anyhow:: anyhow!( "tract: no output tensor" ) ) ?;
219256
220- let view = output. to_array_view :: < f32 > ( )
257+ let view = output
258+ . to_array_view :: < f32 > ( )
221259 . map_err ( |e| anyhow:: anyhow!( "tract: output view: {}" , e) ) ?;
222260
223261 let shape = view. shape ( ) ;
224262 if shape. len ( ) != 3 {
225- anyhow:: bail!( "tract: expected 3D output [batch, seq, dim], got {:?}" , shape) ;
263+ anyhow:: bail!(
264+ "tract: expected 3D output [batch, seq, dim], got {:?}" ,
265+ shape
266+ ) ;
226267 }
227268
228269 let seq_len = shape[ 1 ] ;
@@ -233,11 +274,15 @@ fn run_embed_batch(
233274 let mut sum = vec ! [ 0f32 ; dim] ;
234275 let mut count = 0f32 ;
235276 for j in 0 ..valid_len {
236- for k in 0 ..dim { sum[ k] += view[ [ i, j, k] ] ; }
277+ for k in 0 ..dim {
278+ sum[ k] += view[ [ i, j, k] ] ;
279+ }
237280 count += 1.0 ;
238281 }
239282 if count > 0.0 {
240- for x in & mut sum { * x /= count; }
283+ for x in & mut sum {
284+ * x /= count;
285+ }
241286 }
242287 embeddings. push ( normalize_l2 ( sum) ) ;
243288 }
0 commit comments