summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-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
-rw-r--r--research/lesci/context.go7
6 files changed, 32 insertions, 17 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)
diff --git a/research/lesci/context.go b/research/lesci/context.go
index 8496507..378b4e5 100644
--- a/research/lesci/context.go
+++ b/research/lesci/context.go
@@ -53,8 +53,13 @@ func (e *Experiment) BuildContext(db *sql.DB) error {
}
}
+ cfg := llm.TokenBufferConfig{
+ Window: 1024,
+ Stride: 512,
+ }
+
return AppendRows(db, "context", func(append AppendFunc) error {
- return eval.RunAndCollect("Context", e.data, 1024, 512, func(r []logProb) error {
+ return eval.RunAndCollect("Context", e.data, cfg, func(r []logProb) error {
for _, l := range r {
if err := append([]driver.Value{l.document, l.token, l.value, l.offset}); err != nil {
return err