diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-06 22:59:50 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-06 23:28:04 +0200 |
| commit | b9faf92148c7e6432ff0411783f444b7779fa36a (patch) | |
| tree | 96904cd4cf053b3d272b5f2b73c1141483e33863 | |
| parent | 58bdf1ac7155f4c5076cf4a077a5f39dfdedbac7 (diff) | |
Handle environment initialization on package level
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 16 | ||||
| -rw-r--r-- | gpt2/model.go | 10 | ||||
| -rw-r--r-- | gpt2/model_test.go | 14 | ||||
| -rw-r--r-- | gpt2/onnx.go | 13 | ||||
| -rw-r--r-- | llm/cmd/eval/main.go | 8 |
5 files changed, 46 insertions, 15 deletions
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index 663bdd8..d854519 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -10,12 +10,20 @@ import ( ) func main() { + if err := gpt2.InitializeEnvironment(); err != nil { + log.Fatal(err) + } + prompt := []int64{464, 2068, 7586} generate(prompt) // [-13.483142 -11.277906] score(prompt) // [-13.48314 -11.277912] _ = prompt + + if err := gpt2.DestroyEnvironment(); err != nil { + log.Fatal(err) + } } func generate(prompt []int64) { @@ -35,9 +43,7 @@ func generate(prompt []int64) { fmt.Println(selectLogProbs(logits[:len(logits)-1], prompt[1:])) - if err := m.Destroy(); err != nil { - log.Fatal(err) - } + m.Destroy() } func score(prompt []int64) { @@ -55,9 +61,7 @@ func score(prompt []int64) { fmt.Println(logProbs) - if err := m.Destroy(); err != nil { - log.Fatal(err) - } + m.Destroy() } func selectLogProbs(logits [][]float32, tokens []int64) []float32 { diff --git a/gpt2/model.go b/gpt2/model.go index f492c56..d0d1c2f 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -60,12 +60,6 @@ func IntraOpNumThreads() int { } func (m *Model) Init() error { - ort.SetSharedLibraryPath(SharedLibraryPath()) - - if err := ort.InitializeEnvironment(); err != nil { - return err - } - m.allocator = NewAllocator(m.config, m.withCache, m.withLogits, m.withLogProbs) var options *ort.SessionOptions @@ -99,10 +93,8 @@ func (m *Model) Init() error { return nil } -func (m *Model) Destroy() error { +func (m *Model) Destroy() { m.allocator.Destroy() - - return ort.DestroyEnvironment() } func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error) { diff --git a/gpt2/model_test.go b/gpt2/model_test.go index d197ccc..0072254 100644 --- a/gpt2/model_test.go +++ b/gpt2/model_test.go @@ -9,6 +9,20 @@ import ( "testing" ) +func TestMain(m *testing.M) { + if err := InitializeEnvironment(); err != nil { + log.Fatal(err) + } + + exit := m.Run() + + if err := DestroyEnvironment(); err != nil { + log.Fatal(err) + } + + os.Exit(exit) +} + func fromModel() []float32 { prompt := []int64{464, 2068, 7586, 21831, 18045, 625, 262, 16931, 3290} diff --git a/gpt2/onnx.go b/gpt2/onnx.go new file mode 100644 index 0000000..0da2ae2 --- /dev/null +++ b/gpt2/onnx.go @@ -0,0 +1,13 @@ +package gpt2 + +import ort "github.com/yalue/onnxruntime_go" + +func InitializeEnvironment() error { + ort.SetSharedLibraryPath(SharedLibraryPath()) + + return ort.InitializeEnvironment() +} + +func DestroyEnvironment() error { + return ort.DestroyEnvironment() +} diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go index 78a865e..43752d9 100644 --- a/llm/cmd/eval/main.go +++ b/llm/cmd/eval/main.go @@ -9,8 +9,16 @@ import ( ) func main() { + if err := gpt2.InitializeEnvironment(); err != nil { + log.Fatal(err) + } + perplexity() // logprobs() + + if err := gpt2.DestroyEnvironment(); err != nil { + log.Fatal(err) + } } func data() *dataset.ParquetReader { |
