summaryrefslogtreecommitdiff
path: root/research/lesci
diff options
context:
space:
mode:
Diffstat (limited to 'research/lesci')
-rw-r--r--research/lesci/cmd/lesci/main.go9
-rw-r--r--research/lesci/config.go2
-rw-r--r--research/lesci/context.go75
-rw-r--r--research/lesci/experiment.go10
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 {