1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
|
package lesci
import (
"database/sql"
"fmt"
"log"
"os"
"path/filepath"
_ "github.com/duckdb/duckdb-go/v2"
"go.jknobloc.com/x/dataset"
"go.jknobloc.com/x/llm"
)
type Experiment struct {
model llm.Causal
tokenizer llm.Tokenizer
counterfactual llm.Tokenizer
data dataset.Reader
name string
cutoff int
window int
config Config
options Options
tokenBufferConfig llm.TokenBufferConfig
evaluatorConfig llm.EvaluatorConfig
}
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,
counterfactual: counterfactual,
data: data,
name: name,
cutoff: cutoff,
window: window,
config: cfg,
options: opts,
}
e.tokenBufferConfig = llm.TokenBufferConfig{
Window: 1024,
Stride: 512,
}
e.evaluatorConfig = llm.EvaluatorConfig{
BatchSize: 32,
NumWorkers: 64,
}
return e, nil
}
func (e *Experiment) SetTokenBufferConfig(cfg llm.TokenBufferConfig) {
e.tokenBufferConfig = cfg
}
func (e *Experiment) SetEvaluatorConfig(cfg llm.EvaluatorConfig) {
e.evaluatorConfig = cfg
}
func (e *Experiment) Run() error {
if err := os.MkdirAll(e.name, 0775); err != nil {
log.Fatal(err)
}
dsn := filepath.Join(e.name, "lesci.db")
var db *sql.DB
if database, err := initDatabase(dsn); err != nil {
return err
} else {
db = database
}
defer db.Close()
db.SetMaxOpenConns(1)
fmt.Println(e.name)
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 {
return err
}
if err := e.Peek(db, 20); err != nil {
return err
}
if err := e.Analyze(db); err != nil {
return err
}
if err := e.Plot(db); err != nil {
return err
}
return nil
}
func initDatabase(dsn string) (*sql.DB, error) {
var db *sql.DB
if database, err := sql.Open("duckdb", dsn); err != nil {
return nil, err
} else {
db = database
}
if err := db.Ping(); err != nil {
_ = db.Close()
return nil, err
}
return db, nil
}
|