package semantic import ( "context" "fmt" "math" "sort" "strings" "time" "github.com/kami/maven/internal/router" ) // ScoreLegacy runs the actual router against the corpus and produces the // baseline report. The now parameter is the reference clock for relative // time expressions. func ScoreLegacy(ctx context.Context, r LegacyRouter, exs []RouteExample, now time.Time) LegacyReport { stats := CorpusStatsFrom(exs) rep := LegacyReport{ Stats: stats, Total: len(exs), ByRoute: make(map[SemanticRoute]RouteMetrics), Confusion: make(map[SemanticRoute]map[SemanticRoute]int), } for _, rt := range AllRoutes { rep.Confusion[rt] = make(map[SemanticRoute]int) } for _, e := range exs { input := router.NormalizedInput{Text: e.Text} d, err := r.Route(ctx, input, now) var predicted SemanticRoute var fpHit bool var preroute bool if err != nil { predicted = RouteUncertain } else { // Determine if fast-path resolved this. fpHit = d.Stage == 0 && d.Producer == router.RouteProducerGrammar predicted = IntentToRoute(d.Intent) // Pre-route consumption: command-prohibition grammar emits // IntentAct with Fn=prohibited_act at stage 0. These cases // never reach the general cascade. preroute = d.Stage == 0 && d.Slots.Fn == "prohibited_act" } agree := predicted == e.Route c := LegacyCase{ Example: e, Decision: d, Predicted: predicted, Agree: agree, FastPathHit: fpHit, PrerouteConsumed: preroute, Error: err, } rep.Cases = append(rep.Cases, c) rep.Confusion[e.Route][predicted]++ if agree { rep.Passed++ } if e.Route != RouteAction && predicted == RouteAction { rep.FalseAction++ rep.FalseActionCases = append(rep.FalseActionCases, c) } if e.FastPathResolved || fpHit { if preroute { // Command-prohibition grammar: consumed before cascade. rep.PreRouteTotal++ if agree { rep.PreRoutePassed++ } } else { rep.FastPathTotal++ if agree { rep.FastPathPassed++ } } } else { rep.ResidualTotal++ if agree { rep.ResidualPassed++ } // All residual cases are router-residual (the learned // router would see all of them). rep.RouterResidualTotal++ if agree { rep.RouterResidualPassed++ } } } if rep.Total > 0 { rep.FalseActionRate = float64(rep.FalseAction) / float64(rep.Total) } // Per-route P/R/F1. for _, route := range AllRoutes { tp, fp, fn := 0, 0, 0 for _, got := range AllRoutes { count := rep.Confusion[route][got] if got == route { tp = count } else { fn += count fp += rep.Confusion[got][route] } } rm := RouteMetrics{TP: tp, FP: fp, FN: fn} if tp+fp > 0 { rm.Precision = float64(tp) / float64(tp+fp) } if tp+fn > 0 { rm.Recall = float64(tp) / float64(tp+fn) } if rm.Precision+rm.Recall > 0 { rm.F1 = 2 * rm.Precision * rm.Recall / (rm.Precision + rm.Recall) } rm.F1 = math.Round(rm.F1*1000) / 1000 rm.Precision = math.Round(rm.Precision*1000) / 1000 rm.Recall = math.Round(rm.Recall*1000) / 1000 rep.ByRoute[route] = rm } return rep } // MacroF1 returns the macro-averaged F1 across all routes. func (r LegacyReport) MacroF1() float64 { if len(r.ByRoute) == 0 { return 0 } var sum float64 for _, rm := range r.ByRoute { sum += rm.F1 } return math.Round(sum/float64(len(r.ByRoute))*1000) / 1000 } // String renders the legacy baseline report. func (r LegacyReport) String() string { var b strings.Builder fmt.Fprintf(&b, "legacy baseline: %d/%d (%.1f%%)\n", r.Passed, r.Total, 100*float64(r.Passed)/math.Max(float64(r.Total), 1)) fmt.Fprintf(&b, " macro F1: %.3f\n", r.MacroF1()) fmt.Fprintf(&b, " false-action: %d/%d (%.3f)\n", r.FalseAction, r.Total, r.FalseActionRate) fmt.Fprintf(&b, " fast-path: %d/%d residual: %d/%d\n", r.FastPathPassed, r.FastPathTotal, r.ResidualPassed, r.ResidualTotal) if r.RouterResidualTotal > 0 { fmt.Fprintf(&b, " router-residual: %d/%d (pre-route consumed: %d)\n", r.RouterResidualPassed, r.RouterResidualTotal, r.PreRouteTotal) } fmt.Fprintf(&b, " per-route:\n") routes := make([]SemanticRoute, 0, len(r.ByRoute)) for route := range r.ByRoute { routes = append(routes, route) } sort.Slice(routes, func(i, j int) bool { return routes[i] < routes[j] }) for _, route := range routes { rm := r.ByRoute[route] fmt.Fprintf(&b, " %-15s P=%.3f R=%.3f F1=%.3f (tp=%d fp=%d fn=%d)\n", string(route), rm.Precision, rm.Recall, rm.F1, rm.TP, rm.FP, rm.FN) } fmt.Fprintf(&b, " confusion matrix:\n") fmt.Fprintf(&b, " %-15s", "") for _, g := range routes { fmt.Fprintf(&b, " %12s", string(g)) } fmt.Fprintf(&b, "\n") for _, w := range routes { fmt.Fprintf(&b, " %-15s", string(w)) for _, g := range routes { fmt.Fprintf(&b, " %12d", r.Confusion[w][g]) } fmt.Fprintf(&b, "\n") } if len(r.FalseActionCases) > 0 { fmt.Fprintf(&b, " false-action cases:\n") for _, c := range r.FalseActionCases { decided := "(error)" if c.Error == nil { decided = fmt.Sprintf("%s (%.3f)", c.Decision.Intent, c.Decision.Confidence) } fmt.Fprintf(&b, " %s %q: expected %s, got action (decided %s)\n", c.Example.SourceID, c.Example.Text, c.Example.Route, decided) } } return b.String() } // ContrastFamilies splits the corpus rows by their transform tag and reports // per-family performance against the legacy router. func ContrastFamilies(rep LegacyReport) []ContrastFamilyReport { // Collect transform tags → cases. tagCases := map[string][]LegacyCase{} for _, c := range rep.Cases { for _, tag := range c.Example.Tags { switch tag { case "negation", "question", "reported_speech", "quotation", "hypothetical", "capability_question": tagCases[tag] = append(tagCases[tag], c) } } } var reports []ContrastFamilyReport for _, tag := range []string{ "negation", "question", "reported_speech", "quotation", "hypothetical", "capability_question", } { cases := tagCases[tag] if len(cases) == 0 { continue } cr := ContrastFamilyReport{Transform: tag, Examples: len(cases), Cases: cases} for _, c := range cases { if c.Agree { cr.Correct++ } if c.Example.Route != RouteAction && c.Predicted == RouteAction { cr.FalseAction++ } else if !c.Agree { cr.OtherErrors++ } } reports = append(reports, cr) } return reports } // String renders the contrast family report. func ContrastFamilyReportString(r []ContrastFamilyReport) string { var b strings.Builder fmt.Fprintf(&b, "%-20s %8s %8s %8s %8s\n", "transform", "examples", "correct", "false-act", "other") for _, cr := range r { fmt.Fprintf(&b, "%-20s %8d %8d %8d %8d\n", cr.Transform, cr.Examples, cr.Correct, cr.FalseAction, cr.OtherErrors) } return b.String() }