Forge
csharp72c0d4fb
1using Microsoft.ML.OnnxRuntime;
2using Microsoft.ML.OnnxRuntime.Tensors;
3
4namespace HybridCodebaseIndex.Core.Embeddings;
5
6internal sealed class OnnxEmbeddingProvider : IEmbeddingProvider, IDisposable
7{
8 private readonly InferenceSession _session;
9 private readonly WordPieceTokenizer _tokenizer;
10 private readonly string _inputIdsName;
11 private readonly string _attentionMaskName;
12 private readonly string? _tokenTypeIdsName;
13 private readonly string _outputName;
14 private readonly int _seqLen;
15 private readonly object _gate = new();
16
17 public int Dimension { get; }
18
19 public OnnxEmbeddingProvider(string modelPath, string? vocabPath, bool doLowerCase, int seqLen, bool preferGpu)
20 {
21 if (string.IsNullOrWhiteSpace(modelPath))
22 throw new ArgumentException("embedding_model_path is required for embedding_provider=onnx.", nameof(modelPath));
23 if (string.IsNullOrWhiteSpace(vocabPath))
24 throw new ArgumentException("embedding_vocab_path is required for embedding_provider=onnx.", nameof(vocabPath));
25
26 _seqLen = Math.Clamp(seqLen, 32, 512);
27
28 var opts = new SessionOptions();
29 if (preferGpu)
30 {
31 try
32 {
33 // Optional: CUDA EP needs matching native runtime (GPU package / machine); otherwise catch → CPU.
34 opts.AppendExecutionProvider_CUDA(0);
35 }
36 catch
37 {
38 // fall back to CPU
39 }
40 }
41
42 _session = new InferenceSession(modelPath, opts);
43 _tokenizer = WordPieceTokenizer.FromVocabFile(vocabPath, doLowerCase);
44
45 var inputs = _session.InputMetadata;
46 _inputIdsName = FindKey(inputs, ["input_ids", "inputIds", "input_ids:0"]) ?? inputs.Keys.First();
47 _attentionMaskName = FindKey(inputs, ["attention_mask", "attentionMask"]) ?? inputs.Keys.Skip(1).FirstOrDefault() ?? "attention_mask";
48 _tokenTypeIdsName = FindKey(inputs, ["token_type_ids", "tokenTypeIds"]);
49
50 var outputs = _session.OutputMetadata;
51 _outputName = FindKey(outputs, ["sentence_embedding", "embeddings", "pooled_output", "last_hidden_state"]) ?? outputs.Keys.First();
52
53 // Infer dimension from output metadata when possible
54 var dimsArr = outputs[_outputName].Dimensions.ToArray();
55 // Common: [1, hidden] (pooled) or [1, seq, hidden] (last_hidden_state)
56 Dimension = dimsArr.Length >= 2 && dimsArr[^1] > 0 ? dimsArr[^1] : 384;
57 }
58
59 public ValueTask<float[]> EmbedAsync(string text, CancellationToken cancellationToken)
60 {
61 cancellationToken.ThrowIfCancellationRequested();
62 var (ids, mask, typeIdsArr) = _tokenizer.Encode(text ?? "", _seqLen);
63
64 var inputIds = new DenseTensor<long>(new[] { 1, _seqLen });
65 var attn = new DenseTensor<long>(new[] { 1, _seqLen });
66 var typeIds = _tokenTypeIdsName is null ? null : new DenseTensor<long>(new[] { 1, _seqLen });
67
68 for (var i = 0; i < _seqLen; i++)
69 {
70 inputIds[0, i] = ids[i];
71 attn[0, i] = mask[i];
72 if (typeIds is not null)
73 typeIds[0, i] = typeIdsArr[i];
74 }
75
76 var inputs = new List<NamedOnnxValue>(capacity: typeIds is null ? 2 : 3)
77 {
78 NamedOnnxValue.CreateFromTensor(_inputIdsName, inputIds),
79 NamedOnnxValue.CreateFromTensor(_attentionMaskName, attn),
80 };
81 if (typeIds is not null && _tokenTypeIdsName is not null)
82 inputs.Add(NamedOnnxValue.CreateFromTensor(_tokenTypeIdsName, typeIds));
83
84 lock (_gate)
85 {
86 using var results = _session.Run(inputs);
87 var first = results.First(r => string.Equals(r.Name, _outputName, StringComparison.OrdinalIgnoreCase) || r.Name == _outputName);
88
89 // Try pooled embedding first
90 if (first.Value is DenseTensor<float> pooled && pooled.Rank == 2)
91 {
92 var v = new float[pooled.Dimensions[1]];
93 for (var j = 0; j < v.Length; j++)
94 v[j] = pooled[0, j];
95 NormalizeInPlace(v);
96 return ValueTask.FromResult(v);
97 }
98
99 // Fallback: mean pool last_hidden_state with attention mask
100 var hs = first.AsTensor<float>();
101 if (hs.Rank == 3)
102 {
103 var hidden = hs.Dimensions[2];
104 var v = new float[hidden];
105 double denom = 0;
106 for (var t = 0; t < _seqLen; t++)
107 {
108 if (attn[0, t] == 0)
109 continue;
110 denom += 1;
111 for (var j = 0; j < hidden; j++)
112 v[j] += hs[0, t, j];
113 }
114 if (denom > 0)
115 {
116 var inv = (float)(1.0 / denom);
117 for (var j = 0; j < v.Length; j++)
118 v[j] *= inv;
119 }
120 NormalizeInPlace(v);
121 return ValueTask.FromResult(v);
122 }
123
124 throw new InvalidOperationException($"Unexpected ONNX output shape for '{_outputName}'.");
125 }
126 }
127
128 public void Dispose() => _session.Dispose();
129
130 private static string? FindKey<T>(IReadOnlyDictionary<string, T> dict, string[] candidates)
131 {
132 foreach (var c in candidates)
133 {
134 foreach (var k in dict.Keys)
135 {
136 if (string.Equals(k, c, StringComparison.OrdinalIgnoreCase))
137 return k;
138 }
139 }
140 return null;
141 }
142
143 private static void NormalizeInPlace(float[] v)
144 {
145 double sum = 0;
146 foreach (var x in v)
147 sum += x * x;
148 var norm = Math.Sqrt(sum);
149 if (norm <= 1e-12)
150 return;
151 var inv = (float)(1.0 / norm);
152 for (var i = 0; i < v.Length; i++)
153 v[i] *= inv;
154 }
155}
156
157
View only · write via MCP/CIDE