d52f60c54e
- Add CalendarEvents method to recordingAPI in auth_test.go - Add CalendarEvents method to fakeCore in handlers_test.go Co-Authored-By: opencode <opencode@anthropic.com>
317 lines
6.8 KiB
Go
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)
|