diff options
Diffstat (limited to 'llm')
| -rw-r--r-- | llm/cmd/eval/logprobs.go | 8 | ||||
| -rw-r--r-- | llm/cmd/eval/ppl.go | 8 | ||||
| -rw-r--r-- | llm/evaluator.go | 18 | ||||
| -rw-r--r-- | llm/perplexity_test.go | 2 |
4 files changed, 9 insertions, 27 deletions
diff --git a/llm/cmd/eval/logprobs.go b/llm/cmd/eval/logprobs.go index eaa7459..7c04626 100644 --- a/llm/cmd/eval/logprobs.go +++ b/llm/cmd/eval/logprobs.go @@ -24,15 +24,11 @@ type logProbs []logProb func logprobs() { d := data() + m := model() t := tokenizer() - e := llm.NewEvaluator[logProbs]() - - e.SetTokenizer(t) - e.AddModel(m) - - e.SetCallback(func(j llm.Job, logits [][]float32, tokens []int) logProbs { + e := llm.NewEvaluator(m, t, func(j llm.Job, logits [][]float32, tokens []int) logProbs { probs := selectLogProbs(logits, tokens) r := make(logProbs, len(tokens)) diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go index fdbcff3..0a028d2 100644 --- a/llm/cmd/eval/ppl.go +++ b/llm/cmd/eval/ppl.go @@ -18,15 +18,11 @@ type pplResult struct { func perplexity() { d := data() + m := model() t := tokenizer() - e := llm.NewEvaluator[pplResult]() - - e.SetTokenizer(t) - e.AddModel(m) - - e.SetCallback(func(job llm.Job, logits [][]float32, tokens []int) pplResult { + e := llm.NewEvaluator(m, t, func(job llm.Job, logits [][]float32, tokens []int) pplResult { p, n := llm.NegLogLikelihood(logits, tokens) return pplResult{ diff --git a/llm/evaluator.go b/llm/evaluator.go index 24fe2d4..0eb5418 100644 --- a/llm/evaluator.go +++ b/llm/evaluator.go @@ -20,28 +20,18 @@ type Evaluator[R any] struct { callback func(job Job, logits [][]float32, tokens []int) R } -func NewEvaluator[R any]() *Evaluator[R] { +func NewEvaluator[R any](model Causal, tokenizer Tokenizer, callback func(job Job, logits [][]float32, tokens []int) R) *Evaluator[R] { return &Evaluator[R]{ - models: make([]Causal, 0), + models: []Causal{model}, // TODO multiple devices + tokenizer: tokenizer, batchSize: 1, // TODO arg numWorkers: 4, // TODO arg jobs: make(chan batch, 1024), results: make(chan R, 1024), + callback: callback, } } -func (e *Evaluator[R]) AddModel(model Causal) { - e.models = append(e.models, model) -} - -func (e *Evaluator[R]) SetTokenizer(tokenizer Tokenizer) { - e.tokenizer = tokenizer -} - -func (e *Evaluator[R]) SetCallback(callback func(job Job, logits [][]float32, tokens []int) R) { - e.callback = callback -} - func (e *Evaluator[R]) Results() chan R { return e.results } diff --git a/llm/perplexity_test.go b/llm/perplexity_test.go index 81f10af..fa974a7 100644 --- a/llm/perplexity_test.go +++ b/llm/perplexity_test.go @@ -25,7 +25,7 @@ func TestEvaluator_estimateJobs(t *testing.T) { {tokens: 1, window: 1, stride: 1024, expected: 1}, } - e := NewEvaluator[any]() + e := NewEvaluator[any](nil, nil, nil) for _, tt := range tests { t.Run( |
