diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-06-28 01:28:55 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-07-13 16:31:59 +0200 |
| commit | 0030b589a9752a841ea9fa4dda7b8208b7072bcf (patch) | |
| tree | acf675de38582a36a7b9b398b2f0add6a20356c1 /research/lesci/context.go | |
| parent | 48cac2575c9ba5377ddeba2f17ddb68277e410d9 (diff) | |
Add flags to import context
Diffstat (limited to 'research/lesci/context.go')
| -rw-r--r-- | research/lesci/context.go | 75 |
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{ |
