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
70
Dockerfile
70
Dockerfile
|
|
@ -1,27 +1,59 @@
|
||||||
# Minimal image: ~200-250MB
|
# Go implementation: Minimal image with ONNX Runtime
|
||||||
# Uses pre-converted ONNX model from onnx-community
|
# Multi-stage build: build with Go + ONNX Runtime deps, runtime with minimal Debian
|
||||||
|
|
||||||
FROM python:3.11-slim
|
# Stage 1: Build Go binary
|
||||||
|
FROM golang:1.23-bookworm AS builder
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
# Install runtime dependencies and download model
|
# Install build dependencies: git, g++, make, ca-certificates, curl, libc6-dev
|
||||||
RUN apt-get update && \
|
RUN apt-get update && \
|
||||||
apt-get install -y --no-install-recommends wget && \
|
apt-get install -y --no-install-recommends git g++ make ca-certificates curl libc6-dev && \
|
||||||
pip install --no-cache-dir onnxruntime numpy && \
|
rm -rf /var/lib/apt/lists/*
|
||||||
# Download pre-converted ONNX model and tokenizer files
|
|
||||||
wget -q https://huggingface.co/onnx-community/all-MiniLM-L6-v2-ONNX/resolve/main/onnx/model.onnx -O /app/model.onnx && \
|
|
||||||
wget -q https://huggingface.co/onnx-community/all-MiniLM-L6-v2-ONNX/resolve/main/onnx/model.onnx_data -O /app/model.onnx_data && \
|
|
||||||
wget -q https://huggingface.co/onnx-community/all-MiniLM-L6-v2-ONNX/resolve/main/tokenizer.json -O /app/tokenizer.json && \
|
|
||||||
wget -q https://huggingface.co/onnx-community/all-MiniLM-L6-v2-ONNX/resolve/main/tokenizer_config.json -O /app/tokenizer_config.json && \
|
|
||||||
wget -q https://huggingface.co/onnx-community/all-MiniLM-L6-v2-ONNX/resolve/main/vocab.txt -O /app/vocab.txt && \
|
|
||||||
wget -q https://huggingface.co/onnx-community/all-MiniLM-L6-v2-ONNX/resolve/main/config.json -O /app/config.json && \
|
|
||||||
# Clean up
|
|
||||||
apt-get remove -y wget 2>/dev/null || true && \
|
|
||||||
apt-get clean && \
|
|
||||||
rm -rf /var/lib/apt/lists/* /tmp/* /root/.cache /var/cache/apt/*
|
|
||||||
|
|
||||||
COPY server.py .
|
# Configure git to avoid terminal prompt
|
||||||
|
ENV GIT_TERMINAL_PROMPT=0
|
||||||
|
|
||||||
|
# Copy Go files
|
||||||
|
COPY go.mod go.sum .
|
||||||
|
COPY main.go .
|
||||||
|
|
||||||
|
# Download Go dependencies
|
||||||
|
RUN go mod download 2>&1
|
||||||
|
|
||||||
|
# Download ONNX Runtime shared library for Linux x64
|
||||||
|
RUN curl -sL -o /tmp/onnx.tgz https://github.com/microsoft/onnxruntime/releases/download/v1.27.0/onnxruntime-linux-x64-1.27.0.tgz
|
||||||
|
RUN tar -xzf /tmp/onnx.tgz -C /tmp
|
||||||
|
RUN mkdir -p /usr/local/lib && cp /tmp/onnxruntime-linux-x64-1.27.0/lib/libonnxruntime.so* /usr/local/lib/
|
||||||
|
RUN rm -rf /tmp/onnxruntime-linux-x64-1.27.0 /tmp/onnx.tgz
|
||||||
|
|
||||||
|
# Build Go binary with CGO enabled
|
||||||
|
RUN CGO_ENABLED=1 GOOS=linux GOARCH=amd64 go build -o /app/vector-server main.go 2>&1
|
||||||
|
|
||||||
|
# Stage 2: Runtime
|
||||||
|
FROM debian:bookworm-slim
|
||||||
|
WORKDIR /app
|
||||||
|
|
||||||
|
# Install runtime dependencies: libstdc++, ca-certificates, curl
|
||||||
|
RUN apt-get update && \
|
||||||
|
apt-get install -y --no-install-recommends libstdc++6 ca-certificates curl && \
|
||||||
|
rm -rf /var/lib/apt/lists/*
|
||||||
|
|
||||||
|
# Download model and tokenizer files
|
||||||
|
RUN curl -sL -o /app/model.onnx https://huggingface.co/onnx-community/all-MiniLM-L6-v2-ONNX/resolve/main/onnx/model.onnx && \
|
||||||
|
curl -sL -o /app/model.onnx_data https://huggingface.co/onnx-community/all-MiniLM-L6-v2-ONNX/resolve/main/onnx/model.onnx_data && \
|
||||||
|
curl -sL -o /app/tokenizer.json https://huggingface.co/onnx-community/all-MiniLM-L6-v2-ONNX/resolve/main/tokenizer.json && \
|
||||||
|
apt-get remove -y curl 2>/dev/null || true && \
|
||||||
|
apt-get clean && \
|
||||||
|
rm -rf /var/lib/apt/lists/* /tmp/* /var/tmp/*
|
||||||
|
|
||||||
|
# Copy binary and shared library from builder
|
||||||
|
COPY --from=builder /usr/local/lib/libonnxruntime.so* /usr/local/lib/
|
||||||
|
COPY --from=builder /app/vector-server .
|
||||||
|
|
||||||
|
# Set environment for ONNX Runtime
|
||||||
|
ENV LD_LIBRARY_PATH=/usr/local/lib:${LD_LIBRARY_PATH}
|
||||||
|
ENV ONNXRUNTIME_SHARED_LIBRARY_PATH=/usr/local/lib/libonnxruntime.so
|
||||||
|
|
||||||
EXPOSE 8080
|
EXPOSE 8080
|
||||||
ENV PORT=8080
|
ENV PORT=8080
|
||||||
CMD ["python", "server.py"]
|
CMD ["./vector-server"]
|
||||||
|
|
|
||||||
17
go.mod
Normal file
17
go.mod
Normal file
|
|
@ -0,0 +1,17 @@
|
||||||
|
module vector
|
||||||
|
|
||||||
|
go 1.23
|
||||||
|
|
||||||
|
require (
|
||||||
|
github.com/sugarme/tokenizer v0.3.0
|
||||||
|
github.com/yalue/onnxruntime_go v1.31.0
|
||||||
|
)
|
||||||
|
|
||||||
|
require (
|
||||||
|
github.com/mitchellh/colorstring v0.0.0-20190213212951-d06e56a500db // indirect
|
||||||
|
github.com/rivo/uniseg v0.4.7 // indirect
|
||||||
|
github.com/schollz/progressbar/v2 v2.15.0 // indirect
|
||||||
|
github.com/sugarme/regexpset v0.0.0-20200920021344-4d4ec8eaf93c // indirect
|
||||||
|
golang.org/x/text v0.25.0 // indirect
|
||||||
|
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||||
|
)
|
||||||
30
go.sum
Normal file
30
go.sum
Normal file
|
|
@ -0,0 +1,30 @@
|
||||||
|
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
|
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||||
|
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
|
github.com/emirpasic/gods v1.18.1 h1:FXtiHYKDGKCW2KzwZKx0iC0PQmdlorYgdFG9jPXJ1Bc=
|
||||||
|
github.com/emirpasic/gods v1.18.1/go.mod h1:8tpGGwCnJ5H4r6BWwaV6OrWmMoPhUl5jm/FMNAnJvWQ=
|
||||||
|
github.com/mitchellh/colorstring v0.0.0-20190213212951-d06e56a500db h1:62I3jR2EmQ4l5rM/4FEfDWcRD+abF5XlKShorW5LRoQ=
|
||||||
|
github.com/mitchellh/colorstring v0.0.0-20190213212951-d06e56a500db/go.mod h1:l0dey0ia/Uv7NcFFVbCLtqEBQbrT4OCwCSKTEv6enCw=
|
||||||
|
github.com/patrickmn/go-cache v2.1.0+incompatible h1:HRMgzkcYKYpi3C8ajMPV8OFXaaRUnok+kx1WdO15EQc=
|
||||||
|
github.com/patrickmn/go-cache v2.1.0+incompatible/go.mod h1:3Qf8kWWT7OJRJbdiICTKqZju1ZixQ/KpMGzzAfe6+WQ=
|
||||||
|
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||||
|
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||||
|
github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ=
|
||||||
|
github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88=
|
||||||
|
github.com/schollz/progressbar/v2 v2.15.0 h1:dVzHQ8fHRmtPjD3K10jT3Qgn/+H+92jhPrhmxIJfDz8=
|
||||||
|
github.com/schollz/progressbar/v2 v2.15.0/go.mod h1:UdPq3prGkfQ7MOzZKlDRpYKcFqEMczbD7YmbPgpzKMI=
|
||||||
|
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
||||||
|
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
||||||
|
github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA=
|
||||||
|
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
|
||||||
|
github.com/sugarme/regexpset v0.0.0-20200920021344-4d4ec8eaf93c h1:pwb4kNSHb4K89ymCaN+5lPH/MwnfSVg4rzGDh4d+iy4=
|
||||||
|
github.com/sugarme/regexpset v0.0.0-20200920021344-4d4ec8eaf93c/go.mod h1:2gwkXLWbDGUQWeL3RtpCmcY4mzCtU13kb9UsAg9xMaw=
|
||||||
|
github.com/sugarme/tokenizer v0.3.0 h1:FE8DYbNSz/kSbgEo9l/RjgYHkIJYEdskumitFQBE9FE=
|
||||||
|
github.com/sugarme/tokenizer v0.3.0/go.mod h1:VJ+DLK5ZEZwzvODOWwY0cw+B1dabTd3nCB5HuFCItCc=
|
||||||
|
github.com/yalue/onnxruntime_go v1.31.0 h1:1ln4YW1SFOFfGJZXe3jNOb2JUSt+l2pEneZfV8HdtFA=
|
||||||
|
github.com/yalue/onnxruntime_go v1.31.0/go.mod h1:b4X26A8pekNb1ACJ58wAXgNKeUCGEAQ9dmACut9Sm/4=
|
||||||
|
golang.org/x/text v0.25.0 h1:qVyWApTSYLk/drJRO5mDlNYskwQznZmkpV2c8q9zls4=
|
||||||
|
golang.org/x/text v0.25.0/go.mod h1:WEdwpYrmk1qmdHvhkSTNPm3app7v4rsT8F2UD6+VHIA=
|
||||||
|
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||||
|
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||||
|
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||||
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