router/semantic: e5 embedding cache for linear head experiment (slice 15)

Go tool that computes e5-small embeddings for every corpus row,
assigns frozen/dev split and grouped CV folds, outputs JSON.
This commit is contained in:
2026-09-07 13:51:54 +04:00
parent 6397ea1cc9
commit f8ec77de8a
+273
View File
@@ -0,0 +1,273 @@
// semantic-router-experiment computes embeddings for the semantic coarse-route
// corpus using the deployed multilingual-e5-small ONNX model. It outputs a
// JSON file containing every corpus row with its embedding vector, split
// assignment, and fold membership for grouped cross-validation.
//
// Usage:
//
// MAVEN_ONNX_LIB=/path/to/libonnxruntime.so \
// go run ./cmd/semantic-router-experiment/ -out embeddings.json
package main
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"flag"
"fmt"
"log"
"os"
"sort"
"strings"
"github.com/kami/maven/internal/router"
"github.com/kami/maven/internal/router/semantic"
)
// CachedRow is one corpus row with its precomputed embedding and split metadata.
type CachedRow struct {
Text string `json:"text"`
Route string `json:"route"`
Source string `json:"source"`
SourceID string `json:"source_id"`
SplitGroup string `json:"split_group"`
Tags []string `json:"tags,omitempty"`
FastPathResolved bool `json:"fast_path_resolved"`
RouterResidual *bool `json:"router_residual,omitempty"`
TextHash string `json:"text_hash"`
Embedding []float32 `json:"embedding"`
EmbedderID string `json:"embedder_id"`
FrozenHoldout bool `json:"frozen_holdout"`
CVFold int `json:"cv_fold"`
DevPool bool `json:"dev_pool"`
FamilyID string `json:"family_id"`
}
// ExperimentMeta carries metadata about the experiment run.
type ExperimentMeta struct {
EmbedderID string `json:"embedder_id"`
ModelPath string `json:"model_path"`
TokenizerPath string `json:"tokenizer_path"`
Dimension int `json:"dimension"`
Pooling string `json:"pooling"`
Normalization string `json:"normalization"`
InputTemplate string `json:"input_template"`
TotalExamples int `json:"total_examples"`
FrozenCount int `json:"frozen_count"`
DevCount int `json:"dev_count"`
CVFolds int `json:"cv_folds"`
FoldComposition map[int]FoldStats `json:"fold_composition"`
RouteCounts map[string]int `json:"route_counts"`
SourceCounts map[string]int `json:"source_counts"`
FastPathCount int `json:"fast_path_count"`
ResidualCount int `json:"residual_count"`
HoldoutHash string `json:"holdout_hash"`
}
// FoldStats describes one CV fold.
type FoldStats struct {
EvalCount int `json:"eval_count"`
TrainCount int `json:"train_count"`
Routes map[string]int `json:"eval_routes"`
}
func main() {
outPath := flag.String("out", "embeddings.json", "output JSON path")
folds := flag.Int("folds", 5, "number of CV folds")
flag.Parse()
libPath := os.Getenv("MAVEN_ONNX_LIB")
if libPath == "" {
log.Fatal("MAVEN_ONNX_LIB must be set to the libonnxruntime.so path")
}
// Resolve model paths relative to the module root (cwd when running with go run).
modelPath := "models/embedder/multilingual-e5-small/model_quantized.onnx"
tokPath := "models/embedder/multilingual-e5-small/tokenizer.json"
for _, p := range []string{libPath, modelPath, tokPath} {
if _, err := os.Stat(p); err != nil {
log.Fatalf("missing %s: %v", p, err)
}
}
// Load corpus.
exs, err := semantic.LoadCorpus()
if err != nil {
log.Fatalf("load corpus: %v", err)
}
fmt.Fprintf(os.Stderr, "corpus: %d examples\n", len(exs))
// Initialize embedder.
emb, err := router.NewONNXEmbedder(modelPath, tokPath, libPath)
if err != nil {
log.Fatalf("init embedder: %v", err)
}
defer emb.Close()
fmt.Fprintf(os.Stderr, "embedder: %s (dim=%d)\n", emb.ID(), emb.Dim())
// Compute frozen holdout / dev pool split.
_, devPool, holdoutHash := semantic.FrozenHoldoutSplit(exs)
fmt.Fprintf(os.Stderr, "frozen holdout: hash=%s, dev pool: %d examples\n", holdoutHash, len(devPool))
// Compute grouped CV folds on dev pool only.
cvFolds := semantic.GroupedCVFolds(devPool, *folds)
fmt.Fprintf(os.Stderr, "cv folds: %d\n", len(cvFolds))
// Build a lookup: source_id → cv_fold (from dev pool only).
foldLookup := make(map[string]int)
for _, f := range cvFolds {
for _, e := range f.Eval {
foldLookup[e.SourceID] = f.Fold
}
}
// Build dev set membership lookup.
_, devSet, _ := semantic.FrozenHoldoutSplit(exs)
devIDs := make(map[string]bool)
for _, e := range devSet {
devIDs[e.SourceID] = true
}
// Embed all examples.
ctx := context.Background()
var cached []CachedRow
foldComp := make(map[int]*FoldStats)
for i := 0; i < *folds; i++ {
foldComp[i] = &FoldStats{Routes: make(map[string]int)}
}
routeCounts := make(map[string]int)
sourceCounts := make(map[string]int)
fpCount, resCount := 0, 0
for i, e := range exs {
textHash := sha256.Sum256([]byte(e.Text))
embedding, err := emb.EmbedQuery(ctx, e.Text)
if err != nil {
log.Fatalf("embed row %d (%s): %v", i, e.SourceID, err)
}
inDev := devIDs[e.SourceID]
fold := -1
if inDev {
if f, ok := foldLookup[e.SourceID]; ok {
fold = f
}
}
isFrozen := !inDev
cr := CachedRow{
Text: e.Text,
Route: string(e.Route),
Source: e.Source,
SourceID: e.SourceID,
SplitGroup: e.SplitGroup,
Tags: e.Tags,
FastPathResolved: e.FastPathResolved,
RouterResidual: e.RouterResidual,
TextHash: hex.EncodeToString(textHash[:]),
Embedding: embedding,
EmbedderID: emb.ID(),
FrozenHoldout: isFrozen,
CVFold: fold,
DevPool: inDev,
FamilyID: e.SplitGroup,
}
cached = append(cached, cr)
routeCounts[cr.Route]++
sourceCounts[cr.Source]++
if e.FastPathResolved {
fpCount++
} else {
resCount++
}
if inDev && fold >= 0 {
foldComp[fold].EvalCount++
foldComp[fold].Routes[cr.Route]++
}
}
// Compute train counts per fold.
for i := 0; i < *folds; i++ {
foldComp[i].TrainCount = len(devPool) - foldComp[i].EvalCount
}
// Sort route counts for deterministic output.
sortedRoutes := make([]string, 0, len(routeCounts))
for r := range routeCounts {
sortedRoutes = append(sortedRoutes, r)
}
sort.Strings(sortedRoutes)
sortedRouteCounts := make(map[string]int)
for _, r := range sortedRoutes {
sortedRouteCounts[r] = routeCounts[r]
}
// Build fold stats with sorted keys.
finalFoldComp := make(map[int]FoldStats)
for i := 0; i < *folds; i++ {
finalFoldComp[i] = *foldComp[i]
}
meta := ExperimentMeta{
EmbedderID: emb.ID(),
ModelPath: modelPath,
TokenizerPath: tokPath,
Dimension: emb.Dim(),
Pooling: "mean-pool + L2-normalize",
Normalization: "L2",
InputTemplate: "query: <text>",
TotalExamples: len(exs),
FrozenCount: len(exs) - len(devPool),
DevCount: len(devPool),
CVFolds: *folds,
FoldComposition: finalFoldComp,
RouteCounts: sortedRouteCounts,
SourceCounts: sourceCounts,
FastPathCount: fpCount,
ResidualCount: resCount,
HoldoutHash: holdoutHash,
}
// Output.
output := map[string]any{
"meta": meta,
"examples": cached,
}
data, err := json.MarshalIndent(output, "", " ")
if err != nil {
log.Fatalf("marshal: %v", err)
}
if err := os.WriteFile(*outPath, data, 0644); err != nil {
log.Fatalf("write %s: %v", *outPath, err)
}
// Print summary.
fmt.Fprintf(os.Stderr, "\n=== experiment metadata ===\n")
fmt.Fprintf(os.Stderr, "embedder: %s\n", meta.EmbedderID)
fmt.Fprintf(os.Stderr, "dimension: %d\n", meta.Dimension)
fmt.Fprintf(os.Stderr, "total: %d frozen: %d dev: %d\n", meta.TotalExamples, meta.FrozenCount, meta.DevCount)
fmt.Fprintf(os.Stderr, "fast-path: %d residual: %d\n", meta.FastPathCount, meta.ResidualCount)
fmt.Fprintf(os.Stderr, "routes: %s\n", formatMap(sortedRouteCounts))
fmt.Fprintf(os.Stderr, "fold composition:\n")
for i := 0; i < *folds; i++ {
fs := finalFoldComp[i]
fmt.Fprintf(os.Stderr, " fold %d: eval=%d train=%d routes=%s\n",
i, fs.EvalCount, fs.TrainCount, formatMap(fs.Routes))
}
fmt.Fprintf(os.Stderr, "output: %s (%d bytes)\n", *outPath, len(data))
}
func formatMap(m map[string]int) string {
var parts []string
for k, v := range m {
parts = append(parts, fmt.Sprintf("%s=%d", k, v))
}
sort.Strings(parts)
return "{" + strings.Join(parts, ", ") + "}"
}