diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-10 22:45:40 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-10 22:45:40 +0200 |
| commit | c1598d443f5bad23595de916c32553b620dcbeba (patch) | |
| tree | 9d7f965e837b00bffffdcaeae108586c568263dc | |
| parent | 4386784c1006321c5cf0e1e005b19e0a7d7563ae (diff) | |
Destroy session on model teardown
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 8 | ||||
| -rw-r--r-- | gpt2/model.go | 24 | ||||
| -rw-r--r-- | llm/cmd/eval/logprobs.go | 4 | ||||
| -rw-r--r-- | llm/cmd/eval/ppl.go | 4 |
4 files changed, 27 insertions, 13 deletions
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index a1b3cf8..fc7dec8 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -53,7 +53,9 @@ func generate(prompt []int64) { fmt.Println(selectLogProbs(logits[:len(logits)-1], prompt[1:])) - m.Destroy() + if err := m.Destroy(); err != nil { + log.Fatal(err) + } } func score(prompt []int64) { @@ -81,7 +83,9 @@ func score(prompt []int64) { fmt.Println(logProbs) - m.Destroy() + if err := m.Destroy(); err != nil { + log.Fatal(err) + } } func selectLogProbs(logits [][]float32, tokens []int64) []float32 { diff --git a/gpt2/model.go b/gpt2/model.go index 20dcb71..5fef956 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -12,20 +12,20 @@ import ( ) type Model struct { - name string - deviceID string - config Config - options Options - session *ort.DynamicAdvancedSession - allocator *Allocator + name string + deviceID string + config Config + options Options + session *ort.DynamicAdvancedSession + allocator *Allocator } func NewModel(name string, deviceID string, cfg Config, opts Options) *Model { return &Model{ - name: name, - deviceID: deviceID, - config: cfg, - options: opts, + name: name, + deviceID: deviceID, + config: cfg, + options: opts, } } @@ -89,8 +89,10 @@ func (m *Model) Init() error { return nil } -func (m *Model) Destroy() { +func (m *Model) Destroy() error { m.allocator.Destroy() + + return m.session.Destroy() } func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error) { diff --git a/llm/cmd/eval/logprobs.go b/llm/cmd/eval/logprobs.go index baac805..f39e351 100644 --- a/llm/cmd/eval/logprobs.go +++ b/llm/cmd/eval/logprobs.go @@ -65,6 +65,10 @@ func logprobs() { }); err != nil { log.Fatal(err) } + + if err := m.Destroy(); err != nil { + log.Fatal(err) + } } func prepare(name string) (*sql.Stmt, *sql.DB, error) { diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go index 330e4ab..2394d83 100644 --- a/llm/cmd/eval/ppl.go +++ b/llm/cmd/eval/ppl.go @@ -57,6 +57,10 @@ func perplexity() { ppl := math.Exp(avg) fmt.Println(ppl) + + if err := m.Destroy(); err != nil { + log.Fatal(err) + } } func joined() dataset.Reader { |
