From 1e02a5ec7f7794bb542e99ceb804ed88048b6b54 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Fri, 24 Apr 2026 11:07:01 +0200 Subject: Access artifacts via shelf --- gpt2/cmd/gpt2/main.go | 5 +++-- gpt2/model_test.go | 6 ++++-- 2 files changed, 7 insertions(+), 4 deletions(-) (limited to 'gpt2') diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index 06ec535..b9b75a2 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -7,6 +7,7 @@ import ( "slices" "go.jknobloc.com/x/gpt2" + "go.jknobloc.com/x/shelf" ) func main() { @@ -37,7 +38,7 @@ func generate(prompt []int64) { WithLogProbs: false, } - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", cfg, opts) + m := gpt2.NewModel(shelf.Abs("models/mbpe/gpt2_8192_m000_babylm_v2/model_cache.onnx"), cfg, opts) if err := m.Init(); err != nil { log.Fatal(err) @@ -69,7 +70,7 @@ func score(prompt []int64) { WithLogProbs: true, } - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", cfg, opts) + m := gpt2.NewModel(shelf.Abs("models/mbpe/gpt2_8192_m000_babylm_v2/model_eval.onnx"), cfg, opts) if err := m.Init(); err != nil { log.Fatal(err) diff --git a/gpt2/model_test.go b/gpt2/model_test.go index 051a41a..cd38241 100644 --- a/gpt2/model_test.go +++ b/gpt2/model_test.go @@ -7,6 +7,8 @@ import ( "os" "slices" "testing" + + "go.jknobloc.com/x/shelf" ) func TestMain(m *testing.M) { @@ -46,7 +48,7 @@ func model() *Model { WithLogProbs: false, } - m := NewModel("models/base/model.onnx", DefaultConfig(), opts) + m := NewModel(shelf.Abs("models/gpt2/model_cache.onnx"), DefaultConfig(), opts) if err := m.Init(); err != nil { log.Fatal(err) @@ -58,7 +60,7 @@ func model() *Model { func fromGold() []float32 { shape := []int{1, 9, 50257} - s, err := f32("test/logits.f32", shape[0]*shape[1]*shape[2]) + s, err := f32(shelf.Abs("test/gpt2/logits.f32"), shape[0]*shape[1]*shape[2]) if err != nil { log.Fatal(err) -- cgit v1.3.1