diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-08 21:58:01 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-08 21:58:01 +0200 |
| commit | b5de54cfc298c696fcd090bb7d12e8a6b9a8bfe1 (patch) | |
| tree | 3504cb797e31d3171796a1916c15c8ec7053420b | |
| parent | 59b0d9b9e9e648b3626cc7a17cfc769e6bfa6a9c (diff) | |
Rename default config constructor
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 4 | ||||
| -rw-r--r-- | gpt2/config.go | 2 | ||||
| -rw-r--r-- | gpt2/model_test.go | 2 | ||||
| -rw-r--r-- | llm/cmd/eval/main.go | 2 |
4 files changed, 5 insertions, 5 deletions
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index 472d8b6..bad9639 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -27,7 +27,7 @@ func main() { } func generate(prompt []int64) { - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", "0", gpt2.NewDefaultConfig().WithVocabSize(8193), true, true, false) + m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", "0", gpt2.DefaultConfig().WithVocabSize(8193), true, true, false) if err := m.Init(); err != nil { log.Fatal(err) @@ -47,7 +47,7 @@ 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.NewDefaultConfig().WithVocabSize(8193), false, false, true) + m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", "0", gpt2.DefaultConfig().WithVocabSize(8193), false, false, true) if err := m.Init(); err != nil { log.Fatal(err) diff --git a/gpt2/config.go b/gpt2/config.go index b7cede5..be7db30 100644 --- a/gpt2/config.go +++ b/gpt2/config.go @@ -8,7 +8,7 @@ type Config struct { nPositions int } -func NewDefaultConfig() Config { +func DefaultConfig() Config { return Config{ vocabSize: 50257, nLayers: 12, diff --git a/gpt2/model_test.go b/gpt2/model_test.go index 0072254..a4de628 100644 --- a/gpt2/model_test.go +++ b/gpt2/model_test.go @@ -40,7 +40,7 @@ func fromModel() []float32 { } func model() *Model { - m := NewModel("models/base/model.onnx", "0", NewDefaultConfig(), true, true, false) // TODO check if CUDA is available + m := NewModel("models/base/model.onnx", "0", DefaultConfig(), true, true, false) // 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 b654e49..10c7db2 100644 --- a/llm/cmd/eval/main.go +++ b/llm/cmd/eval/main.go @@ -34,7 +34,7 @@ func data() *dataset.ParquetReader { } func model() *gpt2.Model { - m := gpt2.NewModel("gpt2/models/base/model_eval.onnx", "0", gpt2.NewDefaultConfig(), false, false, true) + m := gpt2.NewModel("gpt2/models/base/model_eval.onnx", "0", gpt2.DefaultConfig(), false, false, true) if err := m.Init(); err != nil { log.Fatal(err) |
