diff options
Diffstat (limited to 'gpt2/allocator.go')
| -rw-r--r-- | gpt2/allocator.go | 32 |
1 files changed, 17 insertions, 15 deletions
diff --git a/gpt2/allocator.go b/gpt2/allocator.go index e65c752..d0ec5ec 100644 --- a/gpt2/allocator.go +++ b/gpt2/allocator.go @@ -7,6 +7,7 @@ import ( ) type Allocator struct { + config Config step int64 inputNames []string outputNames []string @@ -14,20 +15,21 @@ type Allocator struct { withCache bool } -func NewAllocator(withCache bool) *Allocator { +func NewAllocator(config Config, withCache bool) *Allocator { return &Allocator{ + config: config, values: make(map[string]ort.Value), withCache: withCache, } } func (a *Allocator) InputNames() []string { - names := make([]string, 0, 3+2*nLayers) + names := make([]string, 0, 3+2*a.config.nLayers) names = append(names, "input_ids", "position_ids", "attention_mask") if a.withCache { - for i := range nLayers { + for i := range a.config.nLayers { names = append(names, fmt.Sprintf("past_key_values.%d.key", i), fmt.Sprintf("past_key_values.%d.value", i)) } } @@ -36,12 +38,12 @@ func (a *Allocator) InputNames() []string { } func (a *Allocator) OutputNames() []string { - names := make([]string, 0, 1+2*nLayers) + names := make([]string, 0, 1+2*a.config.nLayers) names = append(names, "logits") if a.withCache { - for i := range nLayers { + for i := range a.config.nLayers { names = append(names, fmt.Sprintf("present.%d.key", i), fmt.Sprintf("present.%d.value", i)) } } @@ -76,7 +78,7 @@ func (a *Allocator) initInputs(tokens []int64) error { capacity := 3 if a.withCache { - capacity += 2 * nLayers + capacity += 2 * a.config.nLayers } names := make([]string, 0, capacity) @@ -105,7 +107,7 @@ func (a *Allocator) initInputs(tokens []int64) error { return nil } - for i := range int64(nLayers) { + for i := range int64(a.config.nLayers) { if err := a.pastKeyValues(i, "key", false); err != nil { return err } @@ -128,7 +130,7 @@ func (a *Allocator) initOutputs(tokens []int64) error { capacity := 1 if a.withCache { - capacity += 2 * nLayers + capacity += 2 * a.config.nLayers } names := make([]string, 0, capacity) @@ -145,7 +147,7 @@ func (a *Allocator) initOutputs(tokens []int64) error { return nil } - for i := range int64(nLayers) { + for i := range int64(a.config.nLayers) { if err := a.presentKeyValues(tokens, 0, i, "key", false); err != nil { return err } @@ -183,7 +185,7 @@ func (a *Allocator) Step(token int64) error { return err } - for i := range int64(nLayers) { + for i := range int64(a.config.nLayers) { for _, suffix := range []string{"key", "value"} { if err := a.rotateCache(tokens, i, suffix); err != nil { return err @@ -299,7 +301,7 @@ func (a *Allocator) attentionMask(tokens []int64, start int64, force bool) error } func (a *Allocator) pastKeyValues(i int64, suffix string, force bool) error { - if i > nLayers { + if int(i) > a.config.nLayers { panic("invalid layer index") } @@ -317,7 +319,7 @@ func (a *Allocator) pastKeyValues(i int64, suffix string, force bool) error { _ = a.values[name].Destroy() } - shape := []int64{1, int64(nHeads), 0, int64(headDim)} + shape := []int64{1, int64(a.config.nHeads), 0, int64(a.config.headDim)} if t, err := ort.NewEmptyTensor[float32](shape); err != nil { return err @@ -339,7 +341,7 @@ func (a *Allocator) logits(tokens []int64, force bool) error { _ = a.values[name].Destroy() } - shape := []int64{1, int64(len(tokens)), int64(vocabSize)} + shape := []int64{1, int64(len(tokens)), int64(a.config.vocabSize)} if t, err := ort.NewEmptyTensor[float32](shape); err != nil { return err @@ -351,7 +353,7 @@ func (a *Allocator) logits(tokens []int64, force bool) error { } func (a *Allocator) presentKeyValues(tokens []int64, start, i int64, suffix string, force bool) error { - if i > nLayers { + if int(i) > a.config.nLayers { panic("invalid layer index") } @@ -369,7 +371,7 @@ func (a *Allocator) presentKeyValues(tokens []int64, start, i int64, suffix stri _ = a.values[name].Destroy() } - shape := []int64{1, int64(nHeads), start + int64(len(tokens)), int64(headDim)} + shape := []int64{1, int64(a.config.nHeads), start + int64(len(tokens)), int64(a.config.headDim)} if t, err := ort.NewEmptyTensor[float32](shape); err != nil { return err |
