diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-08 22:14:43 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-08 23:33:38 +0200 |
| commit | a006507a24b93a6af7e3cff8833fb96b48b159d7 (patch) | |
| tree | 8b5585aea514fbaab4f1de6ce2db5da5f2533dbf | |
| parent | 7e48ca5b9143c2b5f1a6eb81f35e90e67e628404 (diff) | |
Refactor model config
| -rw-r--r-- | gpt2/allocator.go | 32 | ||||
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 12 | ||||
| -rw-r--r-- | gpt2/config.go | 26 | ||||
| -rw-r--r-- | gpt2/model.go | 12 |
4 files changed, 42 insertions, 40 deletions
diff --git a/gpt2/allocator.go b/gpt2/allocator.go index 6c34192..3c85410 100644 --- a/gpt2/allocator.go +++ b/gpt2/allocator.go @@ -19,9 +19,9 @@ type Allocator struct { withLogProbs bool } -func NewAllocator(config Config, batchSize int, withCache bool, withLogits bool, withLogProbs bool) *Allocator { +func NewAllocator(cfg Config, batchSize int, withCache bool, withLogits bool, withLogProbs bool) *Allocator { return &Allocator{ - config: config, + config: cfg, batchSize: batchSize, values: make(map[string]ort.Value), withCache: withCache, @@ -34,7 +34,7 @@ func (a *Allocator) InputNames() []string { capacity := 3 if a.withCache { - capacity += 2 * a.config.nLayers + capacity += 2 * a.config.NumLayers } names := make([]string, 0, capacity) @@ -42,7 +42,7 @@ func (a *Allocator) InputNames() []string { names = append(names, "input_ids", "position_ids", "attention_mask") if a.withCache { - for i := range a.config.nLayers { + for i := range a.config.NumLayers { names = append(names, fmt.Sprintf("past_key_values.%d.key", i), fmt.Sprintf("past_key_values.%d.value", i)) } } @@ -62,7 +62,7 @@ func (a *Allocator) OutputNames() []string { } if a.withCache { - capacity += 2 * a.config.nLayers + capacity += 2 * a.config.NumLayers } names := make([]string, 0, capacity) @@ -76,7 +76,7 @@ func (a *Allocator) OutputNames() []string { } if a.withCache { - for i := range a.config.nLayers { + for i := range a.config.NumLayers { names = append(names, fmt.Sprintf("present.%d.key", i), fmt.Sprintf("present.%d.value", i)) } } @@ -120,7 +120,7 @@ func (a *Allocator) initInputs(tokens []int64) error { capacity := 3 if a.withCache { - capacity += 2 * a.config.nLayers + capacity += 2 * a.config.NumLayers } names := make([]string, 0, capacity) @@ -149,7 +149,7 @@ func (a *Allocator) initInputs(tokens []int64) error { return nil } - for i := range int64(a.config.nLayers) { + for i := range int64(a.config.NumLayers) { if err := a.pastKeyValues(i, "key", false); err != nil { return err } @@ -180,7 +180,7 @@ func (a *Allocator) initOutputs(tokens []int64) error { } if a.withCache { - capacity += 2 * a.config.nLayers + capacity += 2 * a.config.NumLayers } names := make([]string, 0, capacity) @@ -207,7 +207,7 @@ func (a *Allocator) initOutputs(tokens []int64) error { return nil } - for i := range int64(a.config.nLayers) { + for i := range int64(a.config.NumLayers) { if err := a.presentKeyValues(0, i, "key", false); err != nil { return err } @@ -259,7 +259,7 @@ func (a *Allocator) Step(token int64) error { } } - for i := range int64(a.config.nLayers) { + for i := range int64(a.config.NumLayers) { for _, suffix := range []string{"key", "value"} { if err := a.rotateCache(i, suffix); err != nil { return err @@ -387,7 +387,7 @@ func (a *Allocator) attentionMask(start int64, force bool) error { } func (a *Allocator) pastKeyValues(i int64, suffix string, force bool) error { - if int(i) > a.config.nLayers { + if int(i) > a.config.NumLayers { panic("invalid layer index") } @@ -405,7 +405,7 @@ func (a *Allocator) pastKeyValues(i int64, suffix string, force bool) error { _ = a.values[name].Destroy() } - shape := []int64{int64(a.batchSize), int64(a.config.nHeads), 0, int64(a.config.headDim)} + shape := []int64{int64(a.batchSize), int64(a.config.NumHeads), 0, int64(a.config.HeadDim)} if t, err := ort.NewEmptyTensor[float32](shape); err != nil { return err @@ -427,7 +427,7 @@ func (a *Allocator) logits(force bool) error { _ = a.values[name].Destroy() } - shape := []int64{int64(a.batchSize), a.sequenceLength, 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 @@ -461,7 +461,7 @@ func (a *Allocator) logProbs(force bool) error { } func (a *Allocator) presentKeyValues(start, i int64, suffix string, force bool) error { - if int(i) > a.config.nLayers { + if int(i) > a.config.NumLayers { panic("invalid layer index") } @@ -479,7 +479,7 @@ func (a *Allocator) presentKeyValues(start, i int64, suffix string, force bool) _ = a.values[name].Destroy() } - shape := []int64{int64(a.batchSize), int64(a.config.nHeads), start + a.sequenceLength, int64(a.config.headDim)} + shape := []int64{int64(a.batchSize), int64(a.config.NumHeads), start + a.sequenceLength, int64(a.config.HeadDim)} if t, err := ort.NewEmptyTensor[float32](shape); err != nil { return err diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index bad9639..be8d011 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -27,7 +27,11 @@ func main() { } func generate(prompt []int64) { - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", "0", gpt2.DefaultConfig().WithVocabSize(8193), true, true, false) + cfg := gpt2.DefaultConfig() + + cfg.VocabSize = 8193 + + m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", "0", cfg, true, true, false) if err := m.Init(); err != nil { log.Fatal(err) @@ -47,7 +51,11 @@ func generate(prompt []int64) { } func score(prompt []int64) { - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", "0", gpt2.DefaultConfig().WithVocabSize(8193), false, false, true) + cfg := gpt2.DefaultConfig() + + cfg.VocabSize = 8193 + + m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", "0", cfg, false, false, true) if err := m.Init(); err != nil { log.Fatal(err) diff --git a/gpt2/config.go b/gpt2/config.go index be7db30..b77578d 100644 --- a/gpt2/config.go +++ b/gpt2/config.go @@ -1,25 +1,19 @@ package gpt2 type Config struct { - vocabSize int - nLayers int - nHeads int - headDim int - nPositions int + VocabSize int + NumLayers int + NumHeads int + HeadDim int + NumPositions int } func DefaultConfig() Config { return Config{ - vocabSize: 50257, - nLayers: 12, - nHeads: 12, - headDim: 64, - nPositions: 1024, + VocabSize: 50257, + NumLayers: 12, + NumHeads: 12, + HeadDim: 64, + NumPositions: 1024, } } - -func (c Config) WithVocabSize(vocabSize int) Config { - c.vocabSize = vocabSize - - return c -} diff --git a/gpt2/model.go b/gpt2/model.go index a5f4d27..0708c7e 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -22,11 +22,11 @@ type Model struct { allocator *Allocator } -func NewModel(name string, deviceID string, config Config, withCache bool, withLogits bool, withLogProbs bool) *Model { +func NewModel(name string, deviceID string, cfg Config, withCache bool, withLogits bool, withLogProbs bool) *Model { return &Model{ name: name, deviceID: deviceID, - config: config, + config: cfg, withCache: withCache, withLogits: withLogits, withLogProbs: withLogProbs, @@ -110,7 +110,7 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in return nil, errors.New("empty prompt") } - if int64(len(prompt))+steps > int64(m.config.nPositions) { + if int64(len(prompt))+steps > int64(m.config.NumPositions) { return nil, errors.New("sequence length exceeds context limit") } @@ -179,13 +179,13 @@ func (m *Model) Score(tokens []int64, batchSize int, logProbs *[]float32) error func (m *Model) logits(output ort.Value) [][]float32 { d := output.(*ort.Tensor[float32]).GetData() - n := len(d) / m.config.vocabSize + n := len(d) / m.config.VocabSize l := make([][]float32, n) for i := range n { - s := i * m.config.vocabSize + s := i * m.config.VocabSize - l[i] = d[s : s+m.config.vocabSize : s+m.config.vocabSize] + l[i] = d[s : s+m.config.VocabSize : s+m.config.VocabSize] } return l |
