Files
Maven/internal/router/onnxembedder.go
T
2026-07-03 00:32:48 +02:00

317 lines
6.8 KiB
Go

package router
import (
"context"
"encoding/json"
"fmt"
"math"
"os"
"strings"
ort "github.com/yalue/onnxruntime_go"
"golang.org/x/text/unicode/norm"
)
const (
padTokenID = 1
unkTokenID = 3
clsTokenID = 0
sepTokenID = 2
maxLength = 128
embedDim = 384
)
type onnxEmbedder struct {
tokenizer *unigramTokenizer
session *ort.DynamicSession[int64, float32]
}
func NewONNXEmbedder(modelPath, tokenizerPath, libPath string) (*onnxEmbedder, error) {
ort.SetSharedLibraryPath(libPath)
if err := ort.InitializeEnvironment(); err != nil {
return nil, fmt.Errorf("onnx: init environment: %w", err)
}
tok, err := newUnigramTokenizer(tokenizerPath)
if err != nil {
return nil, fmt.Errorf("tokenizer: %w", err)
}
session, err := ort.NewDynamicSession[int64, float32](
modelPath,
[]string{"input_ids", "attention_mask", "token_type_ids"},
[]string{"last_hidden_state"},
)
if err != nil {
return nil, fmt.Errorf("onnx: create session: %w", err)
}
return &onnxEmbedder{
tokenizer: tok,
session: session,
}, nil
}
func (e *onnxEmbedder) Dim() int { return embedDim }
func (e *onnxEmbedder) Embed(ctx context.Context, text string) ([]float32, error) {
inputIDs, attentionMask, _ := e.tokenizer.Encode(text)
inputShape := ort.NewShape(1, int64(maxLength))
inputT, err := ort.NewTensor(inputShape, inputIDs)
if err != nil {
return nil, fmt.Errorf("onnx: input tensor: %w", err)
}
defer inputT.Destroy()
maskT, err := ort.NewTensor(inputShape, attentionMask)
if err != nil {
return nil, fmt.Errorf("onnx: mask tensor: %w", err)
}
defer maskT.Destroy()
typeT, err := ort.NewTensor(inputShape, make([]int64, maxLength))
if err != nil {
return nil, fmt.Errorf("onnx: type tensor: %w", err)
}
defer typeT.Destroy()
outputShape := ort.NewShape(1, int64(maxLength), embedDim)
outputT, err := ort.NewTensor(outputShape, make([]float32, maxLength*embedDim))
if err != nil {
return nil, fmt.Errorf("onnx: output tensor: %w", err)
}
defer outputT.Destroy()
if err := e.session.Run(
[]*ort.Tensor[int64]{inputT, maskT, typeT},
[]*ort.Tensor[float32]{outputT},
); err != nil {
return nil, fmt.Errorf("onnx: run: %w", err)
}
emb := meanPool(outputT.GetData(), attentionMask, maxLength, embedDim)
return emb, nil
}
func (e *onnxEmbedder) Close() error {
e.session.Destroy()
return nil
}
func meanPool(hidden []float32, mask []int64, seqLen, dim int) []float32 {
out := make([]float32, dim)
var maskSum float32
for i := 0; i < seqLen; i++ {
if mask[i] == 0 {
continue
}
maskSum++
for j := 0; j < dim; j++ {
out[j] += hidden[i*dim+j]
}
}
if maskSum > 0 {
for j := 0; j < dim; j++ {
out[j] /= maskSum
}
}
var sumSq float64
for _, v := range out {
sumSq += float64(v) * float64(v)
}
if sumSq > 0 {
inv := float32(1.0 / math.Sqrt(sumSq))
for i := range out {
out[i] *= inv
}
}
return out
}
type unigramTokenizer struct {
vocab map[string]vocabEntry
unkScore float64
}
type vocabEntry struct {
id int64
score float64
}
func newUnigramTokenizer(path string) (*unigramTokenizer, error) {
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("read tokenizer.json: %w", err)
}
var raw struct {
Model struct {
Type string `json:"type"`
Vocab json.RawMessage `json:"vocab"`
} `json:"model"`
}
if err := json.Unmarshal(data, &raw); err != nil {
return nil, fmt.Errorf("parse tokenizer.json: %w", err)
}
if raw.Model.Type != "Unigram" {
return nil, fmt.Errorf("unsupported tokenizer type: %s", raw.Model.Type)
}
var rawVocab [][]json.RawMessage
if err := json.Unmarshal(raw.Model.Vocab, &rawVocab); err != nil {
return nil, fmt.Errorf("parse vocab: %w", err)
}
vocab := make(map[string]vocabEntry, len(rawVocab))
var unkScore float64
for _, pair := range rawVocab {
if len(pair) < 2 {
continue
}
var token string
if err := json.Unmarshal(pair[0], &token); err != nil {
continue
}
var score float64
if err := json.Unmarshal(pair[1], &score); err != nil {
continue
}
vocab[token] = vocabEntry{score: score}
}
// Assign IDs based on order
i := int64(0)
for _, pair := range rawVocab {
var token string
if err := json.Unmarshal(pair[0], &token); err != nil {
continue
}
e := vocab[token]
e.id = i
vocab[token] = e
if i == unkTokenID {
unkScore = e.score
}
i++
}
return &unigramTokenizer{vocab: vocab, unkScore: unkScore}, nil
}
func (t *unigramTokenizer) Encode(text string) (inputIDs, attentionMask, tokenTypeIDs []int64) {
tokens := t.tokenize(text)
tokens = append([]int64{clsTokenID}, tokens...)
tokens = append(tokens, sepTokenID)
if len(tokens) > maxLength {
tokens = tokens[:maxLength-1]
tokens = append(tokens, sepTokenID)
}
inputIDs = make([]int64, maxLength)
attentionMask = make([]int64, maxLength)
tokenTypeIDs = make([]int64, maxLength)
for i, id := range tokens {
inputIDs[i] = id
attentionMask[i] = 1
}
return
}
func (t *unigramTokenizer) tokenize(text string) []int64 {
words := preTokenize(text)
var ids []int64
for _, word := range words {
wordIDs := t.encodeWord(word)
ids = append(ids, wordIDs...)
}
return ids
}
type cand struct {
start int
end int
id int64
score float64
}
func (t *unigramTokenizer) encodeWord(word string) []int64 {
runes := []rune(word)
n := len(runes)
if n == 0 {
return nil
}
var candidates []cand
for i := 0; i < n; i++ {
for j := i + 1; j <= n && j-i <= 50; j++ {
sub := string(runes[i:j])
if e, ok := t.vocab[sub]; ok {
candidates = append(candidates, cand{
start: i, end: j, id: e.id, score: e.score,
})
}
}
}
dp := make([]float64, n+1)
prev := make([]int, n+1)
bestID := make([]int64, n+1)
filled := make([]bool, n+1)
dp[0] = 0
filled[0] = true
for i := 1; i <= n; i++ {
bestScore := math.Inf(-1)
bestPrev := -1
bestTokenID := int64(unkTokenID)
for _, c := range candidates {
if c.end == i && filled[c.start] {
candScore := dp[c.start] + c.score
if candScore > bestScore {
bestScore = candScore
bestPrev = c.start
bestTokenID = c.id
}
}
}
if bestScore == math.Inf(-1) {
if filled[i-1] {
dp[i] = dp[i-1] + t.unkScore
prev[i] = i - 1
bestID[i] = unkTokenID
filled[i] = true
}
} else {
dp[i] = bestScore
prev[i] = bestPrev
bestID[i] = bestTokenID
filled[i] = true
}
}
var result []int64
for i := n; i > 0; i = prev[i] {
result = append([]int64{bestID[i]}, result...)
}
// Reverse
for l, r := 0, len(result)-1; l < r; l, r = l+1, r-1 {
result[l], result[r] = result[r], result[l]
}
return result
}
func preTokenize(text string) []string {
text = norm.NFKC.String(text)
text = strings.ToLower(text)
pieces := strings.Fields(text)
out := make([]string, 0, len(pieces))
for _, p := range pieces {
out = append(out, "\u2581"+p)
}
return out
}
var _ Embedder = (*onnxEmbedder)(nil)