summaryrefslogtreecommitdiff
path: root/research/lesci/context.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-06-28 01:28:55 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-07-13 16:31:59 +0200
commit0030b589a9752a841ea9fa4dda7b8208b7072bcf (patch)
treeacf675de38582a36a7b9b398b2f0add6a20356c1 /research/lesci/context.go
parent48cac2575c9ba5377ddeba2f17ddb68277e410d9 (diff)
Add flags to import context
Diffstat (limited to 'research/lesci/context.go')
-rw-r--r--research/lesci/context.go75
1 files changed, 60 insertions, 15 deletions
diff --git a/research/lesci/context.go b/research/lesci/context.go
index 378b4e5..7f98946 100644
--- a/research/lesci/context.go
+++ b/research/lesci/context.go
@@ -16,6 +16,65 @@ type logProb struct {
offset int
}
+var (
+ sqlInsertContext = `
+ INSERT INTO context
+ SELECT uid::INTEGER,
+ token_ids[i]::INTEGER,
+ token_logprob[i]::FLOAT,
+ (i - 1)::INTEGER
+ FROM (SELECT uid, token_ids, token_logprob FROM read_parquet(?) WHERE step = ?) d,
+ UNNEST(range(1, len(token_ids) + 1)) AS t(i)`
+)
+
+func (e *Experiment) ensureContext(db *sql.DB) error {
+ if ok, err := EnsureTable(db, "context", `CREATE TABLE context(uid INTEGER, token INTEGER, logprob FLOAT, pos INTEGER)`); err != nil {
+ return err
+ } else if !ok {
+ fmt.Println("context table not empty")
+
+ if !e.options.ForceContext {
+ fmt.Println("skipping context collection")
+
+ return nil
+ }
+
+ fmt.Println("clearing context")
+
+ if _, err := db.ExecContext(context.Background(), `DELETE FROM context`); err != nil {
+ return err
+ }
+ }
+
+ return nil
+}
+
+func (e *Experiment) ImportContext(db *sql.DB, name string, step int) error {
+ if err := e.ensureContext(db); err != nil {
+ return err
+ }
+
+ if res, err := db.ExecContext(context.Background(), sqlInsertContext, name, step); err != nil {
+ return err
+ } else if n, err := res.RowsAffected(); err == nil {
+ fmt.Printf("imported %d context rows from %s (step %d)\n", n, name, step)
+ }
+
+ if res, err := db.ExecContext(context.Background(), `DELETE FROM context WHERE token = 0 AND logprob = 0`); err != nil {
+ return err
+ } else if n, err := res.RowsAffected(); err == nil {
+ fmt.Printf("dropped %d leading pad rows\n", n)
+ }
+
+ // if res, err := db.ExecContext(context.Background(), `DELETE FROM context WHERE pos >= 2048`); err != nil {
+ // return err
+ // } else if n, err := res.RowsAffected(); err == nil {
+ // fmt.Printf("dropped %d out-of-range rows\n", n)
+ // }
+
+ return nil
+}
+
func (e *Experiment) BuildContext(db *sql.DB) error {
eval := llm.NewEvaluator(e.model, e.tokenizer, func(job llm.Job, logProbs []float32, tokens []int) []logProb {
r := make([]logProb, len(tokens))
@@ -35,22 +94,8 @@ func (e *Experiment) BuildContext(db *sql.DB) error {
NumWorkers: 64,
})
- if ok, err := EnsureTable(db, "context", `CREATE TABLE context(uid INTEGER, token INTEGER, logprob FLOAT, pos INTEGER)`); err != nil {
+ if err := e.ensureContext(db); err != nil {
return err
- } else if !ok {
- fmt.Println("context table not empty")
-
- if !e.options.ForceContext {
- fmt.Println("skipping context collection")
-
- return nil
- }
-
- fmt.Println("clearing context")
-
- if _, err := db.ExecContext(context.Background(), `DELETE FROM context`); err != nil {
- return err
- }
}
cfg := llm.TokenBufferConfig{