diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-04 21:43:38 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-04 21:43:38 +0200 |
| commit | a96c1bf68bf5c477e1d3f105c18346c80d80b47a (patch) | |
| tree | 88e2d5cfb53d870c9febb5b26a1351aea4bb7028 | |
| parent | 758e27cf2b3714588af7db1f5da057b15995ba26 (diff) | |
Add config struct
| -rw-r--r-- | gpt2/allocator.go | 32 | ||||
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 2 | ||||
| -rw-r--r-- | gpt2/config.go | 25 | ||||
| -rw-r--r-- | gpt2/model.go | 22 | ||||
| -rw-r--r-- | gpt2/model_test.go | 2 | ||||
| -rw-r--r-- | llm/cmd/eval/main.go | 2 |
6 files changed, 53 insertions, 32 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 diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index a832ccd..cad3333 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -10,7 +10,7 @@ import ( func main() { prompt := []int64{464, 2068, 7586, 21831} - m := gpt2.NewModel("models/base/model.onnx", "") + m := gpt2.NewModel("gpt2/models/base/model.onnx", "0", gpt2.NewDefaultConfig()) if err := m.Init(); err != nil { log.Fatal(err) diff --git a/gpt2/config.go b/gpt2/config.go new file mode 100644 index 0000000..b7cede5 --- /dev/null +++ b/gpt2/config.go @@ -0,0 +1,25 @@ +package gpt2 + +type Config struct { + vocabSize int + nLayers int + nHeads int + headDim int + nPositions int +} + +func NewDefaultConfig() Config { + return Config{ + vocabSize: 50257, + nLayers: 12, + nHeads: 12, + headDim: 64, + nPositions: 1024, + } +} + +func (c Config) WithVocabSize(vocabSize int) Config { + c.vocabSize = vocabSize + + return c +} diff --git a/gpt2/model.go b/gpt2/model.go index 51889be..6557a61 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -11,25 +11,19 @@ import ( ort "github.com/yalue/onnxruntime_go" ) -const ( - vocabSize = 50257 - nLayers = 12 - nHeads = 12 - headDim = 64 - nPositions = 1024 -) - type Model struct { name string deviceID string + config Config session *ort.DynamicAdvancedSession allocator *Allocator } -func NewModel(name, deviceID string) *Model { +func NewModel(name string, deviceID string, config Config) *Model { return &Model{ name: name, deviceID: deviceID, + config: config, } } @@ -66,7 +60,7 @@ func (m *Model) Init() error { return err } - m.allocator = NewAllocator(true) + m.allocator = NewAllocator(m.config, true) var options *ort.SessionOptions @@ -110,7 +104,7 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in return nil, errors.New("empty prompt") } - if int64(len(prompt))+steps > nPositions { + if int64(len(prompt))+steps > int64(m.config.nPositions) { return nil, errors.New("sequence length exceeds context limit") } @@ -159,13 +153,13 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in func (m *Model) logits(output ort.Value) [][]float32 { d := output.(*ort.Tensor[float32]).GetData() - n := len(d) / vocabSize + n := len(d) / m.config.vocabSize l := make([][]float32, n) for i := range n { - s := i * vocabSize + s := i * m.config.vocabSize - l[i] = d[s : s+vocabSize : s+vocabSize] + l[i] = d[s : s+m.config.vocabSize : s+m.config.vocabSize] } return l diff --git a/gpt2/model_test.go b/gpt2/model_test.go index dcd00e6..18abdb8 100644 --- a/gpt2/model_test.go +++ b/gpt2/model_test.go @@ -26,7 +26,7 @@ func fromModel() []float32 { } func model() *Model { - m := NewModel("models/base/model.onnx", "0") // TODO check if CUDA is available + m := NewModel("models/base/model.onnx", "0", NewDefaultConfig()) // TODO check if CUDA is available if err := m.Init(); err != nil { log.Fatal(err) diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go index bcc6fb4..50448be 100644 --- a/llm/cmd/eval/main.go +++ b/llm/cmd/eval/main.go @@ -26,7 +26,7 @@ func data() *dataset.ParquetReader { } func model() *gpt2.Model { - m := gpt2.NewModel("gpt2/models/base/model.onnx", "0") + m := gpt2.NewModel("gpt2/models/base/model.onnx", "0", gpt2.NewDefaultConfig()) if err := m.Init(); err != nil { log.Fatal(err) |
