summaryrefslogtreecommitdiff
path: root/gpt2/model.go
diff options
context:
space:
mode:
Diffstat (limited to 'gpt2/model.go')
-rw-r--r--gpt2/model.go12
1 files changed, 6 insertions, 6 deletions
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