summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-08 21:58:01 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-08 21:58:01 +0200
commitb5de54cfc298c696fcd090bb7d12e8a6b9a8bfe1 (patch)
tree3504cb797e31d3171796a1916c15c8ec7053420b
parent59b0d9b9e9e648b3626cc7a17cfc769e6bfa6a9c (diff)
Rename default config constructor
-rw-r--r--gpt2/cmd/gpt2/main.go4
-rw-r--r--gpt2/config.go2
-rw-r--r--gpt2/model_test.go2
-rw-r--r--llm/cmd/eval/main.go2
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)