diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-06-27 19:08:39 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-06-27 19:08:39 +0200 |
| commit | e1e4eab2e1a682d633125f3c71cc09c3e95b30d3 (patch) | |
| tree | 92f107c4907f09b0f898351830286e981db87893 /llm | |
| parent | 416fa02619a2111ab3184f083f513bd263fc7cf2 (diff) | |
Add token buffer config
Diffstat (limited to 'llm')
| -rw-r--r-- | llm/cmd/eval/perplexity.go | 7 | ||||
| -rw-r--r-- | llm/evaluator.go | 4 | ||||
| -rw-r--r-- | llm/run.go | 4 | ||||
| -rw-r--r-- | llm/tokenbuffer.go | 15 | ||||
| -rw-r--r-- | llm/tokenbuffer_test.go | 12 |
5 files changed, 26 insertions, 16 deletions
diff --git a/llm/cmd/eval/perplexity.go b/llm/cmd/eval/perplexity.go index ba92c99..dd4dae3 100644 --- a/llm/cmd/eval/perplexity.go +++ b/llm/cmd/eval/perplexity.go @@ -42,7 +42,12 @@ func perplexity() { total := float64(0) n := 0 - if err := e.RunAndCollect("Perplexity", d, 1024, 512, func(r pplResult) error { + cfg := llm.TokenBufferConfig{ + Window: 1024, + Stride: 512, + } + + if err := e.RunAndCollect("Perplexity", d, cfg, func(r pplResult) error { total += r.v n += r.n diff --git a/llm/evaluator.go b/llm/evaluator.go index 8aa4a32..69c4202 100644 --- a/llm/evaluator.go +++ b/llm/evaluator.go @@ -45,7 +45,7 @@ func (e *Evaluator[R]) Results() chan R { return e.results } -func (e *Evaluator[R]) RunAndCollect(title string, data dataset.Reader, window, stride int, callback func(R) error) error { +func (e *Evaluator[R]) RunAndCollect(title string, data dataset.Reader, cfg TokenBufferConfig, callback func(R) error) error { var wg sync.WaitGroup var collectErr error @@ -62,7 +62,7 @@ func (e *Evaluator[R]) RunAndCollect(title string, data dataset.Reader, window, } }() - runErr := e.Run(title, data, window, stride) + runErr := e.Run(title, data, cfg) wg.Wait() @@ -8,7 +8,7 @@ import ( "go.jknobloc.com/x/tui" ) -func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int) error { +func (e *Evaluator[R]) Run(title string, data dataset.Reader, cfg TokenBufferConfig) error { devices := make([]int, len(e.models)) for i := range len(devices) { @@ -47,7 +47,7 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int defer pb.Close() - tb := NewTokenBuffer(e.tokenizer, window, stride) + tb := NewTokenBuffer(e.tokenizer, cfg) tb.SetIncludeTail(false) diff --git a/llm/tokenbuffer.go b/llm/tokenbuffer.go index d19de4a..e2d8a91 100644 --- a/llm/tokenbuffer.go +++ b/llm/tokenbuffer.go @@ -12,22 +12,27 @@ type TokenBuffer struct { includeTail bool } +type TokenBufferConfig struct { + Window int + Stride int +} + type TokenWindow struct { Document int Tokens []int64 Seen int } -func NewTokenBuffer(tokenizer Tokenizer, window, stride int) *TokenBuffer { - if stride > window { +func NewTokenBuffer(tokenizer Tokenizer, cfg TokenBufferConfig) *TokenBuffer { + if cfg.Stride > cfg.Window { panic("stride exceeds window") } return &TokenBuffer{ tokenizer: tokenizer, - window: window, - stride: stride, - buffer: make([]int64, 0, 2*window), + window: cfg.Window, + stride: cfg.Stride, + buffer: make([]int64, 0, 2*cfg.Window), document: -1, position: 0, includeTail: true, diff --git a/llm/tokenbuffer_test.go b/llm/tokenbuffer_test.go index 789caa6..c07c371 100644 --- a/llm/tokenbuffer_test.go +++ b/llm/tokenbuffer_test.go @@ -49,7 +49,7 @@ func TestTokenBuffer_Push(t *testing.T) { fmt.Sprintf("window%d_stride%d", tt.window, tt.stride), func(t *testing.T) { - tb := NewTokenBuffer(byteTokenizer{}, tt.window, tt.stride) + tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: tt.window, Stride: tt.stride}) tb.SetIncludeTail(false) @@ -74,7 +74,7 @@ func TestTokenBuffer_Push(t *testing.T) { } func TestTokenBuffer_PushAccumulates(t *testing.T) { - tb := NewTokenBuffer(byteTokenizer{}, 4, 4) + tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: 4, Stride: 4}) var got [][]int64 @@ -104,7 +104,7 @@ func TestTokenBuffer_PushAccumulates(t *testing.T) { } func TestTokenBuffer_Tail(t *testing.T) { - tb := NewTokenBuffer(byteTokenizer{}, 5, 5) + tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: 5, Stride: 5}) tb.SetIncludeTail(true) @@ -132,7 +132,7 @@ func TestTokenBuffer_Tail(t *testing.T) { } func TestTokenBuffer_Position(t *testing.T) { - tb := NewTokenBuffer(byteTokenizer{}, 4, 2) + tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: 4, Stride: 2}) tb.SetIncludeTail(false) @@ -162,7 +162,7 @@ func TestTokenBuffer_Position(t *testing.T) { } func TestTokenBuffer_DocumentBoundaryYieldsTail(t *testing.T) { - tb := NewTokenBuffer(byteTokenizer{}, 10, 10) + tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: 10, Stride: 10}) tb.SetIncludeTail(true) @@ -194,7 +194,7 @@ func TestTokenBuffer_DocumentBoundaryYieldsTail(t *testing.T) { } func TestTokenBuffer_TailPosition(t *testing.T) { - tb := NewTokenBuffer(byteTokenizer{}, 4, 2) + tb := NewTokenBuffer(byteTokenizer{}, TokenBufferConfig{Window: 4, Stride: 2}) tb.SetIncludeTail(true) |
