summaryrefslogtreecommitdiff
path: root/research/lesci/experiment.go
blob: 53db93193c64f6ed9f3d59bbad3589e51f46f329 (plain)
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
}