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))
for qi := range sample.QA {
qa := &sample.QA[qi]
total++
truncated, _ := BudgetTruncate(allContent, budgetTokens) truncated, _ := BudgetTruncate(allContent, budgetTokens)
contextText := StringListToContent(truncated) 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,35 +196,19 @@ func EvalLegacyLLM(
return results return results
} }
// EvalSeahorseLLM evaluates seahorse retrieval using LLM generation + LLM-as-Judge. // buildSeahorseContext retrieves context for a seahorse QA item.
func EvalSeahorseLLM( func buildSeahorseContext(
ctx context.Context, ctx context.Context,
samples []LocomoSample,
ir *SeahorseIngestResult, ir *SeahorseIngestResult,
sample *LocomoSample,
qa *LocomoQA,
budgetTokens int, budgetTokens int,
client *LLMClient, ) string {
) []EvalResult {
store := ir.Engine.GetRetrieval().Store() store := ir.Engine.GetRetrieval().Store()
retrieval := ir.Engine.GetRetrieval() retrieval := ir.Engine.GetRetrieval()
convID := ir.ConvMap[sample.SampleID]
totalQA := countTotalQA(samples)
results := make([]EvalResult, 0, len(samples))
total := 0
for si := range samples {
sample := &samples[si]
convID, ok := ir.ConvMap[sample.SampleID]
if !ok {
log.Printf("WARN: no conversation ID for sample %s", sample.SampleID)
continue
}
qaResults := make([]QAResult, 0, len(sample.QA))
for qi := range sample.QA {
qa := &sample.QA[qi]
total++
keywords := ExtractKeywords(qa.Question) keywords := ExtractKeywords(qa.Question)
// Search and rank
bestRank := map[int64]float64{} bestRank := map[int64]float64{}
for _, kw := range keywords { for _, kw := range keywords {
searchResults, err := store.SearchMessages(ctx, seahorse.SearchInput{ searchResults, err := store.SearchMessages(ctx, seahorse.SearchInput{
@ -172,7 +217,6 @@ func EvalSeahorseLLM(
Limit: 20, Limit: 20,
}) })
if err != nil { if err != nil {
log.Printf("WARN: search failed for keyword %q: %v", kw, err)
continue continue
} }
for _, sr := range searchResults { for _, sr := range searchResults {
@ -188,8 +232,6 @@ func EvalSeahorseLLM(
for id := range bestRank { for id := range bestRank {
messageIDs = append(messageIDs, id) 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 { sort.Slice(messageIDs, func(i, j int) bool {
return bestRank[messageIDs[i]] < bestRank[messageIDs[j]] return bestRank[messageIDs[i]] < bestRank[messageIDs[j]]
}) })
@ -197,63 +239,83 @@ func EvalSeahorseLLM(
var contentParts []string var contentParts []string
if len(messageIDs) > 0 { if len(messageIDs) > 0 {
expandResult, err := retrieval.ExpandMessages(ctx, messageIDs) expandResult, err := retrieval.ExpandMessages(ctx, messageIDs)
if err != nil { if err == nil {
log.Printf("WARN: expand failed for sample %s: %v", sample.SampleID, err)
} else {
for _, msg := range expandResult.Messages { for _, msg := range expandResult.Messages {
contentParts = append(contentParts, msg.Content) contentParts = append(contentParts, msg.Content)
} }
} }
} }
if len(contentParts) == 0 { if len(contentParts) == 0 {
qaResults = append(qaResults, QAResult{ return ""
}
truncated, _ := BudgetTruncate(contentParts, budgetTokens)
return StringListToContent(truncated)
}
// EvalSeahorseLLM evaluates seahorse retrieval using LLM generation + LLM-as-Judge.
func EvalSeahorseLLM(
ctx context.Context,
samples []LocomoSample,
ir *SeahorseIngestResult,
budgetTokens int,
answerClient, judgeClient *LLMClient,
concurrency int,
) []EvalResult {
if concurrency < 1 {
concurrency = 1
}
totalQA := countTotalQA(samples)
results := make([]EvalResult, 0, len(samples))
for si := range samples {
sample := &samples[si]
if _, ok := ir.ConvMap[sample.SampleID]; !ok {
log.Printf("WARN: no conversation ID for sample %s", sample.SampleID)
continue
}
qaResults := make([]QAResult, len(sample.QA))
evalOne := func(qi int) {
qa := &sample.QA[qi]
contextText := buildSeahorseContext(ctx, ir, sample, qa, budgetTokens)
if contextText == "" {
qaResults[qi] = 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 {
sem := make(chan struct{}, concurrency)
var wg sync.WaitGroup
for qi := range sample.QA {
wg.Add(1)
go func() {
defer wg.Done()
sem <- struct{}{}
defer func() { <-sem }()
evalOne(qi)
}()
} }
wg.Wait()
// Judge the answer; -1.0 = API/parse failure.
if llmAnswer != "" {
var err error
score, err = judgeAnswer(ctx, client, qa.Question, qa.AnswerString(), llmAnswer)
if err != nil {
log.Printf("WARN: LLM judge failed for sample %s Q%d: %v", sample.SampleID, qi, err)
}
}
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

@ -72,15 +72,22 @@ 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

@ -28,6 +28,10 @@ var (
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 {