From f8ec77de8a0c4a5d0d269e9811d1990508777551 Mon Sep 17 00:00:00 2001 From: claude Date: Mon, 7 Sep 2026 13:51:54 +0400 Subject: [PATCH] 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. --- cmd/semantic-router-experiment/main.go | 273 +++++++++++++++++++++++++ 1 file changed, 273 insertions(+) create mode 100644 cmd/semantic-router-experiment/main.go diff --git a/cmd/semantic-router-experiment/main.go b/cmd/semantic-router-experiment/main.go new file mode 100644 index 0000000..cb3286d --- /dev/null +++ b/cmd/semantic-router-experiment/main.go @@ -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: ", + 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, ", ") + "}" +}