summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-04 21:43:38 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-04 21:43:38 +0200
commita96c1bf68bf5c477e1d3f105c18346c80d80b47a (patch)
tree88e2d5cfb53d870c9febb5b26a1351aea4bb7028
parent758e27cf2b3714588af7db1f5da057b15995ba26 (diff)
Add config struct
-rw-r--r--gpt2/allocator.go32
-rw-r--r--gpt2/cmd/gpt2/main.go2
-rw-r--r--gpt2/config.go25
-rw-r--r--gpt2/model.go22
-rw-r--r--gpt2/model_test.go2
-rw-r--r--llm/cmd/eval/main.go2
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)