summaryrefslogtreecommitdiff
path: root/llm
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-06-27 19:08:39 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-06-27 19:08:39 +0200
commite1e4eab2e1a682d633125f3c71cc09c3e95b30d3 (patch)
tree92f107c4907f09b0f898351830286e981db87893 /llm
parent416fa02619a2111ab3184f083f513bd263fc7cf2 (diff)
Add token buffer config
Diffstat (limited to 'llm')
-rw-r--r--llm/cmd/eval/perplexity.go7
-rw-r--r--llm/evaluator.go4
-rw-r--r--llm/run.go4
-rw-r--r--llm/tokenbuffer.go15
-rw-r--r--llm/tokenbuffer_test.go12
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()
diff --git a/llm/run.go b/llm/run.go
index abb4c0d..3be1cb5 100644
--- a/llm/run.go
+++ b/llm/run.go
@@ -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)