iwaitu commited on
Commit
6fa892b
·
verified ·
1 Parent(s): 30b2656

Create README.md

Browse files
Files changed (1) hide show
  1. README.md +60 -0
README.md ADDED
@@ -0,0 +1,60 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ base_model:
4
+ - Qwen/Qwen3-Embedding-0.6B
5
+ ---
6
+
7
+ onnx version.
8
+
9
+
10
+ import numpy as np
11
+ import onnxruntime as ort
12
+ from transformers import AutoTokenizer
13
+
14
+
15
+ class QwenEmbeddingONNX:
16
+ def __init__(self, model_dir, model_path, max_length=512, providers=None):
17
+ self.max_length = max_length
18
+ self.tokenizer = AutoTokenizer.from_pretrained(model_dir)
19
+
20
+ if providers is None:
21
+ providers = ['CPUExecutionProvider']
22
+
23
+ self.session = ort.InferenceSession(
24
+ model_path,
25
+ providers=providers
26
+ )
27
+
28
+ self.input_names = [i.name for i in self.session.get_inputs()]
29
+
30
+ def encode(self, texts, normalize=True):
31
+ if isinstance(texts, str):
32
+ texts = [texts]
33
+
34
+ inputs = self.tokenizer(
35
+ texts,
36
+ padding=True,
37
+ truncation=True,
38
+ max_length=self.max_length,
39
+ return_tensors="np",
40
+ )
41
+
42
+ input_feed = {
43
+ k: inputs[k].astype(np.int64)
44
+ for k in self.input_names
45
+ if k in inputs
46
+ }
47
+
48
+ outputs = self.session.run(None, input_feed)
49
+ embeddings = outputs[0]
50
+
51
+ # 如果输出是 [B, L, H],取第一个 token
52
+ if embeddings.ndim == 3:
53
+ embeddings = embeddings[:, 0, :]
54
+
55
+ # L2 normalize
56
+ if normalize:
57
+ norm = np.linalg.norm(embeddings, axis=1, keepdims=True)
58
+ embeddings = embeddings / (norm + 1e-9)
59
+
60
+ return embeddings