diff --git a/Dockerfile b/Dockerfile index 78218d1..d65938c 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,27 +1,59 @@ -# Minimal image: ~200-250MB -# Uses pre-converted ONNX model from onnx-community +# Go implementation: Minimal image with ONNX Runtime +# 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 -# Install runtime dependencies and download model +# Install build dependencies: git, g++, make, ca-certificates, curl, libc6-dev RUN apt-get update && \ - apt-get install -y --no-install-recommends wget && \ - pip install --no-cache-dir onnxruntime numpy && \ - # 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/* + apt-get install -y --no-install-recommends git g++ make ca-certificates curl libc6-dev && \ + rm -rf /var/lib/apt/lists/* -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 ENV PORT=8080 -CMD ["python", "server.py"] +CMD ["./vector-server"] diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..f65f4c4 --- /dev/null +++ b/go.mod @@ -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 +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..de34701 --- /dev/null +++ b/go.sum @@ -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= diff --git a/main.go b/main.go new file mode 100644 index 0000000..98215cb --- /dev/null +++ b/main.go @@ -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)) +}