package lesci import ( "context" "database/sql" "database/sql/driver" "fmt" "go.jknobloc.com/x/llm" ) type logProb struct { document int token int value float32 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)) for i, token := range tokens { r[i] = logProb{ document: job.Document, token: token, value: logProbs[i], offset: job.Position*512 + job.Seen + i, // TODO refactor } } return r }, llm.EvaluatorConfig{ BatchSize: 32, NumWorkers: 64, }) if err := e.ensureContext(db); err != nil { return err } cfg := llm.TokenBufferConfig{ Window: 1024, Stride: 512, } return AppendRows(db, "context", func(append AppendFunc) 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 } } return nil }) }) }