ollie/embedding/embedding.go

345 lines
8.1 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 is reserved for future batch optimization.
// Currently calls Embed sequentially.
// 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
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]
}
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 "[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
}