diff options
| -rw-r--r-- | research/lesci/cmd/lesci/main.go | 9 | ||||
| -rw-r--r-- | research/lesci/config.go | 2 | ||||
| -rw-r--r-- | research/lesci/context.go | 75 | ||||
| -rw-r--r-- | research/lesci/experiment.go | 10 |
4 files changed, 77 insertions, 19 deletions
diff --git a/research/lesci/cmd/lesci/main.go b/research/lesci/cmd/lesci/main.go index bfabfe8..e9c9a99 100644 --- a/research/lesci/cmd/lesci/main.go +++ b/research/lesci/cmd/lesci/main.go @@ -20,6 +20,9 @@ func main() { outPath := flag.String("o", "", "") + goldData := flag.String("gold-data", "", "gold data") + goldStep := flag.Int("gold-step", 0, "gold step") + dry := flag.Bool("dry", false, "") flag.Parse() @@ -54,16 +57,18 @@ func main() { o := path.Join(shelf.Abs(shelf.Item(*outPath)), path.Base(shelf.Abs(shelf.Item(modelPath))), *chkpt) - exptOpts := lesci.Options{ + expOpts := lesci.Options{ ForceContext: false, ForceExtract: true, + GoldDataName: *goldData, + GoldDataStep: *goldStep, } expCfg := lesci.ConfigDefault() expCfg.ClampRulesBeforeFilter = true - e := must(lesci.NewExperiment(m, t, c, d, o, 50256, 5000, expCfg, exptOpts)) + e := must(lesci.NewExperiment(m, t, c, d, o, 50256, 5000, expCfg, expOpts)) if err := m.Init(); err != nil { log.Fatal(err) diff --git a/research/lesci/config.go b/research/lesci/config.go index 794de0c..f20c0e3 100644 --- a/research/lesci/config.go +++ b/research/lesci/config.go @@ -7,6 +7,8 @@ type Config struct { type Options struct { ForceContext bool ForceExtract bool + GoldDataName string + GoldDataStep int } func ConfigDefault() Config { 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{ diff --git a/research/lesci/experiment.go b/research/lesci/experiment.go index e625301..d01d405 100644 --- a/research/lesci/experiment.go +++ b/research/lesci/experiment.go @@ -62,8 +62,14 @@ func (e *Experiment) Run() error { fmt.Println(e.name) - if err := e.BuildContext(db); err != nil { - return err + if e.options.GoldDataName != "" { + if err := e.ImportContext(db, e.options.GoldDataName, e.options.GoldDataStep); err != nil { + return err + } + } else { + if err := e.BuildContext(db); err != nil { + return err + } } if err := e.ExtractData(db); err != nil { |
