summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-08 22:14:43 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-08 23:33:38 +0200
commita006507a24b93a6af7e3cff8833fb96b48b159d7 (patch)
tree8b5585aea514fbaab4f1de6ce2db5da5f2533dbf
parent7e48ca5b9143c2b5f1a6eb81f35e90e67e628404 (diff)
Refactor model config
-rw-r--r--gpt2/allocator.go32
-rw-r--r--gpt2/cmd/gpt2/main.go12
-rw-r--r--gpt2/config.go26
-rw-r--r--gpt2/model.go12
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