From b9faf92148c7e6432ff0411783f444b7779fa36a Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Mon, 6 Apr 2026 22:59:50 +0200 Subject: Handle environment initialization on package level --- gpt2/cmd/gpt2/main.go | 16 ++++++++++------ 1 file changed, 10 insertions(+), 6 deletions(-) (limited to 'gpt2/cmd') 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 { -- cgit v1.3.1