359 lines
8.4 KiB
Go
359 lines
8.4 KiB
Go
// Package embedding provides text embedding using ONNX-based sentence transformers.
|
|
// It loads an all-MiniLM-L6-v2 model and tokenizer, and provides functions to
|
|
// embed text and compute similarity scores.
|
|
package embedding
|
|
|
|
import (
|
|
"encoding/json"
|
|
"fmt"
|
|
"math"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"sync"
|
|
"unicode"
|
|
|
|
ort "github.com/yalue/onnxruntime_go"
|
|
)
|
|
|
|
// Model holds the ONNX model and tokenizer state.
|
|
type Model struct {
|
|
session *ort.DynamicAdvancedSession
|
|
tokenizer *Tokenizer
|
|
mu sync.Mutex
|
|
}
|
|
|
|
// Vector is a 384-dimensional embedding vector (MiniLM output size).
|
|
type Vector []float32
|
|
|
|
var initOnce sync.Once
|
|
var initErr error
|
|
|
|
// LoadModel loads the ONNX model and tokenizer from the given directory.
|
|
// The directory must contain model.onnx, tokenizer.json, and libonnxruntime.so.
|
|
func LoadModel(modelDir string) (*Model, error) {
|
|
modelPath := filepath.Join(modelDir, "model.onnx")
|
|
tokenizerPath := filepath.Join(modelDir, "tokenizer.json")
|
|
libPath := filepath.Join(modelDir, "libonnxruntime.so")
|
|
|
|
// Initialize ONNX runtime once
|
|
initOnce.Do(func() {
|
|
ort.SetSharedLibraryPath(libPath)
|
|
initErr = ort.InitializeEnvironment()
|
|
})
|
|
if initErr != nil {
|
|
return nil, fmt.Errorf("init onnx environment: %w", initErr)
|
|
}
|
|
|
|
// Load tokenizer
|
|
tok, err := loadTokenizer(tokenizerPath)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("load tokenizer: %w", err)
|
|
}
|
|
|
|
// Create ONNX session
|
|
inputs := []string{"input_ids", "attention_mask", "token_type_ids"}
|
|
outputs := []string{"last_hidden_state"}
|
|
|
|
session, err := ort.NewDynamicAdvancedSession(modelPath, inputs, outputs, nil)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("create onnx session: %w", err)
|
|
}
|
|
|
|
return &Model{
|
|
session: session,
|
|
tokenizer: tok,
|
|
}, nil
|
|
}
|
|
|
|
// Close releases model resources.
|
|
func (m *Model) Close() error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
if m.session != nil {
|
|
return m.session.Destroy()
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// Embed returns the embedding vector for the given text.
|
|
func (m *Model) Embed(text string) (Vector, error) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
|
|
// Tokenize
|
|
inputIDs, attentionMask, tokenTypeIDs := m.tokenizer.Encode(text)
|
|
seqLen := int64(len(inputIDs))
|
|
|
|
// Create input tensors with shape [1, seqLen]
|
|
shape := ort.Shape{1, seqLen}
|
|
|
|
inputIDsTensor, err := ort.NewTensor(shape, toInt64(inputIDs))
|
|
if err != nil {
|
|
return nil, fmt.Errorf("create input_ids tensor: %w", err)
|
|
}
|
|
defer inputIDsTensor.Destroy()
|
|
|
|
attentionTensor, err := ort.NewTensor(shape, toInt64(attentionMask))
|
|
if err != nil {
|
|
return nil, fmt.Errorf("create attention_mask tensor: %w", err)
|
|
}
|
|
defer attentionTensor.Destroy()
|
|
|
|
tokenTypeTensor, err := ort.NewTensor(shape, toInt64(tokenTypeIDs))
|
|
if err != nil {
|
|
return nil, fmt.Errorf("create token_type_ids tensor: %w", err)
|
|
}
|
|
defer tokenTypeTensor.Destroy()
|
|
|
|
// Create output tensor with shape [1, seqLen, 384]
|
|
outputShape := ort.Shape{1, seqLen, 384}
|
|
outputTensor, err := ort.NewEmptyTensor[float32](outputShape)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("create output tensor: %w", err)
|
|
}
|
|
defer outputTensor.Destroy()
|
|
|
|
// Run inference
|
|
err = m.session.Run(
|
|
[]ort.ArbitraryTensor{inputIDsTensor, attentionTensor, tokenTypeTensor},
|
|
[]ort.ArbitraryTensor{outputTensor},
|
|
)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("run inference: %w", err)
|
|
}
|
|
|
|
// Mean pooling: average over sequence length, taking attention mask into account
|
|
output := outputTensor.GetData()
|
|
return meanPool(output, attentionMask, int(seqLen), 384), nil
|
|
}
|
|
|
|
// EmbedBatch embeds multiple texts efficiently.
|
|
func (m *Model) EmbedBatch(texts []string) ([]Vector, error) {
|
|
results := make([]Vector, len(texts))
|
|
for i, text := range texts {
|
|
vec, err := m.Embed(text)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("embed text %d: %w", i, err)
|
|
}
|
|
results[i] = vec
|
|
}
|
|
return results, nil
|
|
}
|
|
|
|
// CosineSimilarity computes the cosine similarity between two vectors.
|
|
func CosineSimilarity(a, b Vector) float32 {
|
|
if len(a) != len(b) {
|
|
return 0
|
|
}
|
|
var dot, normA, normB float64
|
|
for i := range a {
|
|
dot += float64(a[i]) * float64(b[i])
|
|
normA += float64(a[i]) * float64(a[i])
|
|
normB += float64(b[i]) * float64(b[i])
|
|
}
|
|
if normA == 0 || normB == 0 {
|
|
return 0
|
|
}
|
|
return float32(dot / (math.Sqrt(normA) * math.Sqrt(normB)))
|
|
}
|
|
|
|
// meanPool performs mean pooling over the sequence dimension with attention masking.
|
|
func meanPool(output []float32, mask []int, seqLen, hiddenSize int) Vector {
|
|
result := make(Vector, hiddenSize)
|
|
var count float32
|
|
for i := 0; i < seqLen; i++ {
|
|
if mask[i] == 0 {
|
|
continue
|
|
}
|
|
count++
|
|
for j := 0; j < hiddenSize; j++ {
|
|
result[j] += output[i*hiddenSize+j]
|
|
}
|
|
}
|
|
if count > 0 {
|
|
for j := range result {
|
|
result[j] /= count
|
|
}
|
|
}
|
|
// L2 normalize
|
|
var norm float64
|
|
for _, v := range result {
|
|
norm += float64(v) * float64(v)
|
|
}
|
|
norm = math.Sqrt(norm)
|
|
if norm > 0 {
|
|
for j := range result {
|
|
result[j] = float32(float64(result[j]) / norm)
|
|
}
|
|
}
|
|
return result
|
|
}
|
|
|
|
func toInt64(ints []int) []int64 {
|
|
out := make([]int64, len(ints))
|
|
for i, v := range ints {
|
|
out[i] = int64(v)
|
|
}
|
|
return out
|
|
}
|
|
|
|
// --- Tokenizer ---
|
|
|
|
// Tokenizer handles WordPiece tokenization for BERT-style models.
|
|
type Tokenizer struct {
|
|
vocab map[string]int
|
|
maxLen int
|
|
clsID int
|
|
sepID int
|
|
padID int
|
|
unkID int
|
|
lowercase bool
|
|
}
|
|
|
|
// tokenizerJSON is the HuggingFace tokenizers JSON format.
|
|
type tokenizerJSON struct {
|
|
Truncation *struct {
|
|
MaxLength int `json:"max_length"`
|
|
} `json:"truncation"`
|
|
Normalizer *struct {
|
|
Lowercase bool `json:"lowercase"`
|
|
} `json:"normalizer"`
|
|
Model struct {
|
|
Vocab map[string]int `json:"vocab"`
|
|
} `json:"model"`
|
|
AddedTokens []struct {
|
|
ID int `json:"id"`
|
|
Content string `json:"content"`
|
|
} `json:"added_tokens"`
|
|
}
|
|
|
|
func loadTokenizer(path string) (*Tokenizer, error) {
|
|
data, err := os.ReadFile(path)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
var tj tokenizerJSON
|
|
if err := json.Unmarshal(data, &tj); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
tok := &Tokenizer{
|
|
vocab: tj.Model.Vocab,
|
|
maxLen: 128, // default
|
|
unkID: 100, // [UNK]
|
|
clsID: 101, // [CLS]
|
|
sepID: 102, // [SEP]
|
|
padID: 0, // [PAD]
|
|
}
|
|
if tj.Truncation != nil {
|
|
tok.maxLen = tj.Truncation.MaxLength
|
|
}
|
|
if tj.Normalizer != nil {
|
|
tok.lowercase = tj.Normalizer.Lowercase
|
|
}
|
|
|
|
// Override IDs from added_tokens if present
|
|
for _, at := range tj.AddedTokens {
|
|
switch at.Content {
|
|
case "[PAD]":
|
|
tok.padID = at.ID
|
|
case "[UNK]":
|
|
tok.unkID = at.ID
|
|
case "[CLS]":
|
|
tok.clsID = at.ID
|
|
case "[SEP]":
|
|
tok.sepID = at.ID
|
|
}
|
|
}
|
|
|
|
return tok, nil
|
|
}
|
|
|
|
// Encode tokenizes text and returns input_ids, attention_mask, and token_type_ids.
|
|
func (t *Tokenizer) Encode(text string) (inputIDs, attentionMask, tokenTypeIDs []int) {
|
|
if t.lowercase {
|
|
text = strings.ToLower(text)
|
|
}
|
|
|
|
// Basic whitespace + punctuation tokenization
|
|
words := tokenizeBasic(text)
|
|
|
|
// WordPiece tokenization
|
|
var tokens []int
|
|
tokens = append(tokens, t.clsID)
|
|
for _, word := range words {
|
|
wordTokens := t.tokenizeWord(word)
|
|
tokens = append(tokens, wordTokens...)
|
|
}
|
|
tokens = append(tokens, t.sepID)
|
|
|
|
// Truncate if needed (keep [CLS] and [SEP])
|
|
if len(tokens) > t.maxLen {
|
|
tokens = append(tokens[:t.maxLen-1], t.sepID)
|
|
}
|
|
|
|
// Build masks
|
|
seqLen := len(tokens)
|
|
inputIDs = tokens
|
|
attentionMask = make([]int, seqLen)
|
|
tokenTypeIDs = make([]int, seqLen)
|
|
for i := 0; i < seqLen; i++ {
|
|
attentionMask[i] = 1
|
|
tokenTypeIDs[i] = 0
|
|
}
|
|
|
|
return inputIDs, attentionMask, tokenTypeIDs
|
|
}
|
|
|
|
func (t *Tokenizer) tokenizeWord(word string) []int {
|
|
var tokens []int
|
|
remaining := word
|
|
for len(remaining) > 0 {
|
|
found := false
|
|
for end := len(remaining); end > 0; end-- {
|
|
subword := remaining[:end]
|
|
if len(tokens) > 0 {
|
|
subword = "##" + subword
|
|
}
|
|
if id, ok := t.vocab[subword]; ok {
|
|
tokens = append(tokens, id)
|
|
remaining = remaining[end:]
|
|
found = true
|
|
break
|
|
}
|
|
}
|
|
if !found {
|
|
tokens = append(tokens, t.unkID)
|
|
break
|
|
}
|
|
}
|
|
return tokens
|
|
}
|
|
|
|
// tokenizeBasic splits on whitespace and punctuation.
|
|
func tokenizeBasic(text string) []string {
|
|
var words []string
|
|
var current strings.Builder
|
|
for _, r := range text {
|
|
if unicode.IsSpace(r) {
|
|
if current.Len() > 0 {
|
|
words = append(words, current.String())
|
|
current.Reset()
|
|
}
|
|
} else if unicode.IsPunct(r) {
|
|
if current.Len() > 0 {
|
|
words = append(words, current.String())
|
|
current.Reset()
|
|
}
|
|
words = append(words, string(r))
|
|
} else {
|
|
current.WriteRune(r)
|
|
}
|
|
}
|
|
if current.Len() > 0 {
|
|
words = append(words, current.String())
|
|
}
|
|
return words
|
|
}
|