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