summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--bpc/bpc.go22
-rw-r--r--bpc/cmd/bpc/main.go3
-rw-r--r--gpt2/cmd/gpt2/main.go2
-rw-r--r--gpt2/cuda.go36
-rw-r--r--gpt2/model.go39
-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
10 files changed, 413 insertions, 29 deletions
diff --git a/bpc/bpc.go b/bpc/bpc.go
index 67804d3..fd4ab6a 100644
--- a/bpc/bpc.go
+++ b/bpc/bpc.go
@@ -3,26 +3,20 @@ package bpc
import (
"fmt"
"llm"
+ "log"
)
func Run(model llm.Causal, tokenizer llm.Tokenizer) {
- t := tokenizer.Tokenize("The quick brown")
+ e := llm.NewEvaluator()
- logits := make([][]float32, 0)
+ e.AddModel(model)
+ e.SetTokenizer(tokenizer)
- if _, err := model.Generate(toInt64(t), 0, &logits); err != nil {
- // TODO handle
- }
-
- fmt.Println(logits)
-}
-
-func toInt64(s []int) []int64 {
- r := make([]int64, len(s))
+ ppl, err := e.Perplexity("data/shakespeare.txt")
- for i, v := range s {
- r[i] = int64(v)
+ if err != nil {
+ log.Fatal(err)
}
- return r
+ fmt.Printf("\nPerplexity: %.2f\n", ppl)
}
diff --git a/bpc/cmd/bpc/main.go b/bpc/cmd/bpc/main.go
index b6ebef1..8f7fc1e 100644
--- a/bpc/cmd/bpc/main.go
+++ b/bpc/cmd/bpc/main.go
@@ -5,6 +5,7 @@ import (
"gpt2"
"log"
mbpe "mbpe-dyn"
+ "os"
)
func main() {
@@ -24,7 +25,7 @@ func main() {
}
func model() *gpt2.Model {
- return gpt2.NewModel("../gpt2/models/base/model.onnx")
+ return gpt2.NewModel("../gpt2/models/base/model.onnx", os.Getenv("BPC_CUDA_DEVICE_ID"))
}
func tokenizer() *mbpe.Tokenizer {
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go
index 96dabeb..b1e8611 100644
--- a/gpt2/cmd/gpt2/main.go
+++ b/gpt2/cmd/gpt2/main.go
@@ -9,7 +9,7 @@ import (
func main() {
prompt := []int64{464, 2068, 7586, 21831}
- m := gpt2.NewModel("models/base/model.onnx")
+ m := gpt2.NewModel("models/base/model.onnx", "")
if err := m.Init(); err != nil {
log.Fatal(err)
diff --git a/gpt2/cuda.go b/gpt2/cuda.go
new file mode 100644
index 0000000..ae3fbcd
--- /dev/null
+++ b/gpt2/cuda.go
@@ -0,0 +1,36 @@
+package gpt2
+
+import ort "github.com/yalue/onnxruntime_go"
+
+func SessionsOptionsWithCUDADeviceID(deviceID string) (*ort.SessionOptions, error) {
+ var sessionOptions *ort.SessionOptions
+ var cudaProviderOptions *ort.CUDAProviderOptions
+
+ if s, err := ort.NewSessionOptions(); err != nil {
+ return nil, err
+ } else {
+ sessionOptions = s
+ }
+
+ if c, err := ort.NewCUDAProviderOptions(); err != nil {
+ return nil, err
+ } else {
+ cudaProviderOptions = c
+ }
+
+ if err := cudaProviderOptions.Update(map[string]string{
+ "device_id": deviceID,
+ }); err != nil {
+ return nil, err
+ }
+
+ if err := sessionOptions.AppendExecutionProviderCUDA(cudaProviderOptions); err != nil {
+ return nil, err
+ }
+
+ if err := cudaProviderOptions.Destroy(); err != nil {
+ return nil, err
+ }
+
+ return sessionOptions, nil
+}
diff --git a/gpt2/model.go b/gpt2/model.go
index 1b7a9d1..fed1712 100644
--- a/gpt2/model.go
+++ b/gpt2/model.go
@@ -7,6 +7,7 @@ import (
"log"
"math"
"os"
+ "slices"
"sort"
ort "github.com/yalue/onnxruntime_go"
@@ -20,12 +21,14 @@ const (
)
type Model struct {
- name string
+ name string
+ deviceID string
}
-func NewModel(name string) *Model {
+func NewModel(name, deviceID string) *Model {
return &Model{
- name: name,
+ name: name,
+ deviceID: deviceID,
}
}
@@ -63,7 +66,7 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in
out := make([]int64, 0, steps+1)
for step := range context + steps {
- _, _, outputs, err := forward(m.name, token, step, cacheNames, cacheValues)
+ _, _, outputs, err := m.forward(m.name, token, step, cacheNames, cacheValues)
if err != nil {
return nil, err
@@ -96,26 +99,42 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in
return out[:steps], nil
}
-func forward(model string, token int64, position int64, cacheNames []string, cacheValues []ort.Value) (*ort.Tensor[float32], []string, []ort.Value, error) {
+func (m *Model) forward(model string, token int64, position int64, cacheNames []string, cacheValues []ort.Value) (*ort.Tensor[float32], []string, []ort.Value, error) {
inputNames, inputs, _ := initInputs(token, position)
outputNames, outputs, logits, _ := initOutputs(position)
inputNames = append(inputNames, cacheNames...)
inputs = append(inputs, cacheValues...)
+ var options *ort.SessionOptions
+
+ if m.deviceID != "" {
+ if opts, err := SessionsOptionsWithCUDADeviceID(m.deviceID); err != nil {
+ return nil, nil, nil, err
+ } else {
+ options = opts
+ }
+ }
+
session, err := ort.NewAdvancedSession(
model,
inputNames,
outputNames,
inputs,
outputs,
- nil,
+ options,
)
if err != nil {
log.Fatal(err)
}
+ if options != nil {
+ if err := options.Destroy(); err != nil {
+ return nil, nil, nil, err
+ }
+ }
+
defer session.Destroy()
if err := session.Run(); err != nil {
@@ -211,13 +230,7 @@ func initOutputs(position int64) ([]string, []ort.Value, *ort.Tensor[float32], e
}
func softmax(logits []float32) []float32 {
- m := logits[0]
-
- for _, v := range logits {
- if v > m {
- m = v
- }
- }
+ m := slices.Max(logits)
s := float32(0.0)
r := make([]float32, len(logits))
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
+}