summaryrefslogtreecommitdiff
path: root/research
diff options
context:
space:
mode:
Diffstat (limited to 'research')
-rw-r--r--research/lesci/cmd/lesci/main.go11
-rw-r--r--research/lesci/config.go16
-rw-r--r--research/lesci/context.go13
-rw-r--r--research/lesci/experiment.go6
-rw-r--r--research/lesci/extract.go10
-rw-r--r--research/lesci/lesci.go8
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)