From a006507a24b93a6af7e3cff8833fb96b48b159d7 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Wed, 8 Apr 2026 22:14:43 +0200 Subject: Refactor model config --- gpt2/model.go | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) (limited to 'gpt2/model.go') diff --git a/gpt2/model.go b/gpt2/model.go index a5f4d27..0708c7e 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -22,11 +22,11 @@ type Model struct { allocator *Allocator } -func NewModel(name string, deviceID string, config Config, withCache bool, withLogits bool, withLogProbs bool) *Model { +func NewModel(name string, deviceID string, cfg Config, withCache bool, withLogits bool, withLogProbs bool) *Model { return &Model{ name: name, deviceID: deviceID, - config: config, + config: cfg, withCache: withCache, withLogits: withLogits, withLogProbs: withLogProbs, @@ -110,7 +110,7 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in return nil, errors.New("empty prompt") } - if int64(len(prompt))+steps > int64(m.config.nPositions) { + if int64(len(prompt))+steps > int64(m.config.NumPositions) { return nil, errors.New("sequence length exceeds context limit") } @@ -179,13 +179,13 @@ func (m *Model) Score(tokens []int64, batchSize int, logProbs *[]float32) error func (m *Model) logits(output ort.Value) [][]float32 { d := output.(*ort.Tensor[float32]).GetData() - n := len(d) / m.config.vocabSize + n := len(d) / m.config.VocabSize l := make([][]float32, n) for i := range n { - s := i * m.config.vocabSize + s := i * m.config.VocabSize - l[i] = d[s : s+m.config.vocabSize : s+m.config.vocabSize] + l[i] = d[s : s+m.config.VocabSize : s+m.config.VocabSize] } return l -- cgit v1.3.1