// 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 }