diff options
Diffstat (limited to 'gpt2/model.go')
| -rw-r--r-- | gpt2/model.go | 52 |
1 files changed, 42 insertions, 10 deletions
diff --git a/gpt2/model.go b/gpt2/model.go index 6557a61..dd18c80 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -12,18 +12,22 @@ import ( ) type Model struct { - name string - deviceID string - config Config - session *ort.DynamicAdvancedSession - allocator *Allocator + name string + deviceID string + config Config + withCache bool + withLogProbs bool + session *ort.DynamicAdvancedSession + allocator *Allocator } -func NewModel(name string, deviceID string, config Config) *Model { +func NewModel(name string, deviceID string, config Config, withCache bool, withLogProbs bool) *Model { return &Model{ - name: name, - deviceID: deviceID, - config: config, + name: name, + deviceID: deviceID, + config: config, + withCache: withCache, + withLogProbs: withLogProbs, } } @@ -60,7 +64,7 @@ func (m *Model) Init() error { return err } - m.allocator = NewAllocator(m.config, true) + m.allocator = NewAllocator(m.config, m.withCache, m.withLogProbs) var options *ort.SessionOptions @@ -100,6 +104,10 @@ func (m *Model) Destroy() error { } func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error) { + if m.withLogProbs { + panic("generate called on eval model") + } + if len(prompt) == 0 { return nil, errors.New("empty prompt") } @@ -151,6 +159,30 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in return r, nil } +func (m *Model) Score(tokens []int64, logProbs *[]float32) error { + if !m.withLogProbs { + panic("score called on default model") + } + + if err := m.allocator.Init(tokens); err != nil { + return err + } + + if err := m.forward(m.allocator); err != nil { + return err + } + + if logProbs != nil { + _, outputs := m.allocator.Outputs() + + d := outputs[1].(*ort.Tensor[float32]).GetData() + + *logProbs = append(*logProbs, d...) + } + + return nil +} + func (m *Model) logits(output ort.Value) [][]float32 { d := output.(*ort.Tensor[float32]).GetData() n := len(d) / m.config.vocabSize |
