summaryrefslogtreecommitdiff
path: root/research/entropy/context.go
diff options
context:
space:
mode:
Diffstat (limited to 'research/entropy/context.go')
-rw-r--r--research/entropy/context.go53
1 files changed, 53 insertions, 0 deletions
diff --git a/research/entropy/context.go b/research/entropy/context.go
new file mode 100644
index 0000000..3993c03
--- /dev/null
+++ b/research/entropy/context.go
@@ -0,0 +1,53 @@
+package entropy
+
+import (
+ "fmt"
+
+ "go.jknobloc.com/x/dataset"
+ "go.jknobloc.com/x/llm"
+)
+
+type logProb struct {
+ document int
+ token int
+ value float32
+ offset int
+}
+
+func Context(model llm.Causal, tokenizer llm.Tokenizer, data dataset.Reader) error {
+ evaluatorConfig := llm.EvaluatorConfig{
+ BatchSize: 32,
+ NumWorkers: 16,
+ }
+
+ tokenBufferConfig := llm.TokenBufferConfig{
+ Window: 1024,
+ Stride: 512,
+ PadLeft: false,
+ PadRight: false,
+ PadTokenID: 256,
+ }
+
+ eval := llm.NewEvaluator(model, tokenizer, func(job llm.Job, logProbs []float32, tokens []int) []logProb {
+ r := make([]logProb, len(tokens))
+
+ for i, token := range tokens {
+ r[i] = logProb{
+ document: job.Document,
+ token: token,
+ value: logProbs[i],
+ offset: job.Position*tokenBufferConfig.Stride + job.Seen + i,
+ }
+ }
+
+ return r
+ }, evaluatorConfig)
+
+ return eval.RunAndCollect("Context", data, tokenBufferConfig, func(r []logProb) error {
+ for _, l := range r {
+ fmt.Println(l) // TODO implement
+ }
+
+ return nil
+ })
+}