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/run.go | 21 ++++++++++++++------- 1 file changed, 14 insertions(+), 7 deletions(-) (limited to 'llm/run.go') 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) -- cgit v1.3.1