feat: mixed model eval + concurrent QA workers

- Add --judge-model, --judge-api-base, --judge-api-key flags for separate judge model
- Add --concurrency flag (default 1) with semaphore-based goroutine pool
- Add reasoning_content fallback for GLM/DeepSeek style responses
- Prepend /no_think to system prompt for Ollama /v1 compatibility
- Reduce default MaxTokens from 2048 to 512 (answers are 1-3 sentences)
- Extract evalQAWorker and buildSeahorseContext for shared concurrent logic
This commit is contained in:
BeaconCat 2026-04-14 22:27:24 +08:00
parent 1d72a6afa6
commit d9cb8296f0
3 changed files with 251 additions and 150 deletions

View file

@ -8,6 +8,7 @@ import (
"sort" "sort"
"strconv" "strconv"
"strings" "strings"
"sync"
"github.com/sipeed/picoclaw/pkg/seahorse" "github.com/sipeed/picoclaw/pkg/seahorse"
) )
@ -46,7 +47,7 @@ var scoreRe = regexp.MustCompile(`\b([1-5])\b`)
// Returns a score from 0.0 to 1.0, or -1.0 on parse failure. // Returns a score from 0.0 to 1.0, or -1.0 on parse failure.
func judgeAnswer( func judgeAnswer(
ctx context.Context, ctx context.Context,
client *LLMClient, judgeClient *LLMClient,
question, goldAnswer, candidateAnswer string, question, goldAnswer, candidateAnswer string,
) (float64, error) { ) (float64, error) {
userPrompt := fmt.Sprintf( userPrompt := fmt.Sprintf(
@ -54,7 +55,7 @@ func judgeAnswer(
question, goldAnswer, candidateAnswer, question, goldAnswer, candidateAnswer,
) )
response, err := client.Complete(ctx, judgeSystemPrompt, userPrompt) response, err := judgeClient.Complete(ctx, judgeSystemPrompt, userPrompt)
if err != nil { if err != nil {
return -1.0, err return -1.0, err
} }
@ -68,17 +69,80 @@ func judgeAnswer(
return -1.0, nil return -1.0, nil
} }
// qaWork describes one QA evaluation unit.
type qaWork struct {
sampleID string
qaIndex int
globalIndex int
totalQA int
qa *LocomoQA
contextText string
sample *LocomoSample
}
// qaResult collects one QA evaluation output.
type qaResultOut struct {
index int // position in the flat QA list for ordering
result QAResult
answer string
score float64
}
// evalQAWorker processes a single QA item: generate answer + judge score.
func evalQAWorker(
ctx context.Context,
w qaWork,
answerClient, judgeClient *LLMClient,
logPrefix string,
) qaResultOut {
llmAnswer, err := generateAnswer(ctx, answerClient, w.contextText, w.qa.Question)
if err != nil {
log.Printf("WARN: LLM generation failed for sample %s Q%d: %v", w.sampleID, w.qaIndex, err)
llmAnswer = ""
}
score := -1.0
if llmAnswer != "" {
score, err = judgeAnswer(ctx, judgeClient, w.qa.Question, w.qa.AnswerString(), llmAnswer)
if err != nil {
log.Printf("WARN: LLM judge failed for sample %s Q%d: %v", w.sampleID, w.qaIndex, err)
}
}
hitRate := RecallHitRate(w.qa.Evidence, w.sample, w.contextText)
log.Printf("[%s] sample=%s q=%d/%d score=%.2f answer=%q",
logPrefix, w.sampleID, w.globalIndex, w.totalQA, score, truncateStr(llmAnswer, 80))
return qaResultOut{
index: w.globalIndex,
result: QAResult{
Question: w.qa.Question,
Category: w.qa.Category,
GoldAnswer: w.qa.AnswerString(),
TokenF1: score,
HitRate: hitRate,
},
answer: llmAnswer,
score: score,
}
}
// EvalLegacyLLM evaluates legacy store using LLM generation + LLM-as-Judge. // EvalLegacyLLM evaluates legacy store using LLM generation + LLM-as-Judge.
func EvalLegacyLLM( func EvalLegacyLLM(
ctx context.Context, ctx context.Context,
samples []LocomoSample, samples []LocomoSample,
legacy *LegacyStore, legacy *LegacyStore,
budgetTokens int, budgetTokens int,
client *LLMClient, answerClient, judgeClient *LLMClient,
concurrency int,
) []EvalResult { ) []EvalResult {
if concurrency < 1 {
concurrency = 1
}
totalQA := countTotalQA(samples) totalQA := countTotalQA(samples)
results := make([]EvalResult, 0, len(samples)) results := make([]EvalResult, 0, len(samples))
total := 0
for si := range samples { for si := range samples {
sample := &samples[si] sample := &samples[si]
history := legacy.GetHistory(sample.SampleID) history := legacy.GetHistory(sample.SampleID)
@ -88,41 +152,38 @@ func EvalLegacyLLM(
allContent = append(allContent, msg.Content) allContent = append(allContent, msg.Content)
} }
qaResults := make([]QAResult, 0, len(sample.QA)) truncated, _ := BudgetTruncate(allContent, budgetTokens)
for qi := range sample.QA { contextText := StringListToContent(truncated)
qa := &sample.QA[qi]
total++
truncated, _ := BudgetTruncate(allContent, budgetTokens)
contextText := StringListToContent(truncated)
// Generate answer with LLM qaResults := make([]QAResult, len(sample.QA))
llmAnswer, err := generateAnswer(ctx, client, contextText, qa.Question)
if err != nil { if concurrency <= 1 {
log.Printf("WARN: LLM generation failed for sample %s Q%d: %v", sample.SampleID, qi, err) for qi := range sample.QA {
llmAnswer = "" out := evalQAWorker(ctx, qaWork{
sampleID: sample.SampleID, qaIndex: qi,
globalIndex: si*len(sample.QA) + qi + 1, totalQA: totalQA,
qa: &sample.QA[qi], contextText: contextText, sample: sample,
}, answerClient, judgeClient, "legacy-llm")
qaResults[qi] = out.result
} }
} else {
// Judge the answer; -1.0 = API/parse failure. sem := make(chan struct{}, concurrency)
score := -1.0 var wg sync.WaitGroup
if llmAnswer != "" { for qi := range sample.QA {
score, err = judgeAnswer(ctx, client, qa.Question, qa.AnswerString(), llmAnswer) wg.Add(1)
if err != nil { go func() {
log.Printf("WARN: LLM judge failed for sample %s Q%d: %v", sample.SampleID, qi, err) defer wg.Done()
} sem <- struct{}{}
defer func() { <-sem }()
out := evalQAWorker(ctx, qaWork{
sampleID: sample.SampleID, qaIndex: qi,
globalIndex: si*len(sample.QA) + qi + 1, totalQA: totalQA,
qa: &sample.QA[qi], contextText: contextText, sample: sample,
}, answerClient, judgeClient, "legacy-llm")
qaResults[qi] = out.result // safe: each goroutine writes distinct index
}()
} }
wg.Wait()
hitRate := RecallHitRate(qa.Evidence, sample, contextText)
qaResults = append(qaResults, QAResult{
Question: qa.Question,
Category: qa.Category,
GoldAnswer: qa.AnswerString(),
TokenF1: score,
HitRate: hitRate,
})
log.Printf("[legacy-llm] sample=%s q=%d/%d score=%.2f answer=%q",
sample.SampleID, total, totalQA, score, truncateStr(llmAnswer, 80))
} }
results = append(results, EvalResult{ results = append(results, EvalResult{
@ -135,125 +196,126 @@ func EvalLegacyLLM(
return results return results
} }
// buildSeahorseContext retrieves context for a seahorse QA item.
func buildSeahorseContext(
ctx context.Context,
ir *SeahorseIngestResult,
sample *LocomoSample,
qa *LocomoQA,
budgetTokens int,
) string {
store := ir.Engine.GetRetrieval().Store()
retrieval := ir.Engine.GetRetrieval()
convID := ir.ConvMap[sample.SampleID]
keywords := ExtractKeywords(qa.Question)
bestRank := map[int64]float64{}
for _, kw := range keywords {
searchResults, err := store.SearchMessages(ctx, seahorse.SearchInput{
Pattern: kw,
ConversationID: convID,
Limit: 20,
})
if err != nil {
continue
}
for _, sr := range searchResults {
if sr.MessageID > 0 {
if prev, ok := bestRank[sr.MessageID]; !ok || sr.Rank < prev {
bestRank[sr.MessageID] = sr.Rank
}
}
}
}
messageIDs := make([]int64, 0, len(bestRank))
for id := range bestRank {
messageIDs = append(messageIDs, id)
}
sort.Slice(messageIDs, func(i, j int) bool {
return bestRank[messageIDs[i]] < bestRank[messageIDs[j]]
})
var contentParts []string
if len(messageIDs) > 0 {
expandResult, err := retrieval.ExpandMessages(ctx, messageIDs)
if err == nil {
for _, msg := range expandResult.Messages {
contentParts = append(contentParts, msg.Content)
}
}
}
if len(contentParts) == 0 {
return ""
}
truncated, _ := BudgetTruncate(contentParts, budgetTokens)
return StringListToContent(truncated)
}
// EvalSeahorseLLM evaluates seahorse retrieval using LLM generation + LLM-as-Judge. // EvalSeahorseLLM evaluates seahorse retrieval using LLM generation + LLM-as-Judge.
func EvalSeahorseLLM( func EvalSeahorseLLM(
ctx context.Context, ctx context.Context,
samples []LocomoSample, samples []LocomoSample,
ir *SeahorseIngestResult, ir *SeahorseIngestResult,
budgetTokens int, budgetTokens int,
client *LLMClient, answerClient, judgeClient *LLMClient,
concurrency int,
) []EvalResult { ) []EvalResult {
store := ir.Engine.GetRetrieval().Store() if concurrency < 1 {
retrieval := ir.Engine.GetRetrieval() concurrency = 1
}
totalQA := countTotalQA(samples) totalQA := countTotalQA(samples)
results := make([]EvalResult, 0, len(samples)) results := make([]EvalResult, 0, len(samples))
total := 0
for si := range samples { for si := range samples {
sample := &samples[si] sample := &samples[si]
convID, ok := ir.ConvMap[sample.SampleID] if _, ok := ir.ConvMap[sample.SampleID]; !ok {
if !ok {
log.Printf("WARN: no conversation ID for sample %s", sample.SampleID) log.Printf("WARN: no conversation ID for sample %s", sample.SampleID)
continue continue
} }
qaResults := make([]QAResult, 0, len(sample.QA)) qaResults := make([]QAResult, len(sample.QA))
for qi := range sample.QA {
evalOne := func(qi int) {
qa := &sample.QA[qi] qa := &sample.QA[qi]
total++ contextText := buildSeahorseContext(ctx, ir, sample, qa, budgetTokens)
keywords := ExtractKeywords(qa.Question) if contextText == "" {
qaResults[qi] = QAResult{
// Search and rank
bestRank := map[int64]float64{}
for _, kw := range keywords {
searchResults, err := store.SearchMessages(ctx, seahorse.SearchInput{
Pattern: kw,
ConversationID: convID,
Limit: 20,
})
if err != nil {
log.Printf("WARN: search failed for keyword %q: %v", kw, err)
continue
}
for _, sr := range searchResults {
if sr.MessageID > 0 {
if prev, ok := bestRank[sr.MessageID]; !ok || sr.Rank < prev {
bestRank[sr.MessageID] = sr.Rank
}
}
}
}
messageIDs := make([]int64, 0, len(bestRank))
for id := range bestRank {
messageIDs = append(messageIDs, id)
}
// Sort ascending: best (most-negative) rank first.
// BudgetTruncate walks front-to-back, so best-ranked messages are kept.
sort.Slice(messageIDs, func(i, j int) bool {
return bestRank[messageIDs[i]] < bestRank[messageIDs[j]]
})
var contentParts []string
if len(messageIDs) > 0 {
expandResult, err := retrieval.ExpandMessages(ctx, messageIDs)
if err != nil {
log.Printf("WARN: expand failed for sample %s: %v", sample.SampleID, err)
} else {
for _, msg := range expandResult.Messages {
contentParts = append(contentParts, msg.Content)
}
}
}
if len(contentParts) == 0 {
qaResults = append(qaResults, QAResult{
Question: qa.Question, Question: qa.Question,
Category: qa.Category, Category: qa.Category,
GoldAnswer: qa.AnswerString(), GoldAnswer: qa.AnswerString(),
TokenF1: 0.0, TokenF1: 0.0,
HitRate: 0.0, HitRate: 0.0,
}) }
log.Printf("[seahorse-llm] sample=%s q=%d/%d score=0.00 answer=(no context)", log.Printf("[seahorse-llm] sample=%s q=%d/%d score=0.00 answer=(no context)",
sample.SampleID, total, totalQA) sample.SampleID, si*len(sample.QA)+qi+1, totalQA)
continue return
} }
out := evalQAWorker(ctx, qaWork{
sampleID: sample.SampleID, qaIndex: qi,
globalIndex: si*len(sample.QA) + qi + 1, totalQA: totalQA,
qa: qa, contextText: contextText, sample: sample,
}, answerClient, judgeClient, "seahorse-llm")
qaResults[qi] = out.result
}
truncated, _ := BudgetTruncate(contentParts, budgetTokens) if concurrency <= 1 {
contextText := StringListToContent(truncated) for qi := range sample.QA {
evalOne(qi)
// Generate answer with LLM
llmAnswer := ""
score := -1.0
if contextText != "" {
var err error
llmAnswer, err = generateAnswer(ctx, client, contextText, qa.Question)
if err != nil {
log.Printf("WARN: LLM generation failed for sample %s Q%d: %v", sample.SampleID, qi, err)
}
} }
} else {
// Judge the answer; -1.0 = API/parse failure. sem := make(chan struct{}, concurrency)
if llmAnswer != "" { var wg sync.WaitGroup
var err error for qi := range sample.QA {
score, err = judgeAnswer(ctx, client, qa.Question, qa.AnswerString(), llmAnswer) wg.Add(1)
if err != nil { go func() {
log.Printf("WARN: LLM judge failed for sample %s Q%d: %v", sample.SampleID, qi, err) defer wg.Done()
} sem <- struct{}{}
defer func() { <-sem }()
evalOne(qi)
}()
} }
wg.Wait()
hitRate := RecallHitRate(qa.Evidence, sample, contextText)
qaResults = append(qaResults, QAResult{
Question: qa.Question,
Category: qa.Category,
GoldAnswer: qa.AnswerString(),
TokenF1: score,
HitRate: hitRate,
})
log.Printf("[seahorse-llm] sample=%s q=%d/%d score=%.2f answer=%q",
sample.SampleID, total, totalQA, score, truncateStr(llmAnswer, 80))
} }
results = append(results, EvalResult{ results = append(results, EvalResult{

View file

@ -71,16 +71,23 @@ type chatMessage struct {
type chatResponse struct { type chatResponse struct {
Choices []struct { Choices []struct {
Message struct { Message struct {
Content string `json:"content"` Content string `json:"content"`
ReasoningContent string `json:"reasoning_content,omitempty"`
} `json:"message"` } `json:"message"`
} `json:"choices"` } `json:"choices"`
} }
// Complete sends a chat completion request and returns the assistant's reply. // Complete sends a chat completion request and returns the assistant's reply.
func (c *LLMClient) Complete(ctx context.Context, systemPrompt, userPrompt string) (string, error) { func (c *LLMClient) Complete(ctx context.Context, systemPrompt, userPrompt string) (string, error) {
sysContent := systemPrompt
if c.NoThinking && sysContent != "" {
// Prepend /no_think tag — works with Ollama /v1 endpoint and
// Qwen chat templates where the JSON think field is ignored.
sysContent = "/no_think\n" + sysContent
}
messages := []chatMessage{} messages := []chatMessage{}
if systemPrompt != "" { if sysContent != "" {
messages = append(messages, chatMessage{Role: "system", Content: systemPrompt}) messages = append(messages, chatMessage{Role: "system", Content: sysContent})
} }
messages = append(messages, chatMessage{Role: "user", Content: userPrompt}) messages = append(messages, chatMessage{Role: "user", Content: userPrompt})
@ -88,7 +95,7 @@ func (c *LLMClient) Complete(ctx context.Context, systemPrompt, userPrompt strin
Model: c.Model, Model: c.Model,
Messages: messages, Messages: messages,
Temperature: 0.1, Temperature: 0.1,
MaxTokens: 2048, MaxTokens: 512,
} }
if c.NoThinking { if c.NoThinking {
// llama.cpp: chat_template_kwargs // llama.cpp: chat_template_kwargs
@ -180,6 +187,10 @@ func (c *LLMClient) Complete(ctx context.Context, systemPrompt, userPrompt strin
if idx := strings.Index(content, "</think>"); idx >= 0 { if idx := strings.Index(content, "</think>"); idx >= 0 {
content = strings.TrimSpace(content[idx+len("</think>"):]) content = strings.TrimSpace(content[idx+len("</think>"):])
} }
// Fallback: GLM/DeepSeek put thinking output in reasoning_content when thinking is enabled
if content == "" && chatResp.Choices[0].Message.ReasoningContent != "" {
content = strings.TrimSpace(chatResp.Choices[0].Message.ReasoningContent)
}
if content == "" { if content == "" {
return "", fmt.Errorf("empty LLM response") return "", fmt.Errorf("empty LLM response")
} }

View file

@ -16,18 +16,22 @@ import (
) )
var ( var (
flagData string flagData string
flagOut string flagOut string
flagMode string flagMode string
flagBudget int flagBudget int
flagEvalMode string flagEvalMode string
flagAPIBase string flagAPIBase string
flagAPIKey string flagAPIKey string
flagModel string flagModel string
flagNoThinking bool flagNoThinking bool
flagLimit int flagLimit int
flagTimeout int flagTimeout int
flagRetries int flagRetries int
flagJudgeModel string
flagJudgeAPIBase string
flagJudgeAPIKey string
flagConcurrency int
) )
func main() { func main() {
@ -68,6 +72,11 @@ func main() {
evalCmd.Flags().IntVar(&flagLimit, "limit", 0, "max QA questions per sample (0 = all)") evalCmd.Flags().IntVar(&flagLimit, "limit", 0, "max QA questions per sample (0 = all)")
evalCmd.Flags().IntVar(&flagTimeout, "timeout", 120, "HTTP timeout in seconds for LLM requests") evalCmd.Flags().IntVar(&flagTimeout, "timeout", 120, "HTTP timeout in seconds for LLM requests")
evalCmd.Flags().IntVar(&flagRetries, "retries", 3, "max retry attempts for transient LLM errors (timeout/5xx/429)") evalCmd.Flags().IntVar(&flagRetries, "retries", 3, "max retry attempts for transient LLM errors (timeout/5xx/429)")
evalCmd.Flags().StringVar(&flagJudgeModel, "judge-model", "", "model for judge scoring (defaults to --model)")
evalCmd.Flags().
StringVar(&flagJudgeAPIBase, "judge-api-base", "", "API base URL for judge model (defaults to --api-base)")
evalCmd.Flags().StringVar(&flagJudgeAPIKey, "judge-api-key", "", "API key for judge model (defaults to --api-key)")
evalCmd.Flags().IntVar(&flagConcurrency, "concurrency", 1, "number of concurrent QA evaluations")
reportCmd := &cobra.Command{ reportCmd := &cobra.Command{
Use: "report", Use: "report",
@ -96,6 +105,11 @@ func main() {
runCmd.Flags().IntVar(&flagLimit, "limit", 0, "max QA questions per sample (0 = all)") runCmd.Flags().IntVar(&flagLimit, "limit", 0, "max QA questions per sample (0 = all)")
runCmd.Flags().IntVar(&flagTimeout, "timeout", 120, "HTTP timeout in seconds for LLM requests") runCmd.Flags().IntVar(&flagTimeout, "timeout", 120, "HTTP timeout in seconds for LLM requests")
runCmd.Flags().IntVar(&flagRetries, "retries", 3, "max retry attempts for transient LLM errors (timeout/5xx/429)") runCmd.Flags().IntVar(&flagRetries, "retries", 3, "max retry attempts for transient LLM errors (timeout/5xx/429)")
runCmd.Flags().StringVar(&flagJudgeModel, "judge-model", "", "model for judge scoring (defaults to --model)")
runCmd.Flags().
StringVar(&flagJudgeAPIBase, "judge-api-base", "", "API base URL for judge model (defaults to --api-base)")
runCmd.Flags().StringVar(&flagJudgeAPIKey, "judge-api-key", "", "API key for judge model (defaults to --api-key)")
runCmd.Flags().IntVar(&flagConcurrency, "concurrency", 1, "number of concurrent QA evaluations")
rootCmd.AddCommand(ingestCmd, evalCmd, reportCmd, runCmd) rootCmd.AddCommand(ingestCmd, evalCmd, reportCmd, runCmd)
@ -186,14 +200,28 @@ func runEval(cmd *cobra.Command, args []string) error {
default: default:
return fmt.Errorf("invalid --eval-mode %q: must be token or llm", flagEvalMode) return fmt.Errorf("invalid --eval-mode %q: must be token or llm", flagEvalMode)
} }
var llmClient *LLMClient var answerClient, judgeClient *LLMClient
if useLLM { if useLLM {
opts, err := buildLLMOptions() opts, err := buildLLMOptions()
if err != nil { if err != nil {
return err return err
} }
llmClient = NewLLMClient(opts) answerClient = NewLLMClient(opts)
log.Printf("LLM eval mode: model=%s base=%s no-thinking=%v", opts.Model, opts.BaseURL, opts.NoThinking) judgeClient = answerClient // default: same client
if flagJudgeModel != "" {
jOpts := opts // copy base settings
jOpts.Model = flagJudgeModel
if flagJudgeAPIBase != "" {
jOpts.BaseURL = flagJudgeAPIBase
}
if flagJudgeAPIKey != "" {
jOpts.APIKey = flagJudgeAPIKey
}
judgeClient = NewLLMClient(jOpts)
log.Printf("Judge model: model=%s base=%s no-thinking=%v", jOpts.Model, jOpts.BaseURL, jOpts.NoThinking)
}
log.Printf("LLM eval mode: model=%s base=%s no-thinking=%v concurrency=%d",
opts.Model, opts.BaseURL, opts.NoThinking, flagConcurrency)
} }
var tokenResults, llmResults []EvalResult var tokenResults, llmResults []EvalResult
@ -206,7 +234,7 @@ func runEval(cmd *cobra.Command, args []string) error {
legacy.IngestSample(&samples[i]) legacy.IngestSample(&samples[i])
} }
if useLLM { if useLLM {
results := EvalLegacyLLM(ctx, samples, legacy, flagBudget, llmClient) results := EvalLegacyLLM(ctx, samples, legacy, flagBudget, answerClient, judgeClient, flagConcurrency)
llmResults = append(llmResults, results...) llmResults = append(llmResults, results...)
log.Printf("legacy-llm: evaluated %d samples", len(results)) log.Printf("legacy-llm: evaluated %d samples", len(results))
} else { } else {
@ -221,7 +249,7 @@ func runEval(cmd *cobra.Command, args []string) error {
return fmt.Errorf("ingest seahorse: %w", err) return fmt.Errorf("ingest seahorse: %w", err)
} }
if useLLM { if useLLM {
results := EvalSeahorseLLM(ctx, samples, ir, flagBudget, llmClient) results := EvalSeahorseLLM(ctx, samples, ir, flagBudget, answerClient, judgeClient, flagConcurrency)
llmResults = append(llmResults, results...) llmResults = append(llmResults, results...)
log.Printf("seahorse-llm: evaluated %d samples", len(results)) log.Printf("seahorse-llm: evaluated %d samples", len(results))
} else { } else {