diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-03-25 21:50:15 +0100 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-03-25 21:50:15 +0100 |
| commit | a52a3cd0e6d6d999e93877150adf4c4ec19deff0 (patch) | |
| tree | f10b6b1a35cb504d6b467b9e8adc87376282bbfd /llm/perplexity.go | |
| parent | 0c6a2cadf261945558c75256d32fadf569b4a48d (diff) | |
Refactor llm module
Diffstat (limited to 'llm/perplexity.go')
| -rw-r--r-- | llm/perplexity.go | 159 |
1 files changed, 0 insertions, 159 deletions
diff --git a/llm/perplexity.go b/llm/perplexity.go deleted file mode 100644 index df48ff6..0000000 --- a/llm/perplexity.go +++ /dev/null @@ -1,159 +0,0 @@ -package llm - -import ( - "fmt" - "sync" - "time" - - "go.jknobloc.com/x/dataset" - "go.jknobloc.com/x/tui" -) - -func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int) error { - devices := make([]int, len(e.models)) - - for i := range len(devices) { - devices[i] = i - } - - devicePool := newPool(devices...) - - var wg sync.WaitGroup - - for range e.numWorkers { - wg.Add(1) - - go func() { - defer wg.Done() - - for b := range e.jobs { - device := devicePool.Acquire() - - func() { - defer devicePool.Release(device) - - defer func() { - if r := recover(); r != nil { - fmt.Println("HOUSTON") // TODO handle - } - }() - - e.execute(&b, device) - }() - - e.completed.Add(int64(b.Size())) - } - }() - } - - pb := tui.NewProgressBar(title, 20, 0, time.Now()) - - pb.Start(1*time.Second, func() int { - return int(e.completed.Load()) - }) - - defer pb.Close() - - n := 0 - - for d := range data.Texts("text") { - tokens := toInt64(e.tokenizer.Tokenize(d)) - - e.schedule(n, tokens, window, stride, e.batchSize) - - pb.SetTotal(pb.Total() + e.estimateJobs(tokens, window, stride)) - - n++ - } - - close(e.jobs) - - wg.Wait() - - close(e.results) - - return nil -} - -func (e *Evaluator[R]) estimateJobs(tokens []int64, window, stride int) int { - if len(tokens) < window { - return 0 - } - - windows := ((len(tokens) - window) / stride) + 1 - // jobs := (windows + batchSize - 1) / e.batchSize - - return windows -} - -func (e *Evaluator[R]) schedule(uid int, tokens []int64, contextSize, stride, batchSize int) { - b := newBatch(batchSize) - - seen := 1 // first token as context - n := 0 - - // for i := 0; i+contextSize <= len(tokens); i += stride { - for i := 0; i < len(tokens); i += stride { - if b.Size() == batchSize { - e.jobs <- *b - - b = newBatch(batchSize) - } - - j := min(i+contextSize, len(tokens)) - - if j-i < contextSize { - break // don't add jobs with partial windows - } - - // if j < seen { - // break // don't add jobs with no new tokens - // } - - b.AddJob(Job{ - Document: uid, - Position: n, - Tokens: tokens[i:j], - Seen: seen - i, - }) - - seen = i + contextSize - - n++ - } - - if b.Size() > 0 { - e.jobs <- *b - } - - e.scheduled.Add(int64(n)) -} - -func (e *Evaluator[R]) execute(j *batch, device int) { - if j.Size() != 1 { - panic("unimplemented") - } - - job := j.jobs[0] - - if job.Seen < 1 { - panic("empty context") - } - - m := e.models[device] - - logits := make([][]float32, 0, len(job.Tokens)) - - if _, err := m.Generate(job.Tokens, 0, &logits); err != nil { - panic(err) // TODO handle - } - - l := logits[job.Seen-1 : len(logits)-1] - t := toInt(job.Tokens[job.Seen:]) - - r := e.callback(job, l, t) - - e.results <- r - - return -} |
