summaryrefslogtreecommitdiff
path: root/llm
diff options
context:
space:
mode:
Diffstat (limited to 'llm')
-rw-r--r--llm/evaluator.go26
-rw-r--r--llm/job.go19
-rw-r--r--llm/perplexity.go221
-rw-r--r--llm/pool.go53
-rw-r--r--llm/utility.go21
5 files changed, 340 insertions, 0 deletions
diff --git a/llm/evaluator.go b/llm/evaluator.go
new file mode 100644
index 0000000..ccc554e
--- /dev/null
+++ b/llm/evaluator.go
@@ -0,0 +1,26 @@
+package llm
+
+import "sync"
+
+type Evaluator struct {
+ mutex sync.RWMutex
+ models []Causal
+ tokenizer Tokenizer
+ results []float64
+ jobs int
+}
+
+func NewEvaluator() *Evaluator {
+ return &Evaluator{
+ models: make([]Causal, 0),
+ results: make([]float64, 0),
+ }
+}
+
+func (e *Evaluator) AddModel(model Causal) {
+ e.models = append(e.models, model)
+}
+
+func (e *Evaluator) SetTokenizer(tokenizer Tokenizer) {
+ e.tokenizer = tokenizer
+}
diff --git a/llm/job.go b/llm/job.go
new file mode 100644
index 0000000..3e7c5c9
--- /dev/null
+++ b/llm/job.go
@@ -0,0 +1,19 @@
+package llm
+
+type job struct {
+ positions []int
+ tokens [][]int64
+ seen []int
+ results []float64
+ debug [][2]int
+}
+
+func newJob(batchSize int) *job {
+ return &job{
+ positions: make([]int, 0, batchSize),
+ tokens: make([][]int64, 0, batchSize),
+ seen: make([]int, 0, batchSize),
+ results: make([]float64, 0, batchSize),
+ debug: make([][2]int, 0),
+ }
+}
diff --git a/llm/perplexity.go b/llm/perplexity.go
new file mode 100644
index 0000000..dec3945
--- /dev/null
+++ b/llm/perplexity.go
@@ -0,0 +1,221 @@
+package llm
+
+import (
+ "bufio"
+ "context"
+ "log"
+ "math"
+ mbpe "mbpe-dyn"
+ "slices"
+ "sync"
+ "time"
+)
+
+func (e *Evaluator) Perplexity(name string) (float64, error) {
+ tokens := make([]int64, 0)
+
+ if err := mbpe.FromFile(name, func(scanner *bufio.Scanner) error {
+ for scanner.Scan() {
+ line := scanner.Text()
+
+ if err := scanner.Err(); err != nil {
+ return err
+ }
+
+ line += "\n"
+
+ tokens = append(tokens, toInt64(e.tokenizer.Tokenize(line))...)
+ }
+
+ return nil
+ }); err != nil {
+ return 0, err
+ }
+
+ contextSize, stride, batchSize := 64, 32, 1
+
+ if len(tokens) < contextSize {
+ return 0, nil // TODO handle
+ }
+
+ windows := ((len(tokens) - contextSize) / stride) + 1
+ jobs := (windows + batchSize - 1) / batchSize
+
+ pb := mbpe.NewProgressBar("Perplexity", 20, jobs, time.Now())
+
+ ctx, cancel := context.WithCancel(context.Background())
+
+ defer cancel()
+
+ done := make(chan struct{})
+
+ go func(ctx context.Context) {
+ main:
+ for {
+ select {
+ case <-ctx.Done():
+ break main
+ default:
+ time.Sleep(time.Second * 1)
+
+ e.mutex.RLock()
+
+ j := e.jobs
+
+ e.mutex.RUnlock()
+
+ pb.Update(j)
+ pb.Print()
+
+ if j >= jobs {
+ break main
+ }
+ }
+ }
+
+ pb.Finish()
+
+ close(done)
+ }(ctx)
+
+ if err := e.schedule(tokens, contextSize, stride, batchSize); err != nil {
+ log.Fatal(err)
+ }
+
+ total := float64(0)
+
+ for _, nll := range e.results {
+ total += nll
+ }
+
+ average := total / float64(len(e.results)) * float64(contextSize-1)
+
+ return math.Exp(total / average), nil
+}
+
+func (e *Evaluator) schedule(tokens []int64, contextSize, stride, batchSize int) error {
+ jobs := make(chan *job)
+
+ var wg sync.WaitGroup
+
+ devices := make([]int, len(e.models))
+
+ for i := range len(devices) {
+ devices[i] = i
+ }
+
+ devicePool := newPool[int](devices...)
+
+ for d := 0; d < devicePool.Len(); d++ {
+ wg.Add(1)
+
+ go func() {
+ defer wg.Done()
+
+ for j := range jobs {
+ device := devicePool.Acquire()
+
+ e.execute(j, device)
+
+ devicePool.Release(device)
+
+ e.mutex.Lock()
+
+ for _, p := range j.results {
+ e.results = append(e.results, p)
+ }
+
+ e.jobs++
+
+ e.mutex.Unlock()
+ }
+ }()
+ }
+
+ j := newJob(batchSize)
+
+ // 0 to 1023: full logits
+ // 512 to 1535: 1024 upwards
+ // 1024 to 2047: 1536 upwards
+ // ...
+
+ seen := 0
+ n := 0
+
+ // for i := 0; i < len(tokens); i += stride {
+ for i := 0; i+contextSize <= len(tokens); i += stride {
+ // if (len(j.positions) == batchSize) || i+stride > len(tokens) {
+ if len(j.positions) == batchSize {
+ jobs <- j
+
+ j = newJob(batchSize)
+ }
+
+ j.positions = append(j.positions, n)
+ j.tokens = append(j.tokens, tokens[i:min(i+contextSize, len(tokens))]) // TODO verify
+ j.seen = append(j.seen, seen)
+
+ seen = i + contextSize
+ n++
+ }
+
+ if len(j.positions) > 0 {
+ jobs <- j
+ }
+
+ close(jobs)
+
+ wg.Wait()
+
+ return nil
+}
+
+func (e *Evaluator) execute(j *job, device int) {
+ if len(j.positions) != 1 {
+ panic("unimplemented")
+ }
+
+ m := e.models[device]
+
+ logits := make([][]float32, 0, len(j.tokens[0]))
+
+ // fmt.Println("executing job", j.positions[0])
+
+ if _, err := m.Generate(j.tokens[0], 0, &logits); err != nil {
+ panic(err) // TODO handle
+ }
+
+ nll := negLogLikelihood(logits[:len(logits)-1], toInt(j.tokens[0][1:]))
+
+ j.results = append(j.results, nll)
+
+ return
+}
+
+func negLogLikelihood(logits [][]float32, targets []int) float64 {
+ if len(logits) != len(targets) {
+ panic("mismatched input lengths")
+ }
+
+ total := float64(0)
+
+ for i, target := range targets {
+ maxLogit := float64(slices.Max(logits[i]))
+
+ sumExp := float64(0)
+
+ for _, v := range logits[i] {
+ sumExp += math.Exp(float64(v) - maxLogit)
+ }
+
+ logSumExp := maxLogit + math.Log(sumExp)
+
+ targetLogit := float64(logits[i][target])
+
+ logProb := targetLogit - logSumExp
+
+ total -= logProb
+ }
+
+ return total
+}
diff --git a/llm/pool.go b/llm/pool.go
new file mode 100644
index 0000000..294b64f
--- /dev/null
+++ b/llm/pool.go
@@ -0,0 +1,53 @@
+package llm
+
+import "sync"
+
+type pool[K comparable] struct {
+ devices map[K]bool
+ mutex sync.Mutex
+ cond *sync.Cond
+}
+
+func newPool[K comparable](devices ...K) *pool[K] {
+ p := &pool[K]{
+ devices: make(map[K]bool),
+ }
+
+ for _, d := range devices {
+ p.devices[d] = true
+ }
+
+ p.cond = sync.NewCond(&p.mutex)
+
+ return p
+}
+
+func (p *pool[K]) Len() int {
+ return len(p.devices)
+}
+
+func (p *pool[K]) Acquire() K {
+ p.mutex.Lock()
+ defer p.mutex.Unlock()
+
+ for {
+ for k, v := range p.devices {
+ if v {
+ p.devices[k] = false
+
+ return k
+ }
+ }
+
+ p.cond.Wait()
+ }
+}
+
+func (p *pool[K]) Release(device K) {
+ p.mutex.Lock()
+
+ p.devices[device] = true
+
+ p.cond.Signal()
+ p.mutex.Unlock()
+}
diff --git a/llm/utility.go b/llm/utility.go
new file mode 100644
index 0000000..e27ea2e
--- /dev/null
+++ b/llm/utility.go
@@ -0,0 +1,21 @@
+package llm
+
+func toInt64(s []int) []int64 {
+ r := make([]int64, len(s))
+
+ for i, v := range s {
+ r[i] = int64(v)
+ }
+
+ return r
+}
+
+func toInt(s []int64) []int {
+ r := make([]int, len(s))
+
+ for i, v := range s {
+ r[i] = int(v)
+ }
+
+ return r
+}