diff options
| -rw-r--r-- | bpc/bpc.go | 22 | ||||
| -rw-r--r-- | bpc/cmd/bpc/main.go | 3 | ||||
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 2 | ||||
| -rw-r--r-- | gpt2/cuda.go | 36 | ||||
| -rw-r--r-- | gpt2/model.go | 39 | ||||
| -rw-r--r-- | llm/evaluator.go | 26 | ||||
| -rw-r--r-- | llm/job.go | 19 | ||||
| -rw-r--r-- | llm/perplexity.go | 221 | ||||
| -rw-r--r-- | llm/pool.go | 53 | ||||
| -rw-r--r-- | llm/utility.go | 21 |
10 files changed, 413 insertions, 29 deletions
@@ -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 +} |
