summaryrefslogtreecommitdiff
path: root/research
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-06-03 18:04:27 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-06-03 18:04:27 +0200
commitdd984a2b5beb8bc48e32242e1cfabbfbc5353086 (patch)
tree622c49d0f4d99439ea17f794c583f37130427527 /research
parentc7e991d05aa2c9c1fb1ef5b6e39f721d5bec36e3 (diff)
Print counterfactual merge examples
Diffstat (limited to 'research')
-rw-r--r--research/lesci/experiment.go4
-rw-r--r--research/lesci/peek.go63
2 files changed, 67 insertions, 0 deletions
diff --git a/research/lesci/experiment.go b/research/lesci/experiment.go
index 4184199..df30f2c 100644
--- a/research/lesci/experiment.go
+++ b/research/lesci/experiment.go
@@ -66,6 +66,10 @@ func (e *Experiment) Run() error {
return err
}
+ if err := e.Peek(db, 20); err != nil {
+ return err
+ }
+
if err := e.Analyze(db); err != nil {
return err
}
diff --git a/research/lesci/peek.go b/research/lesci/peek.go
new file mode 100644
index 0000000..b98eda5
--- /dev/null
+++ b/research/lesci/peek.go
@@ -0,0 +1,63 @@
+package lesci
+
+import (
+ "context"
+ "database/sql"
+ "fmt"
+
+ "go.jknobloc.com/x/tokenizer/bpe"
+)
+
+var sqlPeek = `
+ WITH pairs AS (
+ SELECT uid, pos, token, logprob,
+ LEAD(token, 1) OVER w AS next_tok
+ FROM context
+ WINDOW w AS (PARTITION BY uid ORDER BY pos)
+ ),
+ matched AS (
+ SELECT p.*, r.c AS merged_as_first
+ FROM pairs p
+ JOIN oov_rules r ON p.token = r.a AND p.next_tok = r.b
+ WHERE r.c >= ?
+ )
+ SELECT
+ m1.token AS tok_a,
+ m1.logprob AS logprob_a,
+ m1.next_tok AS tok_b,
+ m2.logprob AS logprob_b,
+ m1.merged_as_first AS tok_c
+ FROM matched m1
+ JOIN context m2 ON m1.uid = m2.uid AND m1.pos + 1 = m2.pos
+ LIMIT ?`
+
+func (e *Experiment) Peek(db *sql.DB, n int) error {
+ vocab := bpe.Vocab(e.counterfactual.(*bpe.Tokenizer))
+
+ var rows *sql.Rows
+
+ if r, err := db.QueryContext(context.Background(), sqlPeek, e.cutoff, n); err != nil {
+ return err
+ } else {
+ rows = r
+ }
+
+ defer rows.Close()
+
+ for rows.Next() {
+ var idA, idB, idC int
+ var lpA, lpB float32
+
+ if err := rows.Scan(&idA, &lpA, &idB, &lpB, &idC); err != nil {
+ return err
+ }
+
+ tokA := vocab[idA]
+ tokB := vocab[idB]
+ tokC := vocab[idC]
+
+ fmt.Printf("%-20q + %-20q -> %-20q (%.4f + %.4f = %.4f)\n", tokA, tokB, tokC, lpA, lpB, lpA+lpB)
+ }
+
+ return rows.Err()
+}