summaryrefslogtreecommitdiff
path: root/research/lesci
diff options
context:
space:
mode:
Diffstat (limited to 'research/lesci')
-rw-r--r--research/lesci/context.go14
-rw-r--r--research/lesci/experiment.go38
2 files changed, 32 insertions, 20 deletions
diff --git a/research/lesci/context.go b/research/lesci/context.go
index 7f98946..8201e6f 100644
--- a/research/lesci/context.go
+++ b/research/lesci/context.go
@@ -84,27 +84,19 @@ func (e *Experiment) BuildContext(db *sql.DB) error {
document: job.Document,
token: token,
value: logProbs[i],
- offset: job.Position*512 + job.Seen + i, // TODO refactor
+ offset: job.Position*e.tokenBufferConfig.Stride + job.Seen + i,
}
}
return r
- }, llm.EvaluatorConfig{
- BatchSize: 32,
- NumWorkers: 64,
- })
+ }, e.evaluatorConfig)
if err := e.ensureContext(db); err != nil {
return err
}
- cfg := llm.TokenBufferConfig{
- Window: 1024,
- Stride: 512,
- }
-
return AppendRows(db, "context", func(append AppendFunc) error {
- return eval.RunAndCollect("Context", e.data, cfg, func(r []logProb) error {
+ return eval.RunAndCollect("Context", e.data, e.tokenBufferConfig, func(r []logProb) error {
for _, l := range r {
if err := append([]driver.Value{l.document, l.token, l.value, l.offset}); err != nil {
return err
diff --git a/research/lesci/experiment.go b/research/lesci/experiment.go
index d01d405..53db931 100644
--- a/research/lesci/experiment.go
+++ b/research/lesci/experiment.go
@@ -14,15 +14,17 @@ import (
)
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
+ 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) {
@@ -38,9 +40,27 @@ func NewExperiment(model llm.Causal, tokenizer, counterfactual llm.Tokenizer, da
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)