From c87625c71737f41f42f0188ff87b8d9313c8548a Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Sat, 27 Jun 2026 20:24:38 +0200 Subject: Pad token buffer --- llm/cmd/eval/perplexity.go | 8 +++-- llm/job.go | 10 +++--- llm/run.go | 21 ++++++++----- llm/tokenbuffer.go | 37 +++++++++++++++++++--- llm/tokenbuffer_test.go | 78 ++++++++++++++++++++++++++++++++++++++++++++++ 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]) + } + } + }, + ) + } +} -- cgit v1.2.3