Add Go implementation with ONNX Runtime
- Use golang:1.23-bookworm as builder - Use debian:bookworm-slim as runtime - Use github.com/yalue/onnxruntime_go for ONNX inference - Use github.com/sugarme/tokenizer for tokenization - Download ONNX Runtime v1.27.0 shared library - Download model.onnx and model.onnx_data from onnx-community - Support API_SECRET environment variable for authentication - Final image size: ~266 MB Generated by Mistral Vibe. Co-Authored-By: Mistral Vibe <vibe@mistral.ai>
This commit is contained in:
parent
6101df3bfe
commit
21d467a284
4 changed files with 303 additions and 19 deletions
205
main.go
Normal file
205
main.go
Normal file
|
|
@ -0,0 +1,205 @@
|
|||
package main
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"os"
|
||||
|
||||
"github.com/sugarme/tokenizer"
|
||||
"github.com/sugarme/tokenizer/pretrained"
|
||||
"github.com/yalue/onnxruntime_go"
|
||||
)
|
||||
|
||||
var (
|
||||
modelPath = "/app/model.onnx"
|
||||
apiSecret = os.Getenv("API_SECRET")
|
||||
tok *tokenizer.Tokenizer
|
||||
session *onnxruntime_go.DynamicAdvancedSession
|
||||
maxLen = int64(128)
|
||||
embeddingSize = int64(384)
|
||||
)
|
||||
|
||||
func init() {
|
||||
// Initialize ONNX Runtime
|
||||
libPath := os.Getenv("ONNXRUNTIME_SHARED_LIBRARY_PATH")
|
||||
if libPath == "" {
|
||||
libPath = "/usr/local/lib/libonnxruntime.so"
|
||||
}
|
||||
onnxruntime_go.SetSharedLibraryPath(libPath)
|
||||
|
||||
err := onnxruntime_go.InitializeEnvironment()
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to initialize ONNX Runtime: %v", err)
|
||||
}
|
||||
|
||||
// Load tokenizer
|
||||
tok, err = pretrained.FromFile("/app/tokenizer.json")
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to load tokenizer: %v", err)
|
||||
}
|
||||
|
||||
// Load ONNX model
|
||||
session, err = onnxruntime_go.NewDynamicAdvancedSession(
|
||||
modelPath,
|
||||
[]string{"input_ids", "attention_mask", "token_type_ids"},
|
||||
[]string{"last_hidden_state"},
|
||||
nil)
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to load ONNX model: %v", err)
|
||||
}
|
||||
|
||||
log.Println("Model and tokenizer loaded. Starting server...")
|
||||
}
|
||||
|
||||
func encode(text string) ([]int64, []int64, []int64) {
|
||||
// Tokenize using the proper tokenizer
|
||||
inputSeq := tokenizer.NewInputSequence(text)
|
||||
input := tokenizer.NewSingleEncodeInput(inputSeq)
|
||||
encoding, err := tok.Encode(input, true)
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to tokenize: %v", err)
|
||||
}
|
||||
|
||||
inputIds := make([]int64, len(encoding.GetIds()))
|
||||
for i, id := range encoding.GetIds() {
|
||||
inputIds[i] = int64(id)
|
||||
}
|
||||
|
||||
// Truncate or pad
|
||||
paddedIds := make([]int64, maxLen)
|
||||
copy(paddedIds, inputIds)
|
||||
|
||||
attentionMask := make([]int64, maxLen)
|
||||
tokenTypeIds := make([]int64, maxLen)
|
||||
for i := 0; i < int(maxLen); i++ {
|
||||
if i < len(inputIds) {
|
||||
attentionMask[i] = 1
|
||||
} else {
|
||||
attentionMask[i] = 0
|
||||
}
|
||||
tokenTypeIds[i] = 0
|
||||
}
|
||||
|
||||
return paddedIds, attentionMask, tokenTypeIds
|
||||
}
|
||||
|
||||
func vectorHandler(w http.ResponseWriter, r *http.Request) {
|
||||
// Check API secret
|
||||
if apiSecret != "" {
|
||||
auth := r.Header.Get("Authorization")
|
||||
if auth != "Bearer "+apiSecret {
|
||||
http.Error(w, "Unauthorized - Invalid or missing API key", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if r.Method != "POST" {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Text string `json:"text"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
http.Error(w, "Bad request", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
if req.Text == "" {
|
||||
http.Error(w, "Text is required", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
// Encode text
|
||||
inputIds, attentionMask, tokenTypeIds := encode(req.Text)
|
||||
|
||||
// Create input tensors
|
||||
inputShape := onnxruntime_go.NewShape(1, maxLen)
|
||||
|
||||
inputIdsTensor, err := onnxruntime_go.NewTensor(inputShape, inputIds)
|
||||
if err != nil {
|
||||
http.Error(w, fmt.Sprintf("Failed to create input_ids tensor: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
defer inputIdsTensor.Destroy()
|
||||
|
||||
attentionMaskTensor, err := onnxruntime_go.NewTensor(inputShape, attentionMask)
|
||||
if err != nil {
|
||||
http.Error(w, fmt.Sprintf("Failed to create attention_mask tensor: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
defer attentionMaskTensor.Destroy()
|
||||
|
||||
tokenTypeIdsTensor, err := onnxruntime_go.NewTensor(inputShape, tokenTypeIds)
|
||||
if err != nil {
|
||||
http.Error(w, fmt.Sprintf("Failed to create token_type_ids tensor: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
defer tokenTypeIdsTensor.Destroy()
|
||||
|
||||
// Create output tensor
|
||||
outputShape := onnxruntime_go.NewShape(1, maxLen, embeddingSize)
|
||||
outputTensor, err := onnxruntime_go.NewEmptyTensor[float32](outputShape)
|
||||
if err != nil {
|
||||
http.Error(w, fmt.Sprintf("Failed to create output tensor: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
defer outputTensor.Destroy()
|
||||
|
||||
// Run inference
|
||||
inputs := []onnxruntime_go.Value{inputIdsTensor, attentionMaskTensor, tokenTypeIdsTensor}
|
||||
outputs := []onnxruntime_go.Value{outputTensor}
|
||||
|
||||
err = session.Run(inputs, outputs)
|
||||
if err != nil {
|
||||
http.Error(w, fmt.Sprintf("Inference error: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Get embeddings
|
||||
embeddings := outputTensor.GetData()
|
||||
|
||||
// Mean pooling over sequence length (exclude padding)
|
||||
var sum [384]float32
|
||||
count := 0
|
||||
for i := 0; i < int(maxLen); i++ {
|
||||
if attentionMask[i] == 1 {
|
||||
for j := 0; j < int(embeddingSize); j++ {
|
||||
sum[j] += embeddings[i*int(embeddingSize)+j]
|
||||
}
|
||||
count++
|
||||
}
|
||||
}
|
||||
|
||||
var sentenceEmbedding [384]float32
|
||||
if count > 0 {
|
||||
for j := 0; j < int(embeddingSize); j++ {
|
||||
sentenceEmbedding[j] = sum[j] / float32(count)
|
||||
}
|
||||
}
|
||||
|
||||
// Convert to slice for JSON
|
||||
vector := make([]float32, embeddingSize)
|
||||
for i := 0; i < int(embeddingSize); i++ {
|
||||
vector[i] = sentenceEmbedding[i]
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"vector": vector,
|
||||
})
|
||||
}
|
||||
|
||||
func main() {
|
||||
port := os.Getenv("PORT")
|
||||
if port == "" {
|
||||
port = "8080"
|
||||
}
|
||||
|
||||
http.HandleFunc("/vector", vectorHandler)
|
||||
log.Printf("Starting server on port %s...\n", port)
|
||||
log.Fatal(http.ListenAndServe(":"+port, nil))
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue