summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--gpt2/allocator.go98
-rw-r--r--gpt2/cmd/gpt2/main.go2
-rw-r--r--gpt2/model.go6
-rw-r--r--llm/causal.go2
-rw-r--r--llm/run.go30
5 files changed, 82 insertions, 56 deletions
diff --git a/gpt2/allocator.go b/gpt2/allocator.go
index 2609aeb..6c34192 100644
--- a/gpt2/allocator.go
+++ b/gpt2/allocator.go
@@ -7,19 +7,22 @@ import (
)
type Allocator struct {
- config Config
- step int64
- inputNames []string
- outputNames []string
- values map[string]ort.Value
- withCache bool
- withLogits bool
- withLogProbs bool
+ config Config
+ batchSize int
+ sequenceLength int64
+ step int64
+ inputNames []string
+ outputNames []string
+ values map[string]ort.Value
+ withCache bool
+ withLogits bool
+ withLogProbs bool
}
-func NewAllocator(config Config, withCache bool, withLogits bool, withLogProbs bool) *Allocator {
+func NewAllocator(config Config, batchSize int, withCache bool, withLogits bool, withLogProbs bool) *Allocator {
return &Allocator{
config: config,
+ batchSize: batchSize,
values: make(map[string]ort.Value),
withCache: withCache,
withLogits: withLogits,
@@ -81,10 +84,19 @@ func (a *Allocator) OutputNames() []string {
return names
}
+func (a *Allocator) SetBatchSize(n int) {
+ a.batchSize = n
+}
+
func (a *Allocator) Init(tokens []int64) error {
+ if len(tokens)%a.batchSize != 0 {
+ panic("token count is not divisible by batch size")
+ }
+
a.Destroy()
a.values = make(map[string]ort.Value)
+ a.sequenceLength = int64(len(tokens) / a.batchSize)
a.step = 0
if err := a.initInputs(tokens); err != nil {
@@ -99,7 +111,7 @@ func (a *Allocator) Init(tokens []int64) error {
return err
}
- a.step = int64(len(tokens))
+ a.step = a.sequenceLength
return nil
}
@@ -119,13 +131,13 @@ func (a *Allocator) initInputs(tokens []int64) error {
names = append(names, "input_ids")
- if err := a.positionIDs(tokens, 0, false); err != nil {
+ if err := a.positionIDs(0, false); err != nil {
return err
}
names = append(names, "position_ids")
- if err := a.attentionMask(tokens, 0, false); err != nil {
+ if err := a.attentionMask(0, false); err != nil {
return err
}
@@ -174,7 +186,7 @@ func (a *Allocator) initOutputs(tokens []int64) error {
names := make([]string, 0, capacity)
if a.withLogits {
- if err := a.logits(tokens, false); err != nil {
+ if err := a.logits(false); err != nil {
return err
}
@@ -182,7 +194,7 @@ func (a *Allocator) initOutputs(tokens []int64) error {
}
if a.withLogProbs {
- if err := a.logProbs(tokens, false); err != nil {
+ if err := a.logProbs(false); err != nil {
return err
}
@@ -196,13 +208,13 @@ func (a *Allocator) initOutputs(tokens []int64) error {
}
for i := range int64(a.config.nLayers) {
- if err := a.presentKeyValues(tokens, 0, i, "key", false); err != nil {
+ if err := a.presentKeyValues(0, i, "key", false); err != nil {
return err
}
names = append(names, fmt.Sprintf("present.%d.key", i))
- if err := a.presentKeyValues(tokens, 0, i, "value", false); err != nil {
+ if err := a.presentKeyValues(0, i, "value", false); err != nil {
return err
}
@@ -215,35 +227,41 @@ func (a *Allocator) initOutputs(tokens []int64) error {
}
func (a *Allocator) Step(token int64) error {
+ if a.batchSize != 1 {
+ panic("step requires batch size 1")
+ }
+
tokens := []int64{token}
+ a.sequenceLength = 1
+
if err := a.inputIDs(tokens, true); err != nil {
return err
}
- if err := a.positionIDs(tokens, a.step, true); err != nil {
+ if err := a.positionIDs(a.step, true); err != nil {
return err
}
- if err := a.attentionMask(tokens, a.step, true); err != nil {
+ if err := a.attentionMask(a.step, true); err != nil {
return err
}
if a.withLogits {
- if err := a.logits(tokens, true); err != nil {
+ if err := a.logits(true); err != nil {
return err
}
}
if a.withLogProbs {
- if err := a.logProbs(tokens, true); err != nil {
+ if err := a.logProbs(true); err != nil {
return err
}
}
for i := range int64(a.config.nLayers) {
for _, suffix := range []string{"key", "value"} {
- if err := a.rotateCache(tokens, i, suffix); err != nil {
+ if err := a.rotateCache(i, suffix); err != nil {
return err
}
}
@@ -301,7 +319,7 @@ func (a *Allocator) inputIDs(tokens []int64, force bool) error {
_ = a.values[name].Destroy()
}
- shape := []int64{1, int64(len(tokens))}
+ shape := []int64{int64(a.batchSize), a.sequenceLength}
if t, err := ort.NewTensor[int64](shape, tokens); err != nil {
return err
@@ -312,7 +330,7 @@ func (a *Allocator) inputIDs(tokens []int64, force bool) error {
return nil
}
-func (a *Allocator) positionIDs(tokens []int64, start int64, force bool) error {
+func (a *Allocator) positionIDs(start int64, force bool) error {
const name = "position_ids"
if _, ok := a.values[name]; ok {
@@ -323,11 +341,13 @@ func (a *Allocator) positionIDs(tokens []int64, start int64, force bool) error {
_ = a.values[name].Destroy()
}
- data := make([]int64, len(tokens))
- shape := []int64{1, int64(len(tokens))}
+ data := make([]int64, int64(a.batchSize)*a.sequenceLength)
+ shape := []int64{int64(a.batchSize), a.sequenceLength}
- for i := range len(tokens) {
- data[i] = start + int64(i)
+ for b := range int64(a.batchSize) {
+ for s := range a.sequenceLength {
+ data[b*a.sequenceLength+s] = start + s
+ }
}
if t, err := ort.NewTensor[int64](shape, data); err != nil {
@@ -339,7 +359,7 @@ func (a *Allocator) positionIDs(tokens []int64, start int64, force bool) error {
return nil
}
-func (a *Allocator) attentionMask(tokens []int64, start int64, force bool) error {
+func (a *Allocator) attentionMask(start int64, force bool) error {
const name = "attention_mask"
if _, ok := a.values[name]; ok {
@@ -350,8 +370,8 @@ func (a *Allocator) attentionMask(tokens []int64, start int64, force bool) error
_ = a.values[name].Destroy()
}
- data := make([]int64, start+int64(len(tokens)))
- shape := []int64{1, start + int64(len(tokens))}
+ data := make([]int64, int64(a.batchSize)*(start+a.sequenceLength))
+ shape := []int64{int64(a.batchSize), start + a.sequenceLength}
for i := range data {
data[i] = 1
@@ -385,7 +405,7 @@ func (a *Allocator) pastKeyValues(i int64, suffix string, force bool) error {
_ = a.values[name].Destroy()
}
- shape := []int64{1, int64(a.config.nHeads), 0, int64(a.config.headDim)}
+ shape := []int64{int64(a.batchSize), int64(a.config.nHeads), 0, int64(a.config.headDim)}
if t, err := ort.NewEmptyTensor[float32](shape); err != nil {
return err
@@ -396,7 +416,7 @@ func (a *Allocator) pastKeyValues(i int64, suffix string, force bool) error {
return nil
}
-func (a *Allocator) logits(tokens []int64, force bool) error {
+func (a *Allocator) logits(force bool) error {
const name = "logits"
if _, ok := a.values[name]; ok {
@@ -407,7 +427,7 @@ func (a *Allocator) logits(tokens []int64, force bool) error {
_ = a.values[name].Destroy()
}
- shape := []int64{1, int64(len(tokens)), int64(a.config.vocabSize)}
+ shape := []int64{int64(a.batchSize), a.sequenceLength, int64(a.config.vocabSize)}
if t, err := ort.NewEmptyTensor[float32](shape); err != nil {
return err
@@ -418,7 +438,7 @@ func (a *Allocator) logits(tokens []int64, force bool) error {
return nil
}
-func (a *Allocator) logProbs(tokens []int64, force bool) error {
+func (a *Allocator) logProbs(force bool) error {
const name = "token_logprobs"
if _, ok := a.values[name]; ok {
@@ -429,7 +449,7 @@ func (a *Allocator) logProbs(tokens []int64, force bool) error {
_ = a.values[name].Destroy()
}
- shape := []int64{1, int64(len(tokens)) - 1}
+ shape := []int64{int64(a.batchSize), a.sequenceLength - 1}
if t, err := ort.NewEmptyTensor[float32](shape); err != nil {
return err
@@ -440,7 +460,7 @@ func (a *Allocator) logProbs(tokens []int64, force bool) error {
return nil
}
-func (a *Allocator) presentKeyValues(tokens []int64, start, i int64, suffix string, force bool) error {
+func (a *Allocator) presentKeyValues(start, i int64, suffix string, force bool) error {
if int(i) > a.config.nLayers {
panic("invalid layer index")
}
@@ -459,7 +479,7 @@ func (a *Allocator) presentKeyValues(tokens []int64, start, i int64, suffix stri
_ = a.values[name].Destroy()
}
- shape := []int64{1, int64(a.config.nHeads), start + int64(len(tokens)), int64(a.config.headDim)}
+ shape := []int64{int64(a.batchSize), int64(a.config.nHeads), start + a.sequenceLength, int64(a.config.headDim)}
if t, err := ort.NewEmptyTensor[float32](shape); err != nil {
return err
@@ -470,7 +490,7 @@ func (a *Allocator) presentKeyValues(tokens []int64, start, i int64, suffix stri
return nil
}
-func (a *Allocator) rotateCache(tokens []int64, i int64, suffix string) error {
+func (a *Allocator) rotateCache(i int64, suffix string) error {
past := fmt.Sprintf("past_key_values.%d.%s", i, suffix)
present := fmt.Sprintf("present.%d.%s", i, suffix)
@@ -482,7 +502,7 @@ func (a *Allocator) rotateCache(tokens []int64, i int64, suffix string) error {
delete(a.values, present)
- if err := a.presentKeyValues(tokens, a.step, i, suffix, false); err != nil {
+ if err := a.presentKeyValues(a.step, i, suffix, false); err != nil {
return err
}
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go
index d854519..472d8b6 100644
--- a/gpt2/cmd/gpt2/main.go
+++ b/gpt2/cmd/gpt2/main.go
@@ -55,7 +55,7 @@ func score(prompt []int64) {
logProbs := make([]float32, 0, 2)
- if err := m.Score(prompt, &logProbs); err != nil {
+ if err := m.Score(prompt, 1, &logProbs); err != nil {
log.Fatal(err)
}
diff --git a/gpt2/model.go b/gpt2/model.go
index 2b23cdf..a5f4d27 100644
--- a/gpt2/model.go
+++ b/gpt2/model.go
@@ -60,7 +60,7 @@ func IntraOpNumThreads() int {
}
func (m *Model) Init() error {
- m.allocator = NewAllocator(m.config, m.withCache, m.withLogits, m.withLogProbs)
+ m.allocator = NewAllocator(m.config, 1, m.withCache, m.withLogits, m.withLogProbs)
var options *ort.SessionOptions
@@ -153,11 +153,13 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in
return r, nil
}
-func (m *Model) Score(tokens []int64, logProbs *[]float32) error {
+func (m *Model) Score(tokens []int64, batchSize int, logProbs *[]float32) error {
if !m.withLogProbs {
panic("score requires token_logprobs output")
}
+ m.allocator.SetBatchSize(batchSize)
+
if err := m.allocator.Init(tokens); err != nil {
return err
}
diff --git a/llm/causal.go b/llm/causal.go
index b683985..69c3127 100644
--- a/llm/causal.go
+++ b/llm/causal.go
@@ -2,5 +2,5 @@ package llm
type Causal interface {
Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error)
- Score(tokens []int64, logProbs *[]float32) error
+ Score(tokens []int64, batchSize int, logProbs *[]float32) error
}
diff --git a/llm/run.go b/llm/run.go
index 9cdf8c3..21f60eb 100644
--- a/llm/run.go
+++ b/llm/run.go
@@ -113,30 +113,34 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int
}
func (e *Evaluator[R]) execute(j *batch, device int) {
- if j.Size() != 1 {
- panic("unimplemented")
- }
+ b := j.Size()
+
+ s := len(j.jobs[0].Tokens)
- job := j.jobs[0]
+ tokens := make([]int64, b*s)
- if job.Seen < 1 {
- panic("empty context")
+ for i, job := range j.jobs {
+ copy(tokens[i*s:], job.Tokens)
}
m := e.models[device]
- logProbs := make([]float32, 0, len(job.Tokens)-1)
+ logProbs := make([]float32, 0, b*(s-1))
- if err := m.Score(job.Tokens, &logProbs); err != nil {
+ if err := m.Score(tokens, b, &logProbs); err != nil {
panic(err) // TODO handle
}
- l := logProbs[job.Seen-1:]
- t := toInt(job.Tokens[job.Seen:])
+ for i, job := range j.jobs {
+ if job.Seen < 1 {
+ panic("empty context")
+ }
- r := e.callback(job, l, t)
+ l := logProbs[i*(s-1)+job.Seen-1 : (i+1)*(s-1)]
+ t := toInt(job.Tokens[job.Seen:])
- e.results <- r
+ r := e.callback(job, l, t)
- return
+ e.results <- r
+ }
}