diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-06-07 20:26:59 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-06-07 20:26:59 +0200 |
| commit | 7ae0651a55ace3c5d1f651775c1160f6e0b72152 (patch) | |
| tree | 577d906ae8f9c94ca6e435469db19dcbb12f2943 /research | |
| parent | 027bce960feaec195b633b6195c9ff5cbb9f43f9 (diff) | |
Add config to lesci experiment
Diffstat (limited to 'research')
| -rw-r--r-- | research/lesci/cmd/lesci/main.go | 11 | ||||
| -rw-r--r-- | research/lesci/config.go | 16 | ||||
| -rw-r--r-- | research/lesci/context.go | 13 | ||||
| -rw-r--r-- | research/lesci/experiment.go | 6 | ||||
| -rw-r--r-- | research/lesci/extract.go | 10 | ||||
| -rw-r--r-- | research/lesci/lesci.go | 8 |
6 files changed, 57 insertions, 7 deletions
diff --git a/research/lesci/cmd/lesci/main.go b/research/lesci/cmd/lesci/main.go index 6f72283..bfabfe8 100644 --- a/research/lesci/cmd/lesci/main.go +++ b/research/lesci/cmd/lesci/main.go @@ -54,7 +54,16 @@ func main() { o := path.Join(shelf.Abs(shelf.Item(*outPath)), path.Base(shelf.Abs(shelf.Item(modelPath))), *chkpt) - e := must(lesci.NewExperiment(m, t, c, d, o, 50256, 5000)) + exptOpts := lesci.Options{ + ForceContext: false, + ForceExtract: true, + } + + expCfg := lesci.ConfigDefault() + + expCfg.ClampRulesBeforeFilter = true + + e := must(lesci.NewExperiment(m, t, c, d, o, 50256, 5000, expCfg, exptOpts)) if err := m.Init(); err != nil { log.Fatal(err) diff --git a/research/lesci/config.go b/research/lesci/config.go new file mode 100644 index 0000000..794de0c --- /dev/null +++ b/research/lesci/config.go @@ -0,0 +1,16 @@ +package lesci + +type Config struct { + ClampRulesBeforeFilter bool +} + +type Options struct { + ForceContext bool + ForceExtract bool +} + +func ConfigDefault() Config { + return Config{ + ClampRulesBeforeFilter: true, + } +} diff --git a/research/lesci/context.go b/research/lesci/context.go index df9b059..8496507 100644 --- a/research/lesci/context.go +++ b/research/lesci/context.go @@ -1,6 +1,7 @@ package lesci import ( + "context" "database/sql" "database/sql/driver" "fmt" @@ -39,7 +40,17 @@ func (e *Experiment) BuildContext(db *sql.DB) error { } else if !ok { fmt.Println("context table not empty") - return nil + 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 AppendRows(db, "context", func(append AppendFunc) error { diff --git a/research/lesci/experiment.go b/research/lesci/experiment.go index df30f2c..e625301 100644 --- a/research/lesci/experiment.go +++ b/research/lesci/experiment.go @@ -21,9 +21,11 @@ type Experiment struct { name string cutoff int window int + config Config + options Options } -func NewExperiment(model llm.Causal, tokenizer, counterfactual llm.Tokenizer, data dataset.Reader, name string, cutoff, window int) (*Experiment, error) { +func NewExperiment(model llm.Causal, tokenizer, counterfactual llm.Tokenizer, data dataset.Reader, name string, cutoff, window int, cfg Config, opts Options) (*Experiment, error) { e := &Experiment{ model: model, tokenizer: tokenizer, @@ -32,6 +34,8 @@ func NewExperiment(model llm.Causal, tokenizer, counterfactual llm.Tokenizer, da name: name, cutoff: cutoff, window: window, + config: cfg, + options: opts, } return e, nil diff --git a/research/lesci/extract.go b/research/lesci/extract.go index 659cdc3..06c6bd4 100644 --- a/research/lesci/extract.go +++ b/research/lesci/extract.go @@ -17,6 +17,14 @@ func (e *Experiment) ExtractData(db *sql.DB) error { } else if !ok { fmt.Println("oov_rules table not empty") + if !e.options.ForceExtract { + fmt.Println("skipping rule extraction") + + return nil + } + + fmt.Println("clearing oov_rules") + if _, err := db.ExecContext(context.Background(), `DELETE FROM oov_rules`); err != nil { return err } @@ -26,7 +34,7 @@ func (e *Experiment) ExtractData(db *sql.DB) error { rules, valid := Rules(e.counterfactual, merges) - mask := ExtractData(rules, valid, int64(e.cutoff), int64(e.window)) + mask := ExtractData(rules, valid, int64(e.cutoff), int64(e.window), e.config.ClampRulesBeforeFilter) return AppendRows(db, "oov_rules", func(append AppendFunc) error { for i, m := range mask { diff --git a/research/lesci/lesci.go b/research/lesci/lesci.go index 8bae6c6..871cf8f 100644 --- a/research/lesci/lesci.go +++ b/research/lesci/lesci.go @@ -9,16 +9,18 @@ import ( // ExtractData // // https://github.com/pietrolesci/tokenisation-bias/blob/376abc0ed6924986cbaf696ea10fdda71e550e45/notebooks/01_extract_data.ipynb -func ExtractData(rules tensor.Dense[int64], valid []bool, cutoff, window int64) []bool { +func ExtractData(rules tensor.Dense[int64], valid []bool, cutoff, window int64, clampRulesBeforeFilter bool) []bool { shape := rules.Shape() if len(shape) != 2 || shape[0] != len(valid) || shape[1] != 3 { panic("shape mismatch") } - clamped := Window(rules, valid, cutoff, window) + clamped := valid // collect everything for now - // clamped := valid // collect everything for now + if clampRulesBeforeFilter { + clamped = Window(rules, valid, cutoff, window) + } filtered := Filter(rules, clamped, cutoff) oov := OutOfVocab(rules, filtered, cutoff) |
