summaryrefslogtreecommitdiff
path: root/llm
diff options
context:
space:
mode:
Diffstat (limited to 'llm')
-rw-r--r--llm/cmd/eval/perplexity.go8
-rw-r--r--llm/job.go10
-rw-r--r--llm/run.go21
-rw-r--r--llm/tokenbuffer.go37
-rw-r--r--llm/tokenbuffer_test.go78
5 files changed, 136 insertions, 18 deletions
diff --git a/llm/cmd/eval/perplexity.go b/llm/cmd/eval/perplexity.go
index dd4dae3..b84b659 100644
--- a/llm/cmd/eval/perplexity.go
+++ b/llm/cmd/eval/perplexity.go
@@ -43,8 +43,11 @@ func perplexity() {
n := 0
cfg := llm.TokenBufferConfig{
- Window: 1024,
- Stride: 512,
+ Window: 1024,
+ Stride: 512,
+ PadLeft: false, // probably not fine since we don't adjust model inputs for padding
+ PadRight: true, // fine since we remove padded log probs anyway
+ PadTokenID: 50256,
}
if err := e.RunAndCollect("Perplexity", d, cfg, func(r pplResult) error {
@@ -60,6 +63,7 @@ func perplexity() {
ppl := math.Exp(avg)
fmt.Println(ppl)
+ fmt.Printf("%d tokens\n", n)
if err := m.Destroy(); err != nil {
log.Fatal(err)
diff --git a/llm/job.go b/llm/job.go
index d922190..974b905 100644
--- a/llm/job.go
+++ b/llm/job.go
@@ -1,10 +1,12 @@
package llm
type Job struct {
- Document int
- Position int
- Tokens []int64
- Seen int
+ Document int
+ Position int
+ Tokens []int64
+ Seen int
+ PaddingLeft int
+ PaddingRight int
}
type batch struct {
diff --git a/llm/run.go b/llm/run.go
index 3be1cb5..39934dd 100644
--- a/llm/run.go
+++ b/llm/run.go
@@ -49,7 +49,7 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, cfg TokenBufferCon
tb := NewTokenBuffer(e.tokenizer, cfg)
- tb.SetIncludeTail(false)
+ tb.SetIncludeTail(cfg.PadLeft || cfg.PadRight)
b := newBatch(e.batchSize)
@@ -72,10 +72,12 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, cfg TokenBufferCon
}
b.AddJob(Job{
- Document: doc,
- Position: pos,
- Tokens: tokens,
- Seen: seen,
+ Document: doc,
+ Position: pos,
+ Tokens: tokens,
+ Seen: seen,
+ PaddingLeft: w.PaddingLeft,
+ PaddingRight: w.PaddingRight,
})
pos++
@@ -132,8 +134,13 @@ func (e *Evaluator[R]) execute(j *batch, device int) {
panic("empty context")
}
- l := logProbs[i*(s-1)+job.Seen-1 : (i+1)*(s-1)]
- t := toInt(job.Tokens[job.Seen:])
+ l := logProbs[i*(s-1)+job.PaddingLeft+job.Seen-1 : (i+1)*(s-1)]
+ t := toInt(job.Tokens[job.PaddingLeft+job.Seen:])
+
+ if job.PaddingRight > 0 {
+ l = l[:len(l)-job.PaddingRight]
+ t = t[:len(t)-job.PaddingRight]
+ }
r := e.callback(job, l, t)
diff --git a/llm/tokenbuffer.go b/llm/tokenbuffer.go
index e2d8a91..7bd427c 100644
--- a/llm/tokenbuffer.go
+++ b/llm/tokenbuffer.go
@@ -10,17 +10,23 @@ type TokenBuffer struct {
document int
position int
includeTail bool
+ config TokenBufferConfig
}
type TokenBufferConfig struct {
- Window int
- Stride int
+ Window int
+ Stride int
+ PadLeft bool
+ PadRight bool
+ PadTokenID int64
}
type TokenWindow struct {
- Document int
- Tokens []int64
- Seen int
+ Document int
+ Tokens []int64
+ Seen int
+ PaddingLeft int
+ PaddingRight int
}
func NewTokenBuffer(tokenizer Tokenizer, cfg TokenBufferConfig) *TokenBuffer {
@@ -28,6 +34,10 @@ func NewTokenBuffer(tokenizer Tokenizer, cfg TokenBufferConfig) *TokenBuffer {
panic("stride exceeds window")
}
+ if cfg.PadLeft && cfg.PadRight {
+ panic("either pad left or right")
+ }
+
return &TokenBuffer{
tokenizer: tokenizer,
window: cfg.Window,
@@ -36,6 +46,7 @@ func NewTokenBuffer(tokenizer Tokenizer, cfg TokenBufferConfig) *TokenBuffer {
document: -1,
position: 0,
includeTail: true,
+ config: cfg,
}
}
@@ -112,6 +123,22 @@ func (tb *TokenBuffer) Tail() (TokenWindow, bool) {
copy(w.Tokens, tb.buffer)
+ if tb.config.PadLeft || tb.config.PadRight {
+ padding := make([]int64, tb.window-len(w.Tokens))
+
+ for i := range padding {
+ padding[i] = tb.config.PadTokenID
+ }
+
+ if tb.config.PadLeft {
+ w.Tokens = append(padding, w.Tokens...)
+ w.PaddingLeft = len(padding)
+ } else {
+ w.Tokens = append(w.Tokens, padding...)
+ w.PaddingRight = len(padding)
+ }
+ }
+
tb.buffer = tb.buffer[:0]
return w, true
diff --git a/llm/tokenbuffer_test.go b/llm/tokenbuffer_test.go
index c07c371..0e80153 100644
--- a/llm/tokenbuffer_test.go
+++ b/llm/tokenbuffer_test.go
@@ -220,3 +220,81 @@ func TestTokenBuffer_TailPosition(t *testing.T) {
t.Errorf("expected tail position=6 but got %d", tailPosition)
}
}
+
+func TestTokenBuffer_Pad(t *testing.T) {
+ type gold struct {
+ config TokenBufferConfig
+ text string
+ expected [][]int64
+ seen int
+ }
+
+ tests := []gold{
+ {
+ config: TokenBufferConfig{
+ Window: 4,
+ Stride: 2,
+ PadLeft: true,
+ PadRight: false,
+ PadTokenID: 0,
+ },
+ text: "abcdefg",
+ expected: [][]int64{{97, 98, 99, 100}, {99, 100, 101, 102}, {0, 101, 102, 103}},
+ seen: 3,
+ },
+ {
+ config: TokenBufferConfig{
+ Window: 4,
+ Stride: 2,
+ PadLeft: false,
+ PadRight: true,
+ PadTokenID: 0,
+ },
+ text: "abcdefg",
+ expected: [][]int64{{97, 98, 99, 100}, {99, 100, 101, 102}, {101, 102, 103, 0}},
+ seen: 2,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(
+ fmt.Sprintf("window%d_stride%d_%v_%v", tt.config.Window, tt.config.Stride, tt.config.PadLeft, tt.config.PadRight),
+
+ func(t *testing.T) {
+ tb := NewTokenBuffer(byteTokenizer{}, tt.config)
+
+ tb.SetIncludeTail(true)
+
+ var got [][]int64
+
+ for w := range tb.Push(0, tt.text) {
+ got = append(got, w.Tokens)
+ }
+
+ tail, _ := tb.Tail()
+
+ gotSeen := tail.Seen
+
+ if tt.config.PadLeft {
+ gotSeen += tail.PaddingLeft
+ }
+
+ if gotSeen != tt.seen {
+ t.Fatalf("expected seen %d but got %d", tt.seen, gotSeen)
+ }
+
+ got = append(got, tail.Tokens)
+
+ if len(got) != len(tt.expected) {
+ t.Fatalf("expected %d windows but got %d", len(tt.expected), len(got))
+ }
+
+ for i := range got {
+ if !slices.Equal(got[i], tt.expected[i]) {
+ t.Errorf("window %d: expected %v but got %v", i, tt.expected[i], got[i])
+ }
+ }
+ },
+ )
+ }
+}